ratelimit.go
| 1 | // Package ratelimit implements a fixed-window counter per key. |
| 2 | package ratelimit |
| 3 | |
| 4 | import ( |
| 5 | "net" |
| 6 | "net/http" |
| 7 | "strings" |
| 8 | "sync" |
| 9 | "time" |
| 10 | ) |
| 11 | |
| 12 | type bucket struct { |
| 13 | count int |
| 14 | resetAt time.Time |
| 15 | } |
| 16 | |
| 17 | // Limiter allows at most max events per key within a window. |
| 18 | // The window starts at the first event and is not extended by later ones. |
| 19 | type Limiter struct { |
| 20 | max int |
| 21 | window time.Duration |
| 22 | |
| 23 | mu sync.Mutex |
| 24 | buckets map[string]*bucket |
| 25 | nextSweep time.Time |
| 26 | } |
| 27 | |
| 28 | const sweepInterval = time.Minute |
| 29 | |
| 30 | func New(max int, window time.Duration) *Limiter { |
| 31 | return &Limiter{max: max, window: window, buckets: map[string]*bucket{}} |
| 32 | } |
| 33 | |
| 34 | // Allow records an event for key and reports whether it is within the limit. |
| 35 | func (l *Limiter) Allow(key string) bool { |
| 36 | now := time.Now() |
| 37 | l.mu.Lock() |
| 38 | defer l.mu.Unlock() |
| 39 | |
| 40 | // Prune expired buckets so they do not accumulate. Sweeping on call keeps |
| 41 | // the package free of goroutines. |
| 42 | if now.After(l.nextSweep) { |
| 43 | for k, b := range l.buckets { |
| 44 | if now.After(b.resetAt) { |
| 45 | delete(l.buckets, k) |
| 46 | } |
| 47 | } |
| 48 | l.nextSweep = now.Add(sweepInterval) |
| 49 | } |
| 50 | |
| 51 | b := l.buckets[key] |
| 52 | if b == nil || now.After(b.resetAt) { |
| 53 | l.buckets[key] = &bucket{count: 1, resetAt: now.Add(l.window)} |
| 54 | return true |
| 55 | } |
| 56 | if b.count >= l.max { |
| 57 | return false |
| 58 | } |
| 59 | b.count++ |
| 60 | return true |
| 61 | } |
| 62 | |
| 63 | // Blocked reports whether key is over the limit without recording an event. |
| 64 | // Callers that only want to count failures check this first and call Allow |
| 65 | // after a failure. |
| 66 | func (l *Limiter) Blocked(key string) bool { |
| 67 | l.mu.Lock() |
| 68 | defer l.mu.Unlock() |
| 69 | b := l.buckets[key] |
| 70 | return b != nil && time.Now().Before(b.resetAt) && b.count >= l.max |
| 71 | } |
| 72 | |
| 73 | // ClientIP returns the address to rate-limit on. |
| 74 | // With a trusted proxy it takes the first X-Forwarded-For entry, because the |
| 75 | // proxy appends the real client there. Otherwise the header is attacker |
| 76 | // controlled and only the socket address can be believed. |
| 77 | func ClientIP(r *http.Request, trustedProxy bool) string { |
| 78 | if trustedProxy { |
| 79 | first, _, _ := strings.Cut(r.Header.Get("X-Forwarded-For"), ",") |
| 80 | return strings.TrimSpace(first) |
| 81 | } |
| 82 | host, _, err := net.SplitHostPort(r.RemoteAddr) |
| 83 | if err != nil { |
| 84 | return r.RemoteAddr |
| 85 | } |
| 86 | return host |
| 87 | } |
| 88 |