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