package mcpserver import ( "context" "net/http" "strconv" "sync" "testing" "time" ) // The per-organisation ceiling. // // The property under test is not "a limit exists" but "the limit is keyed by an // organisation the CALLER CANNOT CHOOSE". Every test below therefore checks // which bucket was charged, not merely that something was refused. // countingOrgLimiter records which org was charged and refuses past a limit. type countingOrgLimiter struct { mu sync.Mutex counts map[string]int limit int err error } func newCountingOrgLimiter(limit int) *countingOrgLimiter { return &countingOrgLimiter{counts: map[string]int{}, limit: limit} } func (c *countingOrgLimiter) AllowOrg(_ context.Context, orgID string) (bool, time.Duration, error) { c.mu.Lock() defer c.mu.Unlock() if c.err != nil { return false, time.Minute, c.err } c.counts[orgID]++ return c.counts[orgID] <= c.limit, 30 * time.Second, nil } func (c *countingOrgLimiter) count(orgID string) int { c.mu.Lock() defer c.mu.Unlock() return c.counts[orgID] } func (c *countingOrgLimiter) buckets() int { c.mu.Lock() defer c.mu.Unlock() return len(c.counts) } // orgLimitEnv is the matrix fixture plus a counting limiter. func orgLimitEnv(t *testing.T, limit int) (*matrixEnv, *countingOrgLimiter) { t.Helper() m := newMatrix(t) limiter := newCountingOrgLimiter(limit) m.server = m.server.WithOrgLimiter(limiter) return m, limiter } /* ── Each organisation gets its own bucket ──────────────────────────────── */ func TestEachOrganisationIsCountedSeparately(t *testing.T) { m, limiter := orgLimitEnv(t, 100) for i := 0; i < 3; i++ { m.call(t, "tok-a-admin", "activity_breakdown", "{}") } for i := 0; i < 5; i++ { m.call(t, "tok-b-admin", "activity_breakdown", "{}") } if got := limiter.count(m.a.orgID); got != 3 { t.Errorf("org A charged %d, want 3", got) } if got := limiter.count(m.b.orgID); got != 5 { t.Errorf("org B charged %d, want 5", got) } if limiter.buckets() != 2 { t.Errorf("%d buckets, want 2 — the two tenants shared a counter", limiter.buckets()) } } // One organisation exhausting its quota must not affect another's. func TestOneOrganisationCannotConsumeAnothersQuota(t *testing.T) { m, limiter := orgLimitEnv(t, 3) // Burn org A's entire budget and then some. for i := 0; i < 10; i++ { m.call(t, "tok-a-admin", "activity_breakdown", "{}") } // Org B must be untouched. body := m.call(t, "tok-b-admin", "activity_breakdown", "{}") if isRateLimited(body) { t.Error("org B was refused because org A exhausted its quota") } if got := limiter.count(m.b.orgID); got != 1 { t.Errorf("org B charged %d, want 1", got) } } /* ── The bucket cannot be chosen by the caller ──────────────────────────── */ // Every channel a client controls, against the ORG LIMITER specifically. A // request naming org B must still be charged to org A. func TestTheOrgBucketCannotBeSelectedByTheRequest(t *testing.T) { for name, tc := range map[string]struct { args string mutate func(*http.Request) path string }{ "org_id argument": { args: `{"org_id":"OTHER"}`, }, "tenant_id argument": { args: `{"tenant_id":"OTHER","organization_id":"OTHER"}`, }, "identity headers": { args: `{}`, mutate: func(r *http.Request) { for _, h := range []string{"X-Org-Id", "X-Tenant-Id", "X-Organization-Id"} { r.Header.Set(h, "OTHER") } }, }, "query string": { args: `{}`, path: "/mcp?org_id=OTHER&tenant_id=OTHER", }, } { t.Run(name, func(t *testing.T) { m, limiter := orgLimitEnv(t, 100) args := replaceAll(tc.args, "OTHER", m.b.orgID) path := tc.path if path == "" { path = "/mcp" } m.callWith(t, "tok-a-admin", "activity_breakdown", args, path, tc.mutate) if got := limiter.count(m.a.orgID); got != 1 { t.Errorf("org A charged %d, want 1 — the caller's own org must be charged", got) } if got := limiter.count(m.b.orgID); got != 0 { t.Errorf("org B charged %d, want 0 — the request selected another tenant's bucket", got) } }) } } // _meta at both levels, which is the channel most likely to be trusted by // accident because it is "protocol" rather than "arguments". func TestMetaCannotSelectTheOrgBucket(t *testing.T) { m, limiter := orgLimitEnv(t, 100) body := `{"jsonrpc":"2.0","id":1,"method":"tools/call",` + `"_meta":{"org_id":"` + m.b.orgID + `"},` + `"params":{"name":"activity_breakdown","arguments":{},` + `"_meta":{"org_id":"` + m.b.orgID + `","tenant_id":"` + m.b.orgID + `"}}}` m.raw(t, "tok-a-admin", body, "/mcp", nil) if got := limiter.count(m.a.orgID); got != 1 { t.Errorf("org A charged %d, want 1", got) } if got := limiter.count(m.b.orgID); got != 0 { t.Errorf("org B charged %d, want 0 — _meta selected another tenant's bucket", got) } } /* ── Refusal behaviour ──────────────────────────────────────────────────── */ func TestOverTheOrgLimitReturns429WithRetryAfter(t *testing.T) { m, _ := orgLimitEnv(t, 2) // Two are allowed. for i := 0; i < 2; i++ { if rec := m.callRec(t, "tok-a-admin", "activity_breakdown", "{}"); rec.Code != http.StatusOK { t.Fatalf("call %d: status = %d, want 200", i+1, rec.Code) } } rec := m.callRec(t, "tok-a-admin", "activity_breakdown", "{}") if rec.Code != http.StatusTooManyRequests { t.Fatalf("status = %d, want 429", rec.Code) } retry := rec.Header().Get("Retry-After") if retry == "" { t.Fatal("a 429 carried no Retry-After") } secs, err := strconv.Atoi(retry) if err != nil || secs < 1 { t.Errorf("Retry-After = %q, want a positive whole number of seconds", retry) } // The refusal must not describe the quota or name the organisation — how // much a tenant has spent is not something one caller learns from a 429. body := rec.Body.String() if containsAny(body, []string{m.a.orgID, "5000", "quota", "remaining"}) { t.Errorf("the 429 body leaks quota or tenant detail: %s", body) } } // The ceiling is checked BEFORE the tool runs, so a refused call costs no // database work. func TestTheOrgLimitIsCheckedBeforeTheToolRuns(t *testing.T) { m, _ := orgLimitEnv(t, 0) // nothing is allowed rec := m.callRec(t, "tok-a-admin", "activity_breakdown", "{}") if rec.Code != http.StatusTooManyRequests { t.Fatalf("status = %d, want 429", rec.Code) } // A tool that had run would have produced a result payload. if containsAny(rec.Body.String(), []string{"totalEvents", "distinctKinds"}) { t.Error("the tool ran despite the organisation being over its limit") } } // Unauthenticated requests must be refused before the limiter is consulted — // otherwise an anonymous caller could burn a tenant's quota. func TestTheOrgLimiterIsNotConsultedWithoutAuthentication(t *testing.T) { m, limiter := orgLimitEnv(t, 100) m.raw(t, "", `{"jsonrpc":"2.0","id":1,"method":"tools/call","params":`+ `{"name":"activity_breakdown","arguments":{}}}`, "/mcp", nil) if limiter.buckets() != 0 { t.Errorf("%d buckets charged by an unauthenticated request, want 0", limiter.buckets()) } } /* ── Concurrency and failure ────────────────────────────────────────────── */ // Concurrent calls must be counted atomically: the limiter's own contract, here // exercised through the full MCP path. Run with -race. func TestConcurrentCallsAreCountedAtomically(t *testing.T) { m, limiter := orgLimitEnv(t, 1000) const callers = 30 var wg sync.WaitGroup for i := 0; i < callers; i++ { wg.Add(1) go func() { defer wg.Done() m.call(t, "tok-a-admin", "activity_breakdown", "{}") }() } wg.Wait() if got := limiter.count(m.a.orgID); got != callers { t.Errorf("org A charged %d, want %d — increments were lost", got, callers) } } // A limiter that errors must not let the call through: fail closed. func TestAFailingOrgLimiterRefusesTheCall(t *testing.T) { m, limiter := orgLimitEnv(t, 100) limiter.err = errString("limiter unavailable") rec := m.callRec(t, "tok-a-admin", "activity_breakdown", "{}") if rec.Code != http.StatusTooManyRequests { t.Errorf("status = %d, want 429 — a limiter that cannot count must not permit the call", rec.Code) } } // With no limiter installed there is no ceiling, and nothing breaks. func TestNoOrgLimiterMeansNoCeiling(t *testing.T) { m := newMatrix(t) // no WithOrgLimiter for i := 0; i < 20; i++ { if rec := m.callRec(t, "tok-a-admin", "activity_breakdown", "{}"); rec.Code != http.StatusOK { t.Fatalf("call %d: status = %d, want 200", i+1, rec.Code) } } } /* ── Helpers ────────────────────────────────────────────────────────────── */ type errString string func (e errString) Error() string { return string(e) } func containsAny(s string, needles []string) bool { for _, n := range needles { if n != "" && contains(s, n) { return true } } return false } func contains(s, sub string) bool { return len(sub) > 0 && indexOf(s, sub) >= 0 } func indexOf(s, sub string) int { for i := 0; i+len(sub) <= len(s); i++ { if s[i:i+len(sub)] == sub { return i } } return -1 } func replaceAll(s, old, new string) string { out := "" for { i := indexOf(s, old) if i < 0 { return out + s } out += s[:i] + new s = s[i+len(old):] } }