package surf import ( "fmt" "net/http" "net/netip" "strconv" "strings" "sync" "time" "git.golem15.com/golem15/summercms/bouncer" "git.golem15.com/golem15/summercms/pact" ) const tooManyAttemptsBody = `{"message":"Too Many Attempts."}` // Bucket is one named rate-limit definition. Key composes the limiter key // from the request (token id, IP, route param -- D-01's per-bucket rule). type Bucket struct { Name string Max int Decay time.Duration Key func(r *http.Request) string } // BucketProvider is implemented by plugins that declare named buckets (not a // pact interface: it lives in surf and is type-asserted directly in // Assemble/BuildRouter, since pact cannot import surf without a cycle). type BucketProvider interface { Buckets() map[string]Bucket } // FixedWindowLimiter is the concrete rate limiter (distinct from the // pre-existing surf.Limiter interface). It owns the Store, the named-bucket // table, and the inline-throttle parser, and produces the "throttle" // middleware factory's per-route pact.Middleware. type FixedWindowLimiter struct { store Store trusted []netip.Prefix mu sync.Mutex buckets map[string]Bucket owners map[string]string inline map[string]Bucket } // NewFixedWindowLimiter is the ONE constructor signature for this type -- // trusted is required at construction (not a later setter) because both // named-bucket Key closures (registered later via RegisterBucket) and the // inline "N,M" throttle's own key resolver need the same trusted-proxy list. func NewFixedWindowLimiter(store Store, trusted []netip.Prefix) *FixedWindowLimiter { cp := make([]netip.Prefix, len(trusted)) copy(cp, trusted) return &FixedWindowLimiter{ store: store, trusted: cp, buckets: make(map[string]Bucket), owners: make(map[string]string), inline: make(map[string]Bucket), } } // RegisterBucket stores a named bucket. Duplicate names fail. func (l *FixedWindowLimiter) RegisterBucket(pluginID, name string, b Bucket) error { if l == nil { return fmt.Errorf("surf: limiter is nil") } if name == "" { return fmt.Errorf("surf: plugin %q registered empty bucket", pluginID) } l.mu.Lock() defer l.mu.Unlock() if existing, ok := l.owners[name]; ok { return fmt.Errorf("surf: bucket %q already registered by %s", name, existing) } b.Name = name l.buckets[name] = b l.owners[name] = pluginID return nil } // Middleware builds the throttle: factory body. param is either a // registered bucket name or a literal "N,M" pair (inline throttle). func (l *FixedWindowLimiter) Middleware(param string) pact.Middleware { if l == nil { return func(next http.Handler) http.Handler { return next } } b, err := l.resolve(param) return func(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { if err != nil || b.Key == nil || l.store == nil { next.ServeHTTP(w, r) return } key := b.Key(r) if l.store.TooManyAttempts(key, b.Max) { retryAfter := l.store.AvailableIn(key) secs := int(retryAfter / time.Second) w.Header().Set("Content-Type", "application/json") w.Header().Set("Retry-After", strconv.Itoa(secs)) w.Header().Set("X-RateLimit-Reset", strconv.FormatInt(time.Now().Add(retryAfter).Unix(), 10)) w.Header().Set("X-RateLimit-Limit", strconv.Itoa(b.Max)) w.Header().Set("X-RateLimit-Remaining", "0") w.WriteHeader(http.StatusTooManyRequests) _, _ = w.Write([]byte(tooManyAttemptsBody)) return } attempts := l.store.Hit(key, b.Decay) remaining := b.Max - attempts if remaining < 0 { remaining = 0 } w.Header().Set("X-RateLimit-Limit", strconv.Itoa(b.Max)) w.Header().Set("X-RateLimit-Remaining", strconv.Itoa(remaining)) next.ServeHTTP(w, r) }) } } // ValidateThrottle is called once per route at Assemble time so a malformed // inline "N,M" or an unregistered bucket name fails boot instead of the // first live request. func (l *FixedWindowLimiter) ValidateThrottle(param string) error { if l == nil { return fmt.Errorf("surf: limiter is nil") } _, err := l.resolve(param) return err } func (l *FixedWindowLimiter) resolve(param string) (Bucket, error) { l.mu.Lock() defer l.mu.Unlock() if b, ok := l.buckets[param]; ok { return b, nil } if b, ok := l.inline[param]; ok { return b, nil } nStr, mStr, ok := strings.Cut(param, ",") if !ok { return Bucket{}, fmt.Errorf("surf: unknown throttle %q", param) } n, errN := strconv.Atoi(strings.TrimSpace(nStr)) m, errM := strconv.Atoi(strings.TrimSpace(mStr)) if errN != nil || errM != nil || n < 1 || m < 1 { return Bucket{}, fmt.Errorf("surf: malformed throttle %q", param) } trusted := l.trusted b := Bucket{ Name: param, Max: n, Decay: time.Duration(m) * time.Minute, Key: func(r *http.Request) string { if r != nil { if u, ok := bouncer.User(r.Context()); ok { return "u:" + strconv.FormatUint(uint64(u.ID), 10) } } host := "" if r != nil { host = r.Host } return host + "|" + ClientIP(r, trusted) }, } l.inline[param] = b return b, nil }