ratelimit.go 1.5 KB

1234567891011121314151617181920212223242526272829303132333435363738394041424344454647484950515253545556575859606162636465666768697071727374757677
  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. stopCh chan struct{}
  16. }
  17. func NewRateLimiter(r rate.Limit, burst int) *RateLimiter {
  18. rl := &RateLimiter{
  19. visitors: make(map[string]*rate.Limiter),
  20. rate: r,
  21. burst: burst,
  22. lastSeen: make(map[string]time.Time),
  23. stopCh: make(chan struct{}),
  24. }
  25. go rl.cleanup()
  26. return rl
  27. }
  28. func (rl *RateLimiter) Stop() {
  29. close(rl.stopCh)
  30. }
  31. func (rl *RateLimiter) cleanup() {
  32. ticker := time.NewTicker(10 * time.Minute)
  33. defer ticker.Stop()
  34. for {
  35. select {
  36. case <-rl.stopCh:
  37. return
  38. case <-ticker.C:
  39. rl.mu.Lock()
  40. for ip, last := range rl.lastSeen {
  41. if time.Since(last) > 30*time.Minute {
  42. delete(rl.visitors, ip)
  43. delete(rl.lastSeen, ip)
  44. }
  45. }
  46. rl.mu.Unlock()
  47. }
  48. }
  49. }
  50. func (rl *RateLimiter) Middleware() func(http.Handler) http.Handler {
  51. return func(next http.Handler) http.Handler {
  52. return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
  53. ip := GetClientIP(r)
  54. rl.mu.Lock()
  55. defer rl.mu.Unlock()
  56. limiter, ok := rl.visitors[ip]
  57. if !ok {
  58. limiter = rate.NewLimiter(rl.rate, rl.burst)
  59. rl.visitors[ip] = limiter
  60. }
  61. rl.lastSeen[ip] = time.Now()
  62. if !limiter.Allow() {
  63. http.Error(w, `{"error":"rate limit exceeded"}`, http.StatusTooManyRequests)
  64. return
  65. }
  66. next.ServeHTTP(w, r)
  67. })
  68. }
  69. }