internal/httpd/apilimit.go
119 lines · 3071 bytes
1package httpd
2
3import (
4 "net"
5 "net/http"
6 "strconv"
7 "sync"
8 "time"
9)
10
11// apiLimiter is a token bucket per caller, with a separate, smaller budget
12// for writes. Unlike the SSH limiter — which counts only auth failures,
13// because a busy CLI opens many connections legitimately — this counts
14// every request: one HTTP call is one command, and a client looping over a
15// list can issue them far faster than a person can type.
16//
17// Keyed by token hash when authenticated, so one caller's budget follows
18// them across networks, and by IP otherwise, so an unauthenticated flood
19// cannot mint budget by rotating tokens.
20type apiLimiter struct {
21 mu sync.Mutex
22 buckets map[string]*bucket
23 rate float64 // requests per second, sustained
24 burst float64
25 writes float64 // sustained write rate, a fraction of rate
26}
27
28type bucket struct {
29 read, write float64
30 last time.Time
31}
32
33func newAPILimiter(perMinute int) *apiLimiter {
34 if perMinute <= 0 {
35 perMinute = 120
36 }
37 rate := float64(perMinute) / 60
38 return &apiLimiter{
39 buckets: map[string]*bucket{},
40 rate: rate,
41 burst: float64(perMinute),
42 // Writes are rarer and more expensive; a tenth of the read budget
43 // is generous for a client and useless for a scraper.
44 writes: rate / 10,
45 }
46}
47
48// allow reports whether the caller may make this request, and how long to
49// wait if not. write requests draw on both buckets: a write is also a
50// request.
51func (l *apiLimiter) allow(key string, write bool) (bool, time.Duration) {
52 l.mu.Lock()
53 defer l.mu.Unlock()
54 now := time.Now()
55
56 if len(l.buckets) > 4096 {
57 for k, b := range l.buckets {
58 if now.Sub(b.last) > 10*time.Minute {
59 delete(l.buckets, k)
60 }
61 }
62 }
63
64 b := l.buckets[key]
65 if b == nil {
66 b = &bucket{read: l.burst, write: l.burst / 10, last: now}
67 l.buckets[key] = b
68 }
69 elapsed := now.Sub(b.last).Seconds()
70 b.last = now
71 b.read = minf(l.burst, b.read+elapsed*l.rate)
72 b.write = minf(l.burst/10, b.write+elapsed*l.writes)
73
74 if b.read < 1 {
75 return false, retryAfter(1-b.read, l.rate)
76 }
77 if write && b.write < 1 {
78 return false, retryAfter(1-b.write, l.writes)
79 }
80 b.read--
81 if write {
82 b.write--
83 }
84 return true, 0
85}
86
87func retryAfter(deficit, rate float64) time.Duration {
88 if rate <= 0 {
89 return time.Minute
90 }
91 d := time.Duration(deficit / rate * float64(time.Second))
92 if d < time.Second {
93 return time.Second
94 }
95 return d
96}
97
98func minf(a, b float64) float64 {
99 if a < b {
100 return a
101 }
102 return b
103}
104
105// clientIP is the peer address. No forwarded headers are trusted: nothing
106// in front of this process is required to set them, and honouring a
107// client-supplied header would let a caller pick their own bucket.
108func clientIP(r *http.Request) string {
109 if host, _, err := net.SplitHostPort(r.RemoteAddr); err == nil {
110 return host
111 }
112 return r.RemoteAddr
113}
114
115func tooManyRequests(w http.ResponseWriter, wait time.Duration) {
116 w.Header().Set("Retry-After", strconv.Itoa(int(wait.Seconds()+0.5)))
117 apiError(w, http.StatusTooManyRequests,
118 "rate limited; retry in "+strconv.Itoa(int(wait.Seconds()+0.5))+"s")
119}