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
|
// The kernel type-asserts HasConfig (party, before Register), HasCommands
|
||||||
// (generated app main, after Boot), HasMigrations (lagoon migrate), and
|
// (generated app main, after Boot), HasMigrations (lagoon migrate), and
|
||||||
// HasMiddleware/HasMiddlewareFactories/HasRoutes (surf assemble).
|
// 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"
|
"fmt"
|
||||||
"net/http"
|
"net/http"
|
||||||
"strings"
|
"strings"
|
||||||
|
"time"
|
||||||
|
|
||||||
"git.golem15.com/golem15/summercms/backpack"
|
"git.golem15.com/golem15/summercms/backpack"
|
||||||
"git.golem15.com/golem15/summercms/pact"
|
"git.golem15.com/golem15/summercms/pact"
|
||||||
@@ -46,6 +47,7 @@ type Router struct {
|
|||||||
seen map[string]string
|
seen map[string]string
|
||||||
origins []string
|
origins []string
|
||||||
compileErr error
|
compileErr error
|
||||||
|
limiter *FixedWindowLimiter
|
||||||
}
|
}
|
||||||
|
|
||||||
var (
|
var (
|
||||||
@@ -300,7 +302,6 @@ func (r *Router) compile() (http.Handler, error) {
|
|||||||
|
|
||||||
func (r *Router) wrap(rt route) (http.Handler, error) {
|
func (r *Router) wrap(rt route) (http.Handler, error) {
|
||||||
h := constrain(rt.handler, rt.constraints)
|
h := constrain(rt.handler, rt.constraints)
|
||||||
h = noOpLimit(h)
|
|
||||||
h = orgSlot(h)
|
h = orgSlot(h)
|
||||||
for i := len(rt.middleware) - 1; i >= 0; i-- {
|
for i := len(rt.middleware) - 1; i >= 0; i-- {
|
||||||
name := rt.middleware[i]
|
name := rt.middleware[i]
|
||||||
@@ -311,6 +312,11 @@ func (r *Router) wrap(rt route) (http.Handler, error) {
|
|||||||
base, param, hasParam := strings.Cut(name, ":")
|
base, param, hasParam := strings.Cut(name, ":")
|
||||||
if hasParam {
|
if hasParam {
|
||||||
if factory, ok := r.factories[base]; ok {
|
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)
|
h = factory.fn(param)(h)
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
@@ -324,6 +330,18 @@ func (r *Router) wrap(rt route) (http.Handler, error) {
|
|||||||
// Assemble registers plugin middleware and routes, then compiles ServeMux.
|
// Assemble registers plugin middleware and routes, then compiles ServeMux.
|
||||||
func Assemble(app *backpack.App, plugins []party.Plugin) (http.Handler, error) {
|
func Assemble(app *backpack.App, plugins []party.Plugin) (http.Handler, error) {
|
||||||
r := New(corsOrigins(app))
|
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 {
|
for _, p := range plugins {
|
||||||
if hm, ok := p.(pact.HasMiddleware); ok {
|
if hm, ok := p.(pact.HasMiddleware); ok {
|
||||||
for name, fn := range hm.Middlewares() {
|
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 {
|
for _, p := range plugins {
|
||||||
r.BindPlugin(p.ID())
|
r.BindPlugin(p.ID())
|
||||||
if hr, ok := p.(pact.HasRoutes); ok {
|
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 {
|
type Limiter interface {
|
||||||
Wrap(http.Handler) http.Handler
|
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 {
|
func joinPath(prefix, path string) string {
|
||||||
prefix = strings.TrimSuffix(prefix, "/")
|
prefix = strings.TrimSuffix(prefix, "/")
|
||||||
path = strings.TrimSpace(path)
|
path = strings.TrimSpace(path)
|
||||||
|
|||||||
Reference in New Issue
Block a user