Files
krow_backend/go-api/internal/httpserver/mcplimit.go
Aravind f2aa3b3ad8
Some checks failed
CI / fixture (push) Has been cancelled
CI / test (push) Has been cancelled
mcp connection
2026-09-22 10:58:02 +05:30

231 lines
9.1 KiB
Go

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
}