mcp connection
This commit is contained in:
193
go-api/internal/mcpserver/auth.go
Normal file
193
go-api/internal/mcpserver/auth.go
Normal file
@@ -0,0 +1,193 @@
|
||||
package mcpserver
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"net/http"
|
||||
"strings"
|
||||
|
||||
"github.com/krow/krow-backend/go-api/internal/auth"
|
||||
"github.com/krow/krow-backend/go-api/internal/authctx"
|
||||
)
|
||||
|
||||
// Bearer authentication for the MCP surface.
|
||||
//
|
||||
// This file is a SEAM, not an authentication system. It defines the one
|
||||
// question the MCP transport needs answered — "which KROW user does this token
|
||||
// belong to" — and leaves answering it to whatever is plugged in. Phase 3 plugs
|
||||
// in OAuth 2.1 token validation. Nothing here mints, stores, refreshes or
|
||||
// validates a token's contents, because doing any of that now would be
|
||||
// inventing a token format that OAuth then has to replace.
|
||||
//
|
||||
// WHY MCP AUTHENTICATES SEPARATELY FROM THE REST OF THE API
|
||||
//
|
||||
// The cookie middleware in httpserver/auth.go is deliberately not reused, and
|
||||
// this is the most important decision in this file. Mounting MCP behind that
|
||||
// middleware would mean a browser session could authenticate an MCP call: the
|
||||
// middleware puts an Identity in the context, and any handler downstream that
|
||||
// reads the ambient identity would accept it. That is a real vulnerability
|
||||
// rather than a theoretical one — a logged-in user's cookie is sent by the
|
||||
// browser on requests the user did not intend, which is what SameSite exists to
|
||||
// limit and what an MCP endpoint has no business relying on.
|
||||
//
|
||||
// So identity here is PASSED, never ambient. The transport authenticates, and
|
||||
// hands the result to Handle as a parameter. There is no code path in this
|
||||
// package that reads authctx.From on an inbound request, which makes "a cookie
|
||||
// silently authenticated MCP" structurally impossible rather than merely
|
||||
// unintended. See TestCookieCannotAuthenticateMCP.
|
||||
|
||||
/* ── Errors ─────────────────────────────────────────────────────────────── */
|
||||
|
||||
var (
|
||||
// ErrNoAuthenticator is returned when the surface is running without a
|
||||
// token authenticator. It is a configuration fault, and it fails CLOSED:
|
||||
// a deployment that forgot to wire one refuses every call rather than
|
||||
// serving them unauthenticated.
|
||||
ErrNoAuthenticator = errors.New("mcpserver: no token authenticator configured")
|
||||
|
||||
// ErrMissingToken covers an absent or empty Authorization header.
|
||||
ErrMissingToken = errors.New("mcpserver: no bearer token")
|
||||
|
||||
// ErrMalformedToken covers a header this server could not parse as a
|
||||
// bearer credential — a missing scheme, a wrong scheme, an empty value.
|
||||
ErrMalformedToken = errors.New("mcpserver: malformed Authorization header")
|
||||
|
||||
// ErrInvalidToken covers a well-formed token that does not resolve to a
|
||||
// user: unknown, expired, revoked, or issued for something else.
|
||||
//
|
||||
// ONE error for all of those, deliberately. Telling a caller that a token
|
||||
// is "expired" rather than "unknown" confirms it once existed, which is an
|
||||
// oracle over the token space. Same reasoning as tools.Denied().
|
||||
ErrInvalidToken = errors.New("mcpserver: invalid bearer token")
|
||||
)
|
||||
|
||||
/* ── The seam ───────────────────────────────────────────────────────────── */
|
||||
|
||||
// TokenAuthenticator resolves a raw bearer token into a KROW identity.
|
||||
//
|
||||
// Deliberately one method taking a string and returning the SAME
|
||||
// authctx.Identity the cookie path produces. Two things follow from that shape,
|
||||
// and both are the point:
|
||||
//
|
||||
// - There is no second identity model. Everything downstream — the policy
|
||||
// table, the org pre-filter, tools.Context — consumes authctx.Identity and
|
||||
// cannot tell which path produced it, so authorization cannot drift between
|
||||
// the two.
|
||||
// - An implementation cannot report anything except an identity or a failure.
|
||||
// It has no way to return "authenticated, but also here is an org" or any
|
||||
// other channel a caller might trust. The org is inside the identity,
|
||||
// which comes from the user row.
|
||||
//
|
||||
// Phase 3's OAuth implementation of this interface will: hash the presented
|
||||
// token, look it up, check expiry, revocation and audience, load the user, and
|
||||
// build the identity from the USER ROW — never from the token's contents. A
|
||||
// token that carried its own org claim would be a token whose bearer chose
|
||||
// their own tenant.
|
||||
type TokenAuthenticator interface {
|
||||
// Authenticate resolves a raw token, or returns an error.
|
||||
//
|
||||
// Implementations must fail closed and must not distinguish unknown from
|
||||
// expired from revoked in the returned error.
|
||||
Authenticate(ctx context.Context, rawToken string) (authctx.Identity, error)
|
||||
}
|
||||
|
||||
// UserLookup is the subset of the existing user store this package needs.
|
||||
//
|
||||
// Narrowed to one method so an implementation of TokenAuthenticator can re-read
|
||||
// the user on every call — which is what makes suspension take effect on
|
||||
// contact rather than whenever a token happens to lapse. httpserver/auth.go
|
||||
// does exactly this for cookies (see its comment on re-reading the user row),
|
||||
// and the bearer path must not be weaker.
|
||||
//
|
||||
// auth.UserStore already satisfies this.
|
||||
type UserLookup interface {
|
||||
FindByID(ctx context.Context, id string) (auth.User, error)
|
||||
}
|
||||
|
||||
/* ── Header parsing ─────────────────────────────────────────────────────── */
|
||||
|
||||
// bearerToken extracts the credential from an Authorization header.
|
||||
//
|
||||
// Only the Authorization header is consulted. Not a query parameter — the MCP
|
||||
// spec forbids tokens in the URI, and a URI is logged, cached, and put in a
|
||||
// Referer. Not a custom header, not a cookie, not the body. One place, so there
|
||||
// is one thing to reason about.
|
||||
func bearerToken(r *http.Request) (string, error) {
|
||||
header := r.Header.Get("Authorization")
|
||||
if strings.TrimSpace(header) == "" {
|
||||
return "", ErrMissingToken
|
||||
}
|
||||
|
||||
scheme, value, found := strings.Cut(header, " ")
|
||||
if !found {
|
||||
return "", ErrMalformedToken
|
||||
}
|
||||
// Case-insensitive per RFC 7235: "Bearer", "bearer" and "BEARER" are the
|
||||
// same scheme, and rejecting the variants would fail against clients that
|
||||
// are behaving correctly.
|
||||
if !strings.EqualFold(strings.TrimSpace(scheme), "bearer") {
|
||||
return "", ErrMalformedToken
|
||||
}
|
||||
|
||||
token := strings.TrimSpace(value)
|
||||
if token == "" {
|
||||
return "", ErrMalformedToken
|
||||
}
|
||||
// A second space means a second value — "Bearer a b" is not a token, and
|
||||
// accepting the first half would silently authenticate something the
|
||||
// client did not send.
|
||||
if strings.ContainsAny(token, " \t") {
|
||||
return "", ErrMalformedToken
|
||||
}
|
||||
return token, nil
|
||||
}
|
||||
|
||||
// authenticate resolves the request's bearer credential into an identity.
|
||||
//
|
||||
// Every failure returns the same outward answer — 401 with no detail about
|
||||
// which stage failed. The reason is recorded in the log, where the operator is.
|
||||
func (s *Server) authenticate(r *http.Request) (authctx.Identity, error) {
|
||||
token, err := bearerToken(r)
|
||||
if err != nil {
|
||||
return authctx.Identity{}, err
|
||||
}
|
||||
if s.tokens == nil {
|
||||
return authctx.Identity{}, ErrNoAuthenticator
|
||||
}
|
||||
|
||||
identity, err := s.tokens.Authenticate(r.Context(), token)
|
||||
if err != nil {
|
||||
return authctx.Identity{}, ErrInvalidToken
|
||||
}
|
||||
|
||||
// Defence in depth against an authenticator that returns a partially
|
||||
// populated identity. Everything downstream assumes these two are present:
|
||||
// tools/scope.go refuses an empty OrgID, but it should never be asked to,
|
||||
// and a missing UserID would produce a query scoped to nobody.
|
||||
if identity.UserID == "" || identity.OrgID == "" {
|
||||
return authctx.Identity{}, ErrInvalidToken
|
||||
}
|
||||
// A suspended account must not hold a working token. The authenticator is
|
||||
// expected to check this; repeating it here costs nothing and means a
|
||||
// mistake in one implementation is not a live account bypass.
|
||||
if identity.Status != "" && identity.Status != auth.StatusActive {
|
||||
return authctx.Identity{}, ErrInvalidToken
|
||||
}
|
||||
return identity, nil
|
||||
}
|
||||
|
||||
// authFailureReason names the stage that refused, for the log only.
|
||||
func authFailureReason(err error) string {
|
||||
switch {
|
||||
case errors.Is(err, ErrMissingToken):
|
||||
return "missing_token"
|
||||
case errors.Is(err, ErrMalformedToken):
|
||||
return "malformed_header"
|
||||
case errors.Is(err, ErrNoAuthenticator):
|
||||
return "no_authenticator_configured"
|
||||
case errors.Is(err, ErrInvalidToken):
|
||||
return "invalid_token"
|
||||
default:
|
||||
return "error"
|
||||
}
|
||||
}
|
||||
512
go-api/internal/mcpserver/auth_test.go
Normal file
512
go-api/internal/mcpserver/auth_test.go
Normal file
@@ -0,0 +1,512 @@
|
||||
package mcpserver
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"io"
|
||||
"log/slog"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"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"
|
||||
)
|
||||
|
||||
/* ── A test-only authenticator ──────────────────────────────────────────────
|
||||
|
||||
TEST-ONLY. This type lives in a _test.go file and is therefore not compiled
|
||||
into the binary at all — there is no build tag to forget and no flag that
|
||||
could enable it in production. That is deliberate: the one thing worse than
|
||||
having no authentication is having a development authenticator that ships.
|
||||
|
||||
It is a map from token to identity and nothing else. It does not hash, does
|
||||
not expire, does not check an audience and does not consult a database,
|
||||
because none of those are what these tests are testing. What they test is
|
||||
the SEAM: that a token resolves to an identity, that the identity reaches
|
||||
the registry, and that existing authorization then decides the answer.
|
||||
|
||||
Phase 3 replaces this with an OAuth implementation of the same interface.
|
||||
Every test below keeps working, because what they assert is the behaviour of
|
||||
the seam rather than the behaviour of any particular token format. That is
|
||||
the reason for defining the interface before implementing OAuth rather than
|
||||
after. */
|
||||
|
||||
type fakeTokens struct {
|
||||
byToken map[string]authctx.Identity
|
||||
err error
|
||||
}
|
||||
|
||||
func (f *fakeTokens) Authenticate(_ context.Context, raw string) (authctx.Identity, error) {
|
||||
if f.err != nil {
|
||||
return authctx.Identity{}, f.err
|
||||
}
|
||||
id, ok := f.byToken[raw]
|
||||
if !ok {
|
||||
return authctx.Identity{}, ErrInvalidToken
|
||||
}
|
||||
return id, nil
|
||||
}
|
||||
|
||||
const validToken = "test-token-valid"
|
||||
|
||||
func identityFor(orgID, role string) authctx.Identity {
|
||||
return authctx.Identity{
|
||||
UserID: "user-" + role,
|
||||
OrgID: orgID,
|
||||
Email: role + "@example.test",
|
||||
FullName: "Test " + role,
|
||||
Role: role,
|
||||
AccountType: "employer",
|
||||
Status: "active",
|
||||
}
|
||||
}
|
||||
|
||||
// authedServer wires the real registry to a token map.
|
||||
func authedServer(t *testing.T, reg *tools.Registry, tokens map[string]authctx.Identity) *Server {
|
||||
t.Helper()
|
||||
if reg == nil {
|
||||
reg = runtime.DefaultTools(nil, nil)
|
||||
}
|
||||
return New(reg, &fakeTokens{byToken: tokens}, slog.New(slog.NewTextHandler(io.Discard, nil)))
|
||||
}
|
||||
|
||||
// postWith sends a body with an explicit Authorization header value. An empty
|
||||
// header value means the header is not sent at all.
|
||||
func postWith(t *testing.T, s *Server, authHeader, body string) *httptest.ResponseRecorder {
|
||||
t.Helper()
|
||||
req := httptest.NewRequest(http.MethodPost, "/mcp", strings.NewReader(body))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
if authHeader != "" {
|
||||
req.Header.Set("Authorization", authHeader)
|
||||
}
|
||||
rec := httptest.NewRecorder()
|
||||
s.Handler().ServeHTTP(rec, req)
|
||||
return rec
|
||||
}
|
||||
|
||||
const listBody = `{"jsonrpc":"2.0","id":1,"method":"tools/list"}`
|
||||
|
||||
/* ── 1–5. Header and token rejection ────────────────────────────────────── */
|
||||
|
||||
func TestBearerRejection(t *testing.T) {
|
||||
s := authedServer(t, nil, map[string]authctx.Identity{
|
||||
validToken: identityFor("org-a", "admin"),
|
||||
})
|
||||
|
||||
for name, header := range map[string]string{
|
||||
"missing header": "",
|
||||
"no scheme": "abc123",
|
||||
"wrong scheme basic": "Basic dXNlcjpwYXNz",
|
||||
"wrong scheme token": "Token abc123",
|
||||
"empty bearer": "Bearer ",
|
||||
"bearer with only ws": "Bearer ",
|
||||
"two values": "Bearer abc def",
|
||||
"unknown token": "Bearer not-a-real-token",
|
||||
"token with whitespace": "Bearer abc\tdef",
|
||||
} {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
rec := postWith(t, s, header, listBody)
|
||||
|
||||
if rec.Code != http.StatusUnauthorized {
|
||||
t.Fatalf("status = %d, want 401 (body: %s)", rec.Code, rec.Body.String())
|
||||
}
|
||||
// RFC 9728: the client reads the authorization server location from
|
||||
// this header. Without it a compliant client cannot begin discovery.
|
||||
if got := rec.Header().Get("WWW-Authenticate"); !strings.HasPrefix(got, "Bearer") {
|
||||
t.Errorf("WWW-Authenticate = %q, want a Bearer challenge", got)
|
||||
}
|
||||
// The refusal must not say WHICH stage failed. "expired" versus
|
||||
// "unknown" is an oracle over the token space.
|
||||
body := rec.Body.String()
|
||||
for _, leak := range []string{"expired", "revoked", "unknown", "malformed", "not found"} {
|
||||
if strings.Contains(strings.ToLower(body), leak) {
|
||||
t.Errorf("the 401 body distinguishes failure modes (%q): %s", leak, body)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// Case-insensitivity is required by RFC 7235; rejecting "bearer" would fail
|
||||
// against clients that are behaving correctly.
|
||||
func TestBearerSchemeIsCaseInsensitive(t *testing.T) {
|
||||
s := authedServer(t, nil, map[string]authctx.Identity{
|
||||
validToken: identityFor("org-a", "admin"),
|
||||
})
|
||||
for _, scheme := range []string{"Bearer", "bearer", "BEARER", "BeArEr"} {
|
||||
rec := postWith(t, s, scheme+" "+validToken, listBody)
|
||||
if rec.Code != http.StatusOK {
|
||||
t.Errorf("scheme %q: status = %d, want 200", scheme, rec.Code)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// A server with no authenticator wired must refuse everything rather than
|
||||
// serve it unauthenticated. Misconfiguration fails closed.
|
||||
func TestNoAuthenticatorFailsClosed(t *testing.T) {
|
||||
s := New(runtime.DefaultTools(nil, nil), nil, slog.New(slog.NewTextHandler(io.Discard, nil)))
|
||||
rec := postWith(t, s, "Bearer "+validToken, listBody)
|
||||
if rec.Code != http.StatusUnauthorized {
|
||||
t.Fatalf("status = %d, want 401 when no authenticator is configured", rec.Code)
|
||||
}
|
||||
}
|
||||
|
||||
/* ── 6 & 7. Identity and org resolution ─────────────────────────────────── */
|
||||
|
||||
func TestValidBearerReachesToolsList(t *testing.T) {
|
||||
s := authedServer(t, nil, map[string]authctx.Identity{
|
||||
validToken: identityFor("org-a", "admin"),
|
||||
})
|
||||
rec := postWith(t, s, "Bearer "+validToken, listBody)
|
||||
|
||||
var out toolsListResult
|
||||
resultInto(t, rec, &out)
|
||||
if len(out.Tools) != len(wantExposed) {
|
||||
t.Fatalf("exposed %d tools, want %d", len(out.Tools), len(wantExposed))
|
||||
}
|
||||
}
|
||||
|
||||
// An identity missing a tenant must be refused before it reaches a tool.
|
||||
// tools/scope.go would refuse it too, but it should never be asked to: an
|
||||
// authenticator that returns a half-built identity is a bug, not a caller.
|
||||
func TestIncompleteIdentityIsRefused(t *testing.T) {
|
||||
for name, id := range map[string]authctx.Identity{
|
||||
"no org": {UserID: "u1", Role: "admin", Status: "active"},
|
||||
"no user": {OrgID: "org-a", Role: "admin", Status: "active"},
|
||||
"suspended": {UserID: "u1", OrgID: "org-a", Role: "admin", Status: "suspended"},
|
||||
"empty entire": {},
|
||||
} {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
s := authedServer(t, nil, map[string]authctx.Identity{validToken: id})
|
||||
rec := postWith(t, s, "Bearer "+validToken, listBody)
|
||||
if rec.Code != http.StatusUnauthorized {
|
||||
t.Errorf("status = %d, want 401", rec.Code)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
/* ── 8, 9, 10. Identity cannot be overridden ────────────────────────────── */
|
||||
|
||||
// The identity must come from the token and from nothing else. This walks the
|
||||
// channels a client controls and asserts that none of them moves the tenant.
|
||||
func TestIdentityCannotBeOverridden(t *testing.T) {
|
||||
h := testutil.New(t)
|
||||
orgA := freshOrgFor(t, h, "mcp-override-a")
|
||||
orgB := freshOrgFor(t, h, "mcp-override-b")
|
||||
seedActivityRows(t, h, orgA, 5, "a@example.test")
|
||||
seedActivityRows(t, h, orgB, 40, "b@example.test")
|
||||
|
||||
reg := runtime.DefaultTools(h.Pool, nil)
|
||||
s := authedServer(t, reg, map[string]authctx.Identity{
|
||||
validToken: {UserID: "u-a", OrgID: orgA, Email: "a@example.test",
|
||||
Role: "admin", AccountType: "employer", Status: "active"},
|
||||
})
|
||||
|
||||
// Every one of these is a client-controlled channel. None may change which
|
||||
// tenant is read. orgB has 40 rows and orgA has 5, so a successful override
|
||||
// is visible as a total of 40 or 45.
|
||||
attempts := map[string]string{
|
||||
"tool argument": `{"jsonrpc":"2.0","id":1,"method":"tools/call","params":
|
||||
{"name":"activity_breakdown","arguments":{"org_id":"` + orgB + `"}}}`,
|
||||
"jsonrpc meta": `{"jsonrpc":"2.0","id":1,"method":"tools/call","_meta":{"org_id":"` + orgB + `"},
|
||||
"params":{"name":"activity_breakdown","arguments":{}}}`,
|
||||
"params meta": `{"jsonrpc":"2.0","id":1,"method":"tools/call","params":
|
||||
{"name":"activity_breakdown","arguments":{},"_meta":{"org_id":"` + orgB + `"}}}`,
|
||||
"params orgId": `{"jsonrpc":"2.0","id":1,"method":"tools/call","params":
|
||||
{"name":"activity_breakdown","orgId":"` + orgB + `","arguments":{}}}`,
|
||||
"params userId": `{"jsonrpc":"2.0","id":1,"method":"tools/call","params":
|
||||
{"name":"activity_breakdown","userId":"u-b","arguments":{}}}`,
|
||||
}
|
||||
|
||||
for name, body := range attempts {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
rec := postWith(t, s, "Bearer "+validToken, body)
|
||||
total := totalFromActivityBreakdown(t, rec)
|
||||
if total != 5 {
|
||||
t.Errorf("total = %d, want 5 (org A only) — %q moved the tenant", total, name)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// Custom headers naming another tenant must be ignored outright.
|
||||
func TestCustomIdentityHeadersAreIgnored(t *testing.T) {
|
||||
h := testutil.New(t)
|
||||
orgA := freshOrgFor(t, h, "mcp-hdr-a")
|
||||
orgB := freshOrgFor(t, h, "mcp-hdr-b")
|
||||
seedActivityRows(t, h, orgA, 5, "a@example.test")
|
||||
seedActivityRows(t, h, orgB, 40, "b@example.test")
|
||||
|
||||
reg := runtime.DefaultTools(h.Pool, nil)
|
||||
s := authedServer(t, reg, map[string]authctx.Identity{
|
||||
validToken: {UserID: "u-a", OrgID: orgA, Email: "a@example.test",
|
||||
Role: "admin", AccountType: "employer", Status: "active"},
|
||||
})
|
||||
|
||||
body := `{"jsonrpc":"2.0","id":1,"method":"tools/call","params":
|
||||
{"name":"activity_breakdown","arguments":{}}}`
|
||||
req := httptest.NewRequest(http.MethodPost, "/mcp", strings.NewReader(body))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
req.Header.Set("Authorization", "Bearer "+validToken)
|
||||
for _, h := range []string{"X-Org-Id", "X-Organization-Id", "X-User-Id", "X-Tenant-Id", "X-Krow-Org"} {
|
||||
req.Header.Set(h, orgB)
|
||||
}
|
||||
rec := httptest.NewRecorder()
|
||||
s.Handler().ServeHTTP(rec, req)
|
||||
|
||||
if total := totalFromActivityBreakdown(t, rec); total != 5 {
|
||||
t.Errorf("total = %d, want 5 — a custom header moved the tenant", total)
|
||||
}
|
||||
}
|
||||
|
||||
// A cookie must never authenticate MCP.
|
||||
//
|
||||
// This is the test for the decision in auth.go: identity is passed, never
|
||||
// ambient. Even with a valid KROW identity sitting in the request context —
|
||||
// which is exactly what the cookie middleware would put there if this endpoint
|
||||
// were mounted behind it — the call must be refused, because no bearer token
|
||||
// was presented.
|
||||
func TestCookieCannotAuthenticateMCP(t *testing.T) {
|
||||
s := authedServer(t, nil, map[string]authctx.Identity{
|
||||
validToken: identityFor("org-a", "admin"),
|
||||
})
|
||||
|
||||
req := httptest.NewRequest(http.MethodPost, "/mcp", strings.NewReader(listBody))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
// A session cookie, as a browser would send it.
|
||||
req.AddCookie(&http.Cookie{Name: "krow_session", Value: "a-perfectly-valid-session-token"})
|
||||
// AND a fully populated identity in the context, as the cookie middleware
|
||||
// would have placed there. This is the strongest form of the test: even if
|
||||
// somebody mounts MCP behind authenticate(), it must still refuse.
|
||||
ctx := authctx.With(req.Context(), identityFor("org-a", "admin"))
|
||||
req = req.WithContext(ctx)
|
||||
|
||||
rec := httptest.NewRecorder()
|
||||
s.Handler().ServeHTTP(rec, req)
|
||||
|
||||
if rec.Code != http.StatusUnauthorized {
|
||||
t.Fatalf("status = %d, want 401 — a cookie session authenticated an MCP call", rec.Code)
|
||||
}
|
||||
}
|
||||
|
||||
/* ── 11, 12, 13. Authorization still runs ───────────────────────────────── */
|
||||
|
||||
// The point of the whole design: MCP authenticates, and the EXISTING policy
|
||||
// table authorizes. A talent caller must be refused a tool that operators own.
|
||||
func TestExistingAuthorizationStillRuns(t *testing.T) {
|
||||
h := testutil.New(t)
|
||||
org := freshOrgFor(t, h, "mcp-authz")
|
||||
seedActivityRows(t, h, org, 5, "boss@example.test")
|
||||
|
||||
reg := runtime.DefaultTools(h.Pool, nil)
|
||||
s := authedServer(t, reg, map[string]authctx.Identity{
|
||||
"admin-token": {UserID: "u-admin", OrgID: org, Email: "boss@example.test",
|
||||
Role: "admin", AccountType: "employer", Status: "active"},
|
||||
"talent-token": {UserID: "u-talent", OrgID: org, Email: "worker@example.test",
|
||||
Role: "talent", AccountType: "talent", Status: "active"},
|
||||
})
|
||||
|
||||
// `staff` is operators-only in the policy table, and hires_recent reads
|
||||
// job-applications which talent may list only in its own scope.
|
||||
call := `{"jsonrpc":"2.0","id":1,"method":"tools/call","params":
|
||||
{"name":"activity_breakdown","arguments":{}}}`
|
||||
|
||||
adminRec := postWith(t, s, "Bearer admin-token", call)
|
||||
if got := totalFromActivityBreakdown(t, adminRec); got != 5 {
|
||||
t.Errorf("admin total = %d, want 5", got)
|
||||
}
|
||||
|
||||
// A talent caller is authenticated but scoped. Whatever comes back, it must
|
||||
// come back through the policy table rather than around it — the assertion
|
||||
// is that the two roles do NOT get the same answer.
|
||||
talentRec := postWith(t, s, "Bearer talent-token", call)
|
||||
talentTotal := totalFromActivityBreakdownAllowingDenial(t, talentRec)
|
||||
if talentTotal == 5 {
|
||||
t.Error("a talent caller saw the admin's total; authorization did not run")
|
||||
}
|
||||
}
|
||||
|
||||
/* ── 14 & 15. Exclusions hold under authentication ──────────────────────── */
|
||||
|
||||
// Authentication must not become a way to reach a withheld tool. A valid token
|
||||
// is not a key to the write tools.
|
||||
func TestExclusionsHoldForAuthenticatedCallers(t *testing.T) {
|
||||
s := authedServer(t, nil, map[string]authctx.Identity{
|
||||
validToken: identityFor("org-a", "admin"),
|
||||
})
|
||||
|
||||
for _, name := range []string{"assign_worker", "move_application", "knowledge_search"} {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
body := `{"jsonrpc":"2.0","id":1,"method":"tools/call","params":{"name":"` + name + `","arguments":{}}}`
|
||||
rec := postWith(t, s, "Bearer "+validToken, body)
|
||||
|
||||
var out toolsCallResult
|
||||
resultInto(t, rec, &out)
|
||||
if !out.IsError {
|
||||
t.Fatalf("%s was reachable by an authenticated caller", name)
|
||||
}
|
||||
if !strings.Contains(out.Content[0].Text, "mcp.unknown_tool") {
|
||||
t.Errorf("expected mcp.unknown_tool, got: %s", out.Content[0].Text)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// tools/list must not vary by caller in a way that reveals the withheld set.
|
||||
func TestToolsListIsTheSameForEveryRole(t *testing.T) {
|
||||
s := authedServer(t, nil, map[string]authctx.Identity{
|
||||
"admin-token": identityFor("org-a", "admin"),
|
||||
"talent-token": identityFor("org-a", "talent"),
|
||||
})
|
||||
|
||||
names := func(token string) []string {
|
||||
rec := postWith(t, s, "Bearer "+token, listBody)
|
||||
var out toolsListResult
|
||||
resultInto(t, rec, &out)
|
||||
got := make([]string, 0, len(out.Tools))
|
||||
for _, tool := range out.Tools {
|
||||
got = append(got, tool.Name)
|
||||
}
|
||||
return got
|
||||
}
|
||||
|
||||
admin, talent := names("admin-token"), names("talent-token")
|
||||
if strings.Join(admin, ",") != strings.Join(talent, ",") {
|
||||
t.Errorf("tools/list differs by role:\n admin: %v\ntalent: %v", admin, talent)
|
||||
}
|
||||
}
|
||||
|
||||
/* ── The handshake requires a token too ─────────────────────────────────── */
|
||||
|
||||
// REVERSED IN PHASE 4, deliberately.
|
||||
//
|
||||
// This test previously asserted the opposite: that initialize and ping were
|
||||
// reachable without a token, on the reasoning that a client needs somewhere to
|
||||
// start. That was wrong, and the integration test over the mounted route is
|
||||
// what caught it — a client's FIRST request is usually initialize, and
|
||||
// answering it 200 means the client never sees the WWW-Authenticate challenge
|
||||
// that begins the OAuth flow. It believes it is connected and finds out
|
||||
// otherwise at the first real call, with no 401 in hand to discover from.
|
||||
//
|
||||
// Requiring a token everywhere means the first request, whatever it is,
|
||||
// produces the challenge. Nothing is lost: the handshake returns only this
|
||||
// server's name and capabilities, which are of use solely to a client that
|
||||
// means to authenticate.
|
||||
func TestHandshakeMethodsAlsoRequireAuth(t *testing.T) {
|
||||
s := authedServer(t, nil, map[string]authctx.Identity{
|
||||
validToken: identityFor("org-a", "admin"),
|
||||
})
|
||||
|
||||
for name, body := range map[string]string{
|
||||
"initialize": `{"jsonrpc":"2.0","id":1,"method":"initialize","params":{}}`,
|
||||
"ping": `{"jsonrpc":"2.0","id":1,"method":"ping"}`,
|
||||
"tools/list": `{"jsonrpc":"2.0","id":1,"method":"tools/list"}`,
|
||||
"unknown verb": `{"jsonrpc":"2.0","id":1,"method":"resources/list"}`,
|
||||
} {
|
||||
t.Run(name+" without a token", func(t *testing.T) {
|
||||
rec := postWith(t, s, "", body)
|
||||
if rec.Code != http.StatusUnauthorized {
|
||||
t.Errorf("status = %d, want 401 — every method must produce the "+
|
||||
"challenge that starts the OAuth flow", rec.Code)
|
||||
}
|
||||
if !strings.HasPrefix(rec.Header().Get("WWW-Authenticate"), "Bearer") {
|
||||
t.Error("the 401 carries no Bearer challenge")
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
// And with a token, the handshake works — or the test above would pass by
|
||||
// the endpoint being broken rather than by it being guarded.
|
||||
for name, body := range map[string]string{
|
||||
"initialize": `{"jsonrpc":"2.0","id":1,"method":"initialize","params":{}}`,
|
||||
"ping": `{"jsonrpc":"2.0","id":1,"method":"ping"}`,
|
||||
} {
|
||||
t.Run(name+" with a token", func(t *testing.T) {
|
||||
if rec := postWith(t, s, "Bearer "+validToken, body); rec.Code != http.StatusOK {
|
||||
t.Errorf("status = %d, want 200", rec.Code)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// The resource_metadata pointer must be built from configuration, so a 401
|
||||
// tells a client where to look without this package knowing any hostname.
|
||||
func TestChallengeCarriesTheConfiguredResourceMetadata(t *testing.T) {
|
||||
const metadataURL = "https://configured.example.test/.well-known/oauth-protected-resource"
|
||||
s := authedServer(t, nil, map[string]authctx.Identity{}).
|
||||
WithResourceMetadataURL(metadataURL)
|
||||
|
||||
rec := postWith(t, s, "", listBody)
|
||||
challenge := rec.Header().Get("WWW-Authenticate")
|
||||
if !strings.Contains(challenge, `resource_metadata="`+metadataURL+`"`) {
|
||||
t.Errorf("WWW-Authenticate = %q, want it to carry %q", challenge, metadataURL)
|
||||
}
|
||||
}
|
||||
|
||||
/* ── Helpers ────────────────────────────────────────────────────────────── */
|
||||
|
||||
func freshOrgFor(t *testing.T, h *testutil.Harness, slug string) string {
|
||||
t.Helper()
|
||||
var id string
|
||||
if err := h.Pool.QueryRow(context.Background(),
|
||||
`INSERT INTO organizations (name, slug) VALUES ($1, $2) RETURNING id::text`,
|
||||
slug, slug).Scan(&id); err != nil {
|
||||
t.Fatalf("create org %s: %v", slug, err)
|
||||
}
|
||||
return id
|
||||
}
|
||||
|
||||
func seedActivityRows(t *testing.T, h *testutil.Harness, org string, n int, email string) {
|
||||
t.Helper()
|
||||
for i := 0; i < n; i++ {
|
||||
if _, err := h.Pool.Exec(context.Background(),
|
||||
`INSERT INTO user_activity (org_id, event_type, user_email, user_name)
|
||||
VALUES ($1::uuid, 'login', $2, 'Someone')`, org, email); err != nil {
|
||||
t.Fatalf("seed activity: %v", err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// totalFromActivityBreakdown reads the `total` out of a successful tool result.
|
||||
func totalFromActivityBreakdown(t *testing.T, rec *httptest.ResponseRecorder) int {
|
||||
t.Helper()
|
||||
var out toolsCallResult
|
||||
resultInto(t, rec, &out)
|
||||
if out.IsError {
|
||||
t.Fatalf("tool call failed: %s", out.Content[0].Text)
|
||||
}
|
||||
return parseTotal(t, out.Content[0].Text)
|
||||
}
|
||||
|
||||
// totalFromActivityBreakdownAllowingDenial returns -1 when the tool refused,
|
||||
// which is a legitimate authorization outcome rather than a test failure.
|
||||
func totalFromActivityBreakdownAllowingDenial(t *testing.T, rec *httptest.ResponseRecorder) int {
|
||||
t.Helper()
|
||||
var out toolsCallResult
|
||||
resultInto(t, rec, &out)
|
||||
if out.IsError {
|
||||
return -1
|
||||
}
|
||||
return parseTotal(t, out.Content[0].Text)
|
||||
}
|
||||
|
||||
func parseTotal(t *testing.T, raw string) int {
|
||||
t.Helper()
|
||||
// activity_breakdown's own field name, taken from the handler's output
|
||||
// rather than guessed: a wrong name here reads as zero, which would make a
|
||||
// cross-tenant leak look like a pass.
|
||||
var payload struct {
|
||||
Data struct {
|
||||
TotalEvents int `json:"totalEvents"`
|
||||
} `json:"data"`
|
||||
}
|
||||
if err := json.Unmarshal([]byte(raw), &payload); err != nil {
|
||||
t.Fatalf("tool result did not decode: %v\nraw: %s", err, raw)
|
||||
}
|
||||
return payload.Data.TotalEvents
|
||||
}
|
||||
243
go-api/internal/mcpserver/jsonrpc.go
Normal file
243
go-api/internal/mcpserver/jsonrpc.go
Normal file
@@ -0,0 +1,243 @@
|
||||
// Package mcpserver is the Model Context Protocol surface: a second way into
|
||||
// the tool layer, for clients that speak MCP rather than HTTP+cookie.
|
||||
//
|
||||
// It is an ADDITIONAL interface and nothing else. It owns no business logic, no
|
||||
// SQL and no authorization rules. Every call it serves ends up in
|
||||
// tools.Registry.Dispatch — the same entry point the agent loop uses — so a
|
||||
// question asked through MCP is answered by the same handler, under the same
|
||||
// policy table, behind the same org pre-filter as the same question asked by
|
||||
// Owliver. That is the whole design, and the reason this package is small.
|
||||
//
|
||||
// What lives here:
|
||||
//
|
||||
// - JSON-RPC 2.0 framing (this file)
|
||||
// - the three methods MCP needs to be useful: initialize, tools/list,
|
||||
// tools/call (server.go)
|
||||
// - which tools are published, derived from the registry (tools.go)
|
||||
// - bearer authentication, as a seam an OAuth implementation plugs into
|
||||
// (auth.go)
|
||||
// - the Streamable HTTP binding (transport.go)
|
||||
//
|
||||
// Identity is established by this package's own bearer authentication and is
|
||||
// PASSED to the handlers, never read from the ambient request context. That is
|
||||
// what stops a browser cookie from authenticating an MCP call — see auth.go.
|
||||
package mcpserver
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"time"
|
||||
)
|
||||
|
||||
// jsonRPCVersion is the only version this server speaks. A request naming
|
||||
// anything else is malformed rather than merely unsupported: "2.0" is a
|
||||
// constant in the spec, not a negotiation.
|
||||
const jsonRPCVersion = "2.0"
|
||||
|
||||
/* ── Wire types ─────────────────────────────────────────────────────────── */
|
||||
|
||||
// request is one inbound JSON-RPC message.
|
||||
//
|
||||
// ID is json.RawMessage rather than any, because the spec allows a string, a
|
||||
// number or null, and the response MUST echo it back byte-for-byte. Decoding it
|
||||
// into an `any` turns 1 into 1.0 on the way back out, which is a different id to
|
||||
// a client matching responses to requests.
|
||||
type request struct {
|
||||
JSONRPC string `json:"jsonrpc"`
|
||||
ID json.RawMessage `json:"id,omitempty"`
|
||||
Method string `json:"method"`
|
||||
Params json.RawMessage `json:"params,omitempty"`
|
||||
}
|
||||
|
||||
// isNotification reports that no response is expected.
|
||||
//
|
||||
// A notification is a request with no id. The spec is explicit that a server
|
||||
// must not answer one, so the transport drops the response and returns 202.
|
||||
func (r request) isNotification() bool {
|
||||
return len(r.ID) == 0 || string(r.ID) == "null"
|
||||
}
|
||||
|
||||
// response is one outbound JSON-RPC message.
|
||||
//
|
||||
// Result and Error are pointers so exactly one is ever serialised: the spec
|
||||
// forbids both together, and a non-pointer Result would emit `"result":null`
|
||||
// alongside an error.
|
||||
type response struct {
|
||||
JSONRPC string `json:"jsonrpc"`
|
||||
ID json.RawMessage `json:"id,omitempty"`
|
||||
Result any `json:"result,omitempty"`
|
||||
Error *rpcError `json:"error,omitempty"`
|
||||
}
|
||||
|
||||
// rpcError is a JSON-RPC error object.
|
||||
//
|
||||
// retryAfter is NOT serialised. It exists so a rate-limited refusal produced
|
||||
// deep in the handler can reach the transport, which is the only layer that can
|
||||
// set an HTTP status and a Retry-After header. The alternative — returning a
|
||||
// 200 with a JSON-RPC error and no header — would give a client no way to know
|
||||
// how long to wait.
|
||||
type rpcError struct {
|
||||
Code int `json:"code"`
|
||||
Message string `json:"message"`
|
||||
Data any `json:"data,omitempty"`
|
||||
|
||||
retryAfter time.Duration
|
||||
}
|
||||
|
||||
func (e *rpcError) Error() string { return fmt.Sprintf("jsonrpc %d: %s", e.Code, e.Message) }
|
||||
|
||||
/* ── Error codes ────────────────────────────────────────────────────────── */
|
||||
|
||||
// The standard JSON-RPC 2.0 codes. Reserved range is -32768..-32000; anything
|
||||
// this server invents lives outside it.
|
||||
const (
|
||||
codeParseError = -32700
|
||||
codeInvalidRequest = -32600
|
||||
codeMethodNotFound = -32601
|
||||
codeInvalidParams = -32602
|
||||
codeInternalError = -32603
|
||||
)
|
||||
|
||||
// codeUnauthorized is outside the JSON-RPC reserved range (-32768..-32000),
|
||||
// because it is this server's own condition rather than a protocol fault. It
|
||||
// accompanies an HTTP 401: the transport layer carries the authoritative
|
||||
// signal, and this gives a client reading only the JSON-RPC body the same
|
||||
// answer.
|
||||
const codeUnauthorized = -32001
|
||||
|
||||
// codeRateLimited is this server's own condition, outside the reserved range.
|
||||
// It accompanies an HTTP 429 and a Retry-After header.
|
||||
const codeRateLimited = -32002
|
||||
|
||||
// errRateLimited refuses a call that exceeded its organisation's ceiling.
|
||||
//
|
||||
// The message names no number and no organisation. How much quota a tenant has
|
||||
// and how much of it they have spent is not something one caller should learn
|
||||
// from a refusal — it is the same reasoning as the opaque tool denial.
|
||||
func errRateLimited(retryAfter time.Duration) *rpcError {
|
||||
return &rpcError{
|
||||
Code: codeRateLimited,
|
||||
Message: "too many requests for this organisation; retry after the interval in the Retry-After header",
|
||||
retryAfter: retryAfter,
|
||||
}
|
||||
}
|
||||
|
||||
func errParse(detail string) *rpcError {
|
||||
return &rpcError{Code: codeParseError, Message: "invalid JSON", Data: detail}
|
||||
}
|
||||
|
||||
func errInvalidRequest(detail string) *rpcError {
|
||||
return &rpcError{Code: codeInvalidRequest, Message: "invalid JSON-RPC request", Data: detail}
|
||||
}
|
||||
|
||||
func errMethodNotFound(method string) *rpcError {
|
||||
return &rpcError{
|
||||
Code: codeMethodNotFound,
|
||||
Message: "method not found",
|
||||
Data: fmt.Sprintf("this server implements initialize, tools/list and tools/call; it does not implement %q", method),
|
||||
}
|
||||
}
|
||||
|
||||
func errInvalidParams(detail string) *rpcError {
|
||||
return &rpcError{Code: codeInvalidParams, Message: "invalid params", Data: detail}
|
||||
}
|
||||
|
||||
// errInternal deliberately carries no detail.
|
||||
//
|
||||
// An internal failure is the one case where the thing that went wrong is this
|
||||
// server's business and not the caller's: a wrapped database error or a panic
|
||||
// message is reconnaissance. The detail goes to the log, where the operator is.
|
||||
func errInternal() *rpcError {
|
||||
return &rpcError{Code: codeInternalError, Message: "internal error"}
|
||||
}
|
||||
|
||||
/* ── Parsing ────────────────────────────────────────────────────────────── */
|
||||
|
||||
// errBatch marks a batch request, which this server does not accept.
|
||||
//
|
||||
// Rejecting it explicitly rather than failing to parse it is the point: a
|
||||
// client that batches and gets a parse error will retry the same batch, where
|
||||
// one told that batching is unsupported can fall back to sending messages
|
||||
// singly. The current MCP transport binding sends one message per POST, so
|
||||
// nothing a compliant client does requires batching.
|
||||
var errBatch = errors.New("batch requests are not supported")
|
||||
|
||||
// parseRequest decodes one JSON-RPC message and validates its envelope.
|
||||
//
|
||||
// The two are separate returns because they have different fates: a message
|
||||
// that could not be parsed has no id, so its error answers with a null id,
|
||||
// while a message that parsed but is invalid answers with the id it carried.
|
||||
func parseRequest(body []byte) (request, *rpcError) {
|
||||
trimmed := skipSpace(body)
|
||||
if len(trimmed) == 0 {
|
||||
return request{}, errInvalidRequest("the request body was empty")
|
||||
}
|
||||
if trimmed[0] == '[' {
|
||||
return request{}, errInvalidRequest(errBatch.Error())
|
||||
}
|
||||
|
||||
var req request
|
||||
if err := json.Unmarshal(trimmed, &req); err != nil {
|
||||
return request{}, errParse(err.Error())
|
||||
}
|
||||
|
||||
if req.JSONRPC != jsonRPCVersion {
|
||||
return req, errInvalidRequest(fmt.Sprintf(
|
||||
"jsonrpc must be %q, got %q", jsonRPCVersion, req.JSONRPC))
|
||||
}
|
||||
if req.Method == "" {
|
||||
return req, errInvalidRequest("method is required")
|
||||
}
|
||||
// An id, when present, must be a string or a number. Objects and arrays are
|
||||
// forbidden by the spec, and echoing one back would propagate the mistake.
|
||||
if len(req.ID) > 0 && !isValidID(req.ID) {
|
||||
return req, errInvalidRequest("id must be a string, a number or null")
|
||||
}
|
||||
return req, nil
|
||||
}
|
||||
|
||||
// isValidID reports whether a raw id is a string, a number or null.
|
||||
func isValidID(raw json.RawMessage) bool {
|
||||
t := skipSpace(raw)
|
||||
if len(t) == 0 {
|
||||
return false
|
||||
}
|
||||
switch t[0] {
|
||||
case '{', '[':
|
||||
return false
|
||||
}
|
||||
var v any
|
||||
return json.Unmarshal(t, &v) == nil
|
||||
}
|
||||
|
||||
// decodeParams unmarshals params into dst, treating absent params as an empty
|
||||
// object so a method with only optional fields can be called with none.
|
||||
func decodeParams(raw json.RawMessage, dst any) *rpcError {
|
||||
t := skipSpace(raw)
|
||||
if len(t) == 0 || string(t) == "null" {
|
||||
return nil
|
||||
}
|
||||
// Arrays are legal JSON-RPC (positional params) and are not used by MCP,
|
||||
// whose methods all take an object. Saying so beats a confusing type error.
|
||||
if t[0] == '[' {
|
||||
return errInvalidParams("params must be an object; positional params are not supported")
|
||||
}
|
||||
if err := json.Unmarshal(t, dst); err != nil {
|
||||
return errInvalidParams(err.Error())
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func skipSpace(b []byte) []byte {
|
||||
i := 0
|
||||
for i < len(b) {
|
||||
switch b[i] {
|
||||
case ' ', '\t', '\r', '\n':
|
||||
i++
|
||||
default:
|
||||
return b[i:]
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
595
go-api/internal/mcpserver/mcpserver_test.go
Normal file
595
go-api/internal/mcpserver/mcpserver_test.go
Normal file
@@ -0,0 +1,595 @@
|
||||
package mcpserver
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"io"
|
||||
"log/slog"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"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/tools"
|
||||
)
|
||||
|
||||
// The tool registry these tests run against is the REAL one — the same
|
||||
// construction the service uses — with a nil database and a nil retriever.
|
||||
//
|
||||
// That is sound because nothing here dispatches: Phase 1's tools/call refuses
|
||||
// before reaching a handler, and every other assertion is about metadata, which
|
||||
// is fixed at registration. Testing against a hand-built fixture registry would
|
||||
// be testing a copy of the thing under test: the whole claim of this package is
|
||||
// "what MCP publishes is what the registry holds", and a fixture would let that
|
||||
// claim pass while being false of the real set.
|
||||
// Both dependencies are nil: the handlers capture them in closures and nothing
|
||||
// here reaches a handler, so nothing dereferences them. If a future test does
|
||||
// dispatch, this will nil-panic loudly rather than quietly reading a database
|
||||
// it should not have.
|
||||
func testRegistry(t *testing.T) *tools.Registry {
|
||||
t.Helper()
|
||||
return runtime.DefaultTools(nil, nil)
|
||||
}
|
||||
|
||||
// newServer wires the real registry to a test-only token authenticator.
|
||||
//
|
||||
// PHASE 2 NOTE: tools/list and tools/call now require a bearer token, so this
|
||||
// fixture authenticates. Not one assertion in this file changed — only the
|
||||
// fixture gained a credential. The behaviour change is the point of Phase 2 and
|
||||
// is asserted directly in auth_test.go (TestBearerRejection), rather than being
|
||||
// papered over here.
|
||||
func newServer(t *testing.T) *Server {
|
||||
t.Helper()
|
||||
return New(testRegistry(t), &fakeTokens{byToken: map[string]authctx.Identity{
|
||||
validToken: identityFor("org-test", "admin"),
|
||||
}}, slog.New(slog.NewTextHandler(io.Discard, nil)))
|
||||
}
|
||||
|
||||
// post sends one raw body to the MCP endpoint and returns the recorder.
|
||||
func post(t *testing.T, s *Server, body string) *httptest.ResponseRecorder {
|
||||
t.Helper()
|
||||
req := httptest.NewRequest(http.MethodPost, "/mcp", strings.NewReader(body))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
req.Header.Set("Authorization", "Bearer "+validToken)
|
||||
rec := httptest.NewRecorder()
|
||||
s.Handler().ServeHTTP(rec, req)
|
||||
return rec
|
||||
}
|
||||
|
||||
// decode reads a JSON-RPC response, failing the test if the envelope is wrong.
|
||||
func decode(t *testing.T, rec *httptest.ResponseRecorder) response {
|
||||
t.Helper()
|
||||
var resp response
|
||||
if err := json.Unmarshal(rec.Body.Bytes(), &resp); err != nil {
|
||||
t.Fatalf("response was not JSON: %v\nbody: %s", err, rec.Body.String())
|
||||
}
|
||||
if resp.JSONRPC != jsonRPCVersion {
|
||||
t.Fatalf("jsonrpc = %q, want %q", resp.JSONRPC, jsonRPCVersion)
|
||||
}
|
||||
return resp
|
||||
}
|
||||
|
||||
// resultInto re-decodes a successful result into dst.
|
||||
func resultInto(t *testing.T, rec *httptest.ResponseRecorder, dst any) {
|
||||
t.Helper()
|
||||
var envelope struct {
|
||||
Result json.RawMessage `json:"result"`
|
||||
Error *rpcError `json:"error"`
|
||||
}
|
||||
if err := json.Unmarshal(rec.Body.Bytes(), &envelope); err != nil {
|
||||
t.Fatalf("response was not JSON: %v", err)
|
||||
}
|
||||
if envelope.Error != nil {
|
||||
t.Fatalf("expected a result, got error %d: %s", envelope.Error.Code, envelope.Error.Message)
|
||||
}
|
||||
if err := json.Unmarshal(envelope.Result, dst); err != nil {
|
||||
t.Fatalf("result did not decode: %v\nresult: %s", err, envelope.Result)
|
||||
}
|
||||
}
|
||||
|
||||
/* ── 1. initialize ──────────────────────────────────────────────────────── */
|
||||
|
||||
func TestInitializeSucceeds(t *testing.T) {
|
||||
s := newServer(t)
|
||||
rec := post(t, s, `{"jsonrpc":"2.0","id":1,"method":"initialize","params":{
|
||||
"protocolVersion":"2025-06-18",
|
||||
"capabilities":{},
|
||||
"clientInfo":{"name":"test-client","version":"1.0"}}}`)
|
||||
|
||||
if rec.Code != http.StatusOK {
|
||||
t.Fatalf("status = %d, want 200", rec.Code)
|
||||
}
|
||||
if ct := rec.Header().Get("Content-Type"); !strings.HasPrefix(ct, "application/json") {
|
||||
t.Errorf("Content-Type = %q, want application/json", ct)
|
||||
}
|
||||
|
||||
var out initializeResult
|
||||
resultInto(t, rec, &out)
|
||||
|
||||
if out.ProtocolVersion != ProtocolVersion {
|
||||
t.Errorf("protocolVersion = %q, want %q", out.ProtocolVersion, ProtocolVersion)
|
||||
}
|
||||
if out.ServerInfo["name"] != ServerName {
|
||||
t.Errorf("serverInfo.name = %v, want %q", out.ServerInfo["name"], ServerName)
|
||||
}
|
||||
// Exactly one capability, because exactly one is implemented. Advertising
|
||||
// resources or prompts would promise methods that answer method-not-found.
|
||||
if _, ok := out.Capabilities["tools"]; !ok {
|
||||
t.Error("capabilities.tools is missing")
|
||||
}
|
||||
for _, unimplemented := range []string{"resources", "prompts", "sampling", "logging"} {
|
||||
if _, ok := out.Capabilities[unimplemented]; ok {
|
||||
t.Errorf("capabilities advertises %q, which is not implemented", unimplemented)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestInitializeEchoesTheRequestID(t *testing.T) {
|
||||
s := newServer(t)
|
||||
// A string id, to prove ids are echoed verbatim rather than coerced.
|
||||
rec := post(t, s, `{"jsonrpc":"2.0","id":"abc-123","method":"initialize","params":{}}`)
|
||||
resp := decode(t, rec)
|
||||
if string(resp.ID) != `"abc-123"` {
|
||||
t.Errorf("id = %s, want \"abc-123\"", resp.ID)
|
||||
}
|
||||
}
|
||||
|
||||
/* ── 2 & 3. tools/list and exact exposure ───────────────────────────────── */
|
||||
|
||||
// wantExposed is the Phase 0 read-only MVP set, written out in full.
|
||||
//
|
||||
// Deliberately a literal rather than a filter over the registry: a test that
|
||||
// derived the expectation the same way the code does would pass no matter what
|
||||
// the rule became. This is the list a human agreed to, and it is the thing that
|
||||
// should fail when the rule changes.
|
||||
var wantExposed = []string{
|
||||
"activity_breakdown",
|
||||
"activity_signals",
|
||||
"available_workers",
|
||||
"candidates_awaiting",
|
||||
"candidates_quality",
|
||||
"hires_performance",
|
||||
"hires_recent",
|
||||
"open_positions",
|
||||
"operations_risk",
|
||||
"positions_risk",
|
||||
"talent_pool",
|
||||
"workforce_attendance",
|
||||
"workforce_coverage",
|
||||
"workforce_overtime",
|
||||
"workforce_training",
|
||||
"workspace_summary",
|
||||
}
|
||||
|
||||
func TestToolsListExposesExactlyTheReadOnlyMVPSet(t *testing.T) {
|
||||
s := newServer(t)
|
||||
rec := post(t, s, `{"jsonrpc":"2.0","id":2,"method":"tools/list"}`)
|
||||
|
||||
var out toolsListResult
|
||||
resultInto(t, rec, &out)
|
||||
|
||||
got := make([]string, 0, len(out.Tools))
|
||||
for _, tool := range out.Tools {
|
||||
got = append(got, tool.Name)
|
||||
}
|
||||
|
||||
if len(got) != len(wantExposed) {
|
||||
t.Fatalf("exposed %d tools, want %d\n got: %v\nwant: %v",
|
||||
len(got), len(wantExposed), got, wantExposed)
|
||||
}
|
||||
for i := range wantExposed {
|
||||
if got[i] != wantExposed[i] {
|
||||
t.Errorf("tool[%d] = %q, want %q", i, got[i], wantExposed[i])
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestToolsListWorksWithoutParams(t *testing.T) {
|
||||
s := newServer(t)
|
||||
// No params key at all — every exposed tool's arguments are optional, so a
|
||||
// bare list call must work.
|
||||
rec := post(t, s, `{"jsonrpc":"2.0","id":3,"method":"tools/list"}`)
|
||||
var out toolsListResult
|
||||
resultInto(t, rec, &out)
|
||||
if len(out.Tools) == 0 {
|
||||
t.Fatal("tools/list returned nothing")
|
||||
}
|
||||
}
|
||||
|
||||
/* ── 4 & 5. write tools and knowledge_search are not exposed ────────────── */
|
||||
|
||||
func TestNoWriteToolIsExposed(t *testing.T) {
|
||||
s := newServer(t)
|
||||
rec := post(t, s, `{"jsonrpc":"2.0","id":4,"method":"tools/list"}`)
|
||||
var out toolsListResult
|
||||
resultInto(t, rec, &out)
|
||||
|
||||
reg := testRegistry(t)
|
||||
for _, published := range out.Tools {
|
||||
tool, ok := reg.Get(published.Name)
|
||||
if !ok {
|
||||
t.Fatalf("tools/list published %q, which is not in the registry", published.Name)
|
||||
}
|
||||
if tool.Effect != tools.EffectRead {
|
||||
t.Errorf("%q is exposed with effect %q; only read tools may be exposed",
|
||||
tool.Name, tool.Effect)
|
||||
}
|
||||
if tool.RequiresConfirmation {
|
||||
t.Errorf("%q is exposed and requires confirmation, which this surface cannot obtain",
|
||||
tool.Name)
|
||||
}
|
||||
// The annotation must agree with the registry, or a client shows a
|
||||
// person the wrong thing before approving.
|
||||
if published.Annotations == nil || !published.Annotations.ReadOnlyHint {
|
||||
t.Errorf("%q is not annotated readOnlyHint", tool.Name)
|
||||
}
|
||||
if published.Annotations != nil && published.Annotations.DestructiveHint {
|
||||
t.Errorf("%q is annotated destructiveHint", tool.Name)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestNamedExclusionsAreNotExposed(t *testing.T) {
|
||||
s := newServer(t)
|
||||
rec := post(t, s, `{"jsonrpc":"2.0","id":5,"method":"tools/list"}`)
|
||||
var out toolsListResult
|
||||
resultInto(t, rec, &out)
|
||||
|
||||
published := map[string]bool{}
|
||||
for _, tool := range out.Tools {
|
||||
published[tool.Name] = true
|
||||
}
|
||||
// The two writes and the one deferred read. Named explicitly, because these
|
||||
// three are the ones a future change is most likely to let through.
|
||||
for _, forbidden := range []string{"assign_worker", "move_application", "knowledge_search"} {
|
||||
if published[forbidden] {
|
||||
t.Errorf("%q must not be exposed", forbidden)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// The registry must still hold the excluded tools — they are withheld from this
|
||||
// surface, not removed from KROW. A test that only checked absence would pass
|
||||
// if somebody deleted them.
|
||||
func TestExcludedToolsStillExistInTheRegistry(t *testing.T) {
|
||||
reg := testRegistry(t)
|
||||
for _, name := range []string{"assign_worker", "move_application", "knowledge_search"} {
|
||||
if _, ok := reg.Get(name); !ok {
|
||||
t.Errorf("%q is missing from the registry entirely", name)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/* ── 6. schemas come from the real registry ─────────────────────────────── */
|
||||
|
||||
func TestToolSchemaIsSourcedFromTheRegistry(t *testing.T) {
|
||||
s := newServer(t)
|
||||
rec := post(t, s, `{"jsonrpc":"2.0","id":6,"method":"tools/list"}`)
|
||||
var out toolsListResult
|
||||
resultInto(t, rec, &out)
|
||||
|
||||
reg := testRegistry(t)
|
||||
for _, published := range out.Tools {
|
||||
tool, _ := reg.Get(published.Name)
|
||||
|
||||
if published.Description != tool.Description {
|
||||
t.Errorf("%q description differs from the registry's", tool.Name)
|
||||
}
|
||||
if published.InputSchema == nil {
|
||||
t.Fatalf("%q published a nil inputSchema", tool.Name)
|
||||
}
|
||||
// Every exposed tool declares a schema, so the published one must be
|
||||
// the registry's own map and not the empty-schema fallback.
|
||||
if tool.InputSchema == nil {
|
||||
t.Errorf("%q has no InputSchema in the registry; expected every exposed tool to declare one", tool.Name)
|
||||
continue
|
||||
}
|
||||
wantJSON, _ := json.Marshal(tool.InputSchema)
|
||||
gotJSON, _ := json.Marshal(published.InputSchema)
|
||||
if string(wantJSON) != string(gotJSON) {
|
||||
t.Errorf("%q schema differs from the registry's\n got: %s\nwant: %s",
|
||||
tool.Name, gotJSON, wantJSON)
|
||||
}
|
||||
// A published schema must be a JSON Schema object, or a client may
|
||||
// refuse the whole list.
|
||||
if published.InputSchema["type"] != "object" {
|
||||
t.Errorf("%q schema type = %v, want \"object\"", tool.Name, published.InputSchema["type"])
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// No exposed tool may take a tenant or principal identifier as an argument.
|
||||
// An argument is something a client chooses; identity is not the client's to
|
||||
// choose. Enforced here rather than by review, because the cost of missing it
|
||||
// once is cross-tenant access.
|
||||
func TestNoToolAcceptsATenantOrPrincipalArgument(t *testing.T) {
|
||||
s := newServer(t)
|
||||
rec := post(t, s, `{"jsonrpc":"2.0","id":7,"method":"tools/list"}`)
|
||||
var out toolsListResult
|
||||
resultInto(t, rec, &out)
|
||||
|
||||
forbidden := []string{
|
||||
"org_id", "orgid", "organization_id", "organisation_id",
|
||||
"tenant_id", "tenantid", "tenant",
|
||||
"user_id", "userid", "principal", "principal_id",
|
||||
"account_id", "caller", "caller_id", "on_behalf_of", "impersonate",
|
||||
}
|
||||
|
||||
for _, tool := range out.Tools {
|
||||
props, ok := tool.InputSchema["properties"].(map[string]any)
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
for field := range props {
|
||||
lower := strings.ToLower(field)
|
||||
for _, bad := range forbidden {
|
||||
if lower == bad {
|
||||
t.Errorf("%q accepts %q as an argument; identity must come from the "+
|
||||
"authenticated principal, never from the request", tool.Name, field)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/* ── 7 & 8. unknown method, unknown tool ────────────────────────────────── */
|
||||
|
||||
func TestUnknownMethodReturnsMethodNotFound(t *testing.T) {
|
||||
s := newServer(t)
|
||||
rec := post(t, s, `{"jsonrpc":"2.0","id":8,"method":"resources/list"}`)
|
||||
|
||||
if rec.Code != http.StatusOK {
|
||||
t.Errorf("status = %d, want 200 — a JSON-RPC error is still a successful HTTP exchange", rec.Code)
|
||||
}
|
||||
resp := decode(t, rec)
|
||||
if resp.Error == nil {
|
||||
t.Fatal("expected an error")
|
||||
}
|
||||
if resp.Error.Code != codeMethodNotFound {
|
||||
t.Errorf("code = %d, want %d", resp.Error.Code, codeMethodNotFound)
|
||||
}
|
||||
}
|
||||
|
||||
func TestUnknownToolIsReportedAsAToolError(t *testing.T) {
|
||||
s := newServer(t)
|
||||
rec := post(t, s, `{"jsonrpc":"2.0","id":9,"method":"tools/call",
|
||||
"params":{"name":"no_such_tool","arguments":{}}}`)
|
||||
|
||||
var out toolsCallResult
|
||||
resultInto(t, rec, &out)
|
||||
if !out.IsError {
|
||||
t.Error("expected isError on an unknown tool")
|
||||
}
|
||||
if len(out.Content) == 0 || !strings.Contains(out.Content[0].Text, "mcp.unknown_tool") {
|
||||
t.Errorf("expected mcp.unknown_tool, got %+v", out.Content)
|
||||
}
|
||||
}
|
||||
|
||||
// An unexposed tool must be indistinguishable from a non-existent one, or the
|
||||
// error messages become an inventory of what this surface is withholding.
|
||||
func TestUnexposedToolIsIndistinguishableFromUnknown(t *testing.T) {
|
||||
s := newServer(t)
|
||||
|
||||
unknown := post(t, s, `{"jsonrpc":"2.0","id":10,"method":"tools/call","params":{"name":"definitely_not_a_tool"}}`)
|
||||
withheld := post(t, s, `{"jsonrpc":"2.0","id":10,"method":"tools/call","params":{"name":"assign_worker"}}`)
|
||||
|
||||
var a, b toolsCallResult
|
||||
resultInto(t, unknown, &a)
|
||||
resultInto(t, withheld, &b)
|
||||
|
||||
codeOf := func(r toolsCallResult) string {
|
||||
var payload struct {
|
||||
Error struct {
|
||||
Code string `json:"code"`
|
||||
} `json:"error"`
|
||||
}
|
||||
_ = json.Unmarshal([]byte(r.Content[0].Text), &payload)
|
||||
return payload.Error.Code
|
||||
}
|
||||
if codeOf(a) != codeOf(b) {
|
||||
t.Errorf("an unexposed tool is distinguishable from an unknown one: %q vs %q",
|
||||
codeOf(a), codeOf(b))
|
||||
}
|
||||
if a.IsError != b.IsError {
|
||||
t.Error("isError differs between unknown and unexposed tools")
|
||||
}
|
||||
}
|
||||
|
||||
/* ── 9 & 10. malformed JSON and params ──────────────────────────────────── */
|
||||
|
||||
func TestMalformedJSONReturnsParseError(t *testing.T) {
|
||||
s := newServer(t)
|
||||
rec := post(t, s, `{"jsonrpc":"2.0","id":1,"method":`)
|
||||
|
||||
resp := decode(t, rec)
|
||||
if resp.Error == nil || resp.Error.Code != codeParseError {
|
||||
t.Fatalf("want parse error %d, got %+v", codeParseError, resp.Error)
|
||||
}
|
||||
}
|
||||
|
||||
func TestEnvelopeValidation(t *testing.T) {
|
||||
for name, tc := range map[string]struct {
|
||||
body string
|
||||
want int
|
||||
}{
|
||||
"missing jsonrpc": {`{"id":1,"method":"tools/list"}`, codeInvalidRequest},
|
||||
"wrong jsonrpc": {`{"jsonrpc":"1.0","id":1,"method":"tools/list"}`, codeInvalidRequest},
|
||||
"missing method": {`{"jsonrpc":"2.0","id":1}`, codeInvalidRequest},
|
||||
"empty body": {``, codeInvalidRequest},
|
||||
"object id": {`{"jsonrpc":"2.0","id":{"a":1},"method":"tools/list"}`, codeInvalidRequest},
|
||||
"not an object": {`"a string"`, codeParseError},
|
||||
"params as array": {`{"jsonrpc":"2.0","id":1,"method":"tools/call","params":[1,2]}`, codeInvalidParams},
|
||||
"params wrong type": {`{"jsonrpc":"2.0","id":1,"method":"tools/call","params":{"name":42}}`, codeInvalidParams},
|
||||
"missing tool name": {`{"jsonrpc":"2.0","id":1,"method":"tools/call","params":{}}`, codeInvalidParams},
|
||||
} {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
rec := post(t, newServer(t), tc.body)
|
||||
resp := decode(t, rec)
|
||||
if resp.Error == nil {
|
||||
t.Fatalf("expected an error, got %s", rec.Body.String())
|
||||
}
|
||||
if resp.Error.Code != tc.want {
|
||||
t.Errorf("code = %d, want %d (%s)", resp.Error.Code, tc.want, resp.Error.Message)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
/* ── 11. batching is explicitly rejected ────────────────────────────────── */
|
||||
|
||||
func TestBatchRequestsAreExplicitlyRejected(t *testing.T) {
|
||||
s := newServer(t)
|
||||
rec := post(t, s, `[{"jsonrpc":"2.0","id":1,"method":"tools/list"},
|
||||
{"jsonrpc":"2.0","id":2,"method":"tools/list"}]`)
|
||||
|
||||
resp := decode(t, rec)
|
||||
if resp.Error == nil {
|
||||
t.Fatal("a batch must be refused")
|
||||
}
|
||||
if resp.Error.Code != codeInvalidRequest {
|
||||
t.Errorf("code = %d, want %d", resp.Error.Code, codeInvalidRequest)
|
||||
}
|
||||
// The refusal must say WHY, so a client can fall back to sending singly
|
||||
// rather than retrying the same batch forever.
|
||||
detail, _ := resp.Error.Data.(string)
|
||||
if !strings.Contains(detail, "batch") {
|
||||
t.Errorf("the refusal does not mention batching: %q", detail)
|
||||
}
|
||||
}
|
||||
|
||||
/* ── 12. no panic on malformed input ────────────────────────────────────── */
|
||||
|
||||
func TestNoPanicOnHostileInput(t *testing.T) {
|
||||
// Each of these has crashed a hand-written JSON-RPC server somewhere.
|
||||
bodies := []string{
|
||||
``,
|
||||
` `,
|
||||
`null`,
|
||||
`[]`,
|
||||
`[[[[[]]]]]`,
|
||||
`{`,
|
||||
`}`,
|
||||
`{"jsonrpc":"2.0"}`,
|
||||
`{"jsonrpc":null,"id":null,"method":null}`,
|
||||
`{"jsonrpc":"2.0","id":1,"method":"tools/call","params":null}`,
|
||||
`{"jsonrpc":"2.0","id":1,"method":"tools/call","params":{"name":""}}`,
|
||||
`{"jsonrpc":"2.0","id":1,"method":"tools/call","params":{"name":"activity_breakdown","arguments":"not-an-object"}}`,
|
||||
`{"jsonrpc":"2.0","id":1,"method":"tools/call","params":{"name":"activity_breakdown","arguments":[]}}`,
|
||||
`{"jsonrpc":"2.0","id":[1,2,3],"method":"tools/list"}`,
|
||||
`{"jsonrpc":"2.0","id":1,"method":"` + strings.Repeat("A", 10_000) + `"}`,
|
||||
`{"jsonrpc":"2.0","id":1,"method":"tools/list","params":{"cursor":` + strings.Repeat("[", 500) + `}}`,
|
||||
"\x00\x01\x02",
|
||||
}
|
||||
|
||||
s := newServer(t)
|
||||
for i, body := range bodies {
|
||||
// A panic escaping here fails the test by crashing it, which is the
|
||||
// assertion: the handler must answer every one of these.
|
||||
rec := post(t, s, body)
|
||||
if rec.Code < 200 || rec.Code >= 600 {
|
||||
t.Errorf("body %d produced status %d", i, rec.Code)
|
||||
}
|
||||
if rec.Body.Len() > 0 && rec.Code != http.StatusAccepted {
|
||||
var resp response
|
||||
if err := json.Unmarshal(rec.Body.Bytes(), &resp); err != nil {
|
||||
t.Errorf("body %d produced a non-JSON response: %s", i, rec.Body.String())
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/* ── Transport ──────────────────────────────────────────────────────────── */
|
||||
|
||||
func TestOnlyPOSTIsAccepted(t *testing.T) {
|
||||
s := newServer(t)
|
||||
for _, method := range []string{http.MethodGet, http.MethodPut, http.MethodDelete, http.MethodPatch} {
|
||||
req := httptest.NewRequest(method, "/mcp", nil)
|
||||
rec := httptest.NewRecorder()
|
||||
s.Handler().ServeHTTP(rec, req)
|
||||
|
||||
if rec.Code != http.StatusMethodNotAllowed {
|
||||
t.Errorf("%s: status = %d, want 405", method, rec.Code)
|
||||
}
|
||||
if allow := rec.Header().Get("Allow"); allow != http.MethodPost {
|
||||
t.Errorf("%s: Allow = %q, want POST", method, allow)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestNonJSONContentTypeIsRefused(t *testing.T) {
|
||||
s := newServer(t)
|
||||
req := httptest.NewRequest(http.MethodPost, "/mcp",
|
||||
strings.NewReader(`{"jsonrpc":"2.0","id":1,"method":"tools/list"}`))
|
||||
req.Header.Set("Content-Type", "application/x-www-form-urlencoded")
|
||||
rec := httptest.NewRecorder()
|
||||
s.Handler().ServeHTTP(rec, req)
|
||||
|
||||
if rec.Code != http.StatusUnsupportedMediaType {
|
||||
t.Errorf("status = %d, want 415", rec.Code)
|
||||
}
|
||||
}
|
||||
|
||||
func TestOversizedBodyIsRefused(t *testing.T) {
|
||||
s := newServer(t)
|
||||
huge := `{"jsonrpc":"2.0","id":1,"method":"tools/list","params":{"cursor":"` +
|
||||
strings.Repeat("x", MaxRequestBytes+1) + `"}}`
|
||||
rec := post(t, s, huge)
|
||||
|
||||
if rec.Code != http.StatusRequestEntityTooLarge {
|
||||
t.Errorf("status = %d, want 413", rec.Code)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNotificationGetsNoBody(t *testing.T) {
|
||||
s := newServer(t)
|
||||
// No id — a notification. The spec forbids a response.
|
||||
rec := post(t, s, `{"jsonrpc":"2.0","method":"notifications/initialized"}`)
|
||||
|
||||
if rec.Code != http.StatusAccepted {
|
||||
t.Errorf("status = %d, want 202", rec.Code)
|
||||
}
|
||||
if rec.Body.Len() != 0 {
|
||||
t.Errorf("a notification was answered with a body: %s", rec.Body.String())
|
||||
}
|
||||
}
|
||||
|
||||
/* ── Phase 1 boundary ───────────────────────────────────────────────────── */
|
||||
|
||||
// Handle must refuse a nil identity even when called directly.
|
||||
//
|
||||
// The transport answers 401 before this point in the ordinary case, so this is
|
||||
// the SECOND gate: a future caller of Handle that forgets to authenticate must
|
||||
// fail closed rather than dispatch as nobody. It exists to fail loudly if
|
||||
// somebody later supplies a default identity to "make it work".
|
||||
func TestHandleRefusesANilIdentity(t *testing.T) {
|
||||
s := newServer(t)
|
||||
req := request{
|
||||
JSONRPC: jsonRPCVersion,
|
||||
ID: json.RawMessage(`11`),
|
||||
Method: "tools/call",
|
||||
Params: json.RawMessage(`{"name":"activity_breakdown","arguments":{}}`),
|
||||
}
|
||||
result, rpcErr := s.Handle(context.Background(), nil, req)
|
||||
if rpcErr != nil {
|
||||
t.Fatalf("unexpected rpc error: %v", rpcErr)
|
||||
}
|
||||
out, ok := result.(toolsCallResult)
|
||||
if !ok {
|
||||
t.Fatalf("unexpected result type %T", result)
|
||||
}
|
||||
if !out.IsError || !strings.Contains(out.Content[0].Text, "mcp.unauthenticated") {
|
||||
t.Errorf("expected mcp.unauthenticated, got: %+v", out.Content)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPingIsCheap(t *testing.T) {
|
||||
s := newServer(t)
|
||||
rec := post(t, s, `{"jsonrpc":"2.0","id":12,"method":"ping"}`)
|
||||
var out map[string]any
|
||||
resultInto(t, rec, &out)
|
||||
if len(out) != 0 {
|
||||
t.Errorf("ping returned %v, want an empty result", out)
|
||||
}
|
||||
}
|
||||
314
go-api/internal/mcpserver/orglimit_test.go
Normal file
314
go-api/internal/mcpserver/orglimit_test.go
Normal file
@@ -0,0 +1,314 @@
|
||||
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):]
|
||||
}
|
||||
}
|
||||
387
go-api/internal/mcpserver/server.go
Normal file
387
go-api/internal/mcpserver/server.go
Normal file
@@ -0,0 +1,387 @@
|
||||
package mcpserver
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"log/slog"
|
||||
"time"
|
||||
|
||||
"github.com/krow/krow-backend/go-api/internal/authctx"
|
||||
"github.com/krow/krow-backend/go-api/internal/tools"
|
||||
)
|
||||
|
||||
// ProtocolVersion is the MCP revision this server implements.
|
||||
//
|
||||
// Echoed back from initialize. A client asking for a different revision is not
|
||||
// refused: the spec's negotiation is that the server states what it speaks and
|
||||
// the client decides whether it can work with that. Refusing would turn a
|
||||
// version skew into an outage where it is usually a compatible difference.
|
||||
const ProtocolVersion = "2025-06-18"
|
||||
|
||||
// ServerName and ServerVersion identify this implementation to a client.
|
||||
const (
|
||||
ServerName = "krow-mcp"
|
||||
ServerVersion = "0.1.0"
|
||||
)
|
||||
|
||||
// maxToolArgumentBytes bounds one tool call's arguments.
|
||||
//
|
||||
// The same reasoning as runs.go's maxRunRequestBytes: arguments are a handful
|
||||
// of scalars against a schema that sets additionalProperties:false, so anything
|
||||
// large is either a mistake or an attempt to push text past the tool layer into
|
||||
// a prompt. The transport bounds the whole body too; this bounds the part that
|
||||
// reaches a handler.
|
||||
const maxToolArgumentBytes = 64 << 10
|
||||
|
||||
// Server answers MCP methods against the existing tool registry.
|
||||
//
|
||||
// It holds a *tools.Registry and nothing else that matters. There is no second
|
||||
// registry, no adapter table and no per-tool code in this package: what MCP
|
||||
// publishes is what the registry holds, filtered by the rule in tools.go.
|
||||
type Server struct {
|
||||
reg *tools.Registry
|
||||
log *slog.Logger
|
||||
|
||||
// tokens resolves a bearer credential into an identity. Nil means this
|
||||
// surface cannot authenticate anyone, and every call is refused — see
|
||||
// ErrNoAuthenticator. Failing closed is the only safe default for a field
|
||||
// whose absence would otherwise mean "let everyone in".
|
||||
tokens TokenAuthenticator
|
||||
|
||||
// resourceMetadataURL is where a 401 points a client so it can begin
|
||||
// discovery. Empty means the challenge carries no pointer, which is a
|
||||
// valid but less useful 401: a client then has nowhere to look.
|
||||
resourceMetadataURL string
|
||||
|
||||
// orgLimiter bounds tool calls per organisation.
|
||||
//
|
||||
// HERE rather than in the HTTP middleware, and that placement is the whole
|
||||
// point: an organisation is not knowable until the bearer token has been
|
||||
// resolved to a user and that user's row read. A middleware running before
|
||||
// authentication could only key by something the CLIENT supplied, which is
|
||||
// precisely the identity this surface refuses to trust.
|
||||
//
|
||||
// Nil means no per-organisation ceiling, which is the correct default for
|
||||
// a deployment that has not configured one.
|
||||
orgLimiter OrgLimiter
|
||||
}
|
||||
|
||||
// OrgLimiter bounds how much one organisation may ask for.
|
||||
//
|
||||
// Takes an org id that the caller has already established from an authenticated
|
||||
// identity. It cannot be handed anything from a request, because the only
|
||||
// caller is dispatch, which has an authctx.Identity and nothing else.
|
||||
type OrgLimiter interface {
|
||||
// AllowOrg reports whether this organisation may make another call, and
|
||||
// how long until its window rolls over.
|
||||
AllowOrg(ctx context.Context, orgID string) (allowed bool, retryAfter time.Duration, err error)
|
||||
}
|
||||
|
||||
// WithOrgLimiter installs the per-organisation ceiling.
|
||||
func (s *Server) WithOrgLimiter(l OrgLimiter) *Server {
|
||||
s.orgLimiter = l
|
||||
return s
|
||||
}
|
||||
|
||||
// WithResourceMetadataURL sets the RFC 9728 document a 401 points at.
|
||||
//
|
||||
// Supplied by the caller rather than derived here, because this package does
|
||||
// not know its own deployment's URLs and must not invent them. A hardcoded
|
||||
// hostname would be one deployment's identity baked into every other one.
|
||||
func (s *Server) WithResourceMetadataURL(u string) *Server {
|
||||
s.resourceMetadataURL = u
|
||||
return s
|
||||
}
|
||||
|
||||
// challenge builds the WWW-Authenticate header for a 401.
|
||||
//
|
||||
// RFC 9728 section 5.1: the client reads `resource_metadata` from here to find
|
||||
// the protected-resource document, and from there the authorization server.
|
||||
// Without the parameter a compliant client has a 401 and nowhere to go, which
|
||||
// is why this is the difference between "authentication failed" and "here is
|
||||
// how to authenticate".
|
||||
func (s *Server) challenge() string {
|
||||
c := `Bearer realm="` + ServerName + `"`
|
||||
if s.resourceMetadataURL != "" {
|
||||
c += `, resource_metadata="` + s.resourceMetadataURL + `"`
|
||||
}
|
||||
return c
|
||||
}
|
||||
|
||||
// New builds a server over an existing registry.
|
||||
//
|
||||
// The registry is the one the rest of the service already built — the caller
|
||||
// passes runtime.DefaultTools(...)'s result, the same value the HTTP server
|
||||
// uses for its author catalogue. Taking it as a parameter rather than building
|
||||
// one here is what guarantees there is only ever one.
|
||||
func New(reg *tools.Registry, tokens TokenAuthenticator, log *slog.Logger) *Server {
|
||||
if log == nil {
|
||||
log = slog.Default()
|
||||
}
|
||||
return &Server{reg: reg, tokens: tokens, log: log}
|
||||
}
|
||||
|
||||
// Handle dispatches one parsed JSON-RPC request on behalf of an identity.
|
||||
//
|
||||
// The identity is a PARAMETER, not something read from the context, and that is
|
||||
// the security property rather than a style choice. If this function resolved
|
||||
// the caller from ctx, then mounting the endpoint behind the cookie middleware
|
||||
// would make a browser session sufficient to call MCP tools — the middleware
|
||||
// puts an Identity in the context, and this code would find it. Taking it as an
|
||||
// argument means only the MCP transport's own bearer authentication can supply
|
||||
// one. See auth.go.
|
||||
//
|
||||
// A nil identity means unauthenticated. The three handshake methods are allowed
|
||||
// without one; tools/call is not.
|
||||
//
|
||||
// Returns a result or an error, never both. Notifications are handled by the
|
||||
// transport, which discards whatever comes back.
|
||||
func (s *Server) Handle(ctx context.Context, ident *authctx.Identity, req request) (any, *rpcError) {
|
||||
switch req.Method {
|
||||
case "initialize":
|
||||
return s.handleInitialize(req.Params)
|
||||
case "notifications/initialized":
|
||||
// The client telling us it is ready. Nothing to do, and answering is
|
||||
// not required — it arrives as a notification.
|
||||
return map[string]any{}, nil
|
||||
case "ping":
|
||||
// Cheap liveness, defined by the spec as an empty result. Costs nothing
|
||||
// and saves a client from using tools/list as a heartbeat.
|
||||
return map[string]any{}, nil
|
||||
case "tools/list":
|
||||
return s.handleToolsList(req.Params)
|
||||
case "tools/call":
|
||||
return s.handleToolsCall(ctx, ident, req.Params)
|
||||
default:
|
||||
return nil, errMethodNotFound(req.Method)
|
||||
}
|
||||
}
|
||||
|
||||
/* ── initialize ─────────────────────────────────────────────────────────── */
|
||||
|
||||
type initializeParams struct {
|
||||
ProtocolVersion string `json:"protocolVersion"`
|
||||
Capabilities json.RawMessage `json:"capabilities"`
|
||||
ClientInfo struct {
|
||||
Name string `json:"name"`
|
||||
Version string `json:"version"`
|
||||
} `json:"clientInfo"`
|
||||
}
|
||||
|
||||
type initializeResult struct {
|
||||
ProtocolVersion string `json:"protocolVersion"`
|
||||
Capabilities map[string]any `json:"capabilities"`
|
||||
ServerInfo map[string]any `json:"serverInfo"`
|
||||
Instructions string `json:"instructions,omitempty"`
|
||||
}
|
||||
|
||||
// handleInitialize answers the opening handshake.
|
||||
//
|
||||
// Declares exactly one capability, because exactly one is implemented. A server
|
||||
// that advertised resources or prompts here would be promising methods that
|
||||
// answer method-not-found, and a client would reasonably call them.
|
||||
//
|
||||
// listChanged is false: the tool set is fixed at process start by the registry,
|
||||
// so there is no change to notify anyone about.
|
||||
func (s *Server) handleInitialize(raw json.RawMessage) (any, *rpcError) {
|
||||
var p initializeParams
|
||||
if err := decodeParams(raw, &p); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
s.log.Info("mcp initialize",
|
||||
"client_name", p.ClientInfo.Name,
|
||||
"client_version", p.ClientInfo.Version,
|
||||
"client_protocol", p.ProtocolVersion,
|
||||
"server_protocol", ProtocolVersion)
|
||||
|
||||
return initializeResult{
|
||||
ProtocolVersion: ProtocolVersion,
|
||||
Capabilities: map[string]any{
|
||||
"tools": map[string]any{"listChanged": false},
|
||||
},
|
||||
ServerInfo: map[string]any{
|
||||
"name": ServerName,
|
||||
"version": ServerVersion,
|
||||
},
|
||||
Instructions: "Read-only access to KROW workforce and hiring data. " +
|
||||
"Every call is scoped to the authenticated user's organisation and role; " +
|
||||
"results are structured data for you to summarise, not prose.",
|
||||
}, nil
|
||||
}
|
||||
|
||||
/* ── tools/list ─────────────────────────────────────────────────────────── */
|
||||
|
||||
type toolsListResult struct {
|
||||
Tools []mcpTool `json:"tools"`
|
||||
}
|
||||
|
||||
// handleToolsList publishes the exposed tools, straight from the registry.
|
||||
func (s *Server) handleToolsList(raw json.RawMessage) (any, *rpcError) {
|
||||
// Params are optional here (cursor, for pagination this server does not
|
||||
// need), but a malformed object is still worth refusing rather than
|
||||
// ignoring — silently accepting nonsense trains a client to send it.
|
||||
var p struct {
|
||||
Cursor string `json:"cursor,omitempty"`
|
||||
}
|
||||
if err := decodeParams(raw, &p); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
infos := exposed(s.reg)
|
||||
out := make([]mcpTool, 0, len(infos))
|
||||
for _, info := range infos {
|
||||
out = append(out, toMCPTool(info))
|
||||
}
|
||||
return toolsListResult{Tools: out}, nil
|
||||
}
|
||||
|
||||
/* ── tools/call ─────────────────────────────────────────────────────────── */
|
||||
|
||||
type toolsCallParams struct {
|
||||
Name string `json:"name"`
|
||||
Arguments json.RawMessage `json:"arguments,omitempty"`
|
||||
}
|
||||
|
||||
// toolsCallResult is MCP's shape for a tool's output.
|
||||
//
|
||||
// IsError is part of the RESULT, not a JSON-RPC error: a tool that refused is
|
||||
// not a protocol fault, and reporting it as one would deny the model the chance
|
||||
// to read the refusal and do something sensible. It is the same distinction
|
||||
// tools.Result already draws, and runs.go draws for terminations.
|
||||
type toolsCallResult struct {
|
||||
Content []contentBlock `json:"content"`
|
||||
IsError bool `json:"isError,omitempty"`
|
||||
}
|
||||
|
||||
type contentBlock struct {
|
||||
Type string `json:"type"`
|
||||
Text string `json:"text"`
|
||||
}
|
||||
|
||||
func textResult(payload any, isError bool) (toolsCallResult, *rpcError) {
|
||||
encoded, err := json.MarshalIndent(payload, "", " ")
|
||||
if err != nil {
|
||||
return toolsCallResult{}, errInternal()
|
||||
}
|
||||
return toolsCallResult{
|
||||
Content: []contentBlock{{Type: "text", Text: string(encoded)}},
|
||||
IsError: isError,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// handleToolsCall validates a call and dispatches it through the registry.
|
||||
//
|
||||
// The order is: exposure, then bounds, then identity, then dispatch.
|
||||
//
|
||||
// Exposure is checked BEFORE identity on purpose. "There is no such tool here"
|
||||
// does not depend on who is asking, and answering it first means the surface's
|
||||
// tool inventory is not something an attacker can probe by comparing an
|
||||
// authenticated 404 against an unauthenticated 401.
|
||||
//
|
||||
// Authentication is the transport's job and has already happened by the time
|
||||
// this runs; `ident` is nil only when it failed or was never attempted. This
|
||||
// function does not read the ambient context for a caller — see Handle.
|
||||
func (s *Server) handleToolsCall(ctx context.Context, ident *authctx.Identity, raw json.RawMessage) (any, *rpcError) {
|
||||
var p toolsCallParams
|
||||
if err := decodeParams(raw, &p); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if p.Name == "" {
|
||||
return nil, errInvalidParams("name is required")
|
||||
}
|
||||
if len(p.Arguments) > maxToolArgumentBytes {
|
||||
return nil, errInvalidParams("arguments are too large")
|
||||
}
|
||||
|
||||
// Unexposed and unknown are the SAME answer, deliberately. See
|
||||
// isExposedName — distinguishing them inventories what this surface is
|
||||
// hiding.
|
||||
if !isExposedName(s.reg, p.Name) {
|
||||
s.log.Warn("mcp tool call refused", "tool", p.Name, "reason", "not_exposed")
|
||||
return textResult(map[string]any{
|
||||
"error": map[string]any{
|
||||
"code": "mcp.unknown_tool",
|
||||
"message": "there is no tool called " + p.Name + " on this surface",
|
||||
},
|
||||
}, true)
|
||||
}
|
||||
|
||||
// Unauthenticated calls never reach a handler. The transport answers 401
|
||||
// before this point in the ordinary case; this is the second gate, so that
|
||||
// a future caller of Handle that forgets to authenticate fails closed
|
||||
// rather than dispatching as nobody.
|
||||
if ident == nil {
|
||||
s.log.Warn("mcp tool call refused", "tool", p.Name, "reason", "no_identity")
|
||||
return textResult(map[string]any{
|
||||
"error": map[string]any{
|
||||
"code": "mcp.unauthenticated",
|
||||
"message": "this call is not authenticated",
|
||||
},
|
||||
}, true)
|
||||
}
|
||||
|
||||
return s.dispatch(ctx, *ident, p)
|
||||
}
|
||||
|
||||
// dispatch runs the tool through the existing registry.
|
||||
//
|
||||
// This is the only place this package touches the tool layer, and it is four
|
||||
// lines on purpose. Everything that decides what comes back — the policy table,
|
||||
// the org pre-filter, the row scopes, the opaque denial, the truncation — is
|
||||
// inside Dispatch and the handler beneath it, unchanged and unreachable from
|
||||
// here.
|
||||
//
|
||||
// The principal is the caller's, from the context. It is never read from
|
||||
// params: an MCP client that could name its own principal could read anything,
|
||||
// which is the bug I1 exists to prevent.
|
||||
func (s *Server) dispatch(ctx context.Context, identity authctx.Identity, p toolsCallParams) (any, *rpcError) {
|
||||
// The per-organisation ceiling, checked after authentication and before
|
||||
// any work. The org comes from `identity`, which came from the token —
|
||||
// there is no path by which a request can name a different bucket, because
|
||||
// this function is never given anything from the request except the tool
|
||||
// name and its arguments.
|
||||
if s.orgLimiter != nil {
|
||||
allowed, retryAfter, err := s.orgLimiter.AllowOrg(ctx, identity.OrgID)
|
||||
if err != nil {
|
||||
// The limiter has already decided whether a failure permits the
|
||||
// call. Logged without the org's usage, which is not the caller's
|
||||
// business.
|
||||
s.log.Error("org rate limiter unavailable", "error", err)
|
||||
}
|
||||
if !allowed {
|
||||
s.log.Warn("mcp org rate limit exceeded",
|
||||
"org_id", identity.OrgID, "tool", p.Name)
|
||||
return nil, errRateLimited(retryAfter)
|
||||
}
|
||||
}
|
||||
|
||||
args := p.Arguments
|
||||
if len(skipSpace(args)) == 0 {
|
||||
args = json.RawMessage(`{}`)
|
||||
}
|
||||
|
||||
tc := tools.Context{
|
||||
Principal: identity,
|
||||
// No RunID: an MCP call is not an agent run and writes no trajectory.
|
||||
// No KnowledgeSources: there is no spec, which is why knowledge_search
|
||||
// is deferred rather than published — see tools.go.
|
||||
}
|
||||
|
||||
res := s.reg.Dispatch(ctx, tc, p.Name, args)
|
||||
|
||||
s.log.Info("mcp tool call",
|
||||
"tool", p.Name,
|
||||
"user_id", identity.UserID,
|
||||
"org_id", identity.OrgID,
|
||||
"ok", res.Error == nil,
|
||||
"truncated", res.Truncated)
|
||||
|
||||
if res.Error != nil {
|
||||
return textResult(map[string]any{"error": res.Error}, true)
|
||||
}
|
||||
return textResult(map[string]any{
|
||||
"data": res.Data,
|
||||
"truncated": res.Truncated,
|
||||
}, false)
|
||||
}
|
||||
632
go-api/internal/mcpserver/tenant_test.go
Normal file
632
go-api/internal/mcpserver/tenant_test.go
Normal file
@@ -0,0 +1,632 @@
|
||||
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)
|
||||
}
|
||||
145
go-api/internal/mcpserver/tools.go
Normal file
145
go-api/internal/mcpserver/tools.go
Normal file
@@ -0,0 +1,145 @@
|
||||
package mcpserver
|
||||
|
||||
import (
|
||||
"sort"
|
||||
|
||||
"github.com/krow/krow-backend/go-api/internal/tools"
|
||||
)
|
||||
|
||||
// Which tools this surface publishes, and why it is a rule rather than a list.
|
||||
//
|
||||
// The set is DERIVED from the registry on every call, not enumerated. A list
|
||||
// would be a promise someone has to keep: register a write tool tomorrow,
|
||||
// forget to update the list, and it ships to every connected client. A
|
||||
// derivation cannot forget. The only hand-maintained part is `deferred` below,
|
||||
// which names tools held back for a reason other than their effect — and
|
||||
// holding something back is the safe direction to be wrong in.
|
||||
//
|
||||
// Three conditions, all required:
|
||||
//
|
||||
// 1. Effect is read. I4's whole point is that a write is gated; a write
|
||||
// reachable over a surface with no confirmation round-trip is I4 defeated.
|
||||
// 2. RequiresConfirmation is false. Belt and braces: Register already forces
|
||||
// it true for a write, so this catches a READ tool that opted in — some
|
||||
// reads are expensive enough to be worth asking about — which this surface
|
||||
// has no way to ask about yet.
|
||||
// 3. Not in `deferred`.
|
||||
|
||||
// deferred names tools held back for a reason that is not their effect.
|
||||
//
|
||||
// knowledge_search is read-only and still cannot ship. Its corpora come from
|
||||
// tools.Context.KnowledgeSources, which the agent loop fills from the running
|
||||
// agent's SPEC — deliberately, so that which documents may be read is not
|
||||
// something a model can choose. An MCP call has no spec, so the field is empty,
|
||||
// and retrieval refuses an empty source list rather than treating it as "all of
|
||||
// them". The tool would therefore fail every call; publishing it would advertise
|
||||
// a capability that cannot work.
|
||||
//
|
||||
// Making it work is a design decision, not an omission: either the connection
|
||||
// binds to an agent spec whose sources it inherits, or sources are derived from
|
||||
// the caller's org ACL (which needs a reindex). Passing them as a tool argument
|
||||
// is the one option that is ruled out, because that is exactly what the field's
|
||||
// placement in the spec exists to prevent.
|
||||
var deferred = map[string]string{
|
||||
"knowledge_search": "corpora come from an agent spec, which an MCP call does not have",
|
||||
}
|
||||
|
||||
// exposed returns the tools this surface publishes, sorted by name.
|
||||
//
|
||||
// Sorted because tools/list is a set, and a stable order makes it diffable in a
|
||||
// test and in a log.
|
||||
func exposed(reg *tools.Registry) []tools.ToolInfo {
|
||||
out := make([]tools.ToolInfo, 0, 16)
|
||||
for _, info := range reg.Catalogue() {
|
||||
if !isExposable(info) {
|
||||
continue
|
||||
}
|
||||
out = append(out, info)
|
||||
}
|
||||
sort.Slice(out, func(i, j int) bool { return out[i].Name < out[j].Name })
|
||||
return out
|
||||
}
|
||||
|
||||
// isExposable is the rule, in one place, so the list and the check cannot
|
||||
// disagree.
|
||||
func isExposable(info tools.ToolInfo) bool {
|
||||
if info.Effect != string(tools.EffectRead) {
|
||||
return false
|
||||
}
|
||||
if info.RequiresConfirmation {
|
||||
return false
|
||||
}
|
||||
if _, held := deferred[info.Name]; held {
|
||||
return false
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
// isExposedName reports whether a tool may be called by name over this surface.
|
||||
//
|
||||
// tools/call consults this BEFORE the registry, so an unexposed tool answers
|
||||
// exactly as an unknown one does. The alternative — dispatching and letting
|
||||
// authorization refuse — would make "this tool exists but you may not reach it
|
||||
// here" distinguishable from "no such tool", which is an inventory of the
|
||||
// surface's own blind spots.
|
||||
func isExposedName(reg *tools.Registry, name string) bool {
|
||||
t, ok := reg.Get(name)
|
||||
if !ok {
|
||||
return false
|
||||
}
|
||||
return isExposable(tools.ToolInfo{
|
||||
Name: t.Name,
|
||||
Effect: string(t.Effect),
|
||||
RequiresConfirmation: t.RequiresConfirmation,
|
||||
})
|
||||
}
|
||||
|
||||
/* ── MCP shapes ─────────────────────────────────────────────────────────── */
|
||||
|
||||
// mcpTool is one entry in a tools/list result.
|
||||
//
|
||||
// Every field is copied from the registry rather than restated. The annotations
|
||||
// are hints a client may show a person before approving a call; they are
|
||||
// derived from Effect so they cannot contradict it.
|
||||
type mcpTool struct {
|
||||
Name string `json:"name"`
|
||||
Description string `json:"description"`
|
||||
InputSchema map[string]any `json:"inputSchema"`
|
||||
Annotations *annotations `json:"annotations,omitempty"`
|
||||
}
|
||||
|
||||
// annotations are the advisory hints from the MCP tool definition.
|
||||
type annotations struct {
|
||||
ReadOnlyHint bool `json:"readOnlyHint"`
|
||||
DestructiveHint bool `json:"destructiveHint"`
|
||||
}
|
||||
|
||||
// emptySchema is what a tool with no declared schema publishes.
|
||||
//
|
||||
// tools/list requires an inputSchema per tool, and a client given `null` may
|
||||
// reasonably refuse the whole list. An object that accepts nothing is the
|
||||
// honest rendering of "this tool takes no arguments".
|
||||
func emptySchema() map[string]any {
|
||||
return map[string]any{
|
||||
"type": "object",
|
||||
"properties": map[string]any{},
|
||||
"additionalProperties": false,
|
||||
}
|
||||
}
|
||||
|
||||
// toMCPTool converts a registry entry into its wire form.
|
||||
func toMCPTool(info tools.ToolInfo) mcpTool {
|
||||
schema := info.InputSchema
|
||||
if schema == nil {
|
||||
schema = emptySchema()
|
||||
}
|
||||
return mcpTool{
|
||||
Name: info.Name,
|
||||
Description: info.Description,
|
||||
InputSchema: schema,
|
||||
Annotations: &annotations{
|
||||
ReadOnlyHint: info.Effect == string(tools.EffectRead),
|
||||
DestructiveHint: info.Effect == string(tools.EffectWrite),
|
||||
},
|
||||
}
|
||||
}
|
||||
216
go-api/internal/mcpserver/transport.go
Normal file
216
go-api/internal/mcpserver/transport.go
Normal file
@@ -0,0 +1,216 @@
|
||||
package mcpserver
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"io"
|
||||
"net/http"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/krow/krow-backend/go-api/internal/authctx"
|
||||
)
|
||||
|
||||
// The Streamable HTTP binding: one endpoint, one message per POST.
|
||||
//
|
||||
// The current MCP spec defines two standard transports — stdio and Streamable
|
||||
// HTTP — and only the second can serve a client that is not launching a
|
||||
// subprocess. Each message is an HTTP POST to a single MCP endpoint, and the
|
||||
// reply is a JSON object or a request-scoped SSE stream.
|
||||
//
|
||||
// This implementation answers with JSON objects and no stream, which is a
|
||||
// complete implementation of the binding for this server's methods rather than
|
||||
// a shortcut: tools/list is a fixed list, and a tool result is structured data
|
||||
// already bounded by MaxResultBytes. There is nothing to deliver incrementally.
|
||||
// Streaming becomes worth adding if a long-running method is ever exposed.
|
||||
//
|
||||
// STATELESS. No session is minted and no Mcp-Session-Id is required, because
|
||||
// nothing here is worth remembering between calls: every request carries its own
|
||||
// bearer credential and every method is independent. Adding session
|
||||
// state now would be state to expire, to bind to a token, to revalidate and to
|
||||
// leak — for no behaviour this server has.
|
||||
|
||||
// MaxRequestBytes bounds an inbound MCP message.
|
||||
//
|
||||
// Well above any legitimate tools/call — arguments are a few scalars — and far
|
||||
// below anything that would be worth sending here. Enforced with
|
||||
// http.MaxBytesReader so the body is refused as it arrives rather than after it
|
||||
// has been buffered.
|
||||
const MaxRequestBytes = 1 << 20 // 1 MiB
|
||||
|
||||
// Handler returns the HTTP handler for the MCP endpoint.
|
||||
//
|
||||
// Deliberately an http.Handler rather than a registered route: this package
|
||||
// does not know its own path, and the server that mounts it decides where it
|
||||
// lives. Note that it does NOT decide what authenticates it — this handler
|
||||
// authenticates its own callers from the Authorization header, and ignores
|
||||
// whatever middleware sits in front. Mounting it behind the cookie middleware
|
||||
// therefore does not make a cookie sufficient to call it.
|
||||
func (s *Server) Handler() http.Handler {
|
||||
return http.HandlerFunc(s.serveHTTP)
|
||||
}
|
||||
|
||||
func (s *Server) serveHTTP(w http.ResponseWriter, r *http.Request) {
|
||||
// One method. GET is what the binding uses for a server-initiated stream,
|
||||
// which this server does not open; saying 405 with an Allow header is more
|
||||
// use to a client than a 404 that suggests the endpoint is absent.
|
||||
if r.Method != http.MethodPost {
|
||||
w.Header().Set("Allow", http.MethodPost)
|
||||
writeRPCError(w, s, http.StatusMethodNotAllowed, nil,
|
||||
errInvalidRequest("this endpoint accepts POST only"))
|
||||
return
|
||||
}
|
||||
|
||||
// Content-Type is checked rather than assumed. A form post or a stray
|
||||
// upload that happened to be valid JSON would otherwise be processed as a
|
||||
// protocol message.
|
||||
if ct := r.Header.Get("Content-Type"); ct != "" && !isJSONContentType(ct) {
|
||||
writeRPCError(w, s, http.StatusUnsupportedMediaType, nil,
|
||||
errInvalidRequest("Content-Type must be application/json"))
|
||||
return
|
||||
}
|
||||
|
||||
body, err := io.ReadAll(http.MaxBytesReader(w, r.Body, MaxRequestBytes))
|
||||
if err != nil {
|
||||
// MaxBytesReader's error is indistinguishable from a truncated upload
|
||||
// without type assertions that buy nothing here: both mean the body is
|
||||
// unusable, and 413 is the more actionable of the two answers.
|
||||
writeRPCError(w, s, http.StatusRequestEntityTooLarge, nil,
|
||||
errInvalidRequest("the request body was too large or could not be read"))
|
||||
return
|
||||
}
|
||||
|
||||
req, rpcErr := parseRequest(body)
|
||||
if rpcErr != nil {
|
||||
// A parse failure has no usable id, so the response carries the id the
|
||||
// message did parse with — null when it parsed with none. HTTP stays
|
||||
// 200: the transport succeeded and the JSON-RPC error IS the answer.
|
||||
writeRPCError(w, s, http.StatusOK, req.ID, rpcErr)
|
||||
return
|
||||
}
|
||||
|
||||
// Authentication, for every method including the handshake — see
|
||||
// methodRequiresAuth. The 401 below is not merely a refusal: its
|
||||
// WWW-Authenticate header is the first step of the MCP authorization flow,
|
||||
// and a client's very first request is what should produce it.
|
||||
var ident *authctx.Identity
|
||||
if resolved, err := s.authenticate(r); err == nil {
|
||||
ident = &resolved
|
||||
} else if methodRequiresAuth(req.Method) {
|
||||
s.log.Warn("mcp request refused",
|
||||
"method", req.Method, "reason", authFailureReason(err))
|
||||
// WWW-Authenticate is not decoration: RFC 9728 has the client read the
|
||||
// resource-metadata URL from this header to find the authorization
|
||||
// server, and from there where to get a token. It is the difference
|
||||
// between "authentication failed" and "here is how to authenticate".
|
||||
w.Header().Set("WWW-Authenticate", s.challenge())
|
||||
writeRPCError(w, s, http.StatusUnauthorized, req.ID,
|
||||
&rpcError{Code: codeUnauthorized, Message: "authentication required"})
|
||||
return
|
||||
}
|
||||
|
||||
// A panic in a handler must not take the process down or leak a stack into
|
||||
// the response. The registry has its own recover around each tool; this is
|
||||
// the outer net for everything else in this package.
|
||||
result, rpcErr := s.handleRecovered(r, ident, req)
|
||||
|
||||
// A notification gets no response body at all — the spec forbids one.
|
||||
if req.isNotification() {
|
||||
w.WriteHeader(http.StatusAccepted)
|
||||
return
|
||||
}
|
||||
if rpcErr != nil {
|
||||
// A rate-limited refusal is the one JSON-RPC error that also carries an
|
||||
// HTTP status, because 429 and Retry-After are how a client knows to
|
||||
// back off. Everything else is 200 with an error body: the transport
|
||||
// succeeded and the error IS the answer.
|
||||
if rpcErr.Code == codeRateLimited {
|
||||
if rpcErr.retryAfter > 0 {
|
||||
w.Header().Set("Retry-After", retryAfterSeconds(rpcErr.retryAfter))
|
||||
}
|
||||
writeRPCError(w, s, http.StatusTooManyRequests, req.ID, rpcErr)
|
||||
return
|
||||
}
|
||||
writeRPCError(w, s, http.StatusOK, req.ID, rpcErr)
|
||||
return
|
||||
}
|
||||
writeJSON(w, s, http.StatusOK, response{
|
||||
JSONRPC: jsonRPCVersion,
|
||||
ID: req.ID,
|
||||
Result: result,
|
||||
})
|
||||
}
|
||||
|
||||
// methodRequiresAuth reports whether a method may only run for a known caller.
|
||||
//
|
||||
// EVERY method does, including the handshake. An earlier revision left
|
||||
// initialize, ping and notifications/initialized open, on the reasoning that a
|
||||
// client needs somewhere to start — and that was wrong in a way worth
|
||||
// recording, because it is the kind of mistake that looks like helpfulness.
|
||||
//
|
||||
// The MCP authorization flow begins with the client making an MCP request
|
||||
// WITHOUT a token and reading the 401's WWW-Authenticate header. A client's
|
||||
// first request is usually initialize. Answering that one with a cheerful 200
|
||||
// means the client never sees the challenge, believes it is connected, and
|
||||
// discovers otherwise only when the first real call fails — by which time it
|
||||
// has no 401 in hand to discover from. Requiring a token everywhere means the
|
||||
// very first request, whatever it is, produces the challenge that starts the
|
||||
// flow.
|
||||
//
|
||||
// Nothing is lost. The handshake is not information a stranger needs: it
|
||||
// returns this server's name and capabilities, which are only useful to a
|
||||
// client that intends to authenticate anyway.
|
||||
//
|
||||
// Kept as a function rather than inlined because it is the single place that
|
||||
// decision lives, and a future method that genuinely must be open should have
|
||||
// to be written down here to become so.
|
||||
func methodRequiresAuth(method string) bool { return true }
|
||||
|
||||
// handleRecovered runs Handle with a recover, converting a panic into an
|
||||
// internal error whose detail goes to the log and not to the caller.
|
||||
func (s *Server) handleRecovered(r *http.Request, ident *authctx.Identity, req request) (result any, rpcErr *rpcError) {
|
||||
defer func() {
|
||||
if p := recover(); p != nil {
|
||||
s.log.Error("mcp handler panicked", "method", req.Method, "panic", p)
|
||||
result, rpcErr = nil, errInternal()
|
||||
}
|
||||
}()
|
||||
return s.Handle(r.Context(), ident, req)
|
||||
}
|
||||
|
||||
// retryAfterSeconds renders a duration for the Retry-After header, rounded up
|
||||
// and never below one second — "Retry-After: 0" invites an immediate retry,
|
||||
// which is the one thing a limited client must not do.
|
||||
func retryAfterSeconds(d time.Duration) string {
|
||||
secs := int(d.Round(time.Second) / time.Second)
|
||||
if secs < 1 {
|
||||
secs = 1
|
||||
}
|
||||
return strconv.Itoa(secs)
|
||||
}
|
||||
|
||||
// isJSONContentType reports whether a Content-Type header names JSON,
|
||||
// tolerating parameters such as "; charset=utf-8".
|
||||
func isJSONContentType(ct string) bool {
|
||||
media := strings.TrimSpace(strings.SplitN(ct, ";", 2)[0])
|
||||
return strings.EqualFold(media, "application/json")
|
||||
}
|
||||
|
||||
func writeRPCError(w http.ResponseWriter, s *Server, status int, id json.RawMessage, e *rpcError) {
|
||||
writeJSON(w, s, status, response{JSONRPC: jsonRPCVersion, ID: id, Error: e})
|
||||
}
|
||||
|
||||
func writeJSON(w http.ResponseWriter, s *Server, status int, payload response) {
|
||||
encoded, err := json.Marshal(payload)
|
||||
if err != nil {
|
||||
// Encoding our own response failed, so there is nothing safe left to
|
||||
// say in JSON. Log it and send a bare 500.
|
||||
s.log.Error("mcp response could not be encoded", "error", err)
|
||||
http.Error(w, "internal error", http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
w.Header().Set("Content-Type", "application/json; charset=utf-8")
|
||||
w.Header().Set("Cache-Control", "no-store")
|
||||
w.WriteHeader(status)
|
||||
_, _ = w.Write(encoded)
|
||||
}
|
||||
Reference in New Issue
Block a user