ratelimit.go
⎇
Raw
1// Package ratelimit implements a fixed-window counter per key.
2package ratelimit
3
4import (
5 "net"
6 "net/http"
7 "strings"
8 "sync"
9 "time"
10)
11
12type 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.
19type 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
28const sweepInterval = time.Minute
29
30func 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.
35func (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.
66func (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.
77func 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