ratelimit.go 1.3 KB

1234567891011121314151617181920212223242526272829303132333435363738394041424344454647484950515253545556575859606162636465
  1. // Package middleware
  2. package middleware
  3. import (
  4. "net/http"
  5. "sync"
  6. "time"
  7. "golang.org/x/time/rate"
  8. )
  9. type RateLimiter struct {
  10. visitors map[string]*rate.Limiter
  11. mu sync.Mutex
  12. rate rate.Limit
  13. burst int
  14. lastSeen map[string]time.Time
  15. }
  16. func NewRateLimiter(r rate.Limit, burst int) *RateLimiter {
  17. rl := &RateLimiter{
  18. visitors: make(map[string]*rate.Limiter),
  19. rate: r,
  20. burst: burst,
  21. lastSeen: make(map[string]time.Time),
  22. }
  23. go rl.cleanup()
  24. return rl
  25. }
  26. func (rl *RateLimiter) cleanup() {
  27. for {
  28. time.Sleep(10 * time.Minute)
  29. rl.mu.Lock()
  30. for ip, last := range rl.lastSeen {
  31. if time.Since(last) > 30*time.Minute {
  32. delete(rl.visitors, ip)
  33. delete(rl.lastSeen, ip)
  34. }
  35. }
  36. rl.mu.Unlock()
  37. }
  38. }
  39. func (rl *RateLimiter) Middleware() func(http.Handler) http.Handler {
  40. return func(next http.Handler) http.Handler {
  41. return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
  42. ip := r.RemoteAddr
  43. rl.mu.Lock()
  44. limiter, ok := rl.visitors[ip]
  45. if !ok {
  46. limiter = rate.NewLimiter(rl.rate, rl.burst)
  47. rl.visitors[ip] = limiter
  48. }
  49. rl.lastSeen[ip] = time.Now()
  50. rl.mu.Unlock()
  51. if !limiter.Allow() {
  52. http.Error(w, `{"error":"rate limit exceeded"}`, http.StatusTooManyRequests)
  53. return
  54. }
  55. next.ServeHTTP(w, r)
  56. })
  57. }
  58. }