Files
krow_backend/go-api/internal/mcpserver/tenant_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

633 lines
23 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
package mcpserver
import (
"context"
"encoding/json"
"fmt"
"io"
"log/slog"
"net/http"
"net/http/httptest"
"strings"
"testing"
"github.com/jackc/pgx/v5/pgxpool"
"github.com/krow/krow-backend/go-api/internal/authctx"
"github.com/krow/krow-backend/go-api/internal/runtime"
"github.com/krow/krow-backend/go-api/internal/testutil"
"github.com/krow/krow-backend/go-api/internal/tools"
)
// The tenant isolation matrix: every exposed tool, both organisations.
//
// THE METHOD, AND WHY IT IS NOT "CHECK THE ROWS"
//
// Comparing returned rows catches the obvious leak and misses the ones that
// matter. A tool that correctly withholds Org B's records while counting them
// in a total has leaked; so has one whose "no data" answer differs from its
// "no access" answer, or whose error names a record it will not show.
//
// So each tool is called twice — once as Org A, once as Org B — over identical
// but DISTINGUISHABLE data, and the two complete responses are compared as
// text. Any value that differs between tenants must be a value that came from
// that tenant. A number, a name, an id or a flag that crosses is caught
// whatever part of the payload it hides in, including aggregates, counts,
// metadata and error text.
//
// The fixtures are deliberately lopsided — Org B has several times Org A's
// volume — so a leak shows up as a wrong NUMBER, not merely a wrong name. A
// total of 60 where 7 was correct is unmistakable in a way that a missing name
// is not.
/* ── Fixtures ───────────────────────────────────────────────────────────── */
// tenant is one seeded organisation and the identities that can act for it.
type tenant struct {
orgID string
label string
admin authctx.Identity
employer authctx.Identity
talent authctx.Identity
// scale multiplies every seeded row count, so the two tenants' numbers
// cannot coincide by accident.
scale int
}
func seedTenant(t *testing.T, h *testutil.Harness, label string, scale int) tenant {
t.Helper()
ctx := context.Background()
var orgID string
if err := h.Pool.QueryRow(ctx,
`INSERT INTO organizations (name, slug) VALUES ($1, $2) RETURNING id::text`,
"Tenant "+label, "tenant-"+strings.ToLower(label)).Scan(&orgID); err != nil {
t.Fatalf("org %s: %v", label, err)
}
mkUser := func(role, email string) authctx.Identity {
var id string
if err := h.Pool.QueryRow(ctx,
`INSERT INTO users (org_id, email, full_name, role, account_type, status)
VALUES ($1::uuid, $2, $3, $4, $5, 'active') RETURNING id::text`,
orgID, email, "User "+label, role, accountTypeFor(role)).Scan(&id); err != nil {
t.Fatalf("user %s: %v", email, err)
}
return authctx.Identity{
UserID: id, OrgID: orgID, Email: email, FullName: "User " + label,
Role: role, AccountType: accountTypeFor(role), Status: "active",
}
}
tn := tenant{
orgID: orgID, label: label, scale: scale,
admin: mkUser("admin", "admin-"+label+"@tenant.test"),
employer: mkUser("employer", "employer-"+label+"@tenant.test"),
talent: mkUser("talent", "talent-"+label+"@tenant.test"),
}
seedTenantData(t, h, tn)
return tn
}
func accountTypeFor(role string) string {
if role == "talent" {
return "talent"
}
return "employer"
}
// seedTenantData fills every table the 16 tools read.
//
// Every value carries the tenant's label, so a leaked string is identifiable on
// sight rather than by cross-referencing ids.
func seedTenantData(t *testing.T, h *testutil.Harness, tn tenant) {
t.Helper()
ctx := context.Background()
n := tn.scale
exec := func(sql string, args ...any) {
t.Helper()
if _, err := h.Pool.Exec(ctx, sql, args...); err != nil {
t.Fatalf("seed %s: %v", tn.label, err)
}
}
// Activity, SPREAD ACROSS DAYS. activity_signals refuses to call anything
// unusual without at least three days of history, so rows all stamped now
// would make it answer "not enough history" for both tenants — which would
// let the isolation check pass while testing nothing.
for i := 0; i < n*3; i++ {
exec(`INSERT INTO user_activity (org_id, event_type, user_email, user_name, details, created_date)
VALUES ($1::uuid, 'login', $2, $3, $4, now() - ($5::int * interval '1 day'))`,
tn.orgID, "actor-"+tn.label+"@tenant.test", "Actor "+tn.label,
"detail-"+tn.label, i%7)
}
for i := 0; i < n; i++ {
exec(`INSERT INTO user_activity (org_id, event_type, user_email, user_name, details, created_date)
VALUES ($1::uuid, 'create_position', $2, $3, $4, now() - ($5::int * interval '1 day'))`,
tn.orgID, "actor-"+tn.label+"@tenant.test", "Actor "+tn.label,
"created-"+tn.label, i%5)
}
// Postings, and the applications against them.
for i := 0; i < n; i++ {
var postingID string
if err := h.Pool.QueryRow(ctx,
`INSERT INTO job_postings (org_id, title, status, headcount, priority)
VALUES ($1::uuid, $2, 'active', 3, 'normal') RETURNING id::text`,
tn.orgID, fmt.Sprintf("Role-%s-%d", tn.label, i)).Scan(&postingID); err != nil {
t.Fatalf("posting %s: %v", tn.label, err)
}
for j := 0; j < n; j++ {
exec(`INSERT INTO job_applications
(org_id, job_posting_id, applicant_name, email, status, ai_score, job_title)
VALUES ($1::uuid, $2::uuid, $3, $4, $5, $6, $7)`,
tn.orgID, postingID,
fmt.Sprintf("Candidate-%s-%d-%d", tn.label, i, j),
fmt.Sprintf("cand-%s-%d-%d@tenant.test", tn.label, i, j),
[]string{"applied", "ai_screened", "hired"}[j%3],
70+j, fmt.Sprintf("Role-%s-%d", tn.label, i))
}
}
// Staff and worker profiles.
for i := 0; i < n; i++ {
exec(`INSERT INTO staff (org_id, name, email, role, status, hire_date)
VALUES ($1::uuid, $2, $3, $4, 'active', CURRENT_DATE)`,
tn.orgID, fmt.Sprintf("Staff-%s-%d", tn.label, i),
fmt.Sprintf("staff-%s-%d@tenant.test", tn.label, i),
"Server-"+tn.label)
exec(`INSERT INTO worker_profiles
(org_id, full_name, email, krow_score, reliability_score, shifts_completed, status)
VALUES ($1::uuid, $2, $3, 80, 90, 5, 'active')`,
tn.orgID, fmt.Sprintf("Worker-%s-%d", tn.label, i),
fmt.Sprintf("worker-%s-%d@tenant.test", tn.label, i))
}
// Shift records — attendance, overtime and coverage all read these.
for i := 0; i < n*2; i++ {
status := []string{"present", "late", "absent"}[i%3]
actual := 8 + i%4
// shift_records_absence_has_no_hours: the schema refuses an absence
// that logged time, which is the domain rule rather than a quirk —
// honoured here so the fixtures are records the product could produce.
if status == "absent" {
actual = 0
}
daysAgo := i % 14
exec(`INSERT INTO shift_records
(org_id, worker_email, worker_name, role, status, scheduled_hours, actual_hours,
shift_date, scheduled_start, scheduled_end, overtime_hours, created_date)
VALUES ($1::uuid, $2, $3, $4, $5, 8, $6::numeric,
CURRENT_DATE - ($7::int * interval '1 day'),
now() - ($7::int * interval '1 day'),
now() - ($7::int * interval '1 day') + interval '8 hours',
$8::numeric, now())`,
tn.orgID, fmt.Sprintf("worker-%s-%d@tenant.test", tn.label, i%n),
fmt.Sprintf("Worker-%s-%d", tn.label, i%n), "Server-"+tn.label,
status, actual, daysAgo, max(actual-8, 0))
}
// Courses, for workforce_training.
for i := 0; i < n; i++ {
exec(`INSERT INTO courses (org_id, title, status)
VALUES ($1::uuid, $2, 'active')`,
tn.orgID, fmt.Sprintf("Course-%s-%d", tn.label, i))
}
}
// requiredArgs supplies arguments for the tools that cannot be called bare.
//
// Only two need anything. available_workers takes a shift window — it is a
// lookup for "who could work THIS" — and a call without one is an invalid
// input rather than an empty result. Everything else answers a bare {}.
//
// The values are tenant-neutral on purpose: nothing here names an
// organisation, so the only thing that can scope the answer is the token.
var requiredArgs = map[string]string{
"available_workers": `{"starts_at":"2026-09-20T18:00:00Z","ends_at":"2026-09-21T02:00:00Z"}`,
}
/* ── The matrix ─────────────────────────────────────────────────────────── */
// exposedToolNames is the set under test, taken from the registry rather than
// written out, so a tool added to the surface is automatically covered.
func exposedToolNames(reg *tools.Registry) []string {
infos := exposed(reg)
out := make([]string, 0, len(infos))
for _, i := range infos {
out = append(out, i.Name)
}
return out
}
type matrixEnv struct {
server *Server
pool *pgxpool.Pool
a, b tenant
tokens map[string]authctx.Identity
}
func newMatrix(t *testing.T) *matrixEnv {
t.Helper()
h := testutil.New(t)
// Lopsided on purpose: Org B's numbers are several times Org A's, so a
// leaked aggregate is a wrong number rather than a plausible one.
a := seedTenant(t, h, "A", 2)
b := seedTenant(t, h, "B", 5)
tokens := map[string]authctx.Identity{
"tok-a-admin": a.admin,
"tok-a-employer": a.employer,
"tok-a-talent": a.talent,
"tok-b-admin": b.admin,
"tok-b-employer": b.employer,
"tok-b-talent": b.talent,
}
srv := New(runtime.DefaultTools(h.Pool, nil), &fakeTokens{byToken: tokens},
slog.New(slog.NewTextHandler(io.Discard, nil)))
return &matrixEnv{server: srv, pool: h.Pool, a: a, b: b, tokens: tokens}
}
// call invokes one tool and returns the whole response body as text.
func (m *matrixEnv) call(t *testing.T, token, tool string, args string) string {
t.Helper()
return m.callRec(t, token, tool, args).Body.String()
}
// callRec is call, returning the whole recorder so a test can read the status
// and the headers — which is what a 429 assertion needs.
func (m *matrixEnv) callRec(t *testing.T, token, tool, args string) *httptest.ResponseRecorder {
t.Helper()
return m.callWith(t, token, tool, args, "/mcp", nil)
}
// callWith is callRec with a path and a hook for mutating the request, so the
// header- and query-injection tests can drive the same path.
func (m *matrixEnv) callWith(t *testing.T, token, tool, args, path string,
mutate func(*http.Request)) *httptest.ResponseRecorder {
t.Helper()
if args == "" {
args = "{}"
}
body := `{"jsonrpc":"2.0","id":1,"method":"tools/call","params":{"name":"` +
tool + `","arguments":` + args + `}}`
return m.raw(t, token, body, path, mutate)
}
// raw posts an arbitrary JSON-RPC body, for tests that need to shape the
// envelope themselves (_meta injection, for one).
func (m *matrixEnv) raw(t *testing.T, token, body, path string,
mutate func(*http.Request)) *httptest.ResponseRecorder {
t.Helper()
req := httptest.NewRequest(http.MethodPost, path, strings.NewReader(body))
req.Header.Set("Content-Type", "application/json")
if token != "" {
req.Header.Set("Authorization", "Bearer "+token)
}
if mutate != nil {
mutate(req)
}
rec := httptest.NewRecorder()
m.server.Handler().ServeHTTP(rec, req)
return rec
}
// isRateLimited reports whether a response body is the org-ceiling refusal.
func isRateLimited(body string) bool {
return strings.Contains(body, "too many requests for this organisation")
}
// TestTenantIsolationMatrix is the heart of Phase 5.
//
// Sixteen tools × two organisations. For each, the ENTIRE response for Org A is
// searched for every marker belonging to Org B, and the reverse. A marker is
// any string that identifies the other tenant — its label, its names, its
// emails, its org id.
func TestTenantIsolationMatrix(t *testing.T) {
m := newMatrix(t)
names := exposedToolNames(m.server.reg)
if len(names) != 16 {
t.Fatalf("%d exposed tools, want 16 — the matrix must cover all of them", len(names))
}
for _, tool := range names {
t.Run(tool, func(t *testing.T) {
args := requiredArgs[tool]
asA := m.call(t, "tok-a-admin", tool, args)
asB := m.call(t, "tok-b-admin", tool, args)
// Neither response may be an internal failure: a tool that errors
// for both tenants would pass a leak check vacuously.
for label, body := range map[string]string{"A": asA, "B": asB} {
if strings.Contains(body, `"code":-32603`) {
t.Fatalf("org %s: the tool failed internally, so isolation is untested: %s",
label, truncate(body))
}
}
assertNoLeak(t, "A", asA, m.b)
assertNoLeak(t, "B", asB, m.a)
// The two tenants must not produce IDENTICAL payloads. If they do,
// either the tool ignores the tenant entirely (a leak) or it
// returns nothing for both (in which case this test proves nothing
// and should be known to prove nothing).
if asA == asB && !strings.Contains(asA, `"data":null`) {
t.Errorf("both tenants received a byte-identical response; "+
"the tool may not be scoping by organisation at all:\n%s", truncate(asA))
}
})
}
}
// assertNoLeak searches one tenant's response for any trace of the other.
func assertNoLeak(t *testing.T, whose, body string, other tenant) {
t.Helper()
markers := map[string]string{
"organisation id": other.orgID,
"actor email": "actor-" + other.label + "@tenant.test",
"staff name": "Staff-" + other.label,
"worker name": "Worker-" + other.label,
"candidate name": "Candidate-" + other.label,
"posting title": "Role-" + other.label,
"course title": "Course-" + other.label,
"role label": "Server-" + other.label,
"detail text": "detail-" + other.label,
"admin email": other.admin.Email,
"user id": other.admin.UserID,
}
for what, marker := range markers {
if strings.Contains(body, marker) {
t.Errorf("org %s's response contains org %s's %s (%q):\n%s",
whose, other.label, what, marker, truncate(body))
}
}
}
func truncate(s string) string {
if len(s) > 1200 {
return s[:1200] + "… [truncated]"
}
return s
}
/* ── Aggregates and side channels ───────────────────────────────────────── */
// A leak through a NUMBER rather than a name.
//
// Org B has far more of everything. If a tool's totals for Org A are affected
// by Org B's rows, the number will be wrong even though no name crosses. This
// asserts the arithmetic directly against the database rather than against the
// other tenant's response, so it catches a tool that counts everything and
// shows only some.
func TestAggregatesAreScopedToTheTenant(t *testing.T) {
m := newMatrix(t)
ctx := context.Background()
// activity_breakdown reports totalEvents, which must equal exactly this
// tenant's rows and not one more.
for _, tn := range []tenant{m.a, m.b} {
token := "tok-" + strings.ToLower(tn.label) + "-admin"
body := m.call(t, token, "activity_breakdown", "{}")
var want int
if err := m.pool.QueryRow(ctx,
`SELECT count(*) FROM user_activity WHERE org_id = $1::uuid`, tn.orgID).Scan(&want); err != nil {
t.Fatalf("count: %v", err)
}
got := extractInt(t, body, "totalEvents")
if got != want {
t.Errorf("org %s: totalEvents = %d, want %d (this tenant's rows only)",
tn.label, got, want)
}
}
}
// An empty tenant must not be able to infer that another tenant is not empty.
//
// The classic side channel: Org A has no data of some kind, Org B has plenty,
// and the "nothing here" answer differs from the "nothing you may see" answer
// in a way that reveals the difference.
func TestAnEmptyTenantLearnsNothingAboutAFullOne(t *testing.T) {
h := testutil.New(t)
full := seedTenant(t, h, "Full", 6)
empty := seedTenantEmpty(t, h, "Empty")
srv := New(runtime.DefaultTools(h.Pool, nil), &fakeTokens{byToken: map[string]authctx.Identity{
"tok-empty": empty.admin,
"tok-full": full.admin,
}}, slog.New(slog.NewTextHandler(io.Discard, nil)))
m := &matrixEnv{server: srv, pool: h.Pool, a: empty, b: full}
for _, tool := range exposedToolNames(srv.reg) {
t.Run(tool, func(t *testing.T) {
body := m.call(t, "tok-empty", tool, requiredArgs[tool])
// Nothing of the full tenant's may appear.
assertNoLeak(t, "Empty", body, full)
// And no number in the empty tenant's response may match the full
// tenant's scale, which would mean a count escaped its filter.
for _, n := range []string{`:6`, `:36`, `:12`} {
if strings.Contains(strings.ReplaceAll(body, " ", ""), n) &&
strings.Contains(body, "total") {
t.Logf("note: %s contains %s; verify it is not the other tenant's count", tool, n)
}
}
})
}
}
func seedTenantEmpty(t *testing.T, h *testutil.Harness, label string) tenant {
t.Helper()
ctx := context.Background()
var orgID string
if err := h.Pool.QueryRow(ctx,
`INSERT INTO organizations (name, slug) VALUES ($1, $2) RETURNING id::text`,
"Tenant "+label, "tenant-"+strings.ToLower(label)).Scan(&orgID); err != nil {
t.Fatalf("org: %v", err)
}
var id string
if err := h.Pool.QueryRow(ctx,
`INSERT INTO users (org_id, email, full_name, role, account_type, status)
VALUES ($1::uuid, $2, 'Empty Admin', 'admin', 'employer', 'active') RETURNING id::text`,
orgID, "admin-"+label+"@tenant.test").Scan(&id); err != nil {
t.Fatalf("user: %v", err)
}
return tenant{orgID: orgID, label: label, admin: authctx.Identity{
UserID: id, OrgID: orgID, Email: "admin-" + label + "@tenant.test",
Role: "admin", AccountType: "employer", Status: "active",
}}
}
/* ── Role matrix ────────────────────────────────────────────────────────── */
// The three real KROW roles against all 16 tools.
//
// This does NOT assert which tools each role may reach — that is the policy
// table's business and it is the source of truth, not this test. What it
// asserts is the two properties that must hold whatever the policy says:
// a refusal must be opaque, and no role may see another tenant.
func TestRoleMatrixAcrossBothTenants(t *testing.T) {
m := newMatrix(t)
names := exposedToolNames(m.server.reg)
for _, role := range []string{"admin", "employer", "talent"} {
for _, tn := range []struct {
label string
other tenant
}{{"a", m.b}, {"b", m.a}} {
token := "tok-" + tn.label + "-" + role
for _, tool := range names {
t.Run(role+"/"+tn.label+"/"+tool, func(t *testing.T) {
body := m.call(t, token, tool, requiredArgs[tool])
// Whatever the policy decides, the other tenant must not
// appear in the answer — including in a refusal.
assertNoLeak(t, role+"/"+tn.label, body, tn.other)
// A denial must be the single opaque one. A refusal that
// explained itself would describe the shape of what it is
// hiding.
if strings.Contains(body, "tool.denied") {
if !strings.Contains(body, "the caller does not have access to this") {
t.Errorf("a denial carried detail beyond the standard message: %s", truncate(body))
}
}
})
}
}
}
}
/* ── Injection: no request-supplied identity may influence anything ─────── */
// Every channel a client controls, against every tool.
//
// The earlier phases tested this on one tool. Here it is every exposed tool,
// because a single handler that read an argument it should not would be enough.
func TestNoRequestSuppliedIdentityInfluencesAnyTool(t *testing.T) {
m := newMatrix(t)
names := exposedToolNames(m.server.reg)
// Arguments naming the other tenant, in every spelling a caller might try.
hostileArgs := `{"org_id":"` + m.b.orgID + `","organization_id":"` + m.b.orgID +
`","tenant_id":"` + m.b.orgID + `","user_id":"` + m.b.admin.UserID +
`","orgId":"` + m.b.orgID + `","principal":"` + m.b.admin.Email +
`","email":"` + m.b.admin.Email + `","on_behalf_of":"` + m.b.admin.Email + `"}`
for _, tool := range names {
t.Run(tool, func(t *testing.T) {
clean := m.call(t, "tok-a-admin", tool, requiredArgs[tool])
hostile := m.call(t, "tok-a-admin", tool, hostileArgs)
// Whatever the tool does with unknown arguments — ignore them, or
// refuse the call — Org B must not appear.
assertNoLeak(t, "A(hostile args)", hostile, m.b)
// And the answer must not have CHANGED in a way that suggests the
// arguments were honoured. A tool that refuses unknown fields is
// fine; one that returns different DATA is not.
if hostile != clean && !strings.Contains(hostile, "error") {
t.Errorf("hostile arguments changed a successful response:\nclean: %s\nhostile: %s",
truncate(clean), truncate(hostile))
}
})
}
}
// Identity in JSON-RPC metadata, headers and the query string.
func TestIdentityChannelsOutsideArgumentsAreIgnored(t *testing.T) {
m := newMatrix(t)
baseline := m.call(t, "tok-a-admin", "activity_breakdown", "{}")
baselineTotal := extractInt(t, baseline, "totalEvents")
send := func(t *testing.T, mutate func(*http.Request), path string) string {
t.Helper()
body := `{"jsonrpc":"2.0","id":1,"method":"tools/call","_meta":{"org_id":"` + m.b.orgID +
`","user_id":"` + m.b.admin.UserID + `"},"params":{"name":"activity_breakdown",` +
`"arguments":{},"_meta":{"org_id":"` + m.b.orgID + `"}}}`
req := httptest.NewRequest(http.MethodPost, path, strings.NewReader(body))
req.Header.Set("Content-Type", "application/json")
req.Header.Set("Authorization", "Bearer tok-a-admin")
mutate(req)
rec := httptest.NewRecorder()
m.server.Handler().ServeHTTP(rec, req)
return rec.Body.String()
}
t.Run("jsonrpc _meta at both levels", func(t *testing.T) {
got := extractInt(t, send(t, func(*http.Request) {}, "/mcp"), "totalEvents")
if got != baselineTotal {
t.Errorf("totalEvents = %d, want %d — _meta moved the tenant", got, baselineTotal)
}
})
t.Run("identity headers", func(t *testing.T) {
body := send(t, func(r *http.Request) {
for _, h := range []string{
"X-Org-Id", "X-Organization-Id", "X-Tenant-Id", "X-User-Id",
"X-Krow-Org", "X-Krow-User", "X-On-Behalf-Of", "X-Forwarded-User",
} {
r.Header.Set(h, m.b.orgID)
}
}, "/mcp")
if got := extractInt(t, body, "totalEvents"); got != baselineTotal {
t.Errorf("totalEvents = %d, want %d — a header moved the tenant", got, baselineTotal)
}
assertNoLeak(t, "A(headers)", body, m.b)
})
t.Run("query string", func(t *testing.T) {
body := send(t, func(*http.Request) {},
"/mcp?org_id="+m.b.orgID+"&tenant_id="+m.b.orgID+"&user_id="+m.b.admin.UserID)
if got := extractInt(t, body, "totalEvents"); got != baselineTotal {
t.Errorf("totalEvents = %d, want %d — the query string moved the tenant", got, baselineTotal)
}
assertNoLeak(t, "A(query)", body, m.b)
})
}
/* ── Helpers ────────────────────────────────────────────────────────────── */
// extractInt pulls a named integer out of a tools/call response.
func extractInt(t *testing.T, body, field string) int {
t.Helper()
var envelope struct {
Result struct {
Content []struct {
Text string `json:"text"`
} `json:"content"`
IsError bool `json:"isError"`
} `json:"result"`
}
if err := json.Unmarshal([]byte(body), &envelope); err != nil {
t.Fatalf("response: %v\n%s", err, truncate(body))
}
if len(envelope.Result.Content) == 0 {
t.Fatalf("no content: %s", truncate(body))
}
var payload map[string]any
if err := json.Unmarshal([]byte(envelope.Result.Content[0].Text), &payload); err != nil {
t.Fatalf("payload: %v", err)
}
data, ok := payload["data"].(map[string]any)
if !ok {
t.Fatalf("no data object: %s", truncate(body))
}
value, ok := data[field].(float64)
if !ok {
t.Fatalf("no %s in %v", field, data)
}
return int(value)
}