231 lines
9.1 KiB
Go
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
|
|
}
|