package httpserver import ( "context" "net/http" "strconv" "time" "github.com/krow/krow-backend/go-api/internal/ratelimit" ) // Rate limiting for the mounted OAuth and MCP routes. // // A middleware rather than a change inside either package, for one reason: the // SUBJECT of a limit is an HTTP concept. Which IP, which bearer token, which // form field names the client — none of that is knowable from inside // internal/oauth, and handing those packages a request so they could work it // out would put transport details in a layer that has none. // // NOTHING SENSITIVE BECOMES A BUCKET KEY. Every subject goes through // ratelimit.Subject, which hashes it. A bucket naming a bearer token would // write that token into a table and into any log line mentioning the bucket. // limited wraps a handler with one rule, keyed by a subject derived per request. // // The subject function returns "" to mean "not limitable" — no token on the // request, say — and the request passes through. That is correct rather than // lax: a request with no identifiable subject is refused by the handler itself // a moment later, and inventing a shared bucket for all of them would let one // caller exhaust a budget that everybody else then queues behind. func (s *Server) limited(rule ratelimit.Rule, subject func(*http.Request) string, next http.Handler) http.Handler { if s.limiter == nil { return next } return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { raw := subject(r) if raw == "" { next.ServeHTTP(w, r) return } decision, err := s.limiter.Allow(r.Context(), rule, ratelimit.Subject(raw)) if err != nil { // The limiter could not count. It has already decided whether that // permits the request — fail closed by default — and this logs the // fault without naming the subject, which is a hash of a // credential. s.log.Error("rate limiter unavailable", "rule", rule.Name, "allowed", decision.Allowed, "error", err) } // Headers on every response, not only refusals, so a well-behaved // client can slow down before it is refused rather than after. w.Header().Set("RateLimit-Limit", strconv.Itoa(decision.Limit)) w.Header().Set("RateLimit-Remaining", strconv.Itoa(decision.Remaining)) if !decision.Allowed { // Retry-After in seconds, rounded up and never zero — "Retry-After: // 0" invites an immediate retry, which is the one thing a limited // client must not do. The same reasoning as retryAfterSeconds in // ratelimit.go, and the same rounding. w.Header().Set("Retry-After", retryAfterSeconds(decision.RetryAfter)) s.log.Warn("rate limit exceeded", "rule", rule.Name, "path", r.URL.Path) w.Header().Set("Content-Type", "application/json; charset=utf-8") w.WriteHeader(http.StatusTooManyRequests) _, _ = w.Write([]byte(`{"error":"rate_limited",` + `"error_description":"too many requests; retry after the interval in the Retry-After header"}`)) return } next.ServeHTTP(w, r) }) } /* ── Subjects ───────────────────────────────────────────────────────────── */ // byClientAddr keys by the caller's address, for endpoints with no credential. // // A method, not a free function, because the address is no longer a property of // the request alone: resolving it needs the trusted-proxy set, which is wired // onto the server. See clientip.go. func (s *Server) byClientAddr(r *http.Request) string { return s.trust.clientAddr(r) } // byAddrAndUser keys the authorization endpoint by address AND signed-in user. // // Both, because either alone is wrong: by user only, an attacker could exhaust // somebody else's budget by naming them; by address only, an office behind one // NAT shares one person's allowance. func (s *Server) byAddrAndUser(r *http.Request) string { subject := s.trust.clientAddr(r) if identity, ok := (sessionResolver{s}).CurrentUser(r); ok { subject += "|" + identity.UserID } return subject } // byFormClientID keys the token endpoint by the client_id it names. // // Reading a form value means parsing the body, which the handler then parses // again — ParseForm caches on the request, so the second call is free. func (s *Server) byFormClientID(r *http.Request) string { if err := r.ParseForm(); err != nil { return "" } if id := r.PostFormValue("client_id"); id != "" { return id } // No client_id: the handler will refuse it. Fall back to the address so a // caller cannot dodge the limit by omitting the field. return s.trust.clientAddr(r) } // byRefreshFamily keys refresh by the token being presented. // // Keyed by the TOKEN's hash, not the family id, because the family is not // knowable without a database read this middleware has no business doing. The // effect is very nearly the same: a rotation produces a new token and therefore // a new bucket, so the practical limit is per-token-per-window rather than // per-family — which bounds a loop just as well, since a loop presenting the // SAME token is exactly what the limit is for. func byRefreshToken(r *http.Request) string { if err := r.ParseForm(); err != nil { return "" } return r.PostFormValue("refresh_token") } // byBearerToken keys MCP by the presented access token. // // The narrowest identity available on an MCP request, and the right one: it is // one connection from one client for one user. Keying by user would let a // person's second client eat the first's budget; keying by IP would make // Claude's shared egress one bucket for every customer. func byBearerToken(r *http.Request) string { const prefix = "Bearer " header := r.Header.Get("Authorization") if len(header) <= len(prefix) { return "" } // Case-insensitive prefix, matching mcpserver's own parsing. if !equalFoldASCII(header[:len(prefix)], prefix) { return "" } return header[len(prefix):] } func equalFoldASCII(a, b string) bool { if len(a) != len(b) { return false } for i := 0; i < len(a); i++ { ca, cb := a[i], b[i] if 'A' <= ca && ca <= 'Z' { ca += 'a' - 'A' } if 'A' <= cb && cb <= 'Z' { cb += 'a' - 'A' } if ca != cb { return false } } return true } // mcpLimited applies BOTH tool-call limits to the MCP endpoint. // // Two rules stacked rather than one, because they stop different things: the // per-minute rule bounds a spike, and the per-hour rule bounds a slow drain // that would sit under the per-minute rule indefinitely. Checked minute-first // so the cheaper refusal happens earlier. func (s *Server) mcpLimited(next http.Handler) http.Handler { return s.limited(ratelimit.MCPToolCallPerMinute, byBearerToken, s.limited(ratelimit.MCPToolCallPerHour, byBearerToken, next)) } // tokenLimited applies the right rule for the grant type being requested. // // One endpoint, two grant types, two different abuse shapes — so one limit // keyed one way would be wrong for the other. A code exchange is bounded per // client; a refresh is bounded per presented token, so one connection looping // cannot spend a second connection's budget. func (s *Server) tokenLimited(next http.Handler) http.Handler { exchange := s.limited(ratelimit.OAuthToken, s.byFormClientID, next) refresh := s.limited(ratelimit.OAuthRefresh, byRefreshToken, next) return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { // ParseForm caches on the request, so the handler's own call is free. if err := r.ParseForm(); err != nil { next.ServeHTTP(w, r) // let the handler produce the proper error return } if r.PostFormValue("grant_type") == "refresh_token" { refresh.ServeHTTP(w, r) return } exchange.ServeHTTP(w, r) }) } /* ── The per-organisation ceiling ───────────────────────────────────────── */ // orgLimiter adapts the shared limiter to mcpserver.OrgLimiter. // // It is the ONE limit that cannot live in the middleware above, because an // organisation is not knowable until the bearer token has been resolved to a // user and that user's row read. A middleware running before authentication // could only key by something the client supplied — which is exactly the // identity this surface refuses to trust. // // So mcpserver calls this from inside dispatch, after it has an // authctx.Identity, and passes the org id from that identity. This type has no // access to the request and therefore no way to be handed a different one. type orgLimiter struct{ s *Server } // AllowOrg counts one call against the organisation's hourly ceiling. // // The org id is hashed like every other subject. It is not a secret, but the // bucket format is uniform and a uuid in a table of counters is one more place // a tenant identifier exists for no reason. func (o orgLimiter) AllowOrg(ctx context.Context, orgID string) (bool, time.Duration, error) { if o.s.limiter == nil || orgID == "" { // No limiter, or no organisation — the latter cannot happen, because // mcpserver refuses an identity without one before it reaches here. return true, 0, nil } d, err := o.s.limiter.Allow(ctx, ratelimit.MCPPerOrgPerHour, ratelimit.Subject(orgID)) return d.Allowed, d.RetryAfter, err }