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

315 lines
9.4 KiB
Go

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):]
}
}