diff --git a/README.md b/README.md index 77e03f8..c1fdfdb 100644 --- a/README.md +++ b/README.md @@ -58,7 +58,7 @@ Four things here that we did not find elsewhere: | `audit` | One record per call — who asked, what was decided, what came out | Dependencies run strictly downward: `mcpkit` → `mcp`; `doc` → `mcp`; `mcptool` → -`mcp`; `mcptest` → `mcp`; `ratelimit` → `auth`; `audit` → `redact`. Nothing points back up, `mcp` +`mcp`; `mcptest` → `mcp`; `ratelimit` → `auth`, `mcp`; `audit` → `redact`. Nothing points back up, `mcp` depends on nothing, and `auth` is imported by exactly one package — `ratelimit`, which keys its buckets on the principal. The server, the dispatcher and the catalogues never learn who is asking. @@ -133,7 +133,9 @@ curl -s localhost:8075/mcp -H 'Authorization: Bearer demo-token' \ One endpoint. `POST` is a JSON-RPC 2.0 request and always gets a synchronous JSON response; a notification — no `id`, or a null one — is dispatched and -answered with `202` and an empty body. `GET` and `DELETE` are `405` with +answered with `202` and an empty body. Only `notifications/*` go without an id: +any other method without one is a `400`, since its answer would be read by +nobody and its work would run past the rate limiter. `GET` and `DELETE` are `405` with `Allow: POST`: the server pushes nothing and holds no session, and saying so is better than opening a stream that dies. A JSON array in the body — a JSON-RPC batch — is `400`, and a body over 4 MiB is `413`. @@ -274,8 +276,11 @@ whoever knows it. Metrics: `app_mcp_transport_rejected_total{reason}` — `batch`, `parse`, `too_large`, `read_body`, `dispatch`, `method_not_allowed`, `origin`, `host`, -plus the modern-era refusals `header_mismatch`, `bad_version`, `missing_meta`; and -`app_mcp_requests_total{era}`, which is how you find out whether +plus the modern-era refusals `header_mismatch`, `bad_version`, `missing_meta`, and +`repeated_key` — a member named twice, in any case, in the envelope, its params +or their `_meta`, which this transport and the decoder behind it would read +differently — and `missing_id`, a call other than a notification sent without an +id; and `app_mcp_requests_total{era}`, which is how you find out whether anything still speaks the old one. A request refused here reaches no handler and appears in no other series, so without these counters a client that speaks the wrong dialect is invisible. @@ -339,7 +344,38 @@ defer limiter.Stop() Each limit is switched off by a value ≤ 0; with all four off `Middleware` returns your handler unwrapped. Only `POST` is counted — the listening `GET` would hold a concurrency slot for the lifetime of a bridge and drain the bucket -by reconnecting. A denial is `429` with `Retry-After: 10`. +by reconnecting. + +A denial is a `429` carrying a JSON-RPC error to the refused call, so a client +can show it as a failed call with a reason rather than as a server that went +away: + +```json +{"jsonrpc":"2.0","id":7,"error":{"code":-32010,"message":"rate limit: cost_budget", + "data":{"reason":"cost_budget","retry_after":1847,"budget_used":"30m2s","budget_limit":"30m","window":"1h"}}} +``` + +`retry_after` and the `Retry-After` header are the same honest wait: until the +window rolls for the budget, until the next token for the rate, a second for a +concurrency slot. The budget fields are there whenever the budget is on. A body +with no id to answer — a notification, a batch, not JSON, too large to look +into — gets the plain-text `429` instead: a refusal reads at most 64 KiB of a +body, so that shedding load costs less than serving it. + +A help tool or a static resource is what teaches a caller to spend less, and +once the budget is gone it would be the first thing to stop answering. `Exempt` +spares such calls the budget — only the budget: they still take a rate token +and a concurrency slot. + +```go +Exempt: func(method, name string) bool { + return method == "tools/call" && name == "help" +}, +``` + +`name` is what `Mcp-Name` mirrors: the tool, the prompt, the resource URI. With +`Exempt` set and the budget on, the limiter parses every `POST` body; otherwise +only a refused one. The bucket is keyed by `Principal.UserID`, so the limiter has to run inside the authentication middleware. Everything is in memory and resets with the process: @@ -354,16 +390,48 @@ the work says so: ```go start := time.Now() // … one upstream call … -ratelimit.Charge(ctx, time.Since(start)) +ratelimit.ChargeFor(ctx, "grafana", time.Since(start)) ``` +The label is what `app_mcp_ratelimit_charge_seconds` breaks the budget down by, +so "who ate the budget" is a PromQL query rather than an afternoon in the audit +log. It is one series per label: name a target from your own catalogue, never a +value from the request. `Charge(ctx, d)` is the same without a label. + `Charge` is safe to call concurrently and does nothing when limits are off, so a handler never has to ask whether they are. A request that charges nothing is priced by its wall clock. -Metrics: `app_mcp_ratelimit_denied_total{reason}` — `rpm`, `user_concurrent`, -`global_concurrent`, `cost_budget` — and `app_mcp_ratelimit_inflight{scope}` with -`user` and `global`. +A handler can also ask what is left, and put it in its answer, so that a long +investigation narrows its queries before it runs into the ceiling rather than +after: + +```go +if b, ok := ratelimit.Remaining(ctx); ok { + // b.Used, b.Limit, b.Left(), b.ResetAt +} +``` + +It is the caller's finished requests in this window plus what this request has +charged so far — a snapshot, which their other requests still running will +change. `ok` is false with the budget off. + +Metrics: + +| Metric | Labels | +|---|---| +| `app_mcp_ratelimit_denied_total` | `reason`: `rpm`, `user_concurrent`, `global_concurrent`, `cost_budget` | +| `app_mcp_ratelimit_inflight` | `scope`: `user`, `global` | +| `app_mcp_ratelimit_charge_seconds` (histogram, 10ms–60s) | `label`: what you passed to `ChargeFor`, `unlabelled` for `Charge`, `wall_clock` for a request that charged nothing | +| `app_mcp_ratelimit_budget_used_seconds` | `user`: spent in the current window | + +The histogram adds up to the budget: exempt calls and a limiter with the budget +off are not in it. The `user` label is `Principal.UserID` as it is — where your +authentication names a user by e-mail, the e-mail is in `/metrics`. The gauge is +set as requests finish, drops to zero when the caller's next request opens a new +window or within fifteen minutes of the old one ending, and goes with an idle +caller's entry — an alert on `used / limit` fires for someone running out, not +for someone who already has it back. ## mcp and doc diff --git a/era.go b/era.go index 6f956e1..b9e4d85 100644 --- a/era.go +++ b/era.go @@ -158,27 +158,18 @@ func metaString(v *fastjson.Value, key string) string { // and always a 400 — every refusal decided before dispatch is one, which is why // fault carries no status of its own. // -// The id is echoed exactly as it arrived — "error responses MUST include the -// same ID as the request they correspond to" — and a request whose id could not -// be read gets null, which is what JSON-RPC asks for. +// The envelope is zenrpc's, the one every dispatched error goes out in, so a +// refusal here cannot come out in a shape of its own. The id is echoed exactly +// as it arrived — "error responses MUST include the same ID as the request they +// correspond to" — and a request whose id could not be read gets null, which is +// what JSON-RPC asks for. func writeFault(w http.ResponseWriter, id *fastjson.Value, f fault) { - body := struct { - JSONRPC string `json:"jsonrpc"` - ID json.RawMessage `json:"id"` - Error struct { - Code int `json:"code"` - Message string `json:"message"` - Data any `json:"data,omitempty"` - } `json:"error"` - }{JSONRPC: "2.0", ID: json.RawMessage("null")} + var raw *json.RawMessage if id != nil { - body.ID = id.MarshalTo(nil) + b := json.RawMessage(id.MarshalTo(nil)) + raw = &b } - body.Error.Code = f.code - body.Error.Message = f.message - body.Error.Data = f.data - - b, err := json.Marshal(body) + b, err := json.Marshal(zenrpc.NewResponseError(raw, f.code, f.message, f.data)) if err != nil { // a fixed shape plus data the caller supplied http.Error(w, f.message, http.StatusBadRequest) return diff --git a/example/main.go b/example/main.go index 0315831..ea40cee 100644 --- a/example/main.go +++ b/example/main.go @@ -315,7 +315,17 @@ func (helloTool) Call(ctx context.Context, args map[string]any) mcp.ToolCallResu if a.Who == "" { a.Who = "stranger" } + + // A tool that does work upstream says what it cost and where it went: the + // budget is spent by it, and app_mcp_ratelimit_charge_seconds is broken + // down by the label. The label names a target from the service's own + // catalogue, never a value from the request. A greeting costs next to + // nothing; the call is here to show where a real tool puts it. + start := time.Now() + greeting := "hello, " + a.Who + ratelimit.ChargeFor(ctx, "greeter", time.Since(start)) + // A struct rather than a map: the same type the output schema was reflected // from, so the answer cannot drift from what the tool promised. - return mcptool.OKResult(helloResult{Greeting: "hello, " + a.Who}) + return mcptool.OKResult(helloResult{Greeting: greeting}) } diff --git a/example/main_test.go b/example/main_test.go index ef8574e..adf341d 100644 --- a/example/main_test.go +++ b/example/main_test.go @@ -2,15 +2,22 @@ package main import ( "net/http" + "net/http/httptest" + "strings" "testing" + "time" + "github.com/vmkteam/mcpkit" "github.com/vmkteam/mcpkit/doc" "github.com/vmkteam/mcpkit/mcp" "github.com/vmkteam/mcpkit/mcptest" + "github.com/vmkteam/mcpkit/mcptool" + "github.com/vmkteam/mcpkit/ratelimit" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" "github.com/vmkteam/embedlog" + "github.com/vmkteam/zenrpc/v2" ) const testKey = "test-token" @@ -132,3 +139,65 @@ func TestExampleRefusesAnUnauthenticatedCall(t *testing.T) { assert.Equal(t, http.StatusUnauthorized, res.Status) assert.NotEmpty(t, res.Header.Get("WWW-Authenticate"), "a 401 has to say how to authenticate") } + +// A denial has to reach the client as an answer to its call, through the +// transport and in either era — the limiter's own tests stop at a stub +// handler. The example allows sixty calls a minute, so the burst runs out +// within the loop. +func TestExampleRefusalAnswersTheCall(t *testing.T) { + t.Parallel() + for _, era := range []mcptest.Era{mcptest.Modern, mcptest.Legacy} { + t.Run(string(era), func(t *testing.T) { + t.Parallel() + c := serve(t, era) + hello := map[string]any{"name": "hello", "arguments": map[string]any{}} + + var res mcptest.Response + for range 70 { + if res = c.Call(t, "tools/call", hello); res.Status != http.StatusOK { + break + } + } + require.Equal(t, http.StatusTooManyRequests, res.Status, "%s", res.Body) + assert.Equal(t, "application/json", res.Header.Get("Content-Type")) + assert.NotEmpty(t, res.Header.Get("Retry-After")) + require.NotNil(t, res.Error, "%s", res.Body) + assert.Equal(t, mcp.CodeRateLimited, res.Error.Code) + assert.Contains(t, string(res.Error.Data), `"reason":"rpm"`) + }) + } +} + +// An exempt name cannot be borrowed by repeating the key it is read from, in +// the same case or another: the limiter reads the first, the dispatcher runs +// the last. Built here rather +// than through newMCP, because it needs an exemption and a spent budget. +func TestExampleExemptionCannotBeBorrowed(t *testing.T) { + t.Parallel() + zsrv := zenrpc.NewServer(zenrpc.Options{}) + zsrv.RegisterAll(map[string]zenrpc.Invoker{ + mcpkit.NamespaceTools: ToolsService{registry: mcptool.NewRegistry(helloTool{})}, + }) + limiter := ratelimit.New(ratelimit.Config{ + CostBudgetPerHour: time.Nanosecond, // the first call spends it + Exempt: func(_, name string) bool { return name == "help" }, + }) + t.Cleanup(limiter.Stop) + h := limiter.Middleware(mcpkit.NewServer(zsrv, embedlog.Logger{}), embedlog.Logger{}) + + post := func(body string) int { + rec := httptest.NewRecorder() + h.ServeHTTP(rec, httptest.NewRequest(http.MethodPost, "/mcp", strings.NewReader(body))) + return rec.Code + } + hello := `{"jsonrpc":"2.0","id":1,"method":"tools/call","params":{"name":"hello","arguments":{}}}` + require.Equal(t, http.StatusOK, post(hello)) + require.Equal(t, http.StatusTooManyRequests, post(hello), "the budget is spent") + + for _, borrowed := range []string{ + `{"jsonrpc":"2.0","id":2,"method":"tools/call","params":{"name":"help","name":"hello","arguments":{}}}`, + `{"jsonrpc":"2.0","id":2,"method":"tools/call","params":{"name":"help","Name":"hello","arguments":{}}}`, + } { + assert.NotEqual(t, http.StatusOK, post(borrowed), "hello ran under help's exemption: %s", borrowed) + } +} diff --git a/internal/metrics/metrics.go b/internal/metrics/metrics.go index 7b7f1ea..59e82b7 100644 --- a/internal/metrics/metrics.go +++ b/internal/metrics/metrics.go @@ -104,6 +104,24 @@ func (g *Group) Gauge(name, help, label string, values ...string) *prometheus.Ga return gv } +// Histogram declares a histogram with one label, and the series it starts at +// zero. +func (g *Group) Histogram(name, help, label string, buckets []float64, values ...string) *prometheus.HistogramVec { + h := prometheus.NewHistogramVec(prometheus.HistogramOpts{ + Namespace: Namespace, + Subsystem: Subsystem, + Name: name, + Help: help, + Buckets: buckets, + }, []string{label}) + g.add(h, func() { + for _, v := range values { + h.WithLabelValues(v) + } + }) + return h +} + // add collects a metric and how to warm it. It runs while the declaring // package's var block does, which is before anything can reach Register: the // slices are unguarded because at that point there is nobody to guard them diff --git a/internal/metrics/metrics_test.go b/internal/metrics/metrics_test.go index 4f85916..d70c644 100644 --- a/internal/metrics/metrics_test.go +++ b/internal/metrics/metrics_test.go @@ -4,6 +4,7 @@ import ( "sync" "testing" + "github.com/prometheus/client_golang/prometheus" "github.com/prometheus/client_golang/prometheus/testutil" "github.com/stretchr/testify/assert" ) @@ -56,6 +57,17 @@ func TestGaugeStartsItsSeriesAtZero(t *testing.T) { assert.Equal(t, 2, testutil.CollectAndCount(gauge)) } +func TestHistogramStartsItsSeriesAtZero(t *testing.T) { + t.Parallel() + g := NewGroup() + h := g.Histogram("warmed_seconds", "help.", "label", []float64{1, 10}, "a", "b") + + assert.Equal(t, 0, testutil.CollectAndCount(h), "nothing exists before Register") + g.Register() + assert.Equal(t, 2, testutil.CollectAndCount(h), "one series per value, from the start") + assert.Contains(t, h.WithLabelValues("a").(prometheus.Metric).Desc().String(), "app_mcp_warmed_seconds") +} + // Register is called from every entry point that touches a metric, so it is // reached concurrently and reached often. Registering the same collector twice // panics, which is what the Once is for. diff --git a/mcp/meta.go b/mcp/meta.go index bed09fe..f24892a 100644 --- a/mcp/meta.go +++ b/mcp/meta.go @@ -48,6 +48,17 @@ const ( CodeUnsupportedProtocolVersion = -32022 ) +// Error codes this library defines itself. JSON-RPC leaves -32000..-32099 to +// the implementation, the spec has taken -32020 upward (above), and that leaves +// -32000..-32019 to us. Nothing in mcpkit or in the services built on it +// emitted a code from there when the first one was taken, so this list is the +// whole of what is in use. +const ( + // CodeRateLimited is a request the rate limiter turned away. Its data says + // which limit it was and how many seconds until it is worth asking again. + CodeRateLimited = -32010 +) + // The sentinel wrapping a header value that cannot be written as plain ASCII. const ( base64Prefix = "=?base64?" diff --git a/mcptest/mcptest_test.go b/mcptest/mcptest_test.go index 592070d..8a6dd14 100644 --- a/mcptest/mcptest_test.go +++ b/mcptest/mcptest_test.go @@ -161,11 +161,11 @@ func TestOptions(t *testing.T) { t.Run("client identity and capabilities are the caller's to set", func(t *testing.T) { t.Parallel() c, got := recorder(t, - WithClientInfo("ringsrv-test", "2.0"), + WithClientInfo("example-client", "2.0"), WithClientCapabilities(map[string]any{"roots": map[string]any{}})) c.Call(t, "tools/list", nil) meta := got.meta(t) - assert.Equal(t, map[string]any{"name": "ringsrv-test", "version": "2.0"}, meta[mcp.MetaClientInfo]) + assert.Equal(t, map[string]any{"name": "example-client", "version": "2.0"}, meta[mcp.MetaClientInfo]) assert.Contains(t, meta[mcp.MetaClientCapabilities], "roots") }) diff --git a/metrics.go b/metrics.go index 0228071..474bc42 100644 --- a/metrics.go +++ b/metrics.go @@ -22,6 +22,8 @@ const ( reasonHeaderMismatch = "header_mismatch" reasonBadVersion = "bad_version" reasonMissingMeta = "missing_meta" + reasonRepeatedKey = "repeated_key" + reasonMissingID = "missing_id" ) // era labels on app_mcp_requests_total. @@ -41,7 +43,7 @@ var ( "reason", reasonBatch, reasonParse, reasonTooLarge, reasonReadBody, reasonDispatch, reasonBadMethod, reasonOrigin, reasonHost, - reasonHeaderMismatch, reasonBadVersion, reasonMissingMeta, + reasonHeaderMismatch, reasonBadVersion, reasonMissingMeta, reasonRepeatedKey, reasonMissingID, ) // Which era clients actually speak. This is the number that decides when diff --git a/ratelimit/concurrent_test.go b/ratelimit/concurrent_test.go index 71e86fb..3fe6532 100644 --- a/ratelimit/concurrent_test.go +++ b/ratelimit/concurrent_test.go @@ -137,7 +137,7 @@ func TestPerUserConcurrentReleasesEverySlot(t *testing.T) { w.WriteHeader(http.StatusOK) }), embedlog.Logger{}) - call := func() int { + send := func() int { req := httptest.NewRequest(http.MethodPost, "/mcp", strings.NewReader("{}")) req = req.WithContext(auth.NewContext(req.Context(), auth.Principal{UserID: "alice"})) rec := httptest.NewRecorder() @@ -147,14 +147,14 @@ func TestPerUserConcurrentReleasesEverySlot(t *testing.T) { var wg sync.WaitGroup for range burst { - wg.Go(func() { call() }) + wg.Go(func() { send() }) } wg.Wait() assert.LessOrEqual(t, peak.Load(), int64(limit), "PerUserConcurrent exceeded") // Every slot came back: the next request is served, not refused. - assert.Equal(t, http.StatusOK, call(), "a slot was leaked — the caller is throttled forever") + assert.Equal(t, http.StatusOK, send(), "a slot was leaked — the caller is throttled forever") l.mu.Lock() defer l.mu.Unlock() @@ -184,7 +184,7 @@ func TestAcquireAndEvictRunTogether(t *testing.T) { wg.Go(func() { user := "user-" + string(rune('a'+i)) for range 100 { - release, reason := l.acquire(user, time.Now()) + release, reason := l.acquire(user, time.Now(), false) if release == nil { require.NotEmpty(t, reason) continue @@ -205,7 +205,7 @@ func TestEvictSkipsBusyWithoutAConcurrencyLimit(t *testing.T) { l := newLimiter(t, Config{CostBudgetPerHour: time.Hour}) now := time.Now() - release, reason := l.acquire("alice", now) + release, reason := l.acquire("alice", now, false) require.NotNil(t, release, reason) l.evictIdle(now.Add(idleTTL + time.Hour)) @@ -222,3 +222,36 @@ func TestEvictSkipsBusyWithoutAConcurrencyLimit(t *testing.T) { l.evictIdle(now.Add(idleTTL + time.Hour)) assert.NotContains(t, l.entries, "alice", "once released it is idle and goes") } + +// Remaining reads the entry that release is writing, from requests of one +// caller that run at once. Each sees a consistent snapshot — at least its own +// charge, at most everyone's — and once all are done the budget is the sum. +func TestRemainingUnderConcurrency(t *testing.T) { + const callers = 16 + l := newLimiter(t, Config{CostBudgetPerHour: time.Hour}) + + var bad atomic.Int64 + h := l.Middleware(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + Charge(r.Context(), time.Second) + b, ok := Remaining(r.Context()) + if !ok || b.Used < time.Second || b.Used > callers*time.Second { + bad.Add(1) + } + w.WriteHeader(http.StatusOK) + }), embedlog.Logger{}) + + var wg sync.WaitGroup + for range callers { + wg.Go(func() { + req := httptest.NewRequest(http.MethodPost, "/mcp", strings.NewReader("{}")) + req = req.WithContext(auth.NewContext(req.Context(), auth.Principal{UserID: "alice"})) + h.ServeHTTP(httptest.NewRecorder(), req) + }) + } + wg.Wait() + + assert.Zero(t, bad.Load(), "a snapshot outside what the requests could have spent") + b, ok := l.budget("alice", 0, time.Now()) + require.True(t, ok) + assert.Equal(t, callers*time.Second, b.Used) +} diff --git a/ratelimit/cost.go b/ratelimit/cost.go index 2985cd1..9e41e27 100644 --- a/ratelimit/cost.go +++ b/ratelimit/cost.go @@ -1,8 +1,8 @@ package ratelimit // Cost accounting for the limiter. It lives beside the limiter because it is -// the limiter's unit of work, and a handler that reports it needs only the two -// exported functions. +// the limiter's unit of work, and a handler that reports it or asks what is +// left needs only the exported functions here. import ( "context" @@ -24,6 +24,13 @@ type Cost struct { mu sync.Mutex work time.Duration units int + + // Set by the middleware before the handler runs, and zero on an + // accumulator made by WithCost alone: whose budget this is, for Remaining, + // and whether the call is exempt from it. + limiter *Limiter + userID string + exempt bool } // charge is what one served request costs: work against the hourly budget, @@ -44,26 +51,106 @@ func WithCost(ctx context.Context) (context.Context, *Cost) { // Charge records one piece of billable work done while serving ctx. Safe to // call concurrently, and a no-op when ctx carries no accumulator — a handler // must not have to know whether rate limiting is switched on. -func Charge(ctx context.Context, work time.Duration) { +// +// It is ChargeFor without a label. +func Charge(ctx context.Context, work time.Duration) { ChargeFor(ctx, "", work) } + +// ChargeFor is Charge that says what did the work, and the label is what +// app_mcp_ratelimit_charge_seconds breaks the budget down by — the question +// "who ate the budget" answered from Prometheus rather than from an audit log. +// +// The label is per charge, not per request: the calls of one request can go +// to different upstreams, and a label per request would fold together exactly +// the split that is worth seeing. +// +// Each label is a series, so it must be a name from the service's own +// catalogue — "grafana", "postgres" — and never a value from the request: a +// query, a URL, a user. That keeps the cardinality the size of the catalogue. +func ChargeFor(ctx context.Context, label string, work time.Duration) { c, ok := ctx.Value(costKey{}).(*Cost) if !ok { return } c.mu.Lock() - defer c.mu.Unlock() c.work += work c.units++ + c.mu.Unlock() + + if c.observed() { + if label == "" { + label = labelUnlabelled + } + charged().WithLabelValues(label).Observe(work.Seconds()) + } +} + +// Budget is a caller's hourly budget at one moment. +type Budget struct { + Used time.Duration + Limit time.Duration + ResetAt time.Time +} + +// Left is what remains of the budget, never below zero. +func (b Budget) Left() time.Duration { return max(0, b.Limit-b.Used) } + +// Remaining reports the budget of the caller ctx is serving: what their +// finished requests have spent in this window, plus what this request has +// charged so far. ok is false when the budget is off or ctx did not come +// through the middleware. +// +// It is a snapshot. Other requests of the same caller that are still running +// are not in it until they finish, and change it when they do. An exempt call +// charges the budget nothing, so what it charged is not added either. +// +// It is what lets a long investigation narrow its own queries before it runs +// into the ceiling instead of after: the service puts it in its answers. +func Remaining(ctx context.Context) (b Budget, ok bool) { + c, ok := ctx.Value(costKey{}).(*Cost) + if !ok || c.limiter == nil { + return Budget{}, false + } + var own time.Duration + if !c.exempt { + c.mu.Lock() + own = c.work + c.mu.Unlock() + } + return c.limiter.budget(c.userID, own, c.limiter.clock()) } // settle prices the finished request. Nothing charged means nothing billable // ran — an initialize, a tools/list, a refusal before the upstream — and those // are priced by the wall clock, exactly as every request was before one // request could carry many calls. +// +// An exempt call spends no budget whatever it did; its units still count +// against the rate. func (c *Cost) settle(wall time.Duration) charge { c.mu.Lock() - defer c.mu.Unlock() - if c.units == 0 { - return charge{work: wall, units: 1} + ch := charge{work: c.work, units: c.units} + c.mu.Unlock() + + if ch.units == 0 { + ch = charge{work: wall, units: 1} + // Every charged piece was observed as it was charged; only the wall + // clock is left to observe, so the histogram sums to what the budget + // was spent. + if c.observed() { + charged().WithLabelValues(labelWallClock).Observe(wall.Seconds()) + } + } + if c.exempt { + ch.work = 0 } - return charge{work: c.work, units: c.units} + return ch +} + +// observed reports whether a charge made here goes into the histogram: only +// when it goes into a budget. An accumulator made by WithCost alone belongs to +// no limiter, a limiter may have the budget off, and an exempt call spends +// nothing. What it reads is set before the handler runs and never changes, so +// it needs no lock. +func (c *Cost) observed() bool { + return c.limiter != nil && c.limiter.cfg.CostBudgetPerHour > 0 && !c.exempt } diff --git a/ratelimit/exempt_test.go b/ratelimit/exempt_test.go new file mode 100644 index 0000000..ef24b7a --- /dev/null +++ b/ratelimit/exempt_test.go @@ -0,0 +1,121 @@ +package ratelimit + +import ( + "io" + "net/http" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "github.com/vmkteam/embedlog" +) + +const helpCall = `{"jsonrpc":"2.0","id":8,"method":"tools/call","params":{"name":"help","arguments":{}}}` + +func exemptHelp(method, name string) bool { return method == "tools/call" && name == "help" } + +// spendAll is a handler every call of which costs twice the budget, so one +// billed call is enough to spend it. +func spendAll(w http.ResponseWriter, r *http.Request) { + Charge(r.Context(), 2*time.Minute) + w.WriteHeader(http.StatusOK) +} + +// Every MCP tool goes through tools/call, so the method alone cannot tell help +// from a query: the name is what separates them. +func TestExemptPassesWhenTheBudgetIsGone(t *testing.T) { + l := newLimiter(t, Config{CostBudgetPerHour: time.Minute, Exempt: exemptHelp}) + h := l.Middleware(http.HandlerFunc(spendAll), embedlog.Logger{}) + + require.Equal(t, http.StatusOK, post(h, toolCall).Code, "the first call spends the budget") + assert.Equal(t, http.StatusTooManyRequests, post(h, toolCall).Code, "the next one is refused") + assert.Equal(t, http.StatusOK, post(h, helpCall).Code, "help is not") + assert.Equal(t, http.StatusOK, post(h, helpCall).Code, "and never is") +} + +// What an exempt call charges does not reach the budget, however much it is. +func TestExemptChargesNothing(t *testing.T) { + l := newLimiter(t, Config{CostBudgetPerHour: time.Minute, Exempt: exemptHelp}) + h := l.Middleware(http.HandlerFunc(spendAll), embedlog.Logger{}) + + require.Equal(t, http.StatusOK, post(h, helpCall).Code) + require.Equal(t, http.StatusOK, post(h, helpCall).Code) + assert.Equal(t, http.StatusOK, post(h, toolCall).Code, "two help calls left the budget whole") + assert.Equal(t, http.StatusTooManyRequests, post(h, toolCall).Code) +} + +// The budget is all an exemption spares. The rate still counts, or an exempt +// name would be a way to call at any speed. +func TestExemptStillPaysTheRate(t *testing.T) { + l := newLimiter(t, Config{PerUserRPM: 1, CostBudgetPerHour: time.Minute, Exempt: exemptHelp}) + h := l.Middleware(http.HandlerFunc(okHandler), embedlog.Logger{}) + + require.Equal(t, http.StatusOK, post(h, helpCall).Code) + rec := post(h, helpCall) + assert.Equal(t, http.StatusTooManyRequests, rec.Code) + assert.Contains(t, rec.Body.String(), reasonRPM) +} + +// Exempt is asked with what Mcp-Name would carry, for whichever method carries +// one — and not at all about a body there is no call in. +func TestExemptIsAskedAboutTheCall(t *testing.T) { + type asked struct{ method, name string } + var got []asked + l := newLimiter(t, Config{CostBudgetPerHour: time.Minute, Exempt: func(method, name string) bool { + got = append(got, asked{method, name}) + return true + }}) + h := l.Middleware(http.HandlerFunc(okHandler), embedlog.Logger{}) + + post(h, helpCall) + post(h, `{"jsonrpc":"2.0","id":2,"method":"resources/read","params":{"uri":"doc://help"}}`) + post(h, `{"jsonrpc":"2.0","id":3,"method":"tools/list"}`) + post(h, `[`+helpCall+`]`) + post(h, `not json`) + + assert.Equal(t, []asked{ + {"tools/call", "help"}, + {"resources/read", "doc://help"}, + {"tools/list", ""}, + }, got) +} + +// A body the limiter could not read is not exempt, whatever Exempt would have +// said about it. +func TestExemptDoesNotCoverAnUnreadableBody(t *testing.T) { + l := newLimiter(t, Config{CostBudgetPerHour: time.Minute, Exempt: func(string, string) bool { return true }}) + h := l.Middleware(http.HandlerFunc(spendAll), embedlog.Logger{}) + + require.Equal(t, http.StatusOK, post(h, `[`+helpCall+`]`).Code, "billed, since it is not exempt") + assert.Equal(t, http.StatusTooManyRequests, post(h, `[`+helpCall+`]`).Code) +} + +// With Exempt set the limiter reads every body, and the handler behind still +// reads all of it. +func TestExemptLeavesTheBodyWhole(t *testing.T) { + var seen []string + l := newLimiter(t, Config{CostBudgetPerHour: time.Minute, Exempt: exemptHelp}) + h := l.Middleware(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + b, err := io.ReadAll(r.Body) + assert.NoError(t, err) + seen = append(seen, string(b)) + w.WriteHeader(http.StatusOK) + }), embedlog.Logger{}) + + post(h, helpCall) + post(h, toolCall) + assert.Equal(t, []string{helpCall, toolCall}, seen) +} + +// An exempt name cannot be borrowed by a call that repeats the key: the +// dispatcher would run the last one, and the limiter will not guess. +func TestExemptIsNotFooledByARepeatedKey(t *testing.T) { + l := newLimiter(t, Config{CostBudgetPerHour: time.Minute, Exempt: exemptHelp}) + h := l.Middleware(http.HandlerFunc(spendAll), embedlog.Logger{}) + require.Equal(t, http.StatusOK, post(h, toolCall).Code, "the budget is spent") + + borrowed := `{"jsonrpc":"2.0","id":9,"method":"tools/call","params":{"name":"help","name":"api_call","arguments":{}}}` + assert.Equal(t, http.StatusTooManyRequests, post(h, borrowed).Code) + assert.Equal(t, http.StatusOK, post(h, helpCall).Code, "help itself still passes") +} diff --git a/ratelimit/metrics.go b/ratelimit/metrics.go index c9d0798..634cb19 100644 --- a/ratelimit/metrics.go +++ b/ratelimit/metrics.go @@ -20,6 +20,22 @@ const ( scopeGlobal = "global" ) +// The label values of app_mcp_ratelimit_charge_seconds the library writes on +// its own. Everything else is a name a service passed to ChargeFor. +const ( + // labelUnlabelled is Charge, or ChargeFor with no label. + labelUnlabelled = "unlabelled" + // labelWallClock is a request that charged nothing and was priced by its + // wall clock. It is what makes the histogram add up to the budget: without + // it a service that never calls Charge would spend its budget invisibly. + labelWallClock = "wall_clock" +) + +// chargeBuckets run from 10ms to a minute. An upstream call takes anywhere from +// a few milliseconds to a timeout of tens of seconds, and everything worth +// telling apart is in the upper half. +var chargeBuckets = []float64{.01, .025, .05, .1, .25, .5, 1, 2.5, 5, 10, 15, 20, 30, 60} + // Metrics of this package are registered once, on first use, in the default // registry — the same one the service publishes, so registering twice would // panic and the Group's sync.Once is what keeps that from happening. @@ -42,10 +58,39 @@ var ( "scope", scopeUser, scopeGlobal, ) + + chargeSeconds = group.Histogram( + "ratelimit_charge_seconds", + "Work charged against the hourly budget, labelled by what did it.", + "label", + chargeBuckets, + labelUnlabelled, labelWallClock, + ) + + // Not warmed: there is no user to start from until one is served. A series + // appears with a user's first settled request and goes with their entry. + // + // The label is Principal.UserID as it is. Where authentication calls a + // user by e-mail, that e-mail is in /metrics. + budgetUsedGauge = group.Gauge( + "ratelimit_budget_used_seconds", + "Hourly budget each caller has spent in the current window.", + "user", + ) ) func registerMetrics() { group.Register() } +func charged() *prometheus.HistogramVec { + registerMetrics() + return chargeSeconds +} + +func budgetUsed() *prometheus.GaugeVec { + registerMetrics() + return budgetUsedGauge +} + func metric() *prometheus.CounterVec { registerMetrics() return deniedTotal diff --git a/ratelimit/metrics_test.go b/ratelimit/metrics_test.go new file mode 100644 index 0000000..5a0a22a --- /dev/null +++ b/ratelimit/metrics_test.go @@ -0,0 +1,195 @@ +package ratelimit + +import ( + "io" + "net/http" + "strings" + "testing" + "time" + + "github.com/prometheus/client_golang/prometheus" + "github.com/prometheus/client_golang/prometheus/testutil" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "github.com/vmkteam/embedlog" +) + +// The histogram is one per process and every test here writes to it, so these +// read differences on labels of their own, and none of them is parallel. + +type observations struct { + count uint64 + sum float64 + found bool +} + +// histogram reads one series of app_mcp_ratelimit_charge_seconds. +func histogram(t *testing.T, label string) observations { + t.Helper() + registerMetrics() + reg := prometheus.NewRegistry() + require.NoError(t, reg.Register(chargeSeconds)) + families, err := reg.Gather() + require.NoError(t, err) + for _, f := range families { + for _, m := range f.GetMetric() { + for _, l := range m.GetLabel() { + if l.GetName() == "label" && l.GetValue() == label { + h := m.GetHistogram() + return observations{count: h.GetSampleCount(), sum: h.GetSampleSum(), found: true} + } + } + } + } + return observations{} +} + +// The two labels the library writes itself are there before anything is +// charged; the ones a service names cannot be. +func TestChargeHistogramStartsAtZero(t *testing.T) { + assert.True(t, histogram(t, labelUnlabelled).found) + assert.True(t, histogram(t, labelWallClock).found) +} + +// Each charge is one observation under its own label, so the calls of one +// request to different upstreams stay apart. +func TestChargeForIsObservedByLabel(t *testing.T) { + grafana, sentry, unlabelled, wall := histogram(t, "by-label-grafana"), histogram(t, "by-label-sentry"), + histogram(t, labelUnlabelled), histogram(t, labelWallClock) + + l := newLimiter(t, Config{CostBudgetPerHour: time.Hour}) + h := l.Middleware(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + ChargeFor(r.Context(), "by-label-grafana", 3*time.Second) + ChargeFor(r.Context(), "by-label-grafana", 2*time.Second) + ChargeFor(r.Context(), "by-label-sentry", time.Second) + Charge(r.Context(), 500*time.Millisecond) + w.WriteHeader(http.StatusOK) + }), embedlog.Logger{}) + post(h, toolCall) + + got := histogram(t, "by-label-grafana") + assert.Equal(t, grafana.count+2, got.count) + assert.InDelta(t, grafana.sum+5, got.sum, 1e-9) + got = histogram(t, "by-label-sentry") + assert.Equal(t, sentry.count+1, got.count) + assert.InDelta(t, sentry.sum+1, got.sum, 1e-9) + got = histogram(t, labelUnlabelled) + assert.Equal(t, unlabelled.count+1, got.count, "Charge is ChargeFor with no label") + assert.InDelta(t, unlabelled.sum+0.5, got.sum, 1e-9) + assert.Equal(t, wall.count, histogram(t, labelWallClock).count, + "a request that charged anything is not priced by its wall clock") +} + +// The histogram adds up to the budget: what a request charged, and what one +// that charged nothing cost by its wall clock. +func TestChargeHistogramAddsUpToTheBudget(t *testing.T) { + charged0, wall0 := histogram(t, "adds-up"), histogram(t, labelWallClock) + + l := newLimiter(t, Config{CostBudgetPerHour: time.Hour}) + h := l.Middleware(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + body, _ := io.ReadAll(r.Body) + if !strings.Contains(string(body), `"help"`) { + ChargeFor(r.Context(), "adds-up", 4*time.Second) + } + w.WriteHeader(http.StatusOK) + }), embedlog.Logger{}) + post(h, toolCall) // four seconds upstream + post(h, helpCall) // nothing charged: priced by its wall clock + + charged, wall := histogram(t, "adds-up"), histogram(t, labelWallClock) + assert.Equal(t, charged0.count+1, charged.count) + assert.Equal(t, wall0.count+1, wall.count, "the request that charged nothing") + + spent, ok := l.budget(anonymousUser, 0, time.Now()) + require.True(t, ok) + assert.InDelta(t, spent.Used.Seconds(), (charged.sum-charged0.sum)+(wall.sum-wall0.sum), 1e-6) +} + +// Only work that goes into a budget is observed: an exempt call spends none, +// and neither does a limiter with the budget off. +func TestChargeHistogramSkipsWhatSpendsNothing(t *testing.T) { + l := newLimiter(t, Config{CostBudgetPerHour: time.Hour, Exempt: exemptHelp}) + h := l.Middleware(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + ChargeFor(r.Context(), "skips-exempt", time.Second) + w.WriteHeader(http.StatusOK) + }), embedlog.Logger{}) + post(h, helpCall) + assert.False(t, histogram(t, "skips-exempt").found, "an exempt call") + + l = newLimiter(t, Config{PerUserRPM: 10}) + h = l.Middleware(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + ChargeFor(r.Context(), "skips-off", time.Second) + w.WriteHeader(http.StatusOK) + }), embedlog.Logger{}) + post(h, toolCall) + assert.False(t, histogram(t, "skips-off").found, "no budget") +} + +// The gauge is what a caller has spent in the window that is running: set as +// requests settle, zero once the window is over, gone with the entry. +func TestBudgetUsedGauge(t *testing.T) { + const user = "gauge-alice" + l := newLimiter(t, Config{CostBudgetPerHour: time.Hour}) + opened := time.Date(2026, 1, 1, 12, 0, 0, 0, time.UTC) + l.clock = func() time.Time { return opened } + + release, _ := l.acquire(user, opened, false) + require.NotNil(t, release) + release(charge{work: 90 * time.Second, units: 1}) + assert.InDelta(t, 90, testutil.ToFloat64(budgetUsedGauge.WithLabelValues(user)), 1e-9) + + l.evictIdle(opened.Add(30 * time.Minute)) + assert.InDelta(t, 90, testutil.ToFloat64(budgetUsedGauge.WithLabelValues(user)), 1e-9, + "the window is still running") + + l.evictIdle(opened.Add(budgetWindow)) + require.Contains(t, l.entries, user, "an hour idle is not idle enough to go") + assert.Zero(t, testutil.ToFloat64(budgetUsedGauge.WithLabelValues(user)), + "the window is over: the caller has the whole budget back") + + l.evictIdle(opened.Add(idleTTL + time.Minute)) + require.NotContains(t, l.entries, user) + assert.False(t, budgetUsedGauge.DeleteLabelValues(user), "the series went with the entry") +} + +// The gauge follows the window, not the lazy roll of the entry: a caller who +// comes back is shown the fresh window at once, and a request that finished +// after its window ended does not bring the dead one back. +func TestBudgetUsedGaugeFollowsTheWindow(t *testing.T) { + gauge := func(user string) float64 { return testutil.ToFloat64(budgetUsedGauge.WithLabelValues(user)) } + opened := time.Date(2026, 1, 1, 12, 0, 0, 0, time.UTC) + + t.Run("the caller comes back", func(t *testing.T) { + const user = "gauge-back" + l := newLimiter(t, Config{CostBudgetPerHour: time.Hour}) + l.clock = func() time.Time { return opened } + release, _ := l.acquire(user, opened, false) + release(charge{work: 90 * time.Second, units: 1}) + require.InDelta(t, 90, gauge(user), 1e-9) + + held, _ := l.acquire(user, opened.Add(budgetWindow+time.Minute), false) + require.NotNil(t, held) + assert.Zero(t, gauge(user), "the window rolled while the request still runs") + held(charge{}) + }) + + t.Run("a request outlives its window", func(t *testing.T) { + const user = "gauge-late" + l := newLimiter(t, Config{CostBudgetPerHour: time.Hour}) + slow, _ := l.acquire(user, opened, false) + require.NotNil(t, slow) + l.clock = func() time.Time { return opened.Add(budgetWindow + time.Minute) } + slow(charge{work: 90 * time.Second, units: 1}) + assert.Zero(t, gauge(user), "what lands in a window that is over is not shown as spent") + }) + + t.Run("a caller who never settled", func(t *testing.T) { + const user = "gauge-never" + l := newLimiter(t, Config{CostBudgetPerHour: time.Hour}) + held, _ := l.acquire(user, opened, false) + require.NotNil(t, held) + defer held(charge{}) + l.evictIdle(opened.Add(budgetWindow)) + assert.False(t, budgetUsedGauge.DeleteLabelValues(user), "eviction made a series for nobody") + }) +} diff --git a/ratelimit/ratelimit.go b/ratelimit/ratelimit.go index 88b8b1c..b8a4f12 100644 --- a/ratelimit/ratelimit.go +++ b/ratelimit/ratelimit.go @@ -34,6 +34,9 @@ const idleTTL = 2 * time.Hour // evictEvery is how often the eviction loop runs. const evictEvery = 15 * time.Minute +// budgetWindow is the length of the fixed cost-budget window. +const budgetWindow = time.Hour + // Config is the set of limits. Each one is switched off by a value ≤ 0, and // all four off means Disabled: Middleware then returns the handler unwrapped. type Config struct { @@ -41,6 +44,19 @@ type Config struct { PerUserConcurrent int GlobalConcurrent int CostBudgetPerHour time.Duration + + // Exempt names the calls the hourly budget does not apply to. They are + // the cheap ones that teach a caller to spend less — a help tool, a static + // resource — and without this they are the first to stop answering once + // the budget is gone. name is what Mcp-Name mirrors: the tool of + // tools/call, the prompt of prompts/get, the URI of resources/read; empty + // for any other method. + // + // Only the budget. An exempt call still takes a rate token and a + // concurrency slot, or it would be a hole anything could be pushed through + // at any rate. Setting it costs a parse of every POST body while the budget + // is on; nil exempts nothing and reads a body only to answer a refusal. + Exempt func(method, name string) bool } type entry struct { @@ -65,13 +81,16 @@ type Limiter struct { stop chan struct{} stopOnce sync.Once + + // clock is time.Now; a test sets it to say when a request finished. + clock func() time.Time } // New returns a limiter and, unless it is disabled, starts the goroutine that // evicts idle entries. Call Stop when done — in tests especially, where the // goroutine would otherwise outlive the test. func New(cfg Config) *Limiter { - l := &Limiter{cfg: cfg, entries: map[string]*entry{}, stop: make(chan struct{})} + l := &Limiter{cfg: cfg, entries: map[string]*entry{}, stop: make(chan struct{}), clock: time.Now} if cfg.GlobalConcurrent > 0 { l.global = make(chan struct{}, cfg.GlobalConcurrent) } @@ -118,16 +137,36 @@ func (l *Limiter) Middleware(next http.Handler, logger embedlog.Logger) http.Han if p, ok := auth.PrincipalFromContext(r.Context()); ok && p.UserID != "" { userID = p.UserID } - release, reason := l.acquire(userID, time.Now()) + // The body is read here when Exempt needs it, and otherwise only to + // answer a refusal. Exempt is about the budget alone, so with the budget + // off there is nothing to ask it. A body that cannot be read is not a + // call anyone could have exempted, so Exempt is never asked about one. + var ( + c call + read, exempt bool + ) + if l.cfg.Exempt != nil && l.cfg.CostBudgetPerHour > 0 { + var ok bool + c, ok = peek(r, maxPeekBytes) + read, exempt = true, ok && l.cfg.Exempt(c.method, c.name) + } + + now := l.clock() + release, reason := l.acquire(userID, now, exempt) if release == nil { - logger.Error(r.Context(), "mcp rate limit", "user", userID, "reason", reason) - w.Header().Set("Retry-After", "10") - http.Error(w, "rate limit: "+reason, http.StatusTooManyRequests) + if !read { + c, _ = peek(r, maxRefusalPeekBytes) + } + rf := l.refusalFor(userID, reason, now) + logger.Error(r.Context(), "mcp rate limit", rf.logArgs(userID)...) + refuse(w, rf, c.id) return } // The handler prices itself through the accumulator; the wall clock is // only the fallback for a request that did no billable work of its own. + // The accumulator also knows whose budget it is, for Remaining. ctx, cost := WithCost(r.Context()) + cost.limiter, cost.userID, cost.exempt = l, userID, exempt start := time.Now() defer func() { release(cost.settle(time.Since(start))) }() next.ServeHTTP(w, r.WithContext(ctx)) @@ -147,31 +186,41 @@ func (l *Limiter) evictLoop(every time.Duration) { } } -// evictIdle drops entries idle longer than idleTTL that are serving nothing. -// Idempotent — safe to call from tests. +// evictIdle drops entries idle longer than idleTTL that are serving nothing, +// and their budget series with them. Idempotent — safe to call from tests. +// +// An entry whose budget window is over but which is not idle enough to go has +// its series put to zero, if it spent anything to have one. The window rolls +// only when the caller comes back, and until then the gauge would go on showing +// the old one — for up to two hours, with an alert on it firing for someone who +// has the whole budget back. func (l *Limiter) evictIdle(now time.Time) { cutoff := now.Add(-idleTTL) l.mu.Lock() defer l.mu.Unlock() for k, e := range l.entries { - if e.lastSeen.Before(cutoff) && e.inflight == 0 { + switch { + case e.lastSeen.Before(cutoff) && e.inflight == 0: delete(l.entries, k) + budgetUsed().DeleteLabelValues(k) + case e.used > 0 && !now.Before(e.resetAt): + budgetUsed().WithLabelValues(k).Set(0) } } } // entryFor returns this user's bucket, creating it on first sight, and reports -// the cost budget if it is already spent. +// the cost budget if it is already spent — unless the call is exempt from it. // // The budget is read here, before anything else is taken: it is a comparison, // and a call refused by it should cost nothing. -func (l *Limiter) entryFor(userID string, now time.Time) (e *entry, reason string) { +func (l *Limiter) entryFor(userID string, now time.Time, exempt bool) (e *entry, reason string) { l.mu.Lock() defer l.mu.Unlock() e, ok := l.entries[userID] if !ok { - e = &entry{resetAt: now.Add(time.Hour)} + e = &entry{resetAt: now.Add(budgetWindow)} if l.cfg.PerUserRPM > 0 { e.rpm = rate.NewLimiter(rate.Limit(float64(l.cfg.PerUserRPM)/60.0), maxBurst(l.cfg.PerUserRPM)) } @@ -182,17 +231,56 @@ func (l *Limiter) entryFor(userID string, now time.Time) (e *entry, reason strin } // Roll the cost-budget window first, so a user who comes back after an hour // of being rpm-throttled gets the budget reset instead of the stale `used`. - if now.After(e.resetAt) { + // + // At resetAt, not after it: a denial tells the caller to come back in + // resetAt−now seconds, and one who does so to the nanosecond must find the + // window rolled. + if !now.Before(e.resetAt) { + if e.used > 0 { + budgetUsed().WithLabelValues(userID).Set(0) + } e.used = 0 - e.resetAt = now.Add(time.Hour) + e.resetAt = now.Add(budgetWindow) } e.lastSeen = now - if l.cfg.CostBudgetPerHour > 0 && e.used >= l.cfg.CostBudgetPerHour { + if !exempt && l.cfg.CostBudgetPerHour > 0 && e.used >= l.cfg.CostBudgetPerHour { return e, reasonCostBudget } return e, "" } +// budget is userID's budget at now with own added to what is already spent — +// the part of a running request that release has not paid in yet. ok is false +// with the budget off. +// +// The entry is there for as long as a request of the user is: eviction skips +// an entry with anything in flight. +func (l *Limiter) budget(userID string, own time.Duration, now time.Time) (b Budget, ok bool) { + if l.cfg.CostBudgetPerHour <= 0 { + return Budget{}, false + } + l.mu.Lock() + defer l.mu.Unlock() + e, ok := l.entries[userID] + if !ok { + return Budget{}, false + } + return l.budgetOf(e, own, now), true +} + +// budgetOf is the budget of e at now, own added to what is spent. Called with +// l.mu held, so a refusal and Remaining read one number the same way. +func (l *Limiter) budgetOf(e *entry, own time.Duration, now time.Time) Budget { + // A window that is over rolls on the caller's next request, and nothing + // is spent in the one that opens then — what a request still running + // charges lands in the old window and goes with it. The new one ends an + // hour after it opens, so no earlier than an hour from now. + if !now.Before(e.resetAt) { + return Budget{Limit: l.cfg.CostBudgetPerHour, ResetAt: now.Add(budgetWindow)} + } + return Budget{Used: e.used + own, Limit: l.cfg.CostBudgetPerHour, ResetAt: e.resetAt} +} + // takeSlots takes the two concurrency slots and then one rate token, giving // back whatever was already taken when a later limit refuses. // @@ -233,8 +321,10 @@ func (l *Limiter) takeSlots(e *entry, now time.Time) (reason string) { } // acquire reserves the slot. release is non-nil iff reason == "" (allowed). -func (l *Limiter) acquire(userID string, now time.Time) (release func(charge), reason string) { - e, reason := l.entryFor(userID, now) +// An exempt call is not refused for the budget; what it is charged is the +// caller's business, in release. +func (l *Limiter) acquire(userID string, now time.Time, exempt bool) (release func(charge), reason string) { + e, reason := l.entryFor(userID, now, exempt) if reason != "" { metric().WithLabelValues(reason).Inc() return nil, reason @@ -262,10 +352,18 @@ func (l *Limiter) acquire(userID string, now time.Time) (release func(charge), r inflight().WithLabelValues(scopeGlobal).Dec() } inflight().WithLabelValues(scopeUser).Dec() + done := l.clock() l.mu.Lock() e.inflight-- if l.cfg.CostBudgetPerHour > 0 { e.used += c.work + // Under the lock, unlike the other metrics here, because the value + // is state and not an event: set outside it, two releases of one + // caller could land in the wrong order, and a release that lost a + // race with eviction would bring back the series of a caller + // already gone. What the window has spent, as Remaining reads it: + // nothing, once the window is over. + budgetUsed().WithLabelValues(userID).Set(l.budgetOf(e, 0, done).Used.Seconds()) } l.mu.Unlock() // acquire spent one token before the work was known; a request that @@ -275,9 +373,14 @@ func (l *Limiter) acquire(userID string, now time.Time) (release func(charge), r // throttles the next call instead of letting twenty calls cost the same // as one. // - // Priced at `now`, the instant acquire admitted the call, not at - // release: the bucket refills while the work runs, and charging at the - // later instant would refund a slow request part of what it just spent. + // Priced now, when the work is done, and not at the instant acquire + // admitted it. rate.Limiter takes an instant earlier than its last + // event as its new last event, so pricing at admission wound the + // bucket's clock back over every call admitted in the meantime, and + // the next one was credited that stretch a second time: a slow batch + // refunded itself about its own duration in tokens. Pricing later is + // never the cheaper option either — the bucket is capped at its burst, + // so charging after the refill spends at least as much as before it. // // In burst-sized steps, because ReserveN refuses outright when n is over // the bucket's burst and silently charges nothing: a limiter configured @@ -286,7 +389,7 @@ func (l *Limiter) acquire(userID string, now time.Time) (release func(charge), r if e.rpm != nil && c.units > 1 { for left := c.units - 1; left > 0; { step := min(left, e.rpm.Burst()) - e.rpm.ReserveN(now, step) + e.rpm.ReserveN(done, step) left -= step } } diff --git a/ratelimit/ratelimit_test.go b/ratelimit/ratelimit_test.go index a1a6758..e1a2965 100644 --- a/ratelimit/ratelimit_test.go +++ b/ratelimit/ratelimit_test.go @@ -28,16 +28,16 @@ func TestRPM(t *testing.T) { l := newLimiter(t, Config{PerUserRPM: 60}) now := time.Date(2026, 1, 1, 12, 0, 0, 0, time.UTC) for i := range 60 { - release, reason := l.acquire("alice", now) + release, reason := l.acquire("alice", now, false) require.NotNilf(t, release, "call %d denied: %s", i, reason) release(charge{}) } - release, reason := l.acquire("alice", now) + release, reason := l.acquire("alice", now, false) assert.Nil(t, release, "61st call should be denied") assert.Equal(t, reasonRPM, reason) // Other users are independent. - release, reason = l.acquire("bob", now) + release, reason = l.acquire("bob", now, false) require.NotNilf(t, release, "bob blocked by alice's bucket: %s", reason) release(charge{}) } @@ -45,17 +45,17 @@ func TestRPM(t *testing.T) { func TestPerUserConcurrent(t *testing.T) { l := newLimiter(t, Config{PerUserConcurrent: 2}) now := time.Now() - r1, _ := l.acquire("alice", now) - r2, _ := l.acquire("alice", now) + r1, _ := l.acquire("alice", now, false) + r2, _ := l.acquire("alice", now, false) require.NotNil(t, r1) require.NotNil(t, r2) - r3, reason := l.acquire("alice", now) + r3, reason := l.acquire("alice", now, false) assert.Nil(t, r3, "third should be denied") assert.Equal(t, reasonUserConcurrent, reason) r1(charge{}) - r3, _ = l.acquire("alice", now) + r3, _ = l.acquire("alice", now, false) require.NotNil(t, r3, "after release, slot should reopen") r3(charge{}) r2(charge{}) @@ -64,15 +64,15 @@ func TestPerUserConcurrent(t *testing.T) { func TestGlobalConcurrent(t *testing.T) { l := newLimiter(t, Config{GlobalConcurrent: 1}) now := time.Now() - r1, _ := l.acquire("alice", now) + r1, _ := l.acquire("alice", now, false) require.NotNil(t, r1, "alice should pass") - r2, reason := l.acquire("bob", now) + r2, reason := l.acquire("bob", now, false) assert.Nil(t, r2, "bob should hit global") assert.Equal(t, reasonGlobalConcurrent, reason) r1(charge{}) - r2, _ = l.acquire("bob", now) + r2, _ = l.acquire("bob", now, false) require.NotNil(t, r2, "after release, bob should pass") r2(charge{}) } @@ -83,15 +83,15 @@ func TestGlobalDenialReturnsTheUserSlot(t *testing.T) { l := newLimiter(t, Config{PerUserConcurrent: 1, GlobalConcurrent: 1}) now := time.Now() - held, _ := l.acquire("alice", now) + held, _ := l.acquire("alice", now, false) require.NotNil(t, held) - denied, reason := l.acquire("bob", now) + denied, reason := l.acquire("bob", now, false) require.Nil(t, denied) require.Equal(t, reasonGlobalConcurrent, reason) held(charge{}) - r, reason := l.acquire("bob", now) + r, reason := l.acquire("bob", now, false) require.NotNilf(t, r, "bob's own slot was never returned: %s", reason) r(charge{}) } @@ -99,17 +99,17 @@ func TestGlobalDenialReturnsTheUserSlot(t *testing.T) { func TestCostBudget(t *testing.T) { l := newLimiter(t, Config{CostBudgetPerHour: 100 * time.Millisecond}) now := time.Now() - r, _ := l.acquire("alice", now) + r, _ := l.acquire("alice", now, false) r(charge{work: 60 * time.Millisecond, units: 1}) - r, _ = l.acquire("alice", now) + r, _ = l.acquire("alice", now, false) r(charge{work: 50 * time.Millisecond, units: 1}) - r2, reason := l.acquire("alice", now) + r2, reason := l.acquire("alice", now, false) assert.Nil(t, r2, "expected cost budget denial") assert.Equal(t, reasonCostBudget, reason) // After hour rollover the budget resets. - r3, _ := l.acquire("alice", now.Add(time.Hour+time.Minute)) + r3, _ := l.acquire("alice", now.Add(time.Hour+time.Minute), false) require.NotNil(t, r3, "budget didn't reset after hour") r3(charge{}) } @@ -122,12 +122,13 @@ func TestChargePerItem(t *testing.T) { t.Run("rpm counts every item", func(t *testing.T) { l := newLimiter(t, Config{PerUserRPM: 60}) now := time.Date(2026, 1, 1, 12, 0, 0, 0, time.UTC) + l.clock = func() time.Time { return now } for range 3 { - release, reason := l.acquire("alice", now) + release, reason := l.acquire("alice", now, false) require.NotNilf(t, release, "denied too early: %s", reason) release(charge{work: time.Second, units: 20}) } - release, reason := l.acquire("alice", now) + release, reason := l.acquire("alice", now, false) assert.Nil(t, release, "three requests of twenty calls are sixty, the bucket is empty") assert.Equal(t, reasonRPM, reason) }) @@ -137,11 +138,12 @@ func TestChargePerItem(t *testing.T) { t.Run("more calls than the burst is still paid", func(t *testing.T) { l := newLimiter(t, Config{PerUserRPM: 10}) now := time.Date(2026, 1, 1, 12, 0, 0, 0, time.UTC) - release, _ := l.acquire("alice", now) + l.clock = func() time.Time { return now } + release, _ := l.acquire("alice", now, false) require.NotNil(t, release) release(charge{work: time.Second, units: 20}) - next, reason := l.acquire("alice", now) + next, reason := l.acquire("alice", now, false) assert.Nil(t, next, "twenty calls on a 10 RPM limiter must leave the bucket empty") assert.Equal(t, reasonRPM, reason) }) @@ -150,16 +152,43 @@ func TestChargePerItem(t *testing.T) { t.Run("budget counts the sum, not the wall clock", func(t *testing.T) { l := newLimiter(t, Config{CostBudgetPerHour: 100 * time.Millisecond}) now := time.Now() - release, _ := l.acquire("alice", now) + release, _ := l.acquire("alice", now, false) require.NotNil(t, release) release(charge{work: 120 * time.Millisecond, units: 4}) - next, reason := l.acquire("alice", now) + next, reason := l.acquire("alice", now, false) assert.Nil(t, next, "four 30ms calls in one request spend 120ms, not the 30ms they took") assert.Equal(t, reasonCostBudget, reason) }) } +// The rest of a slow request's calls is priced when it finishes. Priced at the +// instant it was admitted, the bucket's clock went back over everything +// admitted meanwhile, and the next caller was handed those seconds again. +func TestSlowRequestDoesNotRefundItself(t *testing.T) { + l := newLimiter(t, Config{PerUserRPM: 60}) // a token a second + opened := time.Date(2026, 1, 1, 12, 0, 0, 0, time.UTC) + for range 60 { + release, _ := l.acquire("alice", opened, false) + require.NotNil(t, release) + release(charge{units: 1}) + } + + slow, _ := l.acquire("alice", opened.Add(time.Second), false) + require.NotNil(t, slow, "the token that came back in the first second") + + finished := opened.Add(11 * time.Second) + l.clock = func() time.Time { return finished } + quick, _ := l.acquire("alice", finished, false) + require.NotNil(t, quick, "ten more came back while the slow one ran; one of them is spent") + quick(charge{units: 1}) + slow(charge{units: 11}) // ten more calls than acquire charged for: nine left, minus ten + + next, reason := l.acquire("alice", finished, false) + assert.Nil(t, next, "the slow request's ten calls were paid for, not refunded") + assert.Equal(t, reasonRPM, reason) +} + // A request that charges nothing is priced by its wall clock, the way every // request was before one request could carry many calls. func TestCostSettleFallsBackToWallClock(t *testing.T) { @@ -186,7 +215,7 @@ func TestChargeWithoutAccumulator(t *testing.T) { func TestEvictIdle(t *testing.T) { l := newLimiter(t, Config{PerUserRPM: 60}) now := time.Now() - r, _ := l.acquire("alice", now) + r, _ := l.acquire("alice", now, false) r(charge{}) require.Contains(t, l.entries, "alice") @@ -202,7 +231,7 @@ func TestEvictIdle(t *testing.T) { func TestEvictSkipsBusy(t *testing.T) { l := newLimiter(t, Config{PerUserConcurrent: 1}) now := time.Now() - r, _ := l.acquire("alice", now) + r, _ := l.acquire("alice", now, false) require.NotNil(t, r) // Don't release — entry holds an in-flight slot. @@ -250,7 +279,7 @@ func TestDisabled(t *testing.T) { l := New(Config{}) assert.True(t, l.Disabled(), "zero config should be Disabled") - release, reason := l.acquire("anyone", time.Now()) + release, reason := l.acquire("anyone", time.Now(), false) require.NotNilf(t, release, "disabled limiter should never deny: reason=%q", reason) release(charge{}) @@ -271,7 +300,7 @@ func TestMiddlewareSkipsNonPOST(t *testing.T) { h := l.Middleware(next, embedlog.Logger{}) // Occupy the single concurrency slot. - release, _ := l.acquire(anonymousUser, time.Now()) + release, _ := l.acquire(anonymousUser, time.Now(), false) require.NotNil(t, release) defer func() { release(charge{}) }() @@ -284,7 +313,7 @@ func TestMiddlewareSkipsNonPOST(t *testing.T) { rec := httptest.NewRecorder() h.ServeHTTP(rec, httptest.NewRequest(http.MethodPost, "/mcp", nil)) assert.Equal(t, http.StatusTooManyRequests, rec.Code, "POST must still be limited") - assert.Equal(t, "10", rec.Header().Get("Retry-After")) + assert.Equal(t, "1", rec.Header().Get("Retry-After"), "a concurrency slot frees up when a neighbour finishes") assert.Contains(t, rec.Body.String(), reasonUserConcurrent) } @@ -342,7 +371,7 @@ func TestMetrics(t *testing.T) { l := newLimiter(t, Config{PerUserConcurrent: 1, GlobalConcurrent: 1}) before := testutil.ToFloat64(inflightGauge.WithLabelValues(scopeUser)) - release, _ := l.acquire("alice", time.Now()) + release, _ := l.acquire("alice", time.Now(), false) require.NotNil(t, release) assert.InDelta(t, before+1, testutil.ToFloat64(inflightGauge.WithLabelValues(scopeUser)), 0) diff --git a/ratelimit/refusal.go b/ratelimit/refusal.go new file mode 100644 index 0000000..5a5acbf --- /dev/null +++ b/ratelimit/refusal.go @@ -0,0 +1,228 @@ +package ratelimit + +// What a refused caller is told. A denial used to be a text/plain 429, which a +// bridge reports to the model as "server unavailable": the model could neither +// wait sensibly nor ask for less, and nobody watching could tell a ban from an +// outage. It is now a JSON-RPC error answering the call that was refused, with +// the limit and the wait in it. + +import ( + "bytes" + "encoding/json" + "io" + "net/http" + "strconv" + "strings" + "time" + + "github.com/vmkteam/mcpkit/mcp" + + "github.com/valyala/fastjson" + "github.com/vmkteam/zenrpc/v2" +) + +// maxPeekBytes caps how much of a body the limiter reads to find the call +// Exempt is asked about. An MCP call is kilobytes; a body over this is never +// exempt, and is refused the old way, without an envelope. +const maxPeekBytes = 1 << 20 // 1 MiB + +// maxRefusalPeekBytes caps what a refusal reads to find the id it answers. A +// refusal is how the limiter sheds load, and reading a megabyte of every +// request it turns away made shedding cost as much as serving; a body over +// this gets the plain 429. +const maxRefusalPeekBytes = 64 << 10 // 64 KiB + +// call is what the limiter reads of a request body. The zero value is what an +// unreadable body leaves: no id to answer, nothing to exempt. +type call struct { + // id is the request id as it arrived, or nil for a notification, which has + // no answer to carry an error in. + id json.RawMessage + method string + // name is what Mcp-Name mirrors: the tool of tools/call, the prompt of + // prompts/get, the URI of resources/read. + name string +} + +// peek reads the JSON-RPC request out of the body and puts the body back, so +// the handler behind reads the same bytes. ok is false when there is no single +// request to read: not JSON, a batch, more than limit bytes of it, or a body +// that repeats a key read here. +// +// The last is what keeps Exempt honest. What the limiter reads has to be what +// the dispatcher runs, and the two readers disagree about a repeated key: +// fastjson takes the first and matches exact bytes, encoding/json — which +// decodes the call for zenrpc — takes the last and ignores case. +// {"name":"help","Name":"db_query"} was exempted as help and run as db_query, +// at no cost to the budget. +func peek(r *http.Request, limit int64) (c call, ok bool) { + head, err := io.ReadAll(io.LimitReader(r.Body, limit+1)) + r.Body = struct { + io.Reader + io.Closer + }{io.MultiReader(bytes.NewReader(head), r.Body), r.Body} + if err != nil || int64(len(head)) > limit { + return call{}, false + } + + var p fastjson.Parser + v, err := p.ParseBytes(head) + if err != nil || v.Type() != fastjson.TypeObject { + return call{}, false + } + if repeated(v, "method") || repeated(v, "params") { + return call{}, false + } + c.method = string(v.GetStringBytes("method")) + if path := mcp.NamePathFor(c.method); path != nil { + last := len(path) - 1 + if repeated(v.Get(path[:last]...), path[last]) { + return call{}, false + } + c.name = string(v.GetStringBytes(path...)) + } + if id := v.Get("id"); id != nil && id.Type() != fastjson.TypeNull { + c.id = id.MarshalTo(nil) + } + return c, true +} + +// repeated reports whether key occurs more than once in the object v, spelt in +// any case: encoding/json folds case, and Unicode's with it, when it matches a +// member to a field. A value that is not an object, or is not there, repeats +// nothing. +func repeated(v *fastjson.Value, key string) bool { + if v == nil || v.Type() != fastjson.TypeObject { + return false + } + kb, n := []byte(key), 0 + v.GetObject().Visit(func(k []byte, _ *fastjson.Value) { + if bytes.EqualFold(k, kb) { + n++ + } + }) + return n > 1 +} + +// refusal is one denial, as the caller is told about it. +type refusal struct { + reason string + retryAfter time.Duration + // budget is set whenever the hourly budget is on, whichever limit refused: + // how much of it is gone is worth knowing on any denial. + budget *Budget +} + +// refusalFor describes the denial acquire has just returned for userID at now. +// +// The wait is worked out per limit, because no single number is true of all +// four. The budget comes back when the window rolls, which can be most of an +// hour away. The rate bucket knows exactly when its next token lands. A +// concurrency slot frees up when a neighbouring request finishes — soon, but +// nothing here knows when, so the answer is a second. +func (l *Limiter) refusalFor(userID, reason string, now time.Time) refusal { + rf := refusal{reason: reason, retryAfter: time.Second} + + l.mu.Lock() + defer l.mu.Unlock() + e, ok := l.entries[userID] + if !ok { // entryFor has just created it; only eviction could take it, after hours idle + return rf + } + if l.cfg.CostBudgetPerHour > 0 { + b := l.budgetOf(e, 0, now) + rf.budget = &b + } + switch reason { + case reasonCostBudget: + rf.retryAfter = rf.budget.ResetAt.Sub(now) + case reasonRPM: + // Read, not reserved: a reservation is a token spent until it is + // cancelled, and the refused call would pay for itself. + missing := 1 - e.rpm.TokensAt(now) + rf.retryAfter = time.Duration(missing / float64(e.rpm.Limit()) * float64(time.Second)) + } + return rf +} + +// retryAfterSeconds is the wait in whole seconds, rounded up and never zero: a +// client told to retry in 0 retries at once and is refused again. +func (rf refusal) retryAfterSeconds() int { + return max(1, int((rf.retryAfter+time.Second-1)/time.Second)) +} + +// logArgs is the denial for the log: enough for whoever is on call to tell a +// ban of an hour from one of a second without asking the caller. +func (rf refusal) logArgs(userID string) []any { + args := []any{"user", userID, "reason", rf.reason, "retry_after", rf.retryAfterSeconds()} + if rf.budget != nil { + args = append(args, "budget_used", formatDuration(rf.budget.Used)) + } + return args +} + +// data is the error's data member: what the model reads to decide whether to +// wait or to ask for less. +func (rf refusal) data() any { + d := struct { + Reason string `json:"reason"` + RetryAfter int `json:"retry_after"` + BudgetUsed string `json:"budget_used,omitempty"` + BudgetLimit string `json:"budget_limit,omitempty"` + Window string `json:"window,omitempty"` + }{Reason: rf.reason, RetryAfter: rf.retryAfterSeconds()} + if rf.budget != nil { + d.BudgetUsed = formatDuration(rf.budget.Used) + d.BudgetLimit = formatDuration(rf.budget.Limit) + d.Window = formatDuration(budgetWindow) + } + return d +} + +// refuse answers a denied request, id being what peek read of it. +// +// A request with an id is answered with a JSON-RPC error to that id, which a +// client can show as a failed call rather than as a dead server. Anything else — +// not JSON, a batch, a notification, a body too large to look into — gets the +// plain 429 it always got: there is no id to answer, and the transport behind +// would refuse most of those anyway. +// +// The envelope is zenrpc's, the one every other JSON-RPC error of the server +// goes out in. The status stays 429 under it: that is what a proxy and a +// retrying transport understand, and Retry-After means something only beside +// it. +func refuse(w http.ResponseWriter, rf refusal, id json.RawMessage) { + w.Header().Set("Retry-After", strconv.Itoa(rf.retryAfterSeconds())) + message := "rate limit: " + rf.reason + + if id == nil { + http.Error(w, message, http.StatusTooManyRequests) + return + } + b, err := json.Marshal(zenrpc.NewResponseError(&id, mcp.CodeRateLimited, message, rf.data())) + if err != nil { // a fixed shape; the id is what fastjson just wrote out + http.Error(w, message, http.StatusTooManyRequests) + return + } + w.Header().Set("Content-Type", "application/json") + w.WriteHeader(http.StatusTooManyRequests) + _, _ = w.Write(b) +} + +// formatDuration writes d the way a person would: 5m rather than 5m0s, 1h +// rather than 1h0m0s. Whole seconds from a second up; below that, milliseconds. +func formatDuration(d time.Duration) string { + if d >= time.Second { + d = d.Round(time.Second) + } else { + d = d.Round(time.Millisecond) + } + s := d.String() + if strings.HasSuffix(s, "m0s") { + s = strings.TrimSuffix(s, "0s") + } + if strings.HasSuffix(s, "h0m") { + s = strings.TrimSuffix(s, "0m") + } + return s +} diff --git a/ratelimit/refusal_test.go b/ratelimit/refusal_test.go new file mode 100644 index 0000000..5143821 --- /dev/null +++ b/ratelimit/refusal_test.go @@ -0,0 +1,390 @@ +package ratelimit + +import ( + "encoding/json" + "io" + "net/http" + "net/http/httptest" + "strconv" + "strings" + "testing" + "time" + + "github.com/vmkteam/mcpkit/mcp" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "github.com/vmkteam/embedlog" +) + +const toolCall = `{"jsonrpc":"2.0","id":7,"method":"tools/call","params":{"name":"api_call","arguments":{}}}` + +type rpcAnswer struct { + JSONRPC string `json:"jsonrpc"` + ID json.RawMessage `json:"id"` + Error struct { + Code int `json:"code"` + Message string `json:"message"` + Data map[string]any `json:"data"` + } `json:"error"` +} + +func post(h http.Handler, body string) *httptest.ResponseRecorder { + rec := httptest.NewRecorder() + h.ServeHTTP(rec, httptest.NewRequest(http.MethodPost, "/mcp", strings.NewReader(body))) + return rec +} + +func okHandler(w http.ResponseWriter, _ *http.Request) { w.WriteHeader(http.StatusOK) } + +// A denial answers the call that was refused, so a client shows it as a failed +// call with a reason — not as a server that went away. Whichever limit said no. +func TestRefusalAnswersTheCall(t *testing.T) { + tests := map[string]struct { + cfg Config + // spend brings the limiter to the point of refusing the next call. + spend func(t *testing.T, l *Limiter, h http.Handler) + }{ + reasonRPM: { + cfg: Config{PerUserRPM: 1}, + spend: func(t *testing.T, _ *Limiter, h http.Handler) { + require.Equal(t, http.StatusOK, post(h, toolCall).Code) + }, + }, + reasonUserConcurrent: { + cfg: Config{PerUserConcurrent: 1}, + spend: func(t *testing.T, l *Limiter, _ http.Handler) { + release, _ := l.acquire(anonymousUser, time.Now(), false) + require.NotNil(t, release) + t.Cleanup(func() { release(charge{}) }) + }, + }, + reasonGlobalConcurrent: { + cfg: Config{GlobalConcurrent: 1}, + spend: func(t *testing.T, l *Limiter, _ http.Handler) { + release, _ := l.acquire("bob", time.Now(), false) + require.NotNil(t, release) + t.Cleanup(func() { release(charge{}) }) + }, + }, + reasonCostBudget: { + cfg: Config{CostBudgetPerHour: time.Minute}, + spend: func(t *testing.T, _ *Limiter, h http.Handler) { + require.Equal(t, http.StatusOK, post(h, toolCall).Code) + }, + }, + } + for reason, tc := range tests { + t.Run(reason, func(t *testing.T) { + l := newLimiter(t, tc.cfg) + h := l.Middleware(http.HandlerFunc(spendAll), embedlog.Logger{}) + tc.spend(t, l, h) + + rec := post(h, toolCall) + require.Equal(t, http.StatusTooManyRequests, rec.Code) + assert.Equal(t, "application/json", rec.Header().Get("Content-Type")) + + var got rpcAnswer + require.NoError(t, json.Unmarshal(rec.Body.Bytes(), &got), rec.Body.String()) + assert.Equal(t, "2.0", got.JSONRPC) + assert.JSONEq(t, "7", string(got.ID), "the id of the refused call is echoed") + assert.Equal(t, mcp.CodeRateLimited, got.Error.Code) + assert.Equal(t, "rate limit: "+reason, got.Error.Message) + assert.Equal(t, reason, got.Error.Data["reason"]) + + retryAfter, ok := got.Error.Data["retry_after"].(float64) + require.True(t, ok, "retry_after is a number of seconds: %v", got.Error.Data) + assert.Equal(t, rec.Header().Get("Retry-After"), strconv.Itoa(int(retryAfter)), + "the header and the body tell the same wait") + }) + } +} + +// What the budget looks like is in the data of a budget denial, in the words a +// person would use. +func TestRefusalCarriesTheBudget(t *testing.T) { + l := newLimiter(t, Config{CostBudgetPerHour: time.Minute}) + h := l.Middleware(http.HandlerFunc(spendAll), embedlog.Logger{}) + require.Equal(t, http.StatusOK, post(h, toolCall).Code) + + var got rpcAnswer + require.NoError(t, json.Unmarshal(post(h, toolCall).Body.Bytes(), &got)) + assert.Equal(t, "2m", got.Error.Data["budget_used"]) + assert.Equal(t, "1m", got.Error.Data["budget_limit"]) + assert.Equal(t, "1h", got.Error.Data["window"]) + assert.InDelta(t, 3600, got.Error.Data["retry_after"], 5, "the window opened a moment ago") +} + +// A limiter without a budget has no budget to report. +func TestRefusalWithoutBudget(t *testing.T) { + l := newLimiter(t, Config{PerUserRPM: 1}) + h := l.Middleware(http.HandlerFunc(okHandler), embedlog.Logger{}) + require.Equal(t, http.StatusOK, post(h, toolCall).Code) + + var got rpcAnswer + require.NoError(t, json.Unmarshal(post(h, toolCall).Body.Bytes(), &got)) + assert.NotContains(t, got.Error.Data, "budget_used") + assert.NotContains(t, got.Error.Data, "window") +} + +// An id is echoed as it arrived: a string stays a string. +func TestRefusalEchoesAStringID(t *testing.T) { + l := newLimiter(t, Config{PerUserRPM: 1}) + h := l.Middleware(http.HandlerFunc(okHandler), embedlog.Logger{}) + require.Equal(t, http.StatusOK, post(h, toolCall).Code) + + var got rpcAnswer + require.NoError(t, json.Unmarshal(post(h, `{"jsonrpc":"2.0","id":"a-7","method":"ping"}`).Body.Bytes(), &got)) + assert.JSONEq(t, `"a-7"`, string(got.ID)) +} + +// Without an id there is nothing to answer, and the denial is the plain 429 it +// always was — with the honest Retry-After all the same. +func TestRefusalWithoutAnIDIsPlain(t *testing.T) { + bodies := map[string]string{ + "not json": `rate me`, + "a batch": `[` + toolCall + `]`, + "notification": `{"jsonrpc":"2.0","method":"notifications/initialized"}`, + "null id": `{"jsonrpc":"2.0","id":null,"method":"ping"}`, + "too big": `{"jsonrpc":"2.0","id":1,"method":"ping","params":{"pad":"` + strings.Repeat("x", maxRefusalPeekBytes) + `"}}`, + "empty": ``, + } + for name, body := range bodies { + t.Run(name, func(t *testing.T) { + l := newLimiter(t, Config{PerUserConcurrent: 1}) + release, _ := l.acquire(anonymousUser, time.Now(), false) + require.NotNil(t, release) + defer release(charge{}) + + rec := post(l.Middleware(http.HandlerFunc(okHandler), embedlog.Logger{}), body) + assert.Equal(t, http.StatusTooManyRequests, rec.Code) + assert.Contains(t, rec.Header().Get("Content-Type"), "text/plain") + assert.Equal(t, "rate limit: "+reasonUserConcurrent+"\n", rec.Body.String()) + assert.Equal(t, "1", rec.Header().Get("Retry-After")) + }) + } +} + +// A refusal reads only as much of a body as an honest call takes: shedding +// load must not cost what serving it would. +func TestRefusalReadsLittleOfTheBody(t *testing.T) { + l := newLimiter(t, Config{PerUserConcurrent: 1}) + release, _ := l.acquire(anonymousUser, time.Now(), false) + require.NotNil(t, release) + defer release(charge{}) + + body := &countingReader{r: strings.NewReader( + `{"jsonrpc":"2.0","id":1,"method":"ping","params":{"pad":"` + strings.Repeat("x", maxPeekBytes) + `"}}`)} + rec := httptest.NewRecorder() + l.Middleware(http.HandlerFunc(okHandler), embedlog.Logger{}). + ServeHTTP(rec, httptest.NewRequest(http.MethodPost, "/mcp", body)) + assert.Equal(t, http.StatusTooManyRequests, rec.Code) + assert.LessOrEqual(t, body.n, maxRefusalPeekBytes+1) +} + +type countingReader struct { + r io.Reader + n int +} + +func (c *countingReader) Read(p []byte) (int, error) { + n, err := c.r.Read(p) + c.n += n + return n, err +} + +// Whatever the limiter reads of a body, the handler behind reads all of it — +// including the part past what the limiter was willing to read. +func TestPeekPutsTheBodyBack(t *testing.T) { + bodies := map[string]string{ + "a call": toolCall, + "too big": `{"jsonrpc":"2.0","id":1,"params":{"pad":"` + strings.Repeat("x", maxPeekBytes) + `"}}`, + "garbage": `rate me`, + } + for name, body := range bodies { + t.Run(name, func(t *testing.T) { + r := httptest.NewRequest(http.MethodPost, "/mcp", strings.NewReader(body)) + peek(r, maxPeekBytes) + rest, err := io.ReadAll(r.Body) + require.NoError(t, err) + assert.Equal(t, body, string(rest)) + assert.NoError(t, r.Body.Close()) + }) + } +} + +func TestPeekReadsTheCall(t *testing.T) { + tests := map[string]struct { + body string + want call + }{ + "tools/call": {toolCall, call{id: json.RawMessage("7"), method: "tools/call", name: "api_call"}}, + "prompts/get": {`{"jsonrpc":"2.0","id":"p","method":"prompts/get","params":{"name":"triage"}}`, + call{id: json.RawMessage(`"p"`), method: "prompts/get", name: "triage"}}, + "resources/read": {`{"jsonrpc":"2.0","id":2,"method":"resources/read","params":{"uri":"doc://help"}}`, + call{id: json.RawMessage("2"), method: "resources/read", name: "doc://help"}}, + "no name": {`{"jsonrpc":"2.0","id":3,"method":"tools/list"}`, call{id: json.RawMessage("3"), method: "tools/list"}}, + "notification": {`{"jsonrpc":"2.0","method":"notifications/initialized"}`, + call{method: "notifications/initialized"}}, + } + for name, tc := range tests { + t.Run(name, func(t *testing.T) { + got, ok := peek(httptest.NewRequest(http.MethodPost, "/mcp", strings.NewReader(tc.body)), maxPeekBytes) + require.True(t, ok) + assert.Equal(t, tc.want, got) + }) + } +} + +// The budget comes back when its window rolls, and the caller is told exactly +// when that is — not ten seconds, after which it would be refused for another +// fifty minutes. +func TestRetryAfterCostBudget(t *testing.T) { + l := newLimiter(t, Config{CostBudgetPerHour: 5 * time.Minute}) + opened := time.Date(2026, 1, 1, 12, 0, 0, 0, time.UTC) + release, _ := l.acquire("alice", opened, false) + require.NotNil(t, release) + release(charge{work: 5 * time.Minute, units: 1}) + + at := opened.Add(5 * time.Minute) + denied, reason := l.acquire("alice", at, false) + require.Nil(t, denied) + require.Equal(t, reasonCostBudget, reason) + + rf := l.refusalFor("alice", reason, at) + assert.Equal(t, 55*time.Minute, rf.retryAfter) + assert.Equal(t, 3300, rf.retryAfterSeconds()) + + denied, reason = l.acquire("alice", at.Add(rf.retryAfter-time.Nanosecond), false) + assert.Nil(t, denied, "a moment early is still refused") + assert.Equal(t, reasonCostBudget, reason) + + release, reason = l.acquire("alice", at.Add(rf.retryAfter), false) + require.NotNilf(t, release, "on the dot, the window has rolled: %s", reason) + release(charge{}) +} + +// The wait for a rate token is when the next one lands, and asking for it does +// not take it. +func TestRetryAfterRPMSpendsNothing(t *testing.T) { + l := newLimiter(t, Config{PerUserRPM: 60}) + now := time.Date(2026, 1, 1, 12, 0, 0, 0, time.UTC) + for range 60 { + release, _ := l.acquire("alice", now, false) + require.NotNil(t, release) + release(charge{}) + } + denied, reason := l.acquire("alice", now, false) + require.Nil(t, denied) + require.Equal(t, reasonRPM, reason) + + rf := l.refusalFor("alice", reason, now) + assert.Equal(t, time.Second, rf.retryAfter, "60 RPM is a token a second") + assert.Equal(t, rf, l.refusalFor("alice", reason, now), "asking twice does not move the wait") + + release, reason := l.acquire("alice", now.Add(rf.retryAfter), false) + require.NotNilf(t, release, "the token the refusal measured is still there: %s", reason) + release(charge{}) +} + +func TestRetryAfterConcurrency(t *testing.T) { + l := newLimiter(t, Config{PerUserConcurrent: 1, GlobalConcurrent: 1}) + now := time.Now() + release, _ := l.acquire("alice", now, false) + require.NotNil(t, release) + defer release(charge{}) + + _, reason := l.acquire("alice", now, false) + require.Equal(t, reasonUserConcurrent, reason) + assert.Equal(t, time.Second, l.refusalFor("alice", reason, now).retryAfter) + + _, reason = l.acquire("bob", now, false) + require.Equal(t, reasonGlobalConcurrent, reason) + assert.Equal(t, time.Second, l.refusalFor("bob", reason, now).retryAfter) +} + +func TestRetryAfterSeconds(t *testing.T) { + t.Parallel() + tests := []struct { + d time.Duration + want int + }{ + {0, 1}, + {time.Millisecond, 1}, + {time.Second, 1}, + {time.Second + 1, 2}, + {1200 * time.Millisecond, 2}, + {55 * time.Minute, 3300}, + } + for _, tc := range tests { + assert.Equalf(t, tc.want, refusal{retryAfter: tc.d}.retryAfterSeconds(), "%s", tc.d) + } +} + +func TestFormatDuration(t *testing.T) { + t.Parallel() + tests := []struct { + d time.Duration + want string + }{ + {0, "0s"}, + {100 * time.Millisecond, "100ms"}, + {5 * time.Minute, "5m"}, + {5*time.Minute + 2*time.Second, "5m2s"}, + {5*time.Minute + 2400*time.Millisecond, "5m2s"}, + {20 * time.Minute, "20m"}, + {time.Hour, "1h"}, + {time.Hour + 5*time.Minute, "1h5m"}, + {time.Hour + 5*time.Second, "1h0m5s"}, + } + for _, tc := range tests { + assert.Equalf(t, tc.want, formatDuration(tc.d), "%s", tc.d) + } +} + +// A key the limiter reads, repeated, is read differently by the dispatcher: +// first wins here, last wins in encoding/json, which also ignores case. Such a +// body is not read at all. +// A repeat anywhere the limiter does not look is none of its business. +func TestPeekRefusesARepeatedKey(t *testing.T) { + repeats := map[string]string{ + "name": `{"jsonrpc":"2.0","id":1,"method":"tools/call","params":{"name":"help","name":"db_query"}}`, + "params": `{"jsonrpc":"2.0","id":1,"method":"tools/call","params":{"name":"help"},"params":{"name":"db_query"}}`, + "method": `{"jsonrpc":"2.0","id":1,"method":"tools/list","method":"tools/call","params":{"name":"db_query"}}`, + "uri": `{"jsonrpc":"2.0","id":1,"method":"resources/read","params":{"uri":"doc://help","uri":"doc://secret"}}`, + // Folded as encoding/json folds them: one member to the dispatcher. + "Name": `{"jsonrpc":"2.0","id":1,"method":"tools/call","params":{"name":"help","Name":"db_query"}}`, + "Params": `{"jsonrpc":"2.0","id":1,"method":"tools/call","params":{"name":"help"},"Params":{"name":"db_query"}}`, + "paramſ": `{"jsonrpc":"2.0","id":1,"method":"tools/call","params":{"name":"help"},"paramſ":{"name":"db_query"}}`, + "Method": `{"jsonrpc":"2.0","id":1,"method":"resources/read","params":{"uri":"help","name":"db_query"},"Method":"tools.call"}`, + } + for name, body := range repeats { + t.Run(name, func(t *testing.T) { + _, ok := peek(httptest.NewRequest(http.MethodPost, "/mcp", strings.NewReader(body)), maxPeekBytes) + assert.False(t, ok) + }) + } + + t.Run("elsewhere", func(t *testing.T) { + body := `{"jsonrpc":"2.0","id":1,"method":"tools/call","params":{"name":"help","arguments":{"q":1,"q":2}}}` + c, ok := peek(httptest.NewRequest(http.MethodPost, "/mcp", strings.NewReader(body)), maxPeekBytes) + require.True(t, ok) + assert.Equal(t, "help", c.name) + }) +} + +// Whoever is on call reads the wait and the spend in the log line, without +// asking the caller what they were told. +func TestRefusalLogArgs(t *testing.T) { + t.Parallel() + spent := refusal{ + reason: reasonCostBudget, + retryAfter: 55 * time.Minute, + budget: &Budget{Used: 21 * time.Minute, Limit: 20 * time.Minute}, + } + assert.Equal(t, []any{"user", "alice", "reason", reasonCostBudget, "retry_after", 3300, "budget_used", "21m"}, + spent.logArgs("alice")) + + rate := refusal{reason: reasonRPM, retryAfter: time.Second} + assert.Equal(t, []any{"user", "alice", "reason", reasonRPM, "retry_after", 1}, rate.logArgs("alice")) +} diff --git a/ratelimit/remaining_test.go b/ratelimit/remaining_test.go new file mode 100644 index 0000000..8291828 --- /dev/null +++ b/ratelimit/remaining_test.go @@ -0,0 +1,135 @@ +package ratelimit + +import ( + "context" + "encoding/json" + "io" + "net/http" + "strings" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" + "github.com/vmkteam/embedlog" +) + +// What Remaining reports is what a refusal reports: one number, read from one +// place, whichever way the caller learns it. +func TestRemainingAgreesWithTheRefusal(t *testing.T) { + var ( + seen Budget + ok bool + ) + l := newLimiter(t, Config{CostBudgetPerHour: time.Minute, Exempt: exemptHelp}) + h := l.Middleware(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + body, _ := io.ReadAll(r.Body) + if strings.Contains(string(body), `"help"`) { + seen, ok = Remaining(r.Context()) + } else { + Charge(r.Context(), 90*time.Second) + } + w.WriteHeader(http.StatusOK) + }), embedlog.Logger{}) + + before := time.Now() + require.Equal(t, http.StatusOK, post(h, toolCall).Code) + + var refused rpcAnswer + require.NoError(t, json.Unmarshal(post(h, toolCall).Body.Bytes(), &refused)) + require.Equal(t, reasonCostBudget, refused.Error.Data["reason"]) + + require.Equal(t, http.StatusOK, post(h, helpCall).Code) + require.True(t, ok) + assert.Equal(t, refused.Error.Data["budget_used"], formatDuration(seen.Used)) + assert.Equal(t, refused.Error.Data["budget_limit"], formatDuration(seen.Limit)) + assert.Equal(t, 90*time.Second, seen.Used) + assert.Equal(t, time.Minute, seen.Limit) + assert.Zero(t, seen.Left(), "overspent is nothing left, not a negative") + assert.WithinRange(t, seen.ResetAt, before.Add(budgetWindow), time.Now().Add(budgetWindow)) +} + +// A running request sees its own charges as it makes them; the next one sees +// them settled. +func TestRemainingCountsThisRequest(t *testing.T) { + var used []time.Duration + l := newLimiter(t, Config{CostBudgetPerHour: time.Hour}) + h := l.Middleware(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + for range 2 { + b, ok := Remaining(r.Context()) + assert.True(t, ok) + used = append(used, b.Used) + Charge(r.Context(), 10*time.Second) + } + w.WriteHeader(http.StatusOK) + }), embedlog.Logger{}) + + post(h, toolCall) + post(h, toolCall) + assert.Equal(t, []time.Duration{0, 10 * time.Second, 20 * time.Second, 30 * time.Second}, used) +} + +// An exempt call spends nothing, so what it charges is not shown as spent. +func TestRemainingLeavesOutAnExemptCall(t *testing.T) { + var seen Budget + l := newLimiter(t, Config{CostBudgetPerHour: time.Hour, Exempt: exemptHelp}) + h := l.Middleware(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + Charge(r.Context(), 10*time.Second) + seen, _ = Remaining(r.Context()) + w.WriteHeader(http.StatusOK) + }), embedlog.Logger{}) + + post(h, helpCall) + assert.Zero(t, seen.Used) + post(h, toolCall) + assert.Equal(t, 10*time.Second, seen.Used, "the help call before it left nothing behind") +} + +// No budget, nothing to report — and a context that never went through the +// middleware has no caller to report on. +func TestRemainingWithoutABudget(t *testing.T) { + ok := true + l := newLimiter(t, Config{PerUserRPM: 10}) + h := l.Middleware(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + _, ok = Remaining(r.Context()) + w.WriteHeader(http.StatusOK) + }), embedlog.Logger{}) + post(h, toolCall) + assert.False(t, ok, "the budget is off") + + _, ok = Remaining(context.Background()) + assert.False(t, ok, "no middleware") + + ctx, _ := WithCost(context.Background()) + _, ok = Remaining(ctx) + assert.False(t, ok, "an accumulator of its own, and no caller behind it") +} + +func TestBudgetLeft(t *testing.T) { + t.Parallel() + assert.Equal(t, 15*time.Second, Budget{Used: 45 * time.Second, Limit: time.Minute}.Left()) + assert.Zero(t, Budget{Used: 90 * time.Second, Limit: time.Minute}.Left()) +} + +// A window that is over is empty before the caller comes back to roll it: +// what a request still running charges lands in the old window and goes with +// it, and the next window ends no earlier than an hour from now. +func TestRemainingOnceTheWindowIsOver(t *testing.T) { + l := newLimiter(t, Config{CostBudgetPerHour: time.Hour}) + opened := time.Date(2026, 1, 1, 12, 0, 0, 0, time.UTC) + release, _ := l.acquire("alice", opened, false) + require.NotNil(t, release) + release(charge{work: 90 * time.Second, units: 1}) + + b, ok := l.budget("alice", 5*time.Second, opened.Add(30*time.Minute)) + require.True(t, ok) + assert.Equal(t, 95*time.Second, b.Used, "finished plus running") + assert.Equal(t, opened.Add(budgetWindow), b.ResetAt) + + over := opened.Add(budgetWindow) + b, ok = l.budget("alice", 5*time.Second, over) + require.True(t, ok) + assert.Zero(t, b.Used) + assert.Equal(t, time.Hour, b.Left()) + assert.Equal(t, over.Add(budgetWindow), b.ResetAt) +} diff --git a/transport.go b/transport.go index f358bda..ea83022 100644 --- a/transport.go +++ b/transport.go @@ -4,11 +4,14 @@ package mcpkit // what it refuses is in the package comment; what follows is the handler. import ( + "bytes" "context" "errors" "io" "net/http" "strings" + "unicode" + "unicode/utf8" "github.com/vmkteam/mcpkit/mcp" @@ -148,7 +151,9 @@ func (s *Server) hostAllowed(r *http.Request) bool { return false } -func (s *Server) handlePOST(w http.ResponseWriter, r *http.Request) { +// readRequest reads the one request a POST carries and parses it, answering +// the refusal itself when there is no such request to read. +func (s *Server) readRequest(w http.ResponseWriter, r *http.Request) (body []byte, v *fastjson.Value, ok bool) { r.Body = http.MaxBytesReader(w, r.Body, s.opts.MaxRequestBytes) body, err := io.ReadAll(r.Body) if err != nil { @@ -156,11 +161,11 @@ func (s *Server) handlePOST(w http.ResponseWriter, r *http.Request) { if errors.As(err, &mbe) { rejected(reasonTooLarge) http.Error(w, "request body too large", http.StatusRequestEntityTooLarge) - return + return nil, nil, false } rejected(reasonReadBody) http.Error(w, "read body: "+err.Error(), http.StatusBadRequest) - return + return nil, nil, false } // One request per POST. MCP dropped JSON-RPC batching in 2025-06-18, and @@ -173,18 +178,50 @@ func (s *Server) handlePOST(w http.ResponseWriter, r *http.Request) { if isBatch(body) { rejected(reasonBatch) http.Error(w, "JSON-RPC batches are not supported: send one request per POST", http.StatusBadRequest) - return + return nil, nil, false } // One parse for the whole path. The three questions this handler asks of the // body — does the method carry a slash, is there an id, what did the client // say on initialize — used to be three more parses of the same bytes on top // of this one. - var p fastjson.Parser - v, err := p.ParseBytes(body) + v, err = fastjson.ParseBytes(body) if err != nil { rejected(reasonParse) http.Error(w, "parse jsonrpc: "+err.Error(), http.StatusBadRequest) + return nil, nil, false + } + + // A member named twice where this handler reads the request has no reading + // that is safe to pick. fastjson, which answers every question asked here, + // takes the first, and by exact bytes; encoding/json, which decodes the call + // for zenrpc, takes the last, and without regard to case. Mcp-Name checked + // against one name while the other runs is the very hole the header exists + // to close. + if repeatsKey(v) { + rejected(reasonRepeatedKey) + http.Error(w, "parse jsonrpc: a member is named twice", http.StatusBadRequest) + return nil, nil, false + } + + return body, v, true +} + +func (s *Server) handlePOST(w http.ResponseWriter, r *http.Request) { + body, v, ok := s.readRequest(w, r) + if !ok { + return + } + + // A notification is a method that asks for no answer, and MCP names every + // one of them notifications/…. Anything else without an id is a call whose + // answer nobody will read, and zenrpc would run it detached, after this + // handler has returned: outside the limiter's concurrency slot and past the + // settling of its budget, so a tool called that way ran for free. "If the + // server cannot accept the input, it MUST return an HTTP error status code." + if isNotification(v) && !bytes.HasPrefix(v.GetStringBytes("method"), []byte("notifications/")) { + rejected(reasonMissingID) + http.Error(w, "a request needs an id: only notifications/* go without one", http.StatusBadRequest) return } @@ -374,6 +411,73 @@ func setProtocolVersion(w http.ResponseWriter, r *http.Request) { } } +// repeatsKey reports whether the request names a member twice in an object this +// handler reads: the envelope, its params, their _meta. Tool arguments are not +// among them — nothing here reads those, and only zenrpc decodes them. +// +// Twice as encoding/json counts, which matches a member to a field without +// regard to case, Unicode folding included: "name" and "Name", or "params" and +// "paramſ", are one member to the decoder behind and two to fastjson. +func repeatsKey(v *fastjson.Value) bool { + for _, o := range [...]*fastjson.Value{v, v.Get("params"), v.Get("params", "_meta")} { + if o != nil && o.Type() == fastjson.TypeObject && repeatsMember(o.GetObject()) { + return true + } + } + return false +} + +// pairwiseMembers is how many members an object can have and still be checked +// by comparing every pair. +const pairwiseMembers = 16 + +// repeatsMember reports whether two members of o fold to the same name. An +// honest object holds a handful of members and is compared pair by pair, on +// fastjson's own bytes, without an allocation. A larger one is folded into a +// set: a megabyte of distinct members compared pairwise was seconds of CPU. +func repeatsMember(o *fastjson.Object) bool { + repeated := false + if o.Len() <= pairwiseMembers { + var buf [pairwiseMembers][]byte + seen := buf[:0] + o.Visit(func(k []byte, _ *fastjson.Value) { + for _, s := range seen { + if bytes.EqualFold(s, k) { + repeated = true + } + } + seen = append(seen, k) + }) + return repeated + } + + seen := make(map[string]struct{}, o.Len()) + o.Visit(func(k []byte, _ *fastjson.Value) { + f := foldName(k) + if _, ok := seen[f]; ok { + repeated = true + } + seen[f] = struct{}{} + }) + return repeated +} + +// foldName spells k the same way for every name bytes.EqualFold holds equal to +// it: each rune becomes the smallest of its fold orbit, as encoding/json folds. +func foldName(k []byte) string { + b := make([]byte, 0, len(k)) + for len(k) > 0 { + r, n := utf8.DecodeRune(k) + k = k[n:] + // SimpleFold walks the orbit upward and wraps around to its smallest. + for next := unicode.SimpleFold(r); next > r; next = unicode.SimpleFold(r) { + r = next + } + b = utf8.AppendRune(b, unicode.SimpleFold(r)) + } + return string(b) +} + // isBatch reports whether the body is a JSON array — a batch, which this // server refuses. // diff --git a/transport_test.go b/transport_test.go index 362efca..3c73922 100644 --- a/transport_test.go +++ b/transport_test.go @@ -8,6 +8,7 @@ import ( "net/http" "net/http/httptest" "os" + "strconv" "strings" "sync" "testing" @@ -18,6 +19,7 @@ import ( "github.com/prometheus/client_golang/prometheus/testutil" "github.com/stretchr/testify/assert" "github.com/stretchr/testify/require" + "github.com/valyala/fastjson" "github.com/vmkteam/embedlog" "github.com/vmkteam/zenrpc/v2" "github.com/vmkteam/zenrpc/v2/smd" @@ -193,6 +195,83 @@ func TestHandlePOST_BrokenJSON(t *testing.T) { assert.Contains(t, rr.Body.String(), "parse jsonrpc") } +// A repeated member is read as its first occurrence here and as its last by +// the decoder behind, so the two would disagree about what the request is — +// which method, which tool, which revision. Refused before anything is +// dispatched. A repeat inside the arguments is the tool's business: nothing +// here reads them. +func TestHandlePOST_RepeatedKey(t *testing.T) { + t.Parallel() + refused := map[string]string{ + "method": `{"jsonrpc":"2.0","id":1,"method":"ping","method":"tools/list"}`, + "id": `{"jsonrpc":"2.0","id":1,"id":2,"method":"ping"}`, + "params": `{"jsonrpc":"2.0","id":1,"method":"tools/call","params":{"name":"a"},"params":{"name":"b"}}`, + "name": `{"jsonrpc":"2.0","id":1,"method":"tools/call","params":{"name":"help","name":"db_query"}}`, + "_meta": `{"jsonrpc":"2.0","id":1,"method":"ping","params":{"_meta":{` + + `"io.modelcontextprotocol/protocolVersion":"2026-07-28","io.modelcontextprotocol/protocolVersion":"2025-06-18"}}}`, + // encoding/json matches a member without regard to case, Unicode folding + // included, so these are one member to it and two to fastjson. + "Name": `{"jsonrpc":"2.0","id":1,"method":"tools/call","params":{"name":"help","Name":"db_query"}}`, + "Method": `{"jsonrpc":"2.0","id":1,"method":"ping","Method":"tools.call"}`, + "ID": `{"jsonrpc":"2.0","id":1,"ID":null,"method":"ping"}`, + "paramſ": `{"jsonrpc":"2.0","id":1,"method":"tools/call","params":{"name":"a"},"paramſ":{"name":"b"}}`, + "many members": `{"jsonrpc":"2.0","id":1,"method":"tools/call","params":{"name":"a",` + + manyMembers(100) + `,"NAME":"b"}}`, + } + for name, body := range refused { + t.Run(name, func(t *testing.T) { + t.Parallel() + srv, root, tools := newTestServer(t, Options{}) + rr := post(t, srv, body) + assert.Equal(t, http.StatusBadRequest, rr.Code) + assert.Contains(t, rr.Body.String(), "named twice") + assert.Empty(t, root.methods(), "nothing is dispatched") + assert.Empty(t, tools.methods(), "nothing is dispatched") + }) + } + + t.Run("inside the arguments", func(t *testing.T) { + t.Parallel() + srv, _, tools := newTestServer(t, Options{}) + rr := post(t, srv, `{"jsonrpc":"2.0","id":1,"method":"tools/call","params":{"name":"a","arguments":{"q":1,"q":2}}}`) + assert.Equal(t, http.StatusOK, rr.Code) + assert.Equal(t, []string{"call"}, tools.methods()) + }) + + t.Run("many members, none repeated", func(t *testing.T) { + t.Parallel() + srv, _, tools := newTestServer(t, Options{}) + rr := post(t, srv, `{"jsonrpc":"2.0","id":1,"method":"tools/call","params":{"name":"a",`+manyMembers(100)+`}}`) + assert.Equal(t, http.StatusOK, rr.Code) + assert.Equal(t, []string{"call"}, tools.methods()) + }) +} + +// A body shaped to make the check quadratic costs one pass over its members: +// a megabyte of distinct ones took seconds of CPU when every member was +// compared with every other. +func TestRepeatsKeyIsLinear(t *testing.T) { + t.Parallel() + v, err := fastjson.Parse(`{"jsonrpc":"2.0","id":1,"method":"tools/call","params":{` + manyMembers(100_000) + `}}`) + require.NoError(t, err) + start := time.Now() + assert.False(t, repeatsKey(v)) + assert.Less(t, time.Since(start), time.Second) +} + +// manyMembers is n distinct members, "m0":0,"m1":0,…, to be spliced into an +// object. +func manyMembers(n int) string { + var b strings.Builder + for i := range n { + if i > 0 { + b.WriteByte(',') + } + b.WriteString(`"m` + strconv.Itoa(i) + `":0`) + } + return b.String() +} + // A request the transport refuses never reaches a handler, so it shows up in no // other series this library publishes. Without this counter a client that sends // batches, or one whose JSON is broken, is visible only in somebody else's @@ -205,7 +284,7 @@ func TestTransportRejectionsAreCounted(t *testing.T) { } t.Run("every reason is published from the start", func(t *testing.T) { - assert.Equal(t, 11, testutil.CollectAndCount(rejectedTotal), + assert.Equal(t, 13, testutil.CollectAndCount(rejectedTotal), "rate() cannot tell 'nothing refused' from 'no data'") assert.Equal(t, 2, testutil.CollectAndCount(requestsTotal), "both eras are counted from zero") }) @@ -222,6 +301,18 @@ func TestTransportRejectionsAreCounted(t *testing.T) { assert.Greater(t, count(reasonParse), before) }) + t.Run("repeated key", func(t *testing.T) { + before := count(reasonRepeatedKey) + post(t, srv, `{"jsonrpc":"2.0","id":1,"method":"ping","method":"tools/list"}`) + assert.Greater(t, count(reasonRepeatedKey), before) + }) + + t.Run("missing id", func(t *testing.T) { + before := count(reasonMissingID) + post(t, srv, `{"jsonrpc":"2.0","method":"tools/call","params":{"name":"a"}}`) + assert.Greater(t, count(reasonMissingID), before) + }) + t.Run("too large", func(t *testing.T) { small, _, _ := newTestServer(t, Options{MaxRequestBytes: 64}) before := count(reasonTooLarge) @@ -380,8 +471,8 @@ func TestNullIDIsANotification(t *testing.T) { srv, _, _ := newTestServer(t, Options{}) for name, body := range map[string]string{ - "null id": `{"jsonrpc":"2.0","id":null,"method":"ping","params":{}}`, - "absent id": `{"jsonrpc":"2.0","method":"ping","params":{}}`, + "null id": `{"jsonrpc":"2.0","id":null,"method":"notifications/initialized","params":{}}`, + "absent id": `{"jsonrpc":"2.0","method":"notifications/initialized","params":{}}`, } { t.Run(name, func(t *testing.T) { t.Parallel() @@ -403,6 +494,30 @@ func TestNullIDIsANotification(t *testing.T) { }) } +// Only a notification goes without an id. A call sent without one would be run +// detached from the request — outside the rate limiter's slot and after its +// budget was settled — and its answer read by nobody. +func TestCallWithoutIDIsRefused(t *testing.T) { + t.Parallel() + for name, body := range map[string]string{ + "absent id": `{"jsonrpc":"2.0","method":"tools/call","params":{"name":"a","arguments":{}}}`, + "null id": `{"jsonrpc":"2.0","id":null,"method":"tools/call","params":{"name":"a","arguments":{}}}`, + "an id in capitals": `{"jsonrpc":"2.0","ID":1,"method":"tools/call","params":{"name":"a","arguments":{}}}`, + "zenrpc's own name": `{"jsonrpc":"2.0","method":"tools.call","params":{"name":"a","arguments":{}}}`, + "ping": `{"jsonrpc":"2.0","method":"ping"}`, + } { + t.Run(name, func(t *testing.T) { + t.Parallel() + srv, root, tools := newTestServer(t, Options{}) + rr := post(t, srv, body) + assert.Equal(t, http.StatusBadRequest, rr.Code) + assert.Contains(t, rr.Body.String(), "needs an id") + assert.Empty(t, root.methods(), "nothing is dispatched") + assert.Empty(t, tools.methods(), "nothing is dispatched") + }) + } +} + // The robustness table, modelled on go-sdk's bad_requests.txtar — which is not // a list of hypotheticals but of panics they actually shipped and fixed (their // issues #194–#197). @@ -420,10 +535,10 @@ func TestMalformedBodiesGetDefinedAnswers(t *testing.T) { body string code int }{ - {"initialize without id", `{"jsonrpc":"2.0","method":"initialize","params":{"protocolVersion":"2025-06-18"}}`, http.StatusAccepted}, + {"initialize without id", `{"jsonrpc":"2.0","method":"initialize","params":{"protocolVersion":"2025-06-18"}}`, http.StatusBadRequest}, {"initialize without params", `{"jsonrpc":"2.0","id":1,"method":"initialize"}`, http.StatusOK}, {"initialize with null params", `{"jsonrpc":"2.0","id":2,"method":"initialize","params":null}`, http.StatusOK}, - {"ping without id", `{"jsonrpc":"2.0","method":"ping"}`, http.StatusAccepted}, + {"ping without id", `{"jsonrpc":"2.0","method":"ping"}`, http.StatusBadRequest}, {"notification carrying an id", `{"jsonrpc":"2.0","id":3,"method":"notifications/initialized"}`, http.StatusOK}, {"notification without one", `{"jsonrpc":"2.0","method":"notifications/initialized"}`, http.StatusAccepted}, {"tools/call without params", `{"jsonrpc":"2.0","id":4,"method":"tools/call"}`, http.StatusOK}, @@ -434,7 +549,7 @@ func TestMalformedBodiesGetDefinedAnswers(t *testing.T) { {"a bare string", `"hello"`, http.StatusOK}, {"a bare number", `42`, http.StatusOK}, {"a string id", `{"jsonrpc":"2.0","id":"abc","method":"ping","params":{}}`, http.StatusOK}, - {"a null id", `{"jsonrpc":"2.0","id":null,"method":"ping","params":{}}`, http.StatusAccepted}, + {"a null id", `{"jsonrpc":"2.0","id":null,"method":"ping","params":{}}`, http.StatusBadRequest}, {"no jsonrpc member", `{"id":1,"method":"ping","params":{}}`, http.StatusOK}, {"no method member", `{"jsonrpc":"2.0","id":1,"params":{}}`, http.StatusOK}, {"deeply nested params", `{"jsonrpc":"2.0","id":1,"method":"ping","params":` + strings.Repeat(`{"a":`, 200) + `1` + strings.Repeat(`}`, 200) + `}`, http.StatusOK},