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