diff --git a/cmd/testserver/main.go b/cmd/testserver/main.go index dc62f78..5268637 100644 --- a/cmd/testserver/main.go +++ b/cmd/testserver/main.go @@ -10,6 +10,7 @@ import ( "github.com/bitechdev/ResolveSpec/pkg/config" "github.com/bitechdev/ResolveSpec/pkg/dbmanager" "github.com/bitechdev/ResolveSpec/pkg/logger" + "github.com/bitechdev/ResolveSpec/pkg/middleware" "github.com/bitechdev/ResolveSpec/pkg/modelregistry" "github.com/bitechdev/ResolveSpec/pkg/server" "github.com/bitechdev/ResolveSpec/pkg/testmodels" @@ -67,8 +68,13 @@ func main() { handler.RegisterModel("public", modelNames[i], model) } + // Queue requests per client (X-Client-Id, Authorization, session, then IP) + // so a burst such as a page load cannot flood the database pool. + queue := middleware.NewClientQueue(middleware.ClientQueueConfig{MaxConcurrent: 10}) + defer queue.Close() + // Setup routes using new SetupMuxRoutes function (without authentication) - resolvespec.SetupMuxRoutes(r, handler, nil) + resolvespec.SetupMuxRoutes(r, handler, middleware.Chain(queue.Middleware)) // Create server manager mgr := server.NewManager() diff --git a/pkg/middleware/README.md b/pkg/middleware/README.md index 51536fd..5bbfb8e 100644 --- a/pkg/middleware/README.md +++ b/pkg/middleware/README.md @@ -5,8 +5,9 @@ HTTP middleware utilities for security and performance. ## Table of Contents 1. [Rate Limiting](#rate-limiting) -2. [Request Size Limits](#request-size-limits) -3. [Input Sanitization](#input-sanitization) +2. [Client Request Queue](#client-request-queue) +3. [Request Size Limits](#request-size-limits) +4. [Input Sanitization](#input-sanitization) --- @@ -381,6 +382,109 @@ func healthHandler(w http.ResponseWriter, r *http.Request) { --- +## Client Request Queue + +`ClientQueue` smooths bursts (for example a page load that fires ~15 requests at once) by +limiting how many requests each client runs concurrently and queueing the rest first-in-first-out. +Unlike the rate limiter it never rejects a request that can be served shortly. + +```go +q := middleware.NewClientQueue(middleware.ClientQueueConfig{ + MaxConcurrent: 4, // running at once, per client + MaxQueue: 50, // waiting, per client; beyond this -> 429 + MaxWait: 30 * time.Second, // waiting too long -> 503 + Retry-After +}) +defer q.Close() + +router.Use(q.Middleware) // gorilla/mux; or wrap any http.Handler +``` + +### Adding it to restheadspec / resolvespec + +Both packages' `SetupMuxRoutes` and `SetupBunRouterRoutes` take one middleware and apply it to every +route, so pass the queue there. Use `Chain` to combine it with auth (first is outermost, so requests +are authenticated before they can take a queue slot): + +```go +q := middleware.NewClientQueue(middleware.ClientQueueConfig{MaxConcurrent: 4}) +defer q.Close() + +mw := middleware.Chain(authMiddleware, q.Middleware) // authMiddleware may be nil + +restheadspec.SetupMuxRoutes(muxRouter, headHandler, mw) +resolvespec.SetupMuxRoutes(muxRouter, resolveHandler, mw) +// or SetupBunRouterRoutes(router, handler, mw) +``` + +Share one `ClientQueue` across packages to give a client a single limit over all of them. + +### Client identification + +The first of these that is present is used, in order: + +1. `X-Client-Id` header +2. `Authorization` header +3. the built-in server session (`security.GetSessionID` from the request context, else the session cookie) +4. the client IP + +Nothing is required from the client: without an id it falls back to the session, then to the IP. +Secrets are hashed and never stored. Ids longer than 128 characters are ignored. + +The session is only in the request context if the `security` auth middleware has run first, so put +auth before the queue with `Chain` (below). Without it the cookie is still used. + +To get one queue per browser tab (rather than per session), have the client send an id generated once +per tab; the server cannot read `sessionStorage` itself: + +```js +const id = sessionStorage.clientId ??= crypto.randomUUID(); +fetch(url, { headers: { "X-Client-Id": id } }); +``` + +### Metrics + +Registered on the default Prometheus registry (the same one `pkg/metrics` and `dbmanager` use), with +no per-client labels so cardinality stays bounded. + +| Metric | Type | Meaning | +|---|---|---| +| `clientqueue_requests_total{result}` | counter | `immediate`, `queued`, `rejected_full`, `timeout`, `canceled` | +| `clientqueue_wait_seconds` | histogram | wait for a slot; 0 for requests that ran immediately | +| `clientqueue_wait_max_seconds` | gauge | longest wait since process start | +| `clientqueue_burst_size` | histogram | peak outstanding requests per client busy period | +| `clientqueue_burst_max` | gauge | largest burst since process start | +| `clientqueue_active` / `clientqueue_queue_depth` | gauge | running / waiting now | +| `clientqueue_clients` | gauge | clients currently tracked | + +A **burst** is one client's busy period: the peak number of its requests running plus waiting between +going from idle to busy and back to idle. A page load firing 15 requests at once is a burst of 15. + +```promql +# average burst size +rate(clientqueue_burst_size_sum[5m]) / rate(clientqueue_burst_size_count[5m]) +# bursts larger than the concurrency limit (need queueing) +1 - (sum(rate(clientqueue_burst_size_bucket{le="10"}[5m])) / sum(rate(clientqueue_burst_size_count[5m]))) +# average and p95 wait +rate(clientqueue_wait_seconds_sum[5m]) / rate(clientqueue_wait_seconds_count[5m]) +histogram_quantile(0.95, sum(rate(clientqueue_wait_seconds_bucket[5m])) by (le)) +# share of requests that had to queue +sum(rate(clientqueue_requests_total{result="queued"}[5m])) / sum(rate(clientqueue_requests_total[5m])) +``` + +Use these to pick `MaxConcurrent`: if the p95 burst is well above it and waits are short, it is doing its +job; if waits grow, the limit is too low or the database is the bottleneck. The `_max` gauges reset on +restart; use the histograms for anything over time. A client that never goes idle produces one long +busy period, so its burst is only recorded when it finally drains. + +### Behaviour + +- Queued requests are dropped if the client disconnects, so abandoned requests never run. +- CORS preflights (`OPTIONS`) and connection upgrades (websockets) bypass the queue. +- Idle clients are forgotten after `IdleTimeout` (default 5m). +- The client id is client-supplied, so a hostile client can rotate ids to get more slots. It smooths + well-behaved clients; it is not an abuse control. Combine it with `RateLimiter` for that. +- Queue time counts against your own request timeouts; keep `MaxWait` below them. + ## Request Size Limits Protect against oversized request bodies with configurable size limits. diff --git a/pkg/middleware/clientqueue.go b/pkg/middleware/clientqueue.go new file mode 100644 index 0000000..fc19f80 --- /dev/null +++ b/pkg/middleware/clientqueue.go @@ -0,0 +1,385 @@ +package middleware + +import ( + "container/list" + "context" + "crypto/sha256" + "encoding/hex" + "errors" + "net/http" + "strconv" + "strings" + "sync" + "time" + + "github.com/prometheus/client_golang/prometheus" + "github.com/prometheus/client_golang/prometheus/promauto" + + "github.com/bitechdev/ResolveSpec/pkg/security" +) + +// DefaultClientIDHeader is the header a client (one per browser tab) sends to +// identify itself to the ClientQueue. +const DefaultClientIDHeader = "X-Client-Id" + +const maxClientIDLen = 128 + +// Client queue metrics. They are aggregated over all clients (no per-client +// labels, to keep cardinality bounded) and over all ClientQueue instances. +// +// A burst is one client's busy period: the peak number of its requests +// outstanding (running plus waiting) between the moment it goes from idle to +// busy and the moment it is idle again. A page load firing 15 requests at +// once is a burst of 15. Average burst is burst_size_sum / burst_size_count; +// the highest burst seen is clientqueue_burst_max. +var ( + queueRequests = promauto.NewCounterVec(prometheus.CounterOpts{ + Name: "clientqueue_requests_total", + Help: "Requests by outcome: immediate (ran without waiting), queued (waited, then ran), rejected_full, timeout, canceled", + }, []string{"result"}) + + queueWait = promauto.NewHistogram(prometheus.HistogramOpts{ + Name: "clientqueue_wait_seconds", + Help: "Time admitted requests spent waiting for a slot (0 for immediate ones)", + Buckets: []float64{0.001, 0.005, 0.01, 0.05, 0.1, 0.25, 0.5, 1, 2.5, 5, 10, 30}, + }) + + queueBurst = promauto.NewHistogram(prometheus.HistogramOpts{ + Name: "clientqueue_burst_size", + Help: "Peak outstanding requests (running + waiting) per client busy period", + Buckets: []float64{1, 2, 3, 4, 5, 8, 10, 15, 20, 30, 50, 100}, + }) + + queueBurstMax = promauto.NewGauge(prometheus.GaugeOpts{ + Name: "clientqueue_burst_max", + Help: "Largest burst seen since the process started", + }) + + queueWaitMax = promauto.NewGauge(prometheus.GaugeOpts{ + Name: "clientqueue_wait_max_seconds", + Help: "Longest queue wait seen since the process started", + }) + + queueActive = promauto.NewGauge(prometheus.GaugeOpts{ + Name: "clientqueue_active", + Help: "Requests currently running through the queue", + }) + + queueDepth = promauto.NewGauge(prometheus.GaugeOpts{ + Name: "clientqueue_queue_depth", + Help: "Requests currently waiting for a slot", + }) + + queueClients = promauto.NewGauge(prometheus.GaugeOpts{ + Name: "clientqueue_clients", + Help: "Clients currently tracked by the queue", + }) + + maxMu sync.Mutex + maxBurst float64 + maxWaitSec float64 +) + +// observeMax raises a high-water gauge when v exceeds the recorded maximum. +func observeMax(g prometheus.Gauge, cur *float64, v float64) { + maxMu.Lock() + if v > *cur { + *cur = v + g.Set(v) + } + maxMu.Unlock() +} + +var ( + errQueueFull = errors.New("client queue full") + errQueueWait = errors.New("client queue wait exceeded") +) + +// ClientQueueConfig configures a ClientQueue. Zero values use the defaults. +type ClientQueueConfig struct { + // MaxConcurrent is how many requests one client may have running at once. + // Default 4. + MaxConcurrent int + // MaxQueue is how many further requests one client may have waiting. + // Requests beyond it are rejected with 429. Default 50. + MaxQueue int + // MaxWait is the longest a request waits for a slot before it is rejected + // with 503. Default 30s. + MaxWait time.Duration + // IdleTimeout is how long a client with no requests is remembered. + // Default 5m. + IdleTimeout time.Duration + // HeaderName is the client id header. Default X-Client-Id. + HeaderName string +} + +func (c *ClientQueueConfig) applyDefaults() { + if c.MaxConcurrent <= 0 { + c.MaxConcurrent = 4 + } + if c.MaxQueue <= 0 { + c.MaxQueue = 50 + } + if c.MaxWait <= 0 { + c.MaxWait = 30 * time.Second + } + if c.IdleTimeout <= 0 { + c.IdleTimeout = 5 * time.Minute + } + if c.HeaderName == "" { + c.HeaderName = DefaultClientIDHeader + } +} + +// ClientQueue limits how many requests each client runs concurrently and +// queues the rest first-in-first-out, smoothing bursts such as a page load +// that fires many requests at once. A client is identified by, in order: +// the client id header, the Authorization header (session), then the IP. +type ClientQueue struct { + cfg ClientQueueConfig + mu sync.Mutex + clients map[string]*queueClient + stop chan struct{} + once sync.Once +} + +type queueClient struct { + active int + waiters list.List // of *queueWaiter + lastUsed time.Time + outstanding int // running + waiting + peak int // highest outstanding in the current busy period +} + +// enter records a request entering the client's outstanding set. +func (c *queueClient) enter() { + c.outstanding++ + if c.outstanding > c.peak { + c.peak = c.outstanding + } +} + +// leave records a request leaving it, closing the burst when the client goes idle. +func (c *queueClient) leave() { + c.outstanding-- + if c.outstanding == 0 { + queueBurst.Observe(float64(c.peak)) + observeMax(queueBurstMax, &maxBurst, float64(c.peak)) + c.peak = 0 + } +} + +type queueWaiter struct { + ready chan struct{} + granted bool + elem *list.Element +} + +// NewClientQueue creates a ClientQueue and starts its idle-client cleanup. +// Call Close to stop it. +func NewClientQueue(cfg ClientQueueConfig) *ClientQueue { + cfg.applyDefaults() + q := &ClientQueue{ + cfg: cfg, + clients: make(map[string]*queueClient), + stop: make(chan struct{}), + } + go q.cleanupRoutine() + return q +} + +// Close stops the cleanup goroutine. Requests in flight are unaffected. +func (q *ClientQueue) Close() { q.once.Do(func() { close(q.stop) }) } + +func (q *ClientQueue) cleanupRoutine() { + interval := q.cfg.IdleTimeout / 2 + if interval < time.Second { + interval = time.Second + } + ticker := time.NewTicker(interval) + defer ticker.Stop() + for { + select { + case <-q.stop: + return + case now := <-ticker.C: + q.evictIdle(now) + } + } +} + +func (q *ClientQueue) evictIdle(now time.Time) { + q.mu.Lock() + defer q.mu.Unlock() + for key, c := range q.clients { + if c.active == 0 && c.waiters.Len() == 0 && now.Sub(c.lastUsed) >= q.cfg.IdleTimeout { + delete(q.clients, key) + queueClients.Dec() + } + } +} + +// clientKey identifies the caller by the first of these that is present: +// the client id header, the Authorization header, the built-in server session +// (from the request context once security auth has run, else the session +// cookie), then the IP. Secrets are hashed so they are never held in memory +// or logs. The session is only in the context if the auth middleware runs +// before this one (see Chain). +func (q *ClientQueue) clientKey(r *http.Request) string { + id := strings.TrimSpace(r.Header.Get(q.cfg.HeaderName)) + if id != "" && len(id) <= maxClientIDLen && !strings.ContainsAny(id, "\r\n\x00") { + return "cid:" + id + } + if a := r.Header.Get("Authorization"); a != "" { + return "auth:" + hashKey(a) + } + if sid, ok := security.GetSessionID(r.Context()); ok && sid != "" { + return "session:" + hashKey(sid) + } + if c := security.GetSessionCookie(r); c != "" { + return "session:" + hashKey(c) + } + return "ip:" + getClientIP(r) +} + +func hashKey(v string) string { + sum := sha256.Sum256([]byte(v)) + return hex.EncodeToString(sum[:16]) +} + +// acquire blocks until key has a free slot, the queue is full, ctx ends, or +// MaxWait passes. On success the caller must call release(key). +func (q *ClientQueue) acquire(ctx context.Context, key string) error { + start := time.Now() + q.mu.Lock() + c := q.clients[key] + if c == nil { + c = &queueClient{} + q.clients[key] = c + queueClients.Inc() + } + c.lastUsed = start + if c.active < q.cfg.MaxConcurrent && c.waiters.Len() == 0 { + c.active++ + c.enter() + q.mu.Unlock() + queueActive.Inc() + queueWait.Observe(0) + queueRequests.WithLabelValues("immediate").Inc() + return nil + } + if c.waiters.Len() >= q.cfg.MaxQueue { + q.mu.Unlock() + queueRequests.WithLabelValues("rejected_full").Inc() + return errQueueFull + } + w := &queueWaiter{ready: make(chan struct{})} + w.elem = c.waiters.PushBack(w) + c.enter() + q.mu.Unlock() + queueDepth.Inc() + + timer := time.NewTimer(q.cfg.MaxWait) + defer timer.Stop() + + var err error + select { + case <-w.ready: + waited := time.Since(start).Seconds() + queueWait.Observe(waited) + observeMax(queueWaitMax, &maxWaitSec, waited) + queueRequests.WithLabelValues("queued").Inc() + return nil + case <-ctx.Done(): + err = ctx.Err() + case <-timer.C: + err = errQueueWait + } + + q.mu.Lock() + if w.granted { + // A slot was handed over as we gave up; pass it on. + q.releaseLocked(key) + } else { + c.waiters.Remove(w.elem) + c.leave() + queueDepth.Dec() + } + q.mu.Unlock() + if errors.Is(err, errQueueWait) { + queueRequests.WithLabelValues("timeout").Inc() + } else { + queueRequests.WithLabelValues("canceled").Inc() + } + return err +} + +func (q *ClientQueue) release(key string) { + q.mu.Lock() + q.releaseLocked(key) + q.mu.Unlock() +} + +// releaseLocked hands the slot to the longest-waiting request, or frees it. +func (q *ClientQueue) releaseLocked(key string) { + c := q.clients[key] + if c == nil { + return + } + c.lastUsed = time.Now() + c.leave() + queueActive.Dec() + if front := c.waiters.Front(); front != nil { + w := c.waiters.Remove(front).(*queueWaiter) + w.granted = true + queueDepth.Dec() + queueActive.Inc() + close(w.ready) + return + } + c.active-- +} + +// Middleware queues requests per client. CORS preflights and connection +// upgrades (websockets) bypass the queue, since they are short or long-lived +// respectively and should not hold or wait for a slot. +func (q *ClientQueue) Middleware(next http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Method == http.MethodOptions || r.Header.Get("Upgrade") != "" { + next.ServeHTTP(w, r) + return + } + + key := q.clientKey(r) + if err := q.acquire(r.Context(), key); err != nil { + switch { + case errors.Is(err, errQueueFull): + w.Header().Set("Retry-After", "1") + http.Error(w, `{"error":"queue_full","message":"Too many queued requests"}`, http.StatusTooManyRequests) + case errors.Is(err, errQueueWait): + w.Header().Set("Retry-After", strconv.Itoa(int(q.cfg.MaxWait.Seconds()))) + http.Error(w, `{"error":"queue_timeout","message":"Timed out waiting in the request queue"}`, http.StatusServiceUnavailable) + } + // Otherwise the client went away; there is no one to answer. + return + } + defer q.release(key) + + next.ServeHTTP(w, r) + }) +} + +// Chain composes middlewares into one, for the single middleware slot of +// restheadspec.SetupMuxRoutes and friends. The first argument is the +// outermost: Chain(auth, q.Middleware) authenticates before a request may +// occupy a queue slot, so unauthenticated traffic cannot fill queues. +func Chain(mws ...func(http.Handler) http.Handler) func(http.Handler) http.Handler { + return func(next http.Handler) http.Handler { + for i := len(mws) - 1; i >= 0; i-- { + if mws[i] != nil { + next = mws[i](next) + } + } + return next + } +} diff --git a/pkg/middleware/clientqueue_test.go b/pkg/middleware/clientqueue_test.go new file mode 100644 index 0000000..974cd36 --- /dev/null +++ b/pkg/middleware/clientqueue_test.go @@ -0,0 +1,333 @@ +package middleware + +import ( + "context" + "net/http" + "net/http/httptest" + "strings" + "sync" + "sync/atomic" + "testing" + "time" + + "github.com/bitechdev/ResolveSpec/pkg/security" + "github.com/prometheus/client_golang/prometheus" + dto "github.com/prometheus/client_model/go" +) + +func newTestQueue(t *testing.T, cfg ClientQueueConfig) *ClientQueue { + q := NewClientQueue(cfg) + t.Cleanup(q.Close) + return q +} + +func TestClientQueueKey(t *testing.T) { + q := newTestQueue(t, ClientQueueConfig{}) + mk := func(id, auth string) *http.Request { + r := httptest.NewRequest("GET", "/", nil) + r.RemoteAddr = "10.0.0.1:1234" + if id != "" { + r.Header.Set("X-Client-Id", id) + } + if auth != "" { + r.Header.Set("Authorization", auth) + } + return r + } + withCtxSession := func(r *http.Request, sid string) *http.Request { + return r.WithContext(context.WithValue(r.Context(), security.SessionIDKey, sid)) + } + + // 1. client id wins over everything else. + r := withCtxSession(mk("tab1", "Bearer a"), "sess") + r.AddCookie(&http.Cookie{Name: "session_token", Value: "cookie"}) + if got := q.clientKey(r); got != "cid:tab1" { + t.Errorf("id: %q", got) + } + // 2. Authorization next, hashed. + got := q.clientKey(withCtxSession(mk("", "Bearer secret"), "sess")) + if !strings.HasPrefix(got, "auth:") || strings.Contains(got, "secret") { + t.Errorf("auth must be hashed: %q", got) + } + // 3. built-in session from the context, then the cookie. + got = q.clientKey(withCtxSession(mk("", ""), "sess")) + if !strings.HasPrefix(got, "session:") || strings.Contains(strings.TrimPrefix(got, "session:"), "sess") { + t.Errorf("context session: %q", got) + } + rc := mk("", "") + rc.AddCookie(&http.Cookie{Name: "session_token", Value: "cookie"}) + if got := q.clientKey(rc); !strings.HasPrefix(got, "session:") { + t.Errorf("cookie session: %q", got) + } + // 4. IP last. + if got := q.clientKey(mk("", "")); got != "ip:10.0.0.1" { + t.Errorf("ip: %q", got) + } + // Oversized ids are ignored and fall through. + if got := q.clientKey(mk(strings.Repeat("x", 200), "")); got != "ip:10.0.0.1" { + t.Errorf("oversized id should be ignored: %q", got) + } + // Tabs with different ids get separate queues; same id shares one. + if q.clientKey(mk("tab1", "Bearer a")) == q.clientKey(mk("tab2", "Bearer a")) { + t.Error("different ids must not share a queue") + } + if q.clientKey(mk("", "Bearer a")) != q.clientKey(mk("", "Bearer a")) { + t.Error("same session without id must share a queue") + } +} + +func TestClientQueueLimitsConcurrencyPerClient(t *testing.T) { + q := newTestQueue(t, ClientQueueConfig{MaxConcurrent: 3}) + var cur, peak atomic.Int32 + h := q.Middleware(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + n := cur.Add(1) + for { + p := peak.Load() + if n <= p || peak.CompareAndSwap(p, n) { + break + } + } + time.Sleep(20 * time.Millisecond) + cur.Add(-1) + })) + + var wg sync.WaitGroup + for i := 0; i < 15; i++ { + wg.Add(1) + go func() { + defer wg.Done() + r := httptest.NewRequest("GET", "/", nil) + r.Header.Set("X-Client-Id", "tab1") + w := httptest.NewRecorder() + h.ServeHTTP(w, r) + if w.Code != http.StatusOK { + t.Errorf("code %d", w.Code) + } + }() + } + wg.Wait() + if p := peak.Load(); p > 3 { + t.Fatalf("peak concurrency %d, want <= 3", p) + } +} + +func TestClientQueueClientsAreIndependent(t *testing.T) { + q := newTestQueue(t, ClientQueueConfig{MaxConcurrent: 1}) + block := make(chan struct{}) + started := make(chan struct{}, 2) + h := q.Middleware(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + started <- struct{}{} + <-block + })) + for _, id := range []string{"a", "b"} { + go func(id string) { + r := httptest.NewRequest("GET", "/", nil) + r.Header.Set("X-Client-Id", id) + h.ServeHTTP(httptest.NewRecorder(), r) + }(id) + } + for i := 0; i < 2; i++ { + select { + case <-started: + case <-time.After(time.Second): + t.Fatal("a second client was blocked by the first") + } + } + close(block) +} + +func TestClientQueueFullAndTimeout(t *testing.T) { + q := newTestQueue(t, ClientQueueConfig{MaxConcurrent: 1, MaxQueue: 1, MaxWait: 50 * time.Millisecond}) + block := make(chan struct{}) + started := make(chan struct{}, 1) + h := q.Middleware(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + started <- struct{}{} + <-block + })) + do := func() *httptest.ResponseRecorder { + r := httptest.NewRequest("GET", "/", nil) + r.Header.Set("X-Client-Id", "tab1") + w := httptest.NewRecorder() + h.ServeHTTP(w, r) + return w + } + + go do() // takes the only slot + <-started + + queued := make(chan *httptest.ResponseRecorder, 1) + go func() { queued <- do() }() + time.Sleep(10 * time.Millisecond) // let it enter the queue + + if w := do(); w.Code != http.StatusTooManyRequests { + t.Errorf("queue full: got %d, want 429", w.Code) + } + if w := <-queued; w.Code != http.StatusServiceUnavailable || w.Header().Get("Retry-After") == "" { + t.Errorf("wait timeout: got %d, want 503 with Retry-After", w.Code) + } + close(block) +} + +func TestClientQueueCancelledWaiterFreesQueue(t *testing.T) { + q := newTestQueue(t, ClientQueueConfig{MaxConcurrent: 1}) + if err := q.acquire(t.Context(), "k"); err != nil { + t.Fatal(err) + } + cctx, ccancel := context.WithCancel(t.Context()) + done := make(chan error, 1) + go func() { done <- q.acquire(cctx, "k") }() + time.Sleep(10 * time.Millisecond) + ccancel() + if err := <-done; err == nil { + t.Fatal("expected cancellation error") + } + q.release("k") + + q.mu.Lock() + c := q.clients["k"] + active, waiting := c.active, c.waiters.Len() + q.mu.Unlock() + if active != 0 || waiting != 0 { + t.Fatalf("leaked state: active=%d waiting=%d", active, waiting) + } +} + +func TestClientQueueBypassesPreflightAndEvictsIdle(t *testing.T) { + q := newTestQueue(t, ClientQueueConfig{MaxConcurrent: 1, IdleTimeout: time.Minute}) + block := make(chan struct{}) + started := make(chan struct{}, 1) + h := q.Middleware(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Method == http.MethodGet { + started <- struct{}{} + <-block + } + })) + go func() { + r := httptest.NewRequest("GET", "/", nil) + r.Header.Set("X-Client-Id", "tab1") + h.ServeHTTP(httptest.NewRecorder(), r) + }() + <-started + + r := httptest.NewRequest("OPTIONS", "/", nil) + r.Header.Set("X-Client-Id", "tab1") + done := make(chan struct{}) + go func() { h.ServeHTTP(httptest.NewRecorder(), r); close(done) }() + select { + case <-done: + case <-time.After(time.Second): + t.Fatal("preflight was queued") + } + close(block) + time.Sleep(20 * time.Millisecond) + + q.evictIdle(time.Now().Add(2 * time.Minute)) + q.mu.Lock() + n := len(q.clients) + q.mu.Unlock() + if n != 0 { + t.Fatalf("idle client not evicted: %d tracked", n) + } +} + +func TestChainOrder(t *testing.T) { + var order []string + tag := func(name string) func(http.Handler) http.Handler { + return func(next http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + order = append(order, name) + next.ServeHTTP(w, r) + }) + } + } + h := Chain(tag("auth"), nil, tag("queue"))(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + order = append(order, "handler") + })) + h.ServeHTTP(httptest.NewRecorder(), httptest.NewRequest("GET", "/", nil)) + if got := len(order); got != 3 || order[0] != "auth" || order[1] != "queue" || order[2] != "handler" { + t.Fatalf("order = %v", order) + } +} + +func gaugeVal(t *testing.T, m prometheus.Metric) float64 { + t.Helper() + var d dto.Metric + if err := m.Write(&d); err != nil { + t.Fatal(err) + } + switch { + case d.Gauge != nil: + return d.Gauge.GetValue() + case d.Counter != nil: + return d.Counter.GetValue() + } + t.Fatal("not a gauge or counter") + return 0 +} + +func histVal(t *testing.T, h prometheus.Histogram) (count uint64, sum float64) { + t.Helper() + var d dto.Metric + if err := h.Write(&d); err != nil { + t.Fatal(err) + } + return d.Histogram.GetSampleCount(), d.Histogram.GetSampleSum() +} + +func TestClientQueueMetrics(t *testing.T) { + q := newTestQueue(t, ClientQueueConfig{MaxConcurrent: 2}) + imm0 := gaugeVal(t, queueRequests.WithLabelValues("immediate")) + que0 := gaugeVal(t, queueRequests.WithLabelValues("queued")) + bc0, bs0 := histVal(t, queueBurst) + wc0, _ := histVal(t, queueWait) + + block := make(chan struct{}) + h := q.Middleware(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { <-block })) + var wg sync.WaitGroup + for i := 0; i < 5; i++ { + wg.Add(1) + go func() { + defer wg.Done() + r := httptest.NewRequest("GET", "/", nil) + r.Header.Set("X-Client-Id", "metrics-tab") + h.ServeHTTP(httptest.NewRecorder(), r) + }() + } + // Wait for 2 running and 3 queued. + deadline := time.Now().Add(2 * time.Second) + for { + q.mu.Lock() + c := q.clients["cid:metrics-tab"] + ok := c != nil && c.active == 2 && c.waiters.Len() == 3 + q.mu.Unlock() + if ok { + break + } + if time.Now().After(deadline) { + t.Fatal("requests did not reach 2 running + 3 queued") + } + time.Sleep(time.Millisecond) + } + if got := gaugeVal(t, queueDepth); got < 3 { + t.Errorf("queue depth = %v, want >= 3", got) + } + close(block) + wg.Wait() + + if got := gaugeVal(t, queueRequests.WithLabelValues("immediate")) - imm0; got != 2 { + t.Errorf("immediate = %v, want 2", got) + } + if got := gaugeVal(t, queueRequests.WithLabelValues("queued")) - que0; got != 3 { + t.Errorf("queued = %v, want 3", got) + } + bc, bs := histVal(t, queueBurst) + if bc-bc0 != 1 || bs-bs0 != 5 { + t.Errorf("burst observations = %d (sum %v), want 1 burst of size 5", bc-bc0, bs-bs0) + } + if wc, _ := histVal(t, queueWait); wc-wc0 != 5 { + t.Errorf("wait observations = %d, want 5", wc-wc0) + } + if got := gaugeVal(t, queueBurstMax); got < 5 { + t.Errorf("burst max = %v, want >= 5", got) + } +}