aboutsummaryrefslogtreecommitdiffstats
path: root/internal/web/ratelimit.go
diff options
context:
space:
mode:
Diffstat (limited to 'internal/web/ratelimit.go')
-rw-r--r--internal/web/ratelimit.go87
1 files changed, 87 insertions, 0 deletions
diff --git a/internal/web/ratelimit.go b/internal/web/ratelimit.go
new file mode 100644
index 0000000..b019a4f
--- /dev/null
+++ b/internal/web/ratelimit.go
@@ -0,0 +1,87 @@
+package web
+
+import (
+ "net"
+ "net/http"
+ "sync"
+ "time"
+)
+
+// limiter is a token bucket per key (a client address, a username) for the
+// few anonymous endpoints worth abusing: login guesses and the search's
+// regex scan. Everything else is left to the reverse proxy's limit_req. It
+// lives in memory — one process, no shared state needed — and forgets a key
+// once its bucket has been full for a while.
+type limiter struct {
+ mu sync.Mutex
+ rate float64 // tokens per second
+ burst float64
+ buckets map[string]*bucket
+ lastSweep time.Time
+ now func() time.Time // tests replace it
+}
+
+type bucket struct {
+ tokens float64
+ at time.Time
+}
+
+const limiterSweep = 10 * time.Minute
+
+func newLimiter(perMinute, burst int) *limiter {
+ return &limiter{rate: float64(perMinute) / 60, burst: float64(burst), buckets: map[string]*bucket{}, now: time.Now}
+}
+
+// allow takes one token for key and says whether there was one.
+func (l *limiter) allow(key string) bool {
+ l.mu.Lock()
+ defer l.mu.Unlock()
+ now := l.now()
+ if l.lastSweep.IsZero() {
+ l.lastSweep = now
+ } else if now.Sub(l.lastSweep) > limiterSweep {
+ l.lastSweep = now
+ for k, b := range l.buckets {
+ if l.fill(b, now) >= l.burst {
+ delete(l.buckets, k)
+ }
+ }
+ }
+ b, ok := l.buckets[key]
+ if !ok {
+ b = &bucket{tokens: l.burst, at: now}
+ l.buckets[key] = b
+ }
+ if l.fill(b, now) < 1 {
+ return false
+ }
+ b.tokens--
+ return true
+}
+
+// fill credits the time since the last visit and returns the balance.
+func (l *limiter) fill(b *bucket, now time.Time) float64 {
+ b.tokens = min(l.burst, b.tokens+now.Sub(b.at).Seconds()*l.rate)
+ b.at = now
+ return b.tokens
+}
+
+// clientIP is the address requests are throttled by: the peer's.
+func clientIP(r *http.Request) string {
+ host, _, err := net.SplitHostPort(r.RemoteAddr)
+ if err != nil {
+ return r.RemoteAddr
+ }
+ return host
+}
+
+// throttle answers 429 and logs when key has no tokens left; false means the
+// request was refused.
+func (s *Server) throttle(w http.ResponseWriter, r *http.Request, l *limiter, key string) bool {
+ if l.allow(key) {
+ return true
+ }
+ logf("throttled %s %s from %s", r.Method, r.URL.Path, clientIP(r))
+ s.fail(w, r, http.StatusTooManyRequests, s.tr(r, "Too many requests. Try again in a minute."))
+ return false
+}