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
This commit is contained in:
@@ -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).
|
||||
|
||||
105
surf/clientip.go
Normal file
105
surf/clientip.go
Normal file
@@ -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
|
||||
}
|
||||
59
surf/clientip_test.go
Normal file
59
surf/clientip_test.go
Normal file
@@ -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
|
||||
}
|
||||
169
surf/limiter.go
Normal file
169
surf/limiter.go
Normal file
@@ -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
|
||||
}
|
||||
115
surf/limiter_store.go
Normal file
115
surf/limiter_store.go
Normal file
@@ -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
|
||||
}
|
||||
260
surf/limiter_test.go
Normal file
260
surf/limiter_test.go
Normal file
@@ -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)
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user