ratelimit.go 1.5 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778
  1. // Package middleware
  2. package middleware
  3. import (
  4. "context"
  5. "net/http"
  6. "sync"
  7. "time"
  8. "golang.org/x/time/rate"
  9. )
  10. type RateLimiter struct {
  11. visitors map[string]*rate.Limiter
  12. mu sync.Mutex
  13. rate rate.Limit
  14. burst int
  15. lastSeen map[string]time.Time
  16. stopCh chan struct{}
  17. }
  18. func NewRateLimiter(r rate.Limit, burst int) *RateLimiter {
  19. rl := &RateLimiter{
  20. visitors: make(map[string]*rate.Limiter),
  21. rate: r,
  22. burst: burst,
  23. lastSeen: make(map[string]time.Time),
  24. stopCh: make(chan struct{}),
  25. }
  26. go rl.cleanup()
  27. return rl
  28. }
  29. func (rl *RateLimiter) Stop() {
  30. close(rl.stopCh)
  31. }
  32. func (rl *RateLimiter) cleanup() {
  33. ticker := time.NewTicker(10 * time.Minute)
  34. defer ticker.Stop()
  35. for {
  36. select {
  37. case <-rl.stopCh:
  38. return
  39. case <-ticker.C:
  40. rl.mu.Lock()
  41. for ip, last := range rl.lastSeen {
  42. if time.Since(last) > 30*time.Minute {
  43. delete(rl.visitors, ip)
  44. delete(rl.lastSeen, ip)
  45. }
  46. }
  47. rl.mu.Unlock()
  48. }
  49. }
  50. }
  51. func (rl *RateLimiter) Middleware() func(http.Handler) http.Handler {
  52. return func(next http.Handler) http.Handler {
  53. return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
  54. ip := GetClientIP(r)
  55. rl.mu.Lock()
  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. rl.mu.Unlock()
  63. if !limiter.Allow() {
  64. http.Error(w, `{"error":"rate limit exceeded"}`, http.StatusTooManyRequests)
  65. return
  66. }
  67. next.ServeHTTP(w, r)
  68. })
  69. }
  70. }