From 68c6ab85c500539817fd5000bc07222504f9432d Mon Sep 17 00:00:00 2001 From: Jakub Zych Date: Sat, 19 Sep 2026 19:35:53 +0200 Subject: [PATCH] feat(06-02): add fixed-window limiter, Store, and ClientIP - MemoryStore mirrors Laravel tooManyAttempts-before-hit first-hit-wins - FixedWindowLimiter + throttle factory; success and 429 rate-limit headers - Trusted-proxy ClientIP; remove noOpLimit; keep Limiter interface seam --- pact/capabilities.go | 2 + surf/clientip.go | 105 +++++++++++++++++ surf/clientip_test.go | 59 ++++++++++ surf/limiter.go | 169 +++++++++++++++++++++++++++ surf/limiter_store.go | 115 +++++++++++++++++++ surf/limiter_test.go | 260 ++++++++++++++++++++++++++++++++++++++++++ surf/router.go | 41 +++++-- 7 files changed, 741 insertions(+), 10 deletions(-) create mode 100644 surf/clientip.go create mode 100644 surf/clientip_test.go create mode 100644 surf/limiter.go create mode 100644 surf/limiter_store.go create mode 100644 surf/limiter_test.go diff --git a/pact/capabilities.go b/pact/capabilities.go index aa2d528..04764cf 100644 --- a/pact/capabilities.go +++ b/pact/capabilities.go @@ -128,3 +128,5 @@ type OptionalMessage interface { // The kernel type-asserts HasConfig (party, before Register), HasCommands // (generated app main, after Boot), HasMigrations (lagoon migrate), and // HasMiddleware/HasMiddlewareFactories/HasRoutes (surf assemble). +// surf.BucketProvider is type-asserted in Assemble (not a pact interface: +// pact cannot import surf without a cycle). diff --git a/surf/clientip.go b/surf/clientip.go new file mode 100644 index 0000000..21a44ec --- /dev/null +++ b/surf/clientip.go @@ -0,0 +1,105 @@ +package surf + +import ( + "net" + "net/http" + "net/netip" + "strings" + + "git.golem15.com/golem15/summercms/compass" +) + +// ClientIP is the single source of client IP for limiter keys (D-04). +// RemoteAddr is used unless it parses as being inside one of trusted; +// in that case the rightmost X-Forwarded-For hop NOT inside any trusted +// prefix is used. An empty trusted list means RemoteAddr only. +func ClientIP(r *http.Request, trusted []netip.Prefix) string { + if r == nil { + return "" + } + remote := parseIP(r.RemoteAddr) + if len(trusted) == 0 || remote == (netip.Addr{}) || !addrTrusted(remote, trusted) { + if remote == (netip.Addr{}) { + return "" + } + return remote.String() + } + xff := r.Header.Get("X-Forwarded-For") + if xff == "" { + return remote.String() + } + hops := strings.Split(xff, ",") + for i := len(hops) - 1; i >= 0; i-- { + hop := strings.TrimSpace(hops[i]) + if hop == "" { + continue + } + addr, err := netip.ParseAddr(hop) + if err != nil { + continue + } + addr = addr.Unmap() + if !addrTrusted(addr, trusted) { + return addr.String() + } + } + return remote.String() +} + +// TrustedProxies reads http.trusted_proxies (a []string of CIDRs) from cfg +// and parses it once into []netip.Prefix. A malformed entry is skipped, not +// fatal (logged by the caller if desired). +func TrustedProxies(cfg *compass.Config) []netip.Prefix { + if cfg == nil { + return nil + } + raw, ok := cfg.Lookup("http.trusted_proxies") + if !ok { + return nil + } + var entries []string + switch v := raw.(type) { + case []string: + entries = v + case []any: + for _, item := range v { + s, _ := item.(string) + if s != "" { + entries = append(entries, s) + } + } + default: + return nil + } + var out []netip.Prefix + for _, e := range entries { + p, err := netip.ParsePrefix(strings.TrimSpace(e)) + if err != nil { + continue + } + out = append(out, p) + } + return out +} + +func parseIP(remoteAddr string) netip.Addr { + host, _, err := net.SplitHostPort(remoteAddr) + if err != nil { + host = remoteAddr + } + host = strings.Trim(host, "[]") + addr, err := netip.ParseAddr(host) + if err != nil { + return netip.Addr{} + } + return addr.Unmap() +} + +func addrTrusted(addr netip.Addr, trusted []netip.Prefix) bool { + for _, p := range trusted { + if p.Contains(addr) { + return true + } + } + return false +} diff --git a/surf/clientip_test.go b/surf/clientip_test.go new file mode 100644 index 0000000..1ff4ce9 --- /dev/null +++ b/surf/clientip_test.go @@ -0,0 +1,59 @@ +package surf + +import ( + "net/http" + "net/http/httptest" + "net/netip" + "testing" +) + +func TestClientIPEmptyTrustedUsesRemoteAddr(t *testing.T) { + req := httptest.NewRequest(http.MethodGet, "/", nil) + req.RemoteAddr = "203.0.113.9:1234" + req.Header.Set("X-Forwarded-For", "198.51.100.1") + got := ClientIP(req, nil) + if got != "203.0.113.9" { + t.Fatalf("got %q", got) + } +} + +func TestClientIPRejectsSpoofedXFF(t *testing.T) { + trusted := []netip.Prefix{mustPrefix("10.0.0.0/8")} + req := httptest.NewRequest(http.MethodGet, "/", nil) + req.RemoteAddr = "203.0.113.9:1234" + req.Header.Set("X-Forwarded-For", "198.51.100.1") + got := ClientIP(req, trusted) + if got != "203.0.113.9" { + t.Fatalf("untrusted RemoteAddr must ignore X-Forwarded-For, got %q", got) + } +} + +func TestClientIPRightmostUntrustedHop(t *testing.T) { + trusted := []netip.Prefix{mustPrefix("10.0.0.0/8")} + req := httptest.NewRequest(http.MethodGet, "/", nil) + req.RemoteAddr = "10.0.0.1:443" + req.Header.Set("X-Forwarded-For", "198.51.100.7, 203.0.113.10, 10.0.0.2") + got := ClientIP(req, trusted) + if got != "203.0.113.10" { + t.Fatalf("rightmost untrusted hop = %q", got) + } +} + +func TestClientIPAllHopsTrustedFallsBack(t *testing.T) { + trusted := []netip.Prefix{mustPrefix("10.0.0.0/8")} + req := httptest.NewRequest(http.MethodGet, "/", nil) + req.RemoteAddr = "10.0.0.1:443" + req.Header.Set("X-Forwarded-For", "10.0.0.8, 10.0.0.9") + got := ClientIP(req, trusted) + if got != "10.0.0.1" { + t.Fatalf("fallback RemoteAddr = %q", got) + } +} + +func mustPrefix(s string) netip.Prefix { + p, err := netip.ParsePrefix(s) + if err != nil { + panic(err) + } + return p +} diff --git a/surf/limiter.go b/surf/limiter.go new file mode 100644 index 0000000..b05afa4 --- /dev/null +++ b/surf/limiter.go @@ -0,0 +1,169 @@ +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 +} diff --git a/surf/limiter_store.go b/surf/limiter_store.go new file mode 100644 index 0000000..21eaeaa --- /dev/null +++ b/surf/limiter_store.go @@ -0,0 +1,115 @@ +package surf + +import ( + "sync" + "time" +) + +// Store mirrors Illuminate\Cache\RateLimiter's hit/tooManyAttempts/ +// availableIn control flow: a fixed window, first-hit-wins (an existing +// unexpired window is never extended), with resetAttempts as a side effect +// of TooManyAttempts observing an expired window. +type Store interface { + Hit(key string, decay time.Duration) (attempts int) + TooManyAttempts(key string, max int) bool + AvailableIn(key string) time.Duration +} + +type counterEntry struct { + count int + resetAt time.Time +} + +// MemoryStore is an in-process, mutex-guarded Store. +type MemoryStore struct { + mu sync.Mutex + entries map[string]*counterEntry + sweep time.Duration + stop chan struct{} +} + +// NewMemoryStore returns an in-process, mutex-guarded Store. sweep controls +// the background expired-entry cleanup interval (memory hygiene only -- +// correctness does not depend on it, since expiry is checked lazily). +// A non-positive sweep disables the background goroutine. +func NewMemoryStore(sweep time.Duration) *MemoryStore { + s := &MemoryStore{ + entries: make(map[string]*counterEntry), + sweep: sweep, + stop: make(chan struct{}), + } + if sweep > 0 { + go s.loop() + } + return s +} + +func (s *MemoryStore) loop() { + ticker := time.NewTicker(s.sweep) + defer ticker.Stop() + for { + select { + case <-ticker.C: + s.purge() + case <-s.stop: + return + } + } +} + +func (s *MemoryStore) purge() { + s.mu.Lock() + defer s.mu.Unlock() + now := time.Now() + for k, e := range s.entries { + if now.After(e.resetAt) { + delete(s.entries, k) + } + } +} + +// Hit increments key's counter, opening a decay window on first hit +// (first-hit-wins: an existing unexpired window is never extended). +func (s *MemoryStore) Hit(key string, decay time.Duration) int { + s.mu.Lock() + defer s.mu.Unlock() + now := time.Now() + e, ok := s.entries[key] + if !ok || now.After(e.resetAt) { + e = &counterEntry{count: 0, resetAt: now.Add(decay)} + s.entries[key] = e + } + e.count++ + return e.count +} + +// TooManyAttempts is true only while count >= max and the window has not +// expired. An expired window is deleted (PHP resetAttempts) and returns false. +func (s *MemoryStore) TooManyAttempts(key string, max int) bool { + s.mu.Lock() + defer s.mu.Unlock() + e, ok := s.entries[key] + if !ok { + return false + } + if time.Now().After(e.resetAt) { + delete(s.entries, key) + return false + } + return e.count >= max +} + +// AvailableIn is the time until the window resets, or 0 if the key is absent. +func (s *MemoryStore) AvailableIn(key string) time.Duration { + s.mu.Lock() + defer s.mu.Unlock() + e, ok := s.entries[key] + if !ok { + return 0 + } + d := e.resetAt.Sub(time.Now()) + if d < 0 { + return 0 + } + return d +} diff --git a/surf/limiter_test.go b/surf/limiter_test.go new file mode 100644 index 0000000..3ad737a --- /dev/null +++ b/surf/limiter_test.go @@ -0,0 +1,260 @@ +package surf + +import ( + "encoding/json" + "net/http" + "net/http/httptest" + "strings" + "testing" + "time" + + "git.golem15.com/golem15/summercms/backpack" + "git.golem15.com/golem15/summercms/bouncer" + "git.golem15.com/golem15/summercms/pact" + "git.golem15.com/golem15/summercms/party" +) + +func TestFixedWindowLimiterMemoryStoreWindow(t *testing.T) { + s := NewMemoryStore(0) + decay := 80 * time.Millisecond + + if c := s.Hit("k", decay); c != 1 { + t.Fatalf("hit 1 = %d", c) + } + if c := s.Hit("k", decay); c != 2 { + t.Fatalf("hit 2 = %d", c) + } + if s.TooManyAttempts("k", 3) { + t.Fatal("TooManyAttempts at count==max-1") + } + if c := s.Hit("k", decay); c != 3 { + t.Fatalf("hit 3 = %d", c) + } + if !s.TooManyAttempts("k", 3) { + t.Fatal("TooManyAttempts at count==max") + } + + time.Sleep(decay + 20*time.Millisecond) + if s.TooManyAttempts("k", 3) { + t.Fatal("TooManyAttempts after window elapsed") + } + if c := s.Hit("k", decay); c != 1 { + t.Fatalf("fresh window hit = %d", c) + } +} + +func TestFixedWindowLimiterMemoryStoreFirstHitWins(t *testing.T) { + s := NewMemoryStore(0) + decay := 200 * time.Millisecond + s.Hit("k", decay) + time.Sleep(120 * time.Millisecond) + s.Hit("k", decay) // must not extend the original window + time.Sleep(100 * time.Millisecond) + if s.TooManyAttempts("k", 1) { + t.Fatal("window was extended by a later hit") + } +} + +func TestFixedWindowLimiterSuccessHeaders(t *testing.T) { + lim := NewFixedWindowLimiter(NewMemoryStore(0), nil) + h := lim.Middleware("3,1")(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.WriteHeader(http.StatusOK) + _, _ = w.Write([]byte(`{"ok":true}`)) + })) + rec := httptest.NewRecorder() + req := httptest.NewRequest(http.MethodGet, "http://example.test/x", nil) + h.ServeHTTP(rec, req) + if rec.Code != http.StatusOK { + t.Fatalf("status = %d", rec.Code) + } + if rec.Header().Get("X-RateLimit-Limit") != "3" { + t.Fatalf("limit = %q", rec.Header().Get("X-RateLimit-Limit")) + } + if rec.Header().Get("X-RateLimit-Remaining") != "2" { + t.Fatalf("remaining = %q", rec.Header().Get("X-RateLimit-Remaining")) + } + if rec.Header().Get("Retry-After") != "" || rec.Header().Get("X-RateLimit-Reset") != "" { + t.Fatal("Retry-After / X-RateLimit-Reset on success") + } +} + +func TestFixedWindowLimiterTooManyAttemptsHeaders(t *testing.T) { + lim := NewFixedWindowLimiter(NewMemoryStore(0), nil) + h := lim.Middleware("1,1")(http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.WriteHeader(http.StatusOK) + })) + req := httptest.NewRequest(http.MethodGet, "http://example.test/x", nil) + ok := httptest.NewRecorder() + h.ServeHTTP(ok, req) + if ok.Code != http.StatusOK { + t.Fatalf("first status = %d", ok.Code) + } + rec := httptest.NewRecorder() + h.ServeHTTP(rec, req) + if rec.Code != http.StatusTooManyRequests { + t.Fatalf("status = %d", rec.Code) + } + if rec.Header().Get("Retry-After") == "" { + t.Fatal("missing Retry-After") + } + if rec.Header().Get("X-RateLimit-Reset") == "" { + t.Fatal("missing X-RateLimit-Reset") + } + if rec.Header().Get("X-RateLimit-Limit") != "1" { + t.Fatalf("limit = %q", rec.Header().Get("X-RateLimit-Limit")) + } + if rec.Header().Get("X-RateLimit-Remaining") != "0" { + t.Fatalf("remaining = %q", rec.Header().Get("X-RateLimit-Remaining")) + } + if rec.Body.String() != `{"message":"Too Many Attempts."}` { + t.Fatalf("body = %q", rec.Body.String()) + } + var payload map[string]any + if err := json.Unmarshal(rec.Body.Bytes(), &payload); err != nil { + t.Fatal(err) + } + if payload["message"] != "Too Many Attempts." { + t.Fatalf("payload = %v", payload) + } +} + +func TestFixedWindowLimiterStackedBuckets(t *testing.T) { + okHandler := http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.WriteHeader(http.StatusOK) + }) + + t.Run("exhaust A", func(t *testing.T) { + lim := NewFixedWindowLimiter(NewMemoryStore(0), nil) + if err := lim.RegisterBucket("demo", "bucket-a", Bucket{ + Max: 1, Decay: time.Minute, + Key: func(*http.Request) string { return "a" }, + }); err != nil { + t.Fatal(err) + } + if err := lim.RegisterBucket("demo", "bucket-b", Bucket{ + Max: 100, Decay: time.Minute, + Key: func(*http.Request) string { return "b" }, + }); err != nil { + t.Fatal(err) + } + h := stackThrottle(lim, okHandler, "bucket-a", "bucket-b") + req := httptest.NewRequest(http.MethodGet, "/x", nil) + first := httptest.NewRecorder() + h.ServeHTTP(first, req) + if first.Code != http.StatusOK { + t.Fatalf("first = %d", first.Code) + } + rec := httptest.NewRecorder() + h.ServeHTTP(rec, req) + if rec.Code != http.StatusTooManyRequests { + t.Fatalf("exhausted A status = %d", rec.Code) + } + }) + + t.Run("exhaust B", func(t *testing.T) { + lim := NewFixedWindowLimiter(NewMemoryStore(0), nil) + if err := lim.RegisterBucket("demo", "bucket-a", Bucket{ + Max: 100, Decay: time.Minute, + Key: func(*http.Request) string { return "a-fresh" }, + }); err != nil { + t.Fatal(err) + } + if err := lim.RegisterBucket("demo", "bucket-b", Bucket{ + Max: 1, Decay: time.Minute, + Key: func(*http.Request) string { return "b-only" }, + }); err != nil { + t.Fatal(err) + } + h := stackThrottle(lim, okHandler, "bucket-a", "bucket-b") + req := httptest.NewRequest(http.MethodGet, "/x", nil) + first := httptest.NewRecorder() + h.ServeHTTP(first, req) + if first.Code != http.StatusOK { + t.Fatalf("first = %d", first.Code) + } + rec := httptest.NewRecorder() + h.ServeHTTP(rec, req) + if rec.Code != http.StatusTooManyRequests { + t.Fatalf("exhausted B status = %d", rec.Code) + } + }) +} + +func TestFixedWindowLimiterInlineThrottleKeys(t *testing.T) { + okHandler := http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.WriteHeader(http.StatusOK) + }) + + t.Run("principals differ", func(t *testing.T) { + lim := NewFixedWindowLimiter(NewMemoryStore(0), nil) + h := lim.Middleware("1,1")(okHandler) + reqA := httptest.NewRequest(http.MethodGet, "http://example.test/x", nil) + reqA.RemoteAddr = "192.0.2.1:1" + reqA = reqA.WithContext(bouncer.WithUser(reqA.Context(), &bouncer.Principal{ID: 1})) + reqB := httptest.NewRequest(http.MethodGet, "http://example.test/x", nil) + reqB.RemoteAddr = "192.0.2.1:1" + reqB = reqB.WithContext(bouncer.WithUser(reqB.Context(), &bouncer.Principal{ID: 2})) + + a1 := httptest.NewRecorder() + h.ServeHTTP(a1, reqA) + a2 := httptest.NewRecorder() + h.ServeHTTP(a2, reqA) + if a2.Code != http.StatusTooManyRequests { + t.Fatalf("user 1 second status = %d", a2.Code) + } + b1 := httptest.NewRecorder() + h.ServeHTTP(b1, reqB) + if b1.Code != http.StatusOK { + t.Fatalf("user 2 should have a distinct key, status = %d", b1.Code) + } + }) + + t.Run("anonymous same IP+Host", func(t *testing.T) { + lim := NewFixedWindowLimiter(NewMemoryStore(0), nil) + h := lim.Middleware("1,1")(okHandler) + req1 := httptest.NewRequest(http.MethodGet, "http://example.test/x", nil) + req1.RemoteAddr = "192.0.2.1:1" + req2 := httptest.NewRequest(http.MethodGet, "http://example.test/x", nil) + req2.RemoteAddr = "192.0.2.1:9" + first := httptest.NewRecorder() + h.ServeHTTP(first, req1) + if first.Code != http.StatusOK { + t.Fatalf("first = %d", first.Code) + } + second := httptest.NewRecorder() + h.ServeHTTP(second, req2) + if second.Code != http.StatusTooManyRequests { + t.Fatalf("anonymous same IP+Host should share key, status = %d", second.Code) + } + }) +} + +func TestFixedWindowLimiterUnknownBucketFailsAssemble(t *testing.T) { + p := assemblePlugin{id: "golem15.demo", use: []string{"throttle:missing"}} + _, err := Assemble(backpack.New(nil), []party.Plugin{p}) + if err == nil || !strings.Contains(err.Error(), "missing") { + t.Fatalf("want unknown throttle in error, got %v", err) + } +} + +func TestFixedWindowLimiterDuplicateBucket(t *testing.T) { + lim := NewFixedWindowLimiter(NewMemoryStore(0), nil) + b := Bucket{Max: 1, Decay: time.Minute, Key: func(*http.Request) string { return "k" }} + if err := lim.RegisterBucket("one", "shared", b); err != nil { + t.Fatal(err) + } + err := lim.RegisterBucket("two", "shared", b) + if err == nil || !strings.Contains(err.Error(), "shared") || !strings.Contains(err.Error(), "one") { + t.Fatalf("got %v", err) + } +} + +func stackThrottle(lim *FixedWindowLimiter, next http.Handler, names ...string) http.Handler { + h := next + for i := len(names) - 1; i >= 0; i-- { + h = lim.Middleware(names[i])(h) + } + return h +} + +var _ pact.Router = (*Router)(nil) diff --git a/surf/router.go b/surf/router.go index 8af14ae..46394b8 100644 --- a/surf/router.go +++ b/surf/router.go @@ -4,6 +4,7 @@ import ( "fmt" "net/http" "strings" + "time" "git.golem15.com/golem15/summercms/backpack" "git.golem15.com/golem15/summercms/pact" @@ -46,6 +47,7 @@ type Router struct { seen map[string]string origins []string compileErr error + limiter *FixedWindowLimiter } var ( @@ -300,7 +302,6 @@ func (r *Router) compile() (http.Handler, error) { func (r *Router) wrap(rt route) (http.Handler, error) { h := constrain(rt.handler, rt.constraints) - h = noOpLimit(h) h = orgSlot(h) for i := len(rt.middleware) - 1; i >= 0; i-- { name := rt.middleware[i] @@ -311,6 +312,11 @@ func (r *Router) wrap(rt route) (http.Handler, error) { base, param, hasParam := strings.Cut(name, ":") if hasParam { if factory, ok := r.factories[base]; ok { + if base == "throttle" && r.limiter != nil { + if err := r.limiter.ValidateThrottle(param); err != nil { + return nil, fmt.Errorf("surf: plugin %q: %w", rt.pluginID, err) + } + } h = factory.fn(param)(h) continue } @@ -324,6 +330,18 @@ func (r *Router) wrap(rt route) (http.Handler, error) { // Assemble registers plugin middleware and routes, then compiles ServeMux. func Assemble(app *backpack.App, plugins []party.Plugin) (http.Handler, error) { r := New(corsOrigins(app)) + trusted := TrustedProxies(nil) + if app != nil { + trusted = TrustedProxies(app.Config) + } + // longest bucket decay is 1 minute; sweep at 2x + lim := NewFixedWindowLimiter(NewMemoryStore(2*time.Minute), trusted) + r.limiter = lim + if err := r.RegisterMiddlewareFactory("surf", "throttle", func(param string) pact.Middleware { + return lim.Middleware(param) + }); err != nil { + return nil, err + } for _, p := range plugins { if hm, ok := p.(pact.HasMiddleware); ok { for name, fn := range hm.Middlewares() { @@ -342,6 +360,15 @@ func Assemble(app *backpack.App, plugins []party.Plugin) (http.Handler, error) { } } } + for _, p := range plugins { + if bp, ok := p.(BucketProvider); ok { + for name, b := range bp.Buckets() { + if err := lim.RegisterBucket(p.ID(), name, b); err != nil { + return nil, err + } + } + } + } for _, p := range plugins { r.BindPlugin(p.ID()) if hr, ok := p.(pact.HasRoutes); ok { @@ -429,19 +456,13 @@ func orgSlot(next http.Handler) http.Handler { }) } -// Limiter wraps handlers. Phase 6 replaces the no-op with named buckets. +// Limiter wraps handlers. It is a retained Phase 3 seam with no implementer +// after Phase 6; rate limiting is provided by FixedWindowLimiter through +// the parameterized throttle middleware. type Limiter interface { Wrap(http.Handler) http.Handler } -type noopLimiter struct{} - -func (noopLimiter) Wrap(next http.Handler) http.Handler { return next } - -func noOpLimit(next http.Handler) http.Handler { - return noopLimiter{}.Wrap(next) -} - func joinPath(prefix, path string) string { prefix = strings.TrimSuffix(prefix, "/") path = strings.TrimSpace(path)