internal/httpd/apilimit.go

v1.1.0
gitbay/internal/httpd/apilimit.go history · blame · raw

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}