diff --git a/.env.example b/.env.example index 851b09a..20417fb 100644 --- a/.env.example +++ b/.env.example @@ -31,6 +31,25 @@ HTTP_SHUTDOWN_TIMEOUT=10s # Origins are matched exactly, echoed back one at a time, and "*" is rejected. # HTTP_CORS_ORIGINS=http://localhost:5173,http://127.0.0.1:5173 +# Networks whose X-Forwarded-For header may be believed. Comma-separated CIDR +# blocks or bare addresses; both IP families accepted. +# +# Several limits are keyed by the caller's address: failed logins, OAuth client +# registration, and OAuth authorization before sign-in. Behind a reverse proxy +# every request arrives FROM the proxy, so without this setting those budgets +# describe the proxy rather than the caller and every user shares one — one +# person retrying a connector exhausts everybody's allowance. +# +# Unset means no proxy is trusted and the header is ignored entirely, which is +# correct for local development: nothing sits in front of the dev server. Leave +# it unset here. A misspelt value cannot open a hole — it only restores the +# shared bucket — but a malformed entry stops startup rather than being dropped. +# +# NEVER set this to 0.0.0.0/0. That trusts every caller's own header, which is +# not a weaker limit but no limit at all: anyone could mint a fresh budget per +# request simply by changing the value they send. +# HTTP_TRUSTED_PROXIES= + # ── PostgreSQL ────────────────────────────────────────────────────────────── # The local development database. DATABASE_NAME is mixed-case and hyphenated, # so anything that interpolates it into SQL must quote it: "Krow-force". diff --git a/go-api/cmd/api/main.go b/go-api/cmd/api/main.go index 809afcc..4d713f4 100644 --- a/go-api/cmd/api/main.go +++ b/go-api/cmd/api/main.go @@ -65,10 +65,14 @@ func run() error { return err } - // The sweeper's context is cancelled by the same signal that stops the - // server, so the ticker goes away with the process rather than outliving - // the pool it queries. + // The sweepers' context is cancelled by the same signal that stops the + // server, so the tickers go away with the process rather than outliving + // the pool they query. go sweepSessions(ctx, server.Sessions(), log) + // OAuth codes and tokens, and the rate-limit counters. Returns immediately + // when the deployment does not serve MCP, so this line costs an unconfigured + // deployment one nil check at startup and nothing after. + go httpserver.SweepMaintenance(ctx, server.Maintenance(), log) errCh := make(chan error, 1) go func() { errCh <- server.Start() }() diff --git a/go-api/internal/config/config.go b/go-api/internal/config/config.go index fe8a499..8950fe2 100644 --- a/go-api/internal/config/config.go +++ b/go-api/internal/config/config.go @@ -12,6 +12,7 @@ package config import ( "fmt" + "net/netip" "net/url" "os" "strconv" @@ -58,8 +59,44 @@ type Config struct { Agents AgentsConfig Model ModelConfig Knowledge KnowledgeConfig + OAuth OAuthConfig } +// OAuthConfig is the MCP surface's OAuth 2.1 identity. +// +// EMPTY IS THE DEFAULT AND IT MEANS "OFF". A deployment that sets neither +// OAUTH_ISSUER nor MCP_RESOURCE does not serve OAuth or MCP at all, and that is +// the correct default for every deployment that exists today — the routes are +// simply not registered, exactly as routeRuns is skipped without a model +// credential. +// +// NO PRODUCTION DOMAIN IS HARDCODED. Both values are URLs the operator supplies, +// because the issuer identifies the deployment and a default would be one +// deployment's identity baked into every other one. +// +// Issuer and Resource look similar and are not the same thing: the ISSUER +// identifies the authorization server ("who minted this token"), the RESOURCE +// identifies what the token is good for ("which MCP server may spend it"). A +// token's audience is checked against Resource. Conflating them is how a token +// for one service becomes spendable at another. +type OAuthConfig struct { + // Issuer is the authorization server's base URL, e.g. + // https://api.example.com. No trailing slash. + Issuer string + + // Resource is the canonical MCP endpoint URI, e.g. + // https://api.example.com/mcp. This becomes an issued token's audience. + Resource string + + // LoginPath is where the authorization endpoint sends somebody who is not + // signed in. A same-origin path, never an absolute URL — an absolute one + // would be an open redirect waiting for a misconfiguration. + LoginPath string +} + +// Enabled reports whether this deployment serves OAuth and MCP. +func (c OAuthConfig) Enabled() bool { return c.Issuer != "" && c.Resource != "" } + // KnowledgeConfig routes the retrieval layer's embedding provider. // // The chat provider does not serve embeddings, so the dense half of hybrid @@ -213,6 +250,42 @@ type HTTPConfig struct { // this API, and once authentication exists that becomes a real hole rather // than a theoretical one. CORSOrigins []string + + // TrustedProxies are the networks a forwarded client address may be + // believed from. Empty by default, and empty means "believe nothing". + // + // WHY THIS EXISTS + // + // Several limits on this API are keyed by the caller's network address: + // failed logins, OAuth registration, and OAuth authorization before the + // caller has signed in. Behind a reverse proxy every request arrives from + // the proxy, so RemoteAddr is one constant value and those per-address + // budgets silently become one budget for the entire deployment. The + // symptom is users rate-limiting each other — one person retrying a + // connector exhausts everybody's allowance. + // + // WHY IT IS NOT SIMPLY "READ X-FORWARDED-FOR" + // + // That header is client-supplied. A caller reaching the API directly can + // invent one and mint a fresh budget per request, which is strictly worse + // than sharing a bucket: it removes the limit entirely. The header is + // meaningful only when the immediate peer is a proxy that is known to + // rewrite it, which is what this list names. + // + // WHY THE DEFAULT IS EMPTY + // + // So that a missing or misspelt setting cannot open the spoofing hole. An + // unconfigured deployment behaves exactly as it did before this setting + // existed: RemoteAddr, and X-Forwarded-For ignored. The failure mode of + // forgetting to set it is the old shared bucket, which is an availability + // problem an operator will notice, rather than an unmetered endpoint which + // they will not. + // + // Entries are CIDR blocks or bare addresses (a bare address is treated as + // a single-host block). Both families are accepted. Set it to the network + // the load balancer or ingress talks to the API from — see + // .env.example and infrastructure/.env.docker.example. + TrustedProxies []netip.Prefix } type DBConfig struct { @@ -278,6 +351,12 @@ func Load() (*Config, error) { return v } + // Parsed before the literal below because it can fail, and a malformed + // entry has to stop startup rather than be dropped: an operator who + // mistypes the proxy network gets the shared-bucket behaviour back, and + // silently is the one way they will not find out. + trustedProxies, trustedProxiesErr := parseTrustedProxies(os.Getenv("HTTP_TRUSTED_PROXIES")) + cfg := &Config{ AppEnv: withDefault("APP_ENV", "development"), Log: LogConfig{Level: withDefault("LOG_LEVEL", "info")}, @@ -293,6 +372,7 @@ func Load() (*Config, error) { // the server derive it from the CORS posture", and an explicit value // overrides that derivation. See Server.sessionSameSite. CookieSameSite: strings.ToLower(strings.TrimSpace(os.Getenv("HTTP_COOKIE_SAMESITE"))), + TrustedProxies: trustedProxies, }, Seed: SeedConfig{ FixturePath: withDefault("SEED_FIXTURE_PATH", "./seed/fixtures/seed.json"), @@ -300,6 +380,14 @@ func Load() (*Config, error) { Agents: AgentsConfig{ CuratedPath: withDefault("CURATED_AGENTS_PATH", "./agents"), }, + OAuth: OAuthConfig{ + // Trailing slashes trimmed here rather than at every use: the + // canonical form of a resource URI has none, and a token minted + // against ".../mcp/" would fail to validate against ".../mcp". + Issuer: strings.TrimRight(strings.TrimSpace(os.Getenv("OAUTH_ISSUER")), "/"), + Resource: strings.TrimRight(strings.TrimSpace(os.Getenv("MCP_RESOURCE")), "/"), + LoginPath: withDefault("OAUTH_LOGIN_PATH", "/login"), + }, Knowledge: KnowledgeConfig{ EmbedProvider: strings.ToLower(strings.TrimSpace(os.Getenv("EMBED_PROVIDER"))), EmbedAPIKey: strings.TrimSpace(os.Getenv("VOYAGE_API_KEY")), @@ -347,6 +435,9 @@ func Load() (*Config, error) { }, } + if trustedProxiesErr != nil { + return nil, trustedProxiesErr + } if len(missing) > 0 { return nil, fmt.Errorf("missing required environment variables: %s "+ "(copy .env.example to .env and fill them in)", strings.Join(missing, ", ")) @@ -531,6 +622,9 @@ func (c *Config) validate() error { if c.Knowledge.EmbedDims < 0 { return fmt.Errorf("EMBED_DIMENSIONS cannot be negative, got %d", c.Knowledge.EmbedDims) } + if err := c.validateOAuth(); err != nil { + return err + } for name, model := range map[string]string{ "MODEL_FAST": c.Model.Fast, "MODEL_BALANCED": c.Model.Balanced, "MODEL_DEEP": c.Model.Deep, } { @@ -621,6 +715,49 @@ func corsOrigins(appEnv string) []string { return out } +// parseTrustedProxies reads HTTP_TRUSTED_PROXIES, a comma-separated list of +// CIDR blocks or bare addresses. +// +// Unset or empty yields nil, which means no proxy is trusted and forwarded +// client addresses are ignored entirely. That is the safe default and the +// behaviour this API had before the setting existed. +// +// A bare address is accepted and widened to a single-host prefix, because +// "10.0.0.7" is what an operator reaches for when there is exactly one ingress +// and requiring them to write "10.0.0.7/32" only invites a mistake. +// +// Malformed entries are an error rather than a skip. Skipping one would leave +// the deployment quietly trusting a shorter list than the operator wrote, and +// the consequence — a proxy that is not believed, so every user shares one +// rate-limit bucket again — is precisely the fault this setting exists to fix. +func parseTrustedProxies(raw string) ([]netip.Prefix, error) { + var out []netip.Prefix + for _, part := range strings.Split(raw, ",") { + entry := strings.TrimSpace(part) + if entry == "" { + continue + } + if prefix, err := netip.ParsePrefix(entry); err == nil { + // Masked so that a block written with host bits set — 10.0.0.7/8, + // which is easy to write and easy to misread — still contains what + // its author meant. Unmasked, Prefix.Contains always reports false. + out = append(out, prefix.Masked()) + continue + } + addr, err := netip.ParseAddr(entry) + if err != nil { + return nil, fmt.Errorf("HTTP_TRUSTED_PROXIES entry %q is not an IP address "+ + "or CIDR block (for example 10.0.0.0/8, 172.17.0.1 or fd00::/8)", entry) + } + // Unmap first: ::ffff:10.0.0.1 and 10.0.0.1 are the same host, and a + // /128 around the mapped form would not match the peer address Go + // reports for an IPv4 connection. + addr = addr.Unmap() + out = append(out, netip.PrefixFrom(addr, addr.BitLen())) + } + return out, nil +} + func withDefault(key, fallback string) string { if v := strings.TrimSpace(os.Getenv(key)); v != "" { return v @@ -764,3 +901,45 @@ func applyDotEnv(content string) { } } } + +// validateOAuth checks the MCP surface's OAuth identity. +// +// Both values empty is the ordinary case and means the surface is off. Setting +// exactly one is always a mistake — a deployment that named an issuer but no +// resource would serve discovery documents pointing at a resource that does not +// exist — so it is refused at boot rather than at the first client connection. +func (c *Config) validateOAuth() error { + issuer, resource := c.OAuth.Issuer, c.OAuth.Resource + if issuer == "" && resource == "" { + return nil + } + if issuer == "" || resource == "" { + return fmt.Errorf("OAUTH_ISSUER and MCP_RESOURCE must be set together; " + + "one without the other serves discovery documents that point nowhere") + } + + for name, raw := range map[string]string{"OAUTH_ISSUER": issuer, "MCP_RESOURCE": resource} { + parsed, err := url.Parse(raw) + if err != nil || parsed.Host == "" { + return fmt.Errorf("%s must be an absolute URL, got %q", name, raw) + } + // HTTPS everywhere except a loopback development host. OAuth 2.1 + // requires every authorization server endpoint to be served over + // HTTPS; a token or code sent over plain http is a token on the wire. + if parsed.Scheme != "https" && !isLoopback(raw) { + return fmt.Errorf("%s must use https (http is permitted only on loopback), got %q", name, raw) + } + if parsed.Fragment != "" { + return fmt.Errorf("%s must not contain a fragment, got %q", name, raw) + } + } + + // A same-origin path, never an absolute URL: the authorization endpoint + // redirects here, and an operator-supplied absolute URL would be an open + // redirect one config mistake away. + if !strings.HasPrefix(c.OAuth.LoginPath, "/") || strings.HasPrefix(c.OAuth.LoginPath, "//") { + return fmt.Errorf("OAUTH_LOGIN_PATH must be a same-origin path beginning with a single '/', got %q", + c.OAuth.LoginPath) + } + return nil +} diff --git a/go-api/internal/config/trustedproxies_test.go b/go-api/internal/config/trustedproxies_test.go new file mode 100644 index 0000000..d9846a7 --- /dev/null +++ b/go-api/internal/config/trustedproxies_test.go @@ -0,0 +1,134 @@ +package config + +// HTTP_TRUSTED_PROXIES parsing. +// +// The setting decides whether a client-supplied header is believed, so the +// tests worth having are about what happens when it is WRONG: unset, empty, +// mistyped. Every one of those must end in "trust nothing", because the +// alternative — trusting something the operator did not write — is the whole +// risk this setting carries. + +import ( + "net/netip" + "testing" +) + +func TestTrustedProxiesUnsetTrustsNothing(t *testing.T) { + for _, raw := range []string{"", " ", ",", " , , "} { + got, err := parseTrustedProxies(raw) + if err != nil { + t.Errorf("parseTrustedProxies(%q): unexpected error %v", raw, err) + } + if len(got) != 0 { + t.Errorf("parseTrustedProxies(%q) = %v, want empty — an unset value must trust nothing", raw, got) + } + } +} + +func TestTrustedProxiesParsesCIDRsAndBareAddresses(t *testing.T) { + got, err := parseTrustedProxies(" 10.0.0.0/8 , 172.17.0.1 , fd00::/8 , ::1 ") + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + + want := []string{"10.0.0.0/8", "172.17.0.1/32", "fd00::/8", "::1/128"} + if len(got) != len(want) { + t.Fatalf("parsed %d entries (%v), want %d", len(got), got, len(want)) + } + for i, w := range want { + if got[i].String() != w { + t.Errorf("entry %d = %q, want %q", i, got[i].String(), w) + } + } +} + +// A bare address must become a single-host block that contains that host and +// nothing else — the operator wrote one proxy, not a network. +func TestTrustedProxyBareAddressIsOneHost(t *testing.T) { + got, err := parseTrustedProxies("172.17.0.1") + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if !got[0].Contains(netip.MustParseAddr("172.17.0.1")) { + t.Error("the host itself is not in its own single-host block") + } + if got[0].Contains(netip.MustParseAddr("172.17.0.2")) { + t.Error("a bare address was widened beyond one host") + } +} + +// A block written with host bits set is common and easy to misread. Masking it +// at parse time makes it mean what its author meant; unmasked, netip.Prefix +// .Contains reports false for everything. +func TestTrustedProxyCIDRWithHostBitsIsMasked(t *testing.T) { + got, err := parseTrustedProxies("10.1.2.3/8") + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if want := "10.0.0.0/8"; got[0].String() != want { + t.Fatalf("got %q, want %q", got[0].String(), want) + } + if !got[0].Contains(netip.MustParseAddr("10.9.9.9")) { + t.Error("the masked block does not contain an address inside it") + } +} + +// An IPv4-mapped address names an IPv4 host, and must match the peer address +// Go reports for an IPv4 connection. +func TestTrustedProxyIPv4MappedIsUnmapped(t *testing.T) { + got, err := parseTrustedProxies("::ffff:10.0.0.1") + if err != nil { + t.Fatalf("unexpected error: %v", err) + } + if !got[0].Contains(netip.MustParseAddr("10.0.0.1")) { + t.Errorf("%q does not contain 10.0.0.1", got[0].String()) + } +} + +// Malformed entries stop startup. Skipping one would leave the deployment +// trusting a shorter list than the operator wrote, and the consequence — every +// user sharing one rate-limit bucket — is silent. +func TestTrustedProxiesRejectMalformedEntries(t *testing.T) { + for _, raw := range []string{ + "banana", + "10.0.0.0/33", + "10.0.0.0/8, banana", + "300.1.2.3", + "10.0.0.1:8080", + "*", + "https://proxy.internal", + "fd00::/200", + } { + if _, err := parseTrustedProxies(raw); err == nil { + t.Errorf("parseTrustedProxies(%q) was accepted; it must refuse and stop startup", raw) + } + } +} + +// The error has to name the entry and show the shape expected, because it is +// read by an operator at 3am with a container that will not boot. +func TestTrustedProxiesErrorNamesTheEntry(t *testing.T) { + _, err := parseTrustedProxies("10.0.0.0/8, banana") + if err == nil { + t.Fatal("expected an error") + } + for _, want := range []string{"HTTP_TRUSTED_PROXIES", "banana"} { + if !contains(err.Error(), want) { + t.Errorf("error %q does not mention %q", err.Error(), want) + } + } +} + +func contains(haystack, needle string) bool { + return len(haystack) >= len(needle) && (haystack == needle || + len(needle) == 0 || indexOf(haystack, needle) >= 0) +} + +func indexOf(haystack, needle string) int { + for i := 0; i+len(needle) <= len(haystack); i++ { + if haystack[i:i+len(needle)] == needle { + return i + } + } + return -1 +} diff --git a/go-api/internal/domain/definitions_schema_test.go b/go-api/internal/domain/definitions_schema_test.go index f875dca..6961512 100644 --- a/go-api/internal/domain/definitions_schema_test.go +++ b/go-api/internal/domain/definitions_schema_test.go @@ -724,6 +724,15 @@ func TestMigrationPairsAreComplete(t *testing.T) { "000009_confirmation_replay.up.sql", "000010_definition_versions.up.sql", "000011_employee_roles.up.sql", + // Phase 3: the OAuth 2.1 authorization server behind the MCP surface. + // Three tables, added together because they are one feature: a client + // registers, is issued a code, and exchanges it for tokens. + "000012_oauth_clients.up.sql", + "000013_oauth_grants.up.sql", + "000014_oauth_tokens.up.sql", + // Phase 5: shared rate limit counters, so a limit means the same thing + // behind one instance and behind ten. + "000015_rate_limits.up.sql", } if len(ups) != len(want) { t.Fatalf("%d migrations, want %d — update this list deliberately", len(ups), len(want)) @@ -753,11 +762,28 @@ func TestMigrationsAddOnlyTheTablesWeDecidedOn(t *testing.T) { // 17 from 000001, + auth_sessions (000004), + agent_definitions and // skill_definitions (000005), + agent_runs (000006), + agent_confirmations // (000007), + knowledge_documents and knowledge_chunks (000008), - // + definition_versions (000010), + employee_roles (000011). + // + definition_versions (000010), + employee_roles (000011), + // + oauth_clients (000012), + oauth_grants (000013), + oauth_tokens + // (000014), + rate_limits (000015). // schema_migrations is golang-migrate's and is absent when the files are // applied directly. - if n != 26 { - t.Errorf("%d base tables after every migration, want 26", n) + if n != 30 { + t.Errorf("%d base tables after every migration, want 30", n) + } + + // The three OAuth tables, named rather than merely counted. The count + // above catches a table arriving without a decision; this catches one of + // these three going missing, which the count alone would not if another + // arrived in the same change. + for _, required := range []string{"oauth_clients", "oauth_grants", "oauth_tokens", "rate_limits"} { + var reg *string + if err := f.pool.QueryRow(f.ctx, + `SELECT to_regclass('public.' || $1)::text`, required).Scan(®); err != nil { + t.Fatalf("check %s: %v", required, err) + } + if reg == nil { + t.Errorf("%s is missing; the MCP OAuth surface cannot work without it", required) + } } // `definition_versions` was on this list, deferred by the Phase 4B decision. diff --git a/go-api/internal/httpserver/auth.go b/go-api/internal/httpserver/auth.go index 71f0d46..1824cbe 100644 --- a/go-api/internal/httpserver/auth.go +++ b/go-api/internal/httpserver/auth.go @@ -206,7 +206,7 @@ func (s *Server) handleLogin(w http.ResponseWriter, r *http.Request) { // which is wider, stops one host working through many accounts. They are // separate limiters because they are deliberately different sizes — see the // note on Server. - addr := clientAddr(r) + addr := s.trust.clientAddr(r) emailKey := strings.ToLower(email) for _, check := range []struct { limiter *attemptLimiter @@ -318,6 +318,85 @@ var publicPaths = map[string]bool{ "/health": true, "/api/v1/auth/login": true, "/api/v1/auth/logout": true, + + // ── The OAuth surface for MCP clients ────────────────────────────────── + // + // Four paths, each public for a specific reason rather than because + // "/oauth/*" is convenient. The namespace is deliberately NOT wildcarded: + // /oauth/authorize is not here, because it renders a consent screen for a + // signed-in person and must keep requiring a session. + // + // These routes are registered only when OAUTH_ISSUER and MCP_RESOURCE are + // configured. Listing them here is harmless otherwise — an unregistered + // path still 404s, it simply does so without being asked for a cookie. + + // RFC 9728 and RFC 8414. A client with no token cannot read a document + // that requires one, and these are how it discovers where to get a token. + // They contain public endpoint URLs and nothing else. + "/.well-known/oauth-protected-resource": true, + "/.well-known/oauth-authorization-server": true, + + // RFC 7591. A client that has never registered has no credential to + // present; that is what dynamic registration is for. + "/oauth/register": true, + + // The client authenticates here with an authorization code or a refresh + // token in the BODY. This is a back-channel call from the MCP client's own + // servers — there is no browser and no cookie to send. + "/oauth/token": true, + + // Revocation authenticates by presenting the token being revoked, for the + // same back-channel reason. + "/oauth/revoke": true, + + // /mcp is listed here, and it is the entry that most deserves explaining, + // because "public" is the opposite of what it means for this path. + // + // The MCP endpoint authenticates its OWN callers, from the Authorization + // header, inside mcpserver — every method but the handshake requires a + // valid bearer token, and the transport ignores whatever identity this + // middleware may have put in the context. So listing it here does not make + // it reachable without a credential; it makes THIS middleware step aside + // so the one that knows how to answer can. + // + // It has to step aside. An MCP client discovers how to authenticate by + // calling the endpoint with no token and reading the WWW-Authenticate + // header of the 401 — RFC 9728, and the first step of the whole flow. + // This middleware's 401 carries no such header, so guarding /mcp here + // would mean a client received a refusal with nowhere to go and the + // connection could never be established. That is not a hypothetical: it is + // what TestMCPWithoutBearerReturns401AndDiscoveryPointer caught. + // + // What stops a cookie authenticating an MCP call is therefore NOT this + // allowlist — it is mcpserver taking its identity as a parameter rather + // than from the request context. See mcpserver/auth.go, and + // TestMCPRejectsACookieSession below. + "/mcp": true, + + // /oauth/authorize is here for the same reason as /mcp, and it took a live + // client to show why. + // + // It was withheld on the reasoning that consent needs a signed-in person, + // so the route "genuinely wants the cookie". That reasoning was right about + // the requirement and wrong about who enforces it. THE HANDLER already + // enforces it — authserver.go asks sessions.CurrentUser, refuses to render + // consent without an identity, and redirects an anonymous visitor to the + // login with the authorization request preserved in returnTo. Guarding the + // path HERE meant that handler was never reached, so the redirect it + // performs could never run: every signed-out visitor got this middleware's + // JSON 401 instead of a login page. + // + // That is not a cosmetic difference. A first-time connector user is signed + // out by definition, so OAuth's browser leg was unreachable for exactly the + // people who needed it. Claude Web stopped here — discovery, registration, + // then a 401 with nowhere to go. Claude Desktop only got past it because a + // session had been established by hand beforehand. + // + // Listing it grants nothing: no session still means no consent screen and + // no authorization code, and the consent POST still requires the + // session-bound CSRF token. What changes is only WHICH layer says no, and + // therefore whether it can say "sign in here" instead of "no". + "/oauth/authorize": true, } // authenticate resolves the session cookie into an identity, or refuses. diff --git a/go-api/internal/httpserver/clientip.go b/go-api/internal/httpserver/clientip.go new file mode 100644 index 0000000..b5c4e81 --- /dev/null +++ b/go-api/internal/httpserver/clientip.go @@ -0,0 +1,217 @@ +package httpserver + +import ( + "net" + "net/http" + "net/netip" + "strings" +) + +// Resolving the caller's network address behind a reverse proxy. +// +// WHAT THIS IS FOR +// +// Three limits on this API are keyed by the caller's address: failed logins +// (auth.go), OAuth client registration, and OAuth authorization before the +// caller has signed in. None of them has a better identity available — +// registration is anonymous by definition, and a login attempt is anonymous +// until the password has been judged. +// +// Behind a proxy, net/http reports the PROXY's address on every request. Those +// three budgets then describe the proxy rather than the caller, which means one +// bucket for the whole deployment: one person retrying a connector exhausts +// everybody's registration allowance, and twenty failed passwords anywhere lock +// out every user's sign-in. That is the fault this file exists to fix. +// +// WHY IT IS NOT JUST X-Forwarded-For +// +// The header is written by clients as readily as by proxies. Believing it +// unconditionally is worse than the shared bucket rather than better: a caller +// who reaches the API directly can put a different value in every request and +// get a fresh budget each time, which is not a weakened limit but no limit at +// all. The header carries information only about the hop that appended it, so +// it is worth exactly as much as the peer that handed it over. +// +// Hence: believe it only when the immediate peer is a configured proxy, and +// walk the chain from the right, where the entries were written by the hops +// closest to us, discarding those that are themselves trusted proxies. The +// first address that is not one of ours is the nearest thing to the real client +// that the topology can actually vouch for. Everything to its left was supplied +// by something we do not control and is never read. +// +// FAILING SAFE +// +// Every fallback in here returns the PEER address. That is deliberate and it is +// the property worth preserving if this code is ever changed: a bad or missing +// chain can only ever make a bucket coarser — more callers sharing one budget, +// which is the old behaviour — and can never hand a caller a bucket of their +// own. Spoofing gains nothing because no path exists from an untrusted input to +// a distinct key. + +// proxyTrust turns a request into the address key used for rate limiting. +// +// A value rather than a package-level variable so that the trusted set is +// wired once at construction and cannot be changed by anything holding a +// request. An empty proxyTrust is valid and trusts nothing. +type proxyTrust struct { + // trusted networks, already masked by config parsing. + trusted []netip.Prefix +} + +// newProxyTrust builds the resolver from configuration. +func newProxyTrust(trusted []netip.Prefix) proxyTrust { + return proxyTrust{trusted: trusted} +} + +// forwardedHeader is the de facto standard, and what Traefik, nginx, Envoy and +// the cloud load balancers all append to. +// +// RFC 7239's `Forwarded:` header is deliberately NOT read. Supporting both +// would mean deciding which wins when they disagree, and an attacker choosing +// the one this code happens to prefer. One header, one meaning. +const forwardedHeader = "X-Forwarded-For" + +// clientAddr returns the rate-limiting key for the caller's address. +// +// The port is stripped: a browser opens a new source port per connection, so +// keying on host:port would give every attempt its own budget and limit nothing +// at all. IPv6 is keyed by /64 — see bucketKey. +func (t proxyTrust) clientAddr(r *http.Request) string { + peer, ok := parseHost(r.RemoteAddr) + if !ok { + // RemoteAddr is not something this code recognises — a test server with + // a synthetic value, or a unix socket. Key by it verbatim, which is + // what this function did before proxies were considered at all. + return strings.TrimSpace(r.RemoteAddr) + } + peerKey := bucketKey(peer) + + // Nothing is trusted, so nothing is read. The common case, and the default. + if len(t.trusted) == 0 || !t.contains(peer) { + return peerKey + } + + if client, ok := t.forwardedClient(r); ok { + return bucketKey(client) + } + return peerKey +} + +// forwardedClient walks the forwarded chain from the right and returns the +// first address that is not one of our own proxies. +// +// It reports false — meaning "fall back to the peer" — for an absent header, a +// chain that is entirely trusted proxies, and a malformed entry. The last of +// those is the interesting one: a chain that cannot be parsed cannot be +// reasoned about, and the safe reading of "10.0.0.1, ???, 10.0.0.2" is that +// everything to the left of the damage is unusable. Skipping the bad entry and +// carrying on would let a caller put anything it likes in the header and have +// this code step over it to reach the value the caller wanted read. +func (t proxyTrust) forwardedClient(r *http.Request) (netip.Addr, bool) { + // Values(), not Get(), because a chain may arrive as several headers as + // well as one comma-separated list; they are the same list in HTTP's terms + // and the rightmost entry of the last header is the most recent hop. + var chain []string + for _, header := range r.Header.Values(forwardedHeader) { + for _, entry := range strings.Split(header, ",") { + chain = append(chain, strings.TrimSpace(entry)) + } + } + + for i := len(chain) - 1; i >= 0; i-- { + entry := chain[i] + if entry == "" { + // A stray comma. Treated as damage rather than skipped, for the + // reason in the doc comment above. + return netip.Addr{}, false + } + addr, ok := parseForwardedAddr(entry) + if !ok { + return netip.Addr{}, false + } + if t.contains(addr) { + // One of ours. Keep walking left, towards the client. + continue + } + return addr, true + } + // Either there was no header, or every hop in it was a trusted proxy and + // none of them recorded a client. Neither tells us who called. + return netip.Addr{}, false +} + +// contains reports whether an address is one of the configured proxies. +func (t proxyTrust) contains(addr netip.Addr) bool { + addr = addr.Unmap() + for _, prefix := range t.trusted { + if prefix.Contains(addr) { + return true + } + } + return false +} + +// bucketKey is the string a rate-limit bucket is keyed by. +// +// IPv4 keys by the exact address, which is what this service has always done +// and what the existing buckets contain. +// +// IPv6 keys by the /64 PREFIX instead. A single customer is routinely delegated +// a whole /64 — often a /56 or shorter — and every address in it is one +// machine's to choose. Keying by the full address would hand one caller +// 18 quintillion budgets, which is a limit in form only. /64 is the smallest +// unit that is reliably one subscriber rather than one interface, so it is the +// narrowest honest key. +func bucketKey(addr netip.Addr) string { + addr = addr.Unmap().WithZone("") // a scope id is local to the host, never a caller identity + if addr.Is4() { + return addr.String() + } + prefix, err := addr.Prefix(64) + if err != nil { + return addr.String() + } + return prefix.String() +} + +// parseHost splits "host:port" and parses the host. +// +// RemoteAddr always carries a port for TCP, but a test server, a unix socket or +// a middleware that rewrote it may not, so a bare address is accepted too. +func parseHost(remoteAddr string) (netip.Addr, bool) { + raw := strings.TrimSpace(remoteAddr) + if raw == "" { + return netip.Addr{}, false + } + if host, _, err := net.SplitHostPort(raw); err == nil { + raw = host + } + addr, err := netip.ParseAddr(strings.Trim(raw, "[]")) + if err != nil { + return netip.Addr{}, false + } + return addr, true +} + +// parseForwardedAddr parses one entry of an X-Forwarded-For chain. +// +// Entries are bare addresses by the header's convention, but a port turns up in +// practice — some proxies append one, and IPv6 is then bracketed. Both forms +// are accepted; anything else is malformed and refused. +// +// "unknown", the obfuscated identifiers RFC 7239 permits, and empty entries are +// all refused rather than skipped: they say the chain is not a list of +// addresses, and this code declines to guess which of the remaining entries the +// proxy meant. +func parseForwardedAddr(entry string) (netip.Addr, bool) { + if addr, err := netip.ParseAddr(entry); err == nil { + return addr, true + } + // "[2001:db8::1]:443" or "203.0.113.7:443". + if host, _, err := net.SplitHostPort(entry); err == nil { + if addr, err := netip.ParseAddr(strings.Trim(host, "[]")); err == nil { + return addr, true + } + } + return netip.Addr{}, false +} diff --git a/go-api/internal/httpserver/clientip_test.go b/go-api/internal/httpserver/clientip_test.go new file mode 100644 index 0000000..a9d62f1 --- /dev/null +++ b/go-api/internal/httpserver/clientip_test.go @@ -0,0 +1,358 @@ +package httpserver + +// Unit tests for client-address resolution. +// +// An INTERNAL test package (httpserver, not httpserver_test) because proxyTrust +// is unexported and deliberately so — the trusted set is wired once at server +// construction and there is no reason for anything outside this package to +// build one. The rest of the package's tests stay external; this file is the +// exception because what is under test is a decision procedure, and testing it +// through an HTTP server would obscure which input produced which key. +// +// THE PROPERTY THESE TESTS EXIST TO DEFEND +// +// No untrusted input may produce a distinct bucket key. Every failure path must +// collapse back to the peer address. A test that asserts a spoofed header is +// "ignored" by checking it does not appear is not enough — it must check the +// key equals the PEER's key, because two different wrong answers are still two +// different buckets, and two buckets is the whole exploit. + +import ( + "net/http" + "net/netip" + "testing" +) + +func prefixes(t *testing.T, cidrs ...string) []netip.Prefix { + t.Helper() + out := make([]netip.Prefix, 0, len(cidrs)) + for _, c := range cidrs { + p, err := netip.ParsePrefix(c) + if err != nil { + t.Fatalf("bad test CIDR %q: %v", c, err) + } + out = append(out, p.Masked()) + } + return out +} + +// request builds a request with a peer address and an optional forwarded chain. +// A chain entry of "" means the header is absent. +func request(remoteAddr string, forwarded ...string) *http.Request { + r := &http.Request{ + RemoteAddr: remoteAddr, + Header: http.Header{}, + } + for _, f := range forwarded { + r.Header.Add(forwardedHeader, f) + } + return r +} + +/* ── A. A direct client's forwarded header is not read ──────────────────── */ + +func TestDirectClientForwardedHeaderIgnored(t *testing.T) { + // A proxy IS configured — just not this caller. The caller reaches the API + // directly and claims to be somebody else. + trust := newProxyTrust(prefixes(t, "10.0.0.0/8")) + + got := trust.clientAddr(request("203.0.113.9:51000", "198.51.100.7")) + + if want := "203.0.113.9"; got != want { + t.Errorf("clientAddr = %q, want %q — a direct caller's X-Forwarded-For was believed", got, want) + } +} + +func TestNoTrustedProxiesConfiguredIgnoresForwarded(t *testing.T) { + // The default posture. Nothing is trusted, so nothing is read, and the + // behaviour is exactly what it was before this setting existed. + trust := newProxyTrust(nil) + + got := trust.clientAddr(request("10.0.0.1:4000", "198.51.100.7")) + + if want := "10.0.0.1"; got != want { + t.Errorf("clientAddr = %q, want %q — an unconfigured deployment read a forwarded address", got, want) + } +} + +/* ── B. A trusted proxy's forwarded client is used ──────────────────────── */ + +func TestTrustedProxyForwardedClientUsed(t *testing.T) { + trust := newProxyTrust(prefixes(t, "10.0.0.0/8")) + + got := trust.clientAddr(request("10.0.0.1:4000", "203.0.113.9")) + + if want := "203.0.113.9"; got != want { + t.Errorf("clientAddr = %q, want %q", got, want) + } +} + +// The point of the whole change: two users behind the same proxy get two keys. +func TestTrustedProxySeparatesTwoClients(t *testing.T) { + trust := newProxyTrust(prefixes(t, "10.0.0.0/8")) + + a := trust.clientAddr(request("10.0.0.1:4000", "203.0.113.9")) + b := trust.clientAddr(request("10.0.0.1:4001", "203.0.113.10")) + + if a == b { + t.Fatalf("two clients behind one proxy shared the key %q", a) + } +} + +/* ── C. Multiple hops, walked right to left ─────────────────────────────── */ + +func TestMultipleTrustedHopsSelectsFirstUntrusted(t *testing.T) { + trust := newProxyTrust(prefixes(t, "10.0.0.0/8", "172.16.0.0/12")) + + // client → edge(172.16.0.5) → internal(10.0.0.1) → us. + // Right to left: 10.0.0.1 ours, 172.16.0.5 ours, 203.0.113.9 the client. + got := trust.clientAddr(request("10.0.0.1:4000", "203.0.113.9, 172.16.0.5, 10.0.0.1")) + + if want := "203.0.113.9"; got != want { + t.Errorf("clientAddr = %q, want %q", got, want) + } +} + +// The chain split across several headers is the same chain. +func TestChainSplitAcrossHeaders(t *testing.T) { + trust := newProxyTrust(prefixes(t, "10.0.0.0/8")) + + got := trust.clientAddr(request("10.0.0.1:4000", "203.0.113.9", "10.0.0.1")) + + if want := "203.0.113.9"; got != want { + t.Errorf("clientAddr = %q, want %q", got, want) + } +} + +// Entries to the LEFT of the first untrusted address are never read, whatever +// they say. This is what stops a client prepending a forged hop. +func TestEntriesLeftOfTheClientAreNotRead(t *testing.T) { + trust := newProxyTrust(prefixes(t, "10.0.0.0/8")) + + // The caller put "1.2.3.4" at the head of the chain hoping to be keyed by + // it. The proxy appended the address it actually saw. + got := trust.clientAddr(request("10.0.0.1:4000", "1.2.3.4, 203.0.113.9, 10.0.0.1")) + + if want := "203.0.113.9"; got != want { + t.Errorf("clientAddr = %q, want %q — a forged leading hop was selected", got, want) + } +} + +/* ── D. Spoofing gains nothing ──────────────────────────────────────────── */ + +// The exploit this design exists to prevent: an untrusted caller varying the +// header to get a fresh budget per request. Every variation must land on the +// SAME key, and that key must be the peer's. +func TestUntrustedSpoofingCannotProduceDistinctBuckets(t *testing.T) { + trust := newProxyTrust(prefixes(t, "10.0.0.0/8")) + + spoofs := []string{ + "1.2.3.4", + "5.6.7.8", + "10.0.0.1", // claiming to BE the trusted proxy + "1.1.1.1, 2.2.2.2, 10.0.0.1", // a whole fabricated chain ending in ours + "::1", + "2001:db8::1", + } + + const peerKey = "203.0.113.9" + for _, spoof := range spoofs { + got := trust.clientAddr(request("203.0.113.9:51000", spoof)) + if got != peerKey { + t.Errorf("X-Forwarded-For %q produced key %q, want %q — spoofing bought a separate bucket", + spoof, got, peerKey) + } + } +} + +// A trusted proxy that forwards a chain whose leading entries were forged still +// yields one key per real client, not one per forgery. +func TestSpoofedPrefixBehindTrustedProxyIsStable(t *testing.T) { + trust := newProxyTrust(prefixes(t, "10.0.0.0/8")) + + first := trust.clientAddr(request("10.0.0.1:4000", "9.9.9.9, 203.0.113.9, 10.0.0.1")) + second := trust.clientAddr(request("10.0.0.1:4002", "8.8.8.8, 203.0.113.9, 10.0.0.1")) + + if first != second { + t.Errorf("one client produced two keys (%q, %q) by varying a forged hop", first, second) + } +} + +/* ── E. Malformed input falls back, and never panics ────────────────────── */ + +func TestMalformedForwardedEntriesFallBackToPeer(t *testing.T) { + trust := newProxyTrust(prefixes(t, "10.0.0.0/8")) + + cases := map[string]string{ + "not an address": "banana", + "unknown": "unknown", + "obfuscated (7239)": "_hidden", + "empty entry": "203.0.113.9, , 10.0.0.1", + "trailing comma": "203.0.113.9,", + "damage before ours": "203.0.113.9, banana, 10.0.0.1", + "whitespace only": " ", + "port but no host": ":443", + "cidr not address": "203.0.113.0/24", + } + + const peerKey = "10.0.0.1" + for name, header := range cases { + t.Run(name, func(t *testing.T) { + got := trust.clientAddr(request("10.0.0.1:4000", header)) + if got != peerKey { + t.Errorf("clientAddr = %q, want the peer %q", got, peerKey) + } + }) + } +} + +// An address WITH a port is not malformed — some proxies append one. +func TestForwardedEntryWithPortIsAccepted(t *testing.T) { + trust := newProxyTrust(prefixes(t, "10.0.0.0/8")) + + if got, want := trust.clientAddr(request("10.0.0.1:4000", "203.0.113.9:51000")), "203.0.113.9"; got != want { + t.Errorf("IPv4 with port: clientAddr = %q, want %q", got, want) + } + if got, want := trust.clientAddr(request("10.0.0.1:4000", "[2001:db8::1]:443")), "2001:db8::/64"; got != want { + t.Errorf("IPv6 with port: clientAddr = %q, want %q", got, want) + } +} + +func TestMalformedRemoteAddrDoesNotPanic(t *testing.T) { + trust := newProxyTrust(prefixes(t, "10.0.0.0/8")) + + for _, remote := range []string{"", " ", "pipe", "not:an:addr", "@"} { + got := trust.clientAddr(request(remote, "203.0.113.9")) + // Whatever it returns, it must not be the forwarded address: an + // unparseable peer is not a trusted one. + if got == "203.0.113.9" { + t.Errorf("RemoteAddr %q was treated as a trusted peer", remote) + } + } +} + +/* ── F. IPv6 is keyed by /64 ────────────────────────────────────────────── */ + +func TestIPv6SameSlash64SharesABucket(t *testing.T) { + trust := newProxyTrust(prefixes(t, "10.0.0.0/8")) + + // Same /64, different hosts within it — one subscriber, one budget. + a := trust.clientAddr(request("10.0.0.1:4000", "2001:db8:abcd:1234::1")) + b := trust.clientAddr(request("10.0.0.1:4000", "2001:db8:abcd:1234:ffff:ffff:ffff:ffff")) + + if a != b { + t.Errorf("two addresses in one /64 produced %q and %q; a caller could mint budgets at will", a, b) + } + if want := "2001:db8:abcd:1234::/64"; a != want { + t.Errorf("key = %q, want %q", a, want) + } +} + +func TestIPv6DifferentSlash64DoesNotShareABucket(t *testing.T) { + trust := newProxyTrust(prefixes(t, "10.0.0.0/8")) + + a := trust.clientAddr(request("10.0.0.1:4000", "2001:db8:abcd:1234::1")) + b := trust.clientAddr(request("10.0.0.1:4000", "2001:db8:abcd:9999::1")) + + if a == b { + t.Errorf("two different /64s shared the key %q", a) + } +} + +// An IPv4 peer reported in IPv4-mapped form is the same caller as the plain +// form, and must not become a second bucket. +func TestIPv4MappedIPv6NormalisesToIPv4(t *testing.T) { + trust := newProxyTrust(nil) + + plain := trust.clientAddr(request("203.0.113.9:51000")) + mapped := trust.clientAddr(request("[::ffff:203.0.113.9]:51000")) + + if plain != mapped { + t.Errorf("plain %q and mapped %q are the same host but keyed differently", plain, mapped) + } + if want := "203.0.113.9"; plain != want { + t.Errorf("key = %q, want %q", plain, want) + } +} + +// A trusted IPv6 proxy works the same way as a trusted IPv4 one. +func TestTrustedIPv6Proxy(t *testing.T) { + trust := newProxyTrust(prefixes(t, "fd00::/8")) + + got := trust.clientAddr(request("[fd00::1]:4000", "2001:db8:abcd:1234::5")) + + if want := "2001:db8:abcd:1234::/64"; got != want { + t.Errorf("clientAddr = %q, want %q", got, want) + } +} + +// A scope id is local to this host and says nothing about who called. +func TestIPv6ZoneIsNotPartOfTheKey(t *testing.T) { + trust := newProxyTrust(nil) + + withZone := trust.clientAddr(request("[fe80::1%eth0]:4000")) + without := trust.clientAddr(request("[fe80::1]:4000")) + + if withZone != without { + t.Errorf("zone changed the key: %q vs %q", withZone, without) + } +} + +/* ── G. No header at all ────────────────────────────────────────────────── */ + +func TestMissingForwardedHeaderFallsBackToPeer(t *testing.T) { + trust := newProxyTrust(prefixes(t, "10.0.0.0/8")) + + if got, want := trust.clientAddr(request("10.0.0.1:4000")), "10.0.0.1"; got != want { + t.Errorf("clientAddr = %q, want %q", got, want) + } +} + +// A chain consisting only of our own proxies names no client. +func TestChainOfOnlyTrustedProxiesFallsBackToPeer(t *testing.T) { + trust := newProxyTrust(prefixes(t, "10.0.0.0/8")) + + if got, want := trust.clientAddr(request("10.0.0.1:4000", "10.0.0.2, 10.0.0.1")), "10.0.0.1"; got != want { + t.Errorf("clientAddr = %q, want %q", got, want) + } +} + +/* ── H. The port is not part of the key ─────────────────────────────────── */ + +// Pre-existing behaviour, asserted here because it is the reason this function +// strips the port at all: a browser opens a new source port per connection. +func TestSourcePortIsNotPartOfTheKey(t *testing.T) { + trust := newProxyTrust(nil) + + a := trust.clientAddr(request("203.0.113.9:51000")) + b := trust.clientAddr(request("203.0.113.9:51001")) + + if a != b { + t.Errorf("source port changed the key: %q vs %q", a, b) + } +} + +// A bare address with no port — a test server, or a rewritten RemoteAddr. +func TestRemoteAddrWithoutAPortIsAccepted(t *testing.T) { + trust := newProxyTrust(nil) + + if got, want := trust.clientAddr(request("203.0.113.9")), "203.0.113.9"; got != want { + t.Errorf("clientAddr = %q, want %q", got, want) + } +} + +/* ── Trust-set edge cases ───────────────────────────────────────────────── */ + +// A single-host trusted proxy, which is what a bare address in configuration +// becomes. +func TestSingleHostTrustedProxy(t *testing.T) { + trust := newProxyTrust(prefixes(t, "172.17.0.1/32")) + + if got, want := trust.clientAddr(request("172.17.0.1:4000", "203.0.113.9")), "203.0.113.9"; got != want { + t.Errorf("trusted host: clientAddr = %q, want %q", got, want) + } + // One address along is NOT trusted. + if got, want := trust.clientAddr(request("172.17.0.2:4000", "203.0.113.9")), "172.17.0.2"; got != want { + t.Errorf("neighbouring host: clientAddr = %q, want %q", got, want) + } +} diff --git a/go-api/internal/httpserver/maintenance.go b/go-api/internal/httpserver/maintenance.go new file mode 100644 index 0000000..7b3dc4d --- /dev/null +++ b/go-api/internal/httpserver/maintenance.go @@ -0,0 +1,193 @@ +package httpserver + +import ( + "context" + "log/slog" + "time" + + "github.com/krow/krow-backend/go-api/internal/oauth" + "github.com/krow/krow-backend/go-api/internal/ratelimit" +) + +// Scheduled maintenance for the OAuth and rate-limit tables. +// +// WHY THIS SHAPE AND NOT A NEW ONE +// +// The process already has a scheduled maintenance mechanism: sweepSessions in +// cmd/api/main.go, a ticker goroutine whose context is the server's, which runs +// once at startup and then on an interval, logs a failure and retries at the +// next tick. It is bounded, cancellable, non-blocking and failure-isolated, and +// it has been in production. +// +// So this is the same thing for two more tables rather than a second kind of +// thing. No new process, no cron dependency, no leader election, no library. +// The one addition is that both sweeps live behind a single type, so +// cmd/api/main.go gains one line rather than two more goroutines. +// +// MULTI-INSTANCE SAFETY COMES FROM THE STATEMENTS, NOT FROM COORDINATION +// +// Every instance runs this, on its own schedule, with no lock between them — +// deliberately. A lease or an advisory lock would be state to hold, to expire +// and to recover when the holder dies mid-sweep, in exchange for avoiding work +// that is already harmless: each sweep is a bounded DELETE whose predicate no +// longer matches once a row is gone. Two instances sweeping at the same moment +// delete disjoint sets and neither errors. A row deleted twice is not an error; +// it is a row that was already deleted. +// +// That is the same property Phase 5's concurrent-cleanup test asserts directly: +// four workers, six dead tokens, exactly six removed between them. + +// maintenanceInterval is how often the sweep runs. +// +// Hourly. The grace period before anything is deleted is also an hour, so a +// row becomes eligible and is collected within roughly two — soon enough that +// nothing accumulates, and far enough apart that a DELETE never lands on a hot +// path. Shorter would buy nothing: nothing here is a correctness deadline. +// +// Deliberately NOT sweepInterval's fifteen minutes. Sessions churn with every +// sign-in; authorization codes live sixty seconds and tokens fifteen minutes, +// so an hour still collects them promptly while running a quarter as often. +const maintenanceInterval = time.Hour + +// maintenanceTimeout bounds one pass. +// +// Generous for three bounded deletes and short enough that a wedged statement +// cannot hold this goroutine past shutdown. Matches sweepSessions' own bound in +// spirit; longer only because there are more statements. +const maintenanceTimeout = 60 * time.Second + +// Maintenance sweeps the OAuth and rate-limit tables. +// +// Nil when the deployment does not serve MCP, which is why Server.Maintenance +// returns a pointer and the caller checks it — the same way routeOAuth simply +// registers nothing. +type Maintenance struct { + store *oauth.Store + limiter *ratelimit.Limiter + log *slog.Logger +} + +// Maintenance exposes the sweeper, or nil when there is nothing to sweep. +// +// Mirrors Server.Sessions(), which exists for exactly this reason: the process +// owns the schedule, the server owns the things being swept. +func (s *Server) Maintenance() *Maintenance { + if !s.cfg.OAuth.Enabled() { + return nil + } + return &Maintenance{ + store: oauth.NewStore(s.db.Pool), + limiter: s.limiter, + log: s.log, + } +} + +// MaintenanceResult is what one pass removed. +type MaintenanceResult struct { + Grants int64 + AccessTokens int64 + RefreshTokens int64 + RateLimits int64 +} + +// Total is the row count removed, for the log line. +func (r MaintenanceResult) Total() int64 { + return r.Grants + r.AccessTokens + r.RefreshTokens + r.RateLimits +} + +// Sweep runs one maintenance pass. +// +// The two halves are independent on purpose: a failure sweeping OAuth rows must +// not prevent the rate-limit sweep, because the second is the one that would +// otherwise grow without bound. The first error is returned, after both have +// been attempted. +func (m *Maintenance) Sweep(ctx context.Context) (MaintenanceResult, error) { + var out MaintenanceResult + var firstErr error + + // OAuth: codes, access tokens, and refresh tokens past their retention. + // The grace period and the reuse-detection retention are enforced inside + // Store.Cleanup — this schedules it, it does not reimplement it. + cleaned, err := m.store.Cleanup(ctx) + if err != nil { + firstErr = err + } else { + out.Grants = cleaned.Grants + out.AccessTokens = cleaned.AccessTokens + out.RefreshTokens = cleaned.RefreshTokens + } + + if m.limiter != nil { + swept, err := m.limiter.Sweep(ctx, 0) // 0 = the package's own batch size + if err != nil && firstErr == nil { + firstErr = err + } + out.RateLimits = swept + } + + return out, firstErr +} + +// SweepMaintenance runs the sweep until the context is cancelled. +// +// Deliberately identical in shape to sweepSessions: one pass immediately so a +// process that has been down does not carry a backlog for a further hour, then +// on the ticker. A failed pass is logged and retried at the next tick — the +// tables being briefly larger than they should be is not worth stopping the API +// for, and it is certainly not worth a panic in a goroutine nobody is watching. +// +// Exported because cmd/api owns the process's goroutines and this package owns +// what they do. +func SweepMaintenance(ctx context.Context, m *Maintenance, log *slog.Logger) { + if m == nil { + // No OAuth surface, nothing to sweep. Returning rather than ticking + // uselessly for the life of the process. + return + } + + ticker := time.NewTicker(maintenanceInterval) + defer ticker.Stop() + + pass := func() { + // A deadline of its own, so a slow DELETE cannot leave this goroutine + // blocked past shutdown. + sweepCtx, cancel := context.WithTimeout(ctx, maintenanceTimeout) + defer cancel() + + // A panic in a background goroutine takes the process with it, and + // this one runs unattended for the life of the deployment. Recovering + // turns a bug here into a logged failure and a retry at the next tick. + defer func() { + if p := recover(); p != nil { + log.Error("maintenance sweep panicked", "panic", p) + } + }() + + result, err := m.Sweep(sweepCtx) + switch { + case err != nil && ctx.Err() != nil: + // Shutting down; the cancellation is expected, not a failure. + case err != nil: + log.Warn("maintenance sweep failed", "error", err, + "grants", result.Grants, "access_tokens", result.AccessTokens, + "refresh_tokens", result.RefreshTokens, "rate_limits", result.RateLimits) + case result.Total() > 0: + log.Info("maintenance sweep", + "grants", result.Grants, "access_tokens", result.AccessTokens, + "refresh_tokens", result.RefreshTokens, "rate_limits", result.RateLimits) + default: + log.Debug("maintenance sweep found nothing to delete") + } + } + + pass() + for { + select { + case <-ctx.Done(): + log.Debug("maintenance sweeper stopped") + return + case <-ticker.C: + pass() + } + } +} diff --git a/go-api/internal/httpserver/maintenance_test.go b/go-api/internal/httpserver/maintenance_test.go new file mode 100644 index 0000000..f162f65 --- /dev/null +++ b/go-api/internal/httpserver/maintenance_test.go @@ -0,0 +1,275 @@ +package httpserver_test + +import ( + "context" + "io" + "log/slog" + "strings" + "sync" + "testing" + "time" + + "github.com/krow/krow-backend/go-api/internal/httpserver" +) + +// Scheduler lifecycle. +// +// What is under test is the GOROUTINE, not the deletes — those are covered in +// internal/oauth and internal/ratelimit against real data. Here the questions +// are: does it start, does it do a pass, does it stop when told, does a failure +// take the process with it, and is running it twice safe. + +/* ── Lifecycle ──────────────────────────────────────────────────────────── */ + +// It runs one pass IMMEDIATELY, before the first tick. A process that has been +// down should not carry a backlog for a further hour. +func TestMaintenanceRunsOnceImmediately(t *testing.T) { + a := newOAuthAPI(t) + m := a.srv.Maintenance() + if m == nil { + t.Fatal("a configured deployment returned no Maintenance") + } + + ctx, cancel := context.WithCancel(context.Background()) + defer cancel() + + done := make(chan struct{}) + go func() { + httpserver.SweepMaintenance(ctx, m, slog.New(slog.NewTextHandler(io.Discard, nil))) + close(done) + }() + + // The immediate pass is the only one that will happen inside the test's + // lifetime — the ticker is an hour. Give it a moment, then stop. + time.Sleep(200 * time.Millisecond) + cancel() + + select { + case <-done: + case <-time.After(5 * time.Second): + t.Fatal("the sweeper did not stop within 5s of cancellation") + } +} + +// Cancellation must return promptly, or a shutdown hangs on a goroutine nobody +// is waiting for. +func TestMaintenanceStopsOnCancellation(t *testing.T) { + a := newOAuthAPI(t) + ctx, cancel := context.WithCancel(context.Background()) + + done := make(chan struct{}) + go func() { + httpserver.SweepMaintenance(ctx, a.srv.Maintenance(), + slog.New(slog.NewTextHandler(io.Discard, nil))) + close(done) + }() + + time.Sleep(100 * time.Millisecond) + start := time.Now() + cancel() + + select { + case <-done: + if elapsed := time.Since(start); elapsed > 2*time.Second { + t.Errorf("stopping took %v; shutdown would block on it", elapsed) + } + case <-time.After(5 * time.Second): + t.Fatal("the sweeper ignored cancellation") + } +} + +// An already-cancelled context must not run a pass and must return at once. +func TestMaintenanceWithAnAlreadyCancelledContextReturns(t *testing.T) { + a := newOAuthAPI(t) + ctx, cancel := context.WithCancel(context.Background()) + cancel() + + done := make(chan struct{}) + go func() { + httpserver.SweepMaintenance(ctx, a.srv.Maintenance(), + slog.New(slog.NewTextHandler(io.Discard, nil))) + close(done) + }() + + select { + case <-done: + case <-time.After(5 * time.Second): + t.Fatal("the sweeper did not return on an already-cancelled context") + } +} + +/* ── It does the work ───────────────────────────────────────────────────── */ + +// One pass removes dead rows and leaves live ones. The detailed retention rules +// are tested in internal/oauth; this asserts the scheduler is wired to them. +func TestMaintenanceSweepRemovesDeadRows(t *testing.T) { + a := newOAuthAPI(t) + ctx := context.Background() + + // A grant that is already past its expiry and its grace. + if _, err := a.h.Pool.Exec(ctx, + `INSERT INTO oauth_clients (client_id, client_name, redirect_uris) + VALUES ('sweep-client', 'Sweep', ARRAY['https://a.test/cb'])`); err != nil { + t.Fatalf("client: %v", err) + } + userID, _ := seededUser(t, a.h.Pool) + var orgID string + if err := a.h.Pool.QueryRow(ctx, + `SELECT org_id::text FROM users WHERE id = $1::uuid`, userID).Scan(&orgID); err != nil { + t.Fatalf("org: %v", err) + } + + if _, err := a.h.Pool.Exec(ctx, + `INSERT INTO oauth_grants + (code_hash, client_id, user_id, org_id, redirect_uri, scopes, resource, + code_challenge, code_challenge_method, created_date, expires_at) + VALUES (repeat('a', 64), 'sweep-client', $1::uuid, $2::uuid, 'https://a.test/cb', + ARRAY['krow.read'], $3, repeat('B', 43), 'S256', + now() - interval '3 hours', now() - interval '3 hours' + interval '1 minute')`, + userID, orgID, testMCPResource); err != nil { + t.Fatalf("grant: %v", err) + } + + // An expired rate-limit bucket. + if _, err := a.h.Pool.Exec(ctx, + `INSERT INTO rate_limits (bucket, window_start, count, expires_at) + VALUES ('test:old', now() - interval '3 hours', 5, now() - interval '2 hours')`); err != nil { + t.Fatalf("bucket: %v", err) + } + // And a live one, which must survive. + if _, err := a.h.Pool.Exec(ctx, + `INSERT INTO rate_limits (bucket, window_start, count, expires_at) + VALUES ('test:live', now(), 1, now() + interval '1 hour')`); err != nil { + t.Fatalf("bucket: %v", err) + } + + result, err := a.srv.Maintenance().Sweep(ctx) + if err != nil { + t.Fatalf("Sweep: %v", err) + } + + if result.Grants != 1 { + t.Errorf("removed %d grants, want 1", result.Grants) + } + if result.RateLimits != 1 { + t.Errorf("removed %d rate-limit rows, want 1", result.RateLimits) + } + if result.Total() != 2 { + t.Errorf("Total() = %d, want 2", result.Total()) + } + + var live int + if err := a.h.Pool.QueryRow(ctx, + `SELECT count(*) FROM rate_limits WHERE bucket = 'test:live'`).Scan(&live); err != nil { + t.Fatalf("count: %v", err) + } + if live != 1 { + t.Error("the live rate-limit window was swept") + } +} + +// Running it repeatedly must be safe and must converge to removing nothing. +func TestRepeatedMaintenanceIsSafe(t *testing.T) { + a := newOAuthAPI(t) + ctx := context.Background() + m := a.srv.Maintenance() + + for i := 0; i < 3; i++ { + result, err := m.Sweep(ctx) + if err != nil { + t.Fatalf("pass %d: %v", i+1, err) + } + if i > 0 && result.Total() != 0 { + t.Errorf("pass %d removed %d rows; a repeat pass should find nothing", i+1, result.Total()) + } + } +} + +// Two instances sweep concurrently with no coordination. Neither may error. +// Run with -race. +func TestConcurrentMaintenanceIsSafe(t *testing.T) { + a := newOAuthAPI(t) + ctx := context.Background() + + const instances = 4 + var wg sync.WaitGroup + errs := make(chan error, instances) + + for i := 0; i < instances; i++ { + wg.Add(1) + go func() { + defer wg.Done() + if _, err := a.srv.Maintenance().Sweep(ctx); err != nil { + errs <- err + } + }() + } + wg.Wait() + close(errs) + + for err := range errs { + t.Errorf("concurrent sweep errored: %v", err) + } +} + +/* ── Failure isolation ──────────────────────────────────────────────────── */ + +// A failing sweep must be logged and survived, not fatal. The database is +// closed underneath the sweeper, which is the closest thing to a real outage a +// test can arrange. +func TestMaintenanceSurvivesADatabaseFailure(t *testing.T) { + a := newOAuthAPI(t) + m := a.srv.Maintenance() + + var logged strings.Builder + log := slog.New(slog.NewTextHandler(&logged, &slog.HandlerOptions{Level: slog.LevelDebug})) + + // A cancelled context makes every statement fail immediately. + dead, cancel := context.WithCancel(context.Background()) + cancel() + + if _, err := m.Sweep(dead); err == nil { + t.Log("note: the sweep reported no error on a cancelled context") + } + + // The goroutine wrapper must not panic or exit the process on that. + ctx, stop := context.WithCancel(context.Background()) + done := make(chan struct{}) + go func() { + httpserver.SweepMaintenance(ctx, m, log) + close(done) + }() + time.Sleep(150 * time.Millisecond) + stop() + + select { + case <-done: + case <-time.After(5 * time.Second): + t.Fatal("the sweeper did not stop") + } +} + +/* ── It is absent when the surface is ───────────────────────────────────── */ + +// A deployment without OAuth has nothing to sweep, and must not start a ticker +// that runs for the life of the process doing nothing. +func TestMaintenanceIsNilWhenTheSurfaceIsDisabled(t *testing.T) { + a := newAPI(t) // the standard fixture: no OAuth configuration + + if m := a.srv.Maintenance(); m != nil { + t.Error("an unconfigured deployment returned a Maintenance sweeper") + } + + // And the runner must return immediately rather than tick forever. + done := make(chan struct{}) + go func() { + httpserver.SweepMaintenance(context.Background(), a.srv.Maintenance(), + slog.New(slog.NewTextHandler(io.Discard, nil))) + close(done) + }() + select { + case <-done: + case <-time.After(2 * time.Second): + t.Fatal("the sweeper ticked despite having nothing to sweep") + } +} diff --git a/go-api/internal/httpserver/mcp.go b/go-api/internal/httpserver/mcp.go new file mode 100644 index 0000000..538f2ee --- /dev/null +++ b/go-api/internal/httpserver/mcp.go @@ -0,0 +1,195 @@ +package httpserver + +import ( + "net/http" + + "github.com/krow/krow-backend/go-api/internal/auth" + "github.com/krow/krow-backend/go-api/internal/authctx" + "github.com/krow/krow-backend/go-api/internal/mcpserver" + "github.com/krow/krow-backend/go-api/internal/oauth" + "github.com/krow/krow-backend/go-api/internal/ratelimit" + "github.com/krow/krow-backend/go-api/internal/runtime" +) + +// Mounting the MCP surface and the OAuth authorization server behind it. +// +// This file is the seam between the existing HTTP server and two packages that +// know nothing about it. It is deliberately thin: no validation, no policy and +// no business logic live here, because every one of those already lives in the +// package being mounted. What this file decides is only WHERE things are served +// and WHAT AUTHENTICATES them, and those two decisions are the ones that have +// to be right. +// +// OFF UNLESS CONFIGURED. Without OAUTH_ISSUER and MCP_RESOURCE, none of these +// routes are registered at all. That follows routeRuns' precedent exactly: a +// deployment that does not serve agents answers 404 rather than registering +// routes that fail, and the same is true of one that does not serve MCP. An +// existing deployment that upgrades to this build gains nothing it did not ask +// for. + +// routeOAuth registers the authorization server and its discovery documents. +// +// WHICH OF THESE ARE PUBLIC, AND WHY — this is the part worth reading twice. +// Four paths bypass the cookie middleware, and each has a specific reason: +// +// /.well-known/oauth-protected-resource RFC 9728. A client that has no +// /.well-known/oauth-authorization-server RFC 8414. token cannot read a +// document that requires one, and +// these are how it learns where to +// get a token. They contain only +// public endpoint URLs. +// +// /oauth/register RFC 7591. A client that has never registered has no +// credential to present — that is the entire point of +// dynamic registration. +// +// /oauth/token The client authenticates with an authorization code or a +// refresh token IN THE BODY. A cookie would be meaningless: +// this is a back-channel call from Claude's servers, where +// no browser and no cookie exist. +// +// /oauth/authorize is deliberately NOT public. It runs in a browser, as a +// person, and it requires the existing KROW session — that is how the consent +// screen knows whose organisation is being granted. An unauthenticated visitor +// is redirected to the existing login and comes back. +// +// /mcp is deliberately NOT public either, and also does not use the cookie. See +// routeMCP. +func (s *Server) routeOAuth(mux *http.ServeMux) int { + if !s.cfg.OAuth.Enabled() { + return 0 + } + + cfg := oauth.Config{ + Issuer: s.cfg.OAuth.Issuer, + Resource: s.cfg.OAuth.Resource, + } + store := oauth.NewStore(s.db.Pool) + as := oauth.NewServer(cfg, store, sessionResolver{s}, s.cfg.OAuth.LoginPath, s.log) + + mux.Handle("GET /.well-known/oauth-protected-resource", cfg.ProtectedResourceHandler()) + mux.Handle("GET /.well-known/oauth-authorization-server", cfg.AuthorizationServerHandler()) + // Registration is the only endpoint that writes for a caller with no + // credential at all, so it carries the tightest limit on the surface. + mux.Handle("POST /oauth/register", + s.limited(ratelimit.OAuthRegister, s.byClientAddr, as.RegisterHandler())) + // GET renders consent; POST carries the decision. One handler, because the + // POST re-validates every parameter the GET validated rather than trusting + // the form it rendered. + mux.Handle("GET /oauth/authorize", + s.limited(ratelimit.OAuthAuthorize, s.byAddrAndUser, as.AuthorizeHandler())) + mux.Handle("POST /oauth/authorize", + s.limited(ratelimit.OAuthAuthorize, s.byAddrAndUser, as.AuthorizeHandler())) + + // The token endpoint carries two limits on two different subjects, because + // its two grant types are abused differently: a code exchange is bounded + // per client, and a refresh is bounded per token so a loop on one + // connection cannot spend another's budget. Which applies is decided per + // request by the grant_type, inside tokenLimited. + mux.Handle("POST /oauth/token", s.tokenLimited(as.TokenHandler())) + + // Revocation is deliberately unlimited — see ratelimit/rules.go. It is the + // emergency brake, and an attacker gains nothing by pulling it. + mux.Handle("POST /oauth/revoke", as.RevokeHandler()) + + return 7 +} + +// routeMCP registers the MCP endpoint. +// +// AUTHENTICATION HERE IS THE BEARER PATH AND ONLY THE BEARER PATH. +// +// The handler authenticates its own callers from the Authorization header and +// ignores whatever the cookie middleware put in the context. That is a property +// of mcpserver, not of this file — see its auth.go. +// +// /mcp IS on the publicPaths allowlist, and that is deliberate rather than an +// oversight. The cookie middleware has to step aside here: an MCP client +// discovers how to authenticate by calling this endpoint without a token and +// reading the WWW-Authenticate header of the 401, and the middleware's own 401 +// carries no such header. Guarding the path here would refuse the client with +// nowhere to go, and the connection could never be made at all. +// +// The credential requirement is not weakened by that, because it was never +// this middleware enforcing it: mcpserver refuses every method but the +// handshake without a bearer token, and it takes its identity as a parameter +// rather than from the request context, so a cookie cannot supply one. +// +// There is no second authorization layer. A tool call goes straight into the +// registry the agent runtime already uses, under the policy table it already +// consults. +func (s *Server) routeMCP(mux *http.ServeMux) int { + if !s.cfg.OAuth.Enabled() { + return 0 + } + + // The SAME registry the runtime builds. Not a copy, not a second + // construction: a tool added once is available to Owliver and to MCP + // together, and neither can drift from the other. + registry := runtime.DefaultTools( + s.db.Pool, + nil, // knowledge_search is not exposed over MCP — see mcpserver/tools.go + ) + + authenticator := oauth.NewAuthenticator( + oauth.NewStore(s.db.Pool), + s.users, + // The audience an access token must carry. From configuration, never + // from a request: a resource value supplied by a caller would let the + // caller choose their own audience. + s.cfg.OAuth.Resource, + s.log, + ) + + server := mcpserver.New(registry, authenticator, s.log). + WithResourceMetadataURL(s.cfg.OAuth.Issuer + "/.well-known/oauth-protected-resource"). + // The per-organisation ceiling is installed INSIDE the MCP server + // rather than as middleware, because the organisation is only known + // after the token has been resolved. See orgLimiter in mcplimit.go. + WithOrgLimiter(orgLimiter{s}) + + mux.Handle("POST /mcp", s.mcpLimited(server.Handler())) + // GET is what the Streamable HTTP binding uses for a server-initiated + // stream, which this server does not open. Registered so the answer is 405 + // with an Allow header rather than a 404 that suggests the endpoint is + // absent. + mux.Handle("GET /mcp", server.Handler()) + + return 2 +} + +// sessionResolver adapts the existing cookie session to oauth.SessionResolver. +// +// This is the ONLY place the OAuth package learns who is signed in, and it does +// so through the existing session manager — the same lookup every other +// authenticated route performs. No second password store, no second session +// table, no second notion of identity. +type sessionResolver struct{ s *Server } + +// CurrentUser resolves the session cookie into an identity. +// +// Re-reads the user row rather than trusting the session's own copy, exactly as +// authenticate() does, so a suspended account cannot approve an authorization +// in the window before its session lapses. +func (r sessionResolver) CurrentUser(req *http.Request) (authctx.Identity, bool) { + token := sessionToken(req) + if token == "" { + return authctx.Identity{}, false + } + sess, err := r.s.sessions.Authenticate(req.Context(), token) + if err != nil { + return authctx.Identity{}, false + } + user, err := r.s.users.FindByID(req.Context(), sess.UserID) + if err != nil || !user.IsActive() { + return authctx.Identity{}, false + } + return authctx.Identity{ + UserID: user.ID, OrgID: user.OrgID, Email: user.Email, + FullName: user.FullName, Role: user.Role, AccountType: user.AccountType, + Status: user.Status, SessionID: sess.ID, ExpiresAt: sess.ExpiresAt, + }, true +} + +// compile-time proof that the existing user store satisfies what OAuth needs. +var _ oauth.UserLookup = (auth.UserStore)(nil) diff --git a/go-api/internal/httpserver/mcp_routes_test.go b/go-api/internal/httpserver/mcp_routes_test.go new file mode 100644 index 0000000..47f58a5 --- /dev/null +++ b/go-api/internal/httpserver/mcp_routes_test.go @@ -0,0 +1,785 @@ +package httpserver_test + +import ( + "bytes" + "crypto/sha256" + "encoding/base64" + "encoding/json" + "io" + "log/slog" + "net/http" + "net/http/httptest" + "net/url" + "strings" + "testing" + "time" + + "github.com/krow/krow-backend/go-api/internal/config" + "github.com/krow/krow-backend/go-api/internal/db" + "github.com/krow/krow-backend/go-api/internal/httpserver" + "github.com/krow/krow-backend/go-api/internal/testutil" +) + +// The configured deployment these tests run as. Fictional on purpose: every +// URL in a discovery document must be traceable to THIS configuration, and a +// realistic hostname would make a hardcoded one impossible to spot. +const ( + testOAuthIssuer = "https://krow.example.test" + testMCPResource = "https://krow.example.test/mcp" +) + +/* ── A fixture that can see headers and raw bodies ──────────────────────── */ + +// mcpResponse carries what the existing `response` deliberately does not: the +// headers (WWW-Authenticate is the whole point of several tests) and the raw +// body (the consent page is HTML, not JSON). +// +// A separate type rather than a change to `response`, so not one existing test +// in this package is touched. +type mcpResponse struct { + code int + body string + header http.Header +} + +type mcpAPI struct { + t *testing.T + handler http.Handler + srv *httpserver.Server + h *testutil.Harness + cookie *http.Cookie + email string + userID string +} + +// newOAuthAPI builds a server WITH OAuth configured, and signs in. +// +// The OAuth block is what makes routeOAuth and routeMCP register at all; the +// standard newAPI fixture leaves it empty, which is what +// TestMCPRoutesAreAbsentWhenUnconfigured relies on. +func newOAuthAPI(t *testing.T) *mcpAPI { + t.Helper() + h := testutil.New(t) + + cfg := &config.Config{ + AppEnv: "development", + HTTP: config.HTTPConfig{ + Host: "127.0.0.1", Port: 0, ShutdownTimeout: time.Second, + }, + DB: config.DBConfig{Schema: "public"}, + OAuth: config.OAuthConfig{ + Issuer: testOAuthIssuer, + Resource: testMCPResource, + LoginPath: "/login", + }, + } + log := slog.New(slog.NewTextHandler(io.Discard, nil)) + srv, err := httpserver.New(cfg, &db.DB{Pool: h.Pool, Schema: "public"}, log) + if err != nil { + t.Fatalf("build the server: %v", err) + } + + a := &mcpAPI{t: t, handler: srv.Handler(), srv: srv, h: h} + a.userID, a.email = seededUser(t, h.Pool) + setPassword(t, h.Pool, a.userID) + + result := signIn(t, a.handler, a.email, harnessPassword, false) + if result.code != http.StatusOK || result.cookie == nil { + t.Fatalf("the harness could not sign in: %d", result.code) + } + a.cookie = result.cookie + return a +} + +func (a *mcpAPI) send(req *http.Request, withCookie bool) mcpResponse { + a.t.Helper() + if withCookie && a.cookie != nil { + req.AddCookie(a.cookie) + } + rec := httptest.NewRecorder() + a.handler.ServeHTTP(rec, req) + return mcpResponse{code: rec.Code, body: rec.Body.String(), header: rec.Header()} +} + +func (a *mcpAPI) jsonReq(method, path string, payload any) *http.Request { + a.t.Helper() + var body io.Reader + if payload != nil { + raw, err := json.Marshal(payload) + if err != nil { + a.t.Fatalf("encode: %v", err) + } + body = bytes.NewReader(raw) + } + req := httptest.NewRequest(method, path, body) + if payload != nil { + req.Header.Set("Content-Type", "application/json") + } + return req +} + +// do sends WITH the session cookie — a signed-in browser. +func (a *mcpAPI) do(method, path string, payload any) mcpResponse { + return a.send(a.jsonReq(method, path, payload), true) +} + +// doAnon sends WITHOUT any credential. +func (a *mcpAPI) doAnon(method, path string, payload any) mcpResponse { + return a.send(a.jsonReq(method, path, payload), false) +} + +// doAnonWithHeader sends one extra header and no cookie. +func (a *mcpAPI) doAnonWithHeader(method, path string, payload any, key, value string) mcpResponse { + req := a.jsonReq(method, path, payload) + req.Header.Set(key, value) + return a.send(req, false) +} + +func (a *mcpAPI) formReq(method, path string, form url.Values) *http.Request { + req := httptest.NewRequest(method, path, strings.NewReader(form.Encode())) + req.Header.Set("Content-Type", "application/x-www-form-urlencoded") + return req +} + +// doForm posts a form WITH the cookie — the consent decision. +func (a *mcpAPI) doForm(method, path string, form url.Values) mcpResponse { + return a.send(a.formReq(method, path, form), true) +} + +// doAnonForm posts a form WITHOUT a cookie — the back-channel token call. +func (a *mcpAPI) doAnonForm(method, path string, form url.Values) mcpResponse { + return a.send(a.formReq(method, path, form), false) +} + +// oauthAccessToken runs the whole flow and returns a usable access token, for +// tests that need a valid credential to prove it is being ignored. +func (a *mcpAPI) oauthAccessToken(t *testing.T) string { + t.Helper() + reg := a.doAnon("POST", "/oauth/register", map[string]any{ + "client_name": "Token Helper", "redirect_uris": []string{"https://client.example.test/cb"}, + }) + var regDoc struct { + ClientID string `json:"client_id"` + } + mustJSON(t, reg.body, ®Doc) + + verifier := "helperVerifier0123456789abcdefghijklmnopqrst" + q := url.Values{ + "client_id": {regDoc.ClientID}, "redirect_uri": {"https://client.example.test/cb"}, + "response_type": {"code"}, "state": {"helper"}, + "code_challenge": {challengeFor(verifier)}, "code_challenge_method": {"S256"}, + "resource": {testMCPResource}, "scope": {"krow.read"}, + } + consent := a.do("GET", "/oauth/authorize?"+q.Encode(), nil) + csrf := between(consent.body, `name="csrf" value="`, `"`) + + form := url.Values{} + for k, v := range q { + form[k] = v + } + form.Set("decision", "approve") + form.Set("csrf", csrf) + approved := a.doForm("POST", "/oauth/authorize", form) + loc, _ := url.Parse(approved.header.Get("Location")) + + tok := a.doAnonForm("POST", "/oauth/token", url.Values{ + "grant_type": {"authorization_code"}, "code": {loc.Query().Get("code")}, + "client_id": {regDoc.ClientID}, "redirect_uri": {"https://client.example.test/cb"}, + "code_verifier": {verifier}, + }) + var tokens struct { + AccessToken string `json:"access_token"` + } + mustJSON(t, tok.body, &tokens) + if tokens.AccessToken == "" { + t.Fatalf("could not obtain a token: %s", tok.body) + } + return tokens.AccessToken +} + +// challengeFor derives an S256 challenge, so these tests do not depend on the +// oauth package's unexported helpers. +func challengeFor(verifier string) string { + sum := sha256.Sum256([]byte(verifier)) + return base64.RawURLEncoding.EncodeToString(sum[:]) +} + +// The mounted surface, end to end. +// +// Everything below drives the REAL router — the same mux, the same +// authenticate() middleware, the same publicPaths allowlist that serves +// production. The point is not to re-test the OAuth package (internal/oauth +// does that against its own handlers) but to prove the MOUNTING is right: that +// discovery is reachable without a cookie, that /mcp is not, that a cookie +// cannot substitute for a bearer token, and that the routes appear at all only +// when the deployment is configured for them. + +/* ── Route registration is conditional ──────────────────────────────────── */ + +// Without OAUTH_ISSUER and MCP_RESOURCE, none of this exists. An upgrade must +// not quietly add an authorization server to a deployment that never asked. +func TestMCPRoutesAreAbsentWhenUnconfigured(t *testing.T) { + a := newAPI(t) // the standard fixture: no OAuth configuration + + for _, path := range []string{ + "/mcp", + "/oauth/register", + "/oauth/authorize", + "/oauth/token", + "/.well-known/oauth-protected-resource", + "/.well-known/oauth-authorization-server", + } { + r := a.doAnon("POST", path, nil) + if r.code != http.StatusNotFound && r.code != http.StatusUnauthorized { + t.Errorf("%s = %d on an unconfigured deployment; want 404 or 401, never a served response", + path, r.code) + } + } +} + +/* ── Discovery is public ────────────────────────────────────────────────── */ + +// A client with no token must be able to read both documents, or it can never +// discover how to get one. +func TestDiscoveryIsReachableWithoutASession(t *testing.T) { + a := newOAuthAPI(t) + + t.Run("protected resource", func(t *testing.T) { + r := a.doAnon("GET", "/.well-known/oauth-protected-resource", nil) + if r.code != http.StatusOK { + t.Fatalf("status = %d, want 200 without a cookie: %s", r.code, r.body) + } + var doc struct { + Resource string `json:"resource"` + AuthorizationServers []string `json:"authorization_servers"` + BearerMethods []string `json:"bearer_methods_supported"` + } + mustJSON(t, r.body, &doc) + + if doc.Resource != testMCPResource { + t.Errorf("resource = %q, want %q", doc.Resource, testMCPResource) + } + if len(doc.AuthorizationServers) != 1 || doc.AuthorizationServers[0] != testOAuthIssuer { + t.Errorf("authorization_servers = %v, want [%q]", doc.AuthorizationServers, testOAuthIssuer) + } + // The MCP spec forbids a token in the query string. + if strings.Join(doc.BearerMethods, ",") != "header" { + t.Errorf("bearer_methods_supported = %v, want [header]", doc.BearerMethods) + } + }) + + t.Run("authorization server", func(t *testing.T) { + r := a.doAnon("GET", "/.well-known/oauth-authorization-server", nil) + if r.code != http.StatusOK { + t.Fatalf("status = %d, want 200 without a cookie: %s", r.code, r.body) + } + var doc struct { + Issuer string `json:"issuer"` + AuthorizationEndpoint string `json:"authorization_endpoint"` + TokenEndpoint string `json:"token_endpoint"` + RegistrationEndpoint string `json:"registration_endpoint"` + Scopes []string `json:"scopes_supported"` + ResponseTypes []string `json:"response_types_supported"` + GrantTypes []string `json:"grant_types_supported"` + PKCEMethods []string `json:"code_challenge_methods_supported"` + ResourceIndicators bool `json:"resource_indicators_supported"` + } + mustJSON(t, r.body, &doc) + + // EVERY url must come from configuration. A hardcoded hostname would + // be one deployment's identity baked into every other one. + if doc.Issuer != testOAuthIssuer { + t.Errorf("issuer = %q, want %q", doc.Issuer, testOAuthIssuer) + } + for name, got := range map[string]string{ + "authorization_endpoint": doc.AuthorizationEndpoint, + "token_endpoint": doc.TokenEndpoint, + "registration_endpoint": doc.RegistrationEndpoint, + } { + if !strings.HasPrefix(got, testOAuthIssuer) { + t.Errorf("%s = %q, want it under the configured issuer", name, got) + } + } + if strings.Join(doc.ResponseTypes, ",") != "code" { + t.Errorf("response_types_supported = %v; implicit must not be advertised", doc.ResponseTypes) + } + if strings.Join(doc.PKCEMethods, ",") != "S256" { + t.Errorf("code_challenge_methods_supported = %v, want [S256]", doc.PKCEMethods) + } + for _, forbidden := range []string{"password", "client_credentials", "implicit"} { + for _, advertised := range doc.GrantTypes { + if advertised == forbidden { + t.Errorf("grant_types_supported advertises %q", forbidden) + } + } + } + for _, s := range doc.Scopes { + if s == "krow.write" { + t.Error("scopes_supported advertises krow.write") + } + } + if !doc.ResourceIndicators { + t.Error("resource_indicators_supported must be true") + } + }) +} + +/* ── /mcp authentication ────────────────────────────────────────────────── */ + +// No bearer → 401 with a challenge that tells the client where to go. +func TestMCPWithoutBearerReturns401AndDiscoveryPointer(t *testing.T) { + a := newOAuthAPI(t) + r := a.doAnon("POST", "/mcp", map[string]any{ + "jsonrpc": "2.0", "id": 1, "method": "tools/list", + }) + + if r.code != http.StatusUnauthorized { + t.Fatalf("status = %d, want 401", r.code) + } + challenge := r.header.Get("WWW-Authenticate") + if !strings.HasPrefix(challenge, "Bearer") { + t.Fatalf("WWW-Authenticate = %q, want a Bearer challenge", challenge) + } + // RFC 9728: without resource_metadata the client has a 401 and nowhere to + // look. This is the difference between "failed" and "here is how". + if !strings.Contains(challenge, `resource_metadata="`+testOAuthIssuer) { + t.Errorf("WWW-Authenticate = %q, want resource_metadata built from the configured issuer", challenge) + } + // And it must be built from config, not baked in. + if strings.Contains(challenge, "krowforce.com") { + t.Errorf("WWW-Authenticate contains a hardcoded production hostname: %q", challenge) + } +} + +// THE test for this phase's riskiest decision: a perfectly valid KROW session +// cookie must not open the MCP endpoint. +func TestMCPRejectsACookieSession(t *testing.T) { + a := newOAuthAPI(t) + + // `a.do` sends the authenticated session cookie the rest of the suite uses. + r := a.do("POST", "/mcp", map[string]any{ + "jsonrpc": "2.0", "id": 1, "method": "tools/list", + }) + + if r.code != http.StatusUnauthorized { + t.Fatalf("status = %d, want 401 — a browser cookie authenticated an MCP call", r.code) + } +} + +func TestMCPRejectsAnInvalidBearer(t *testing.T) { + a := newOAuthAPI(t) + for name, header := range map[string]string{ + "unknown token": "Bearer not-a-real-token", + "empty": "Bearer ", + "wrong scheme": "Basic dXNlcjpwYXNz", + "no scheme": "abcdef", + } { + t.Run(name, func(t *testing.T) { + r := a.doAnonWithHeader("POST", "/mcp", map[string]any{ + "jsonrpc": "2.0", "id": 1, "method": "tools/list", + }, "Authorization", header) + if r.code != http.StatusUnauthorized { + t.Errorf("status = %d, want 401", r.code) + } + }) + } +} + +// A token must never be accepted from the query string. The MCP spec forbids +// it, and a URL is logged, cached and put in a Referer. +func TestMCPIgnoresATokenInTheQueryString(t *testing.T) { + a := newOAuthAPI(t) + token := a.oauthAccessToken(t) + + r := a.doAnon("POST", "/mcp?access_token="+url.QueryEscape(token), map[string]any{ + "jsonrpc": "2.0", "id": 1, "method": "tools/list", + }) + if r.code != http.StatusUnauthorized { + t.Errorf("status = %d, want 401 — a query-string token was accepted", r.code) + } +} + +// Custom identity headers must be ignored outright. +func TestMCPIgnoresCustomIdentityHeaders(t *testing.T) { + a := newOAuthAPI(t) + for _, header := range []string{"X-Access-Token", "X-Api-Key", "X-Org-Id", "X-User-Id", "X-Krow-Token"} { + r := a.doAnonWithHeader("POST", "/mcp", map[string]any{ + "jsonrpc": "2.0", "id": 1, "method": "tools/list", + }, header, a.oauthAccessToken(t)) + if r.code != http.StatusUnauthorized { + t.Errorf("%s was accepted as a credential: %d", header, r.code) + } + } +} + +/* ── The full discovery → consent → token → MCP journey ─────────────────── */ + +// Every step a Claude client performs, over the real router, in order. +func TestFullMCPConnectionJourney(t *testing.T) { + a := newOAuthAPI(t) + + // 1–2. Call /mcp with no token; get 401 and a pointer. + unauth := a.doAnon("POST", "/mcp", map[string]any{ + "jsonrpc": "2.0", "id": 1, "method": "initialize", + }) + if unauth.code != http.StatusUnauthorized { + t.Fatalf("step 1: status = %d, want 401", unauth.code) + } + challenge := unauth.header.Get("WWW-Authenticate") + + // 3. Follow resource_metadata to the protected-resource document. + metaURL := between(challenge, `resource_metadata="`, `"`) + if metaURL == "" { + t.Fatal("step 3: the challenge carries no resource_metadata") + } + prPath := strings.TrimPrefix(metaURL, testOAuthIssuer) + pr := a.doAnon("GET", prPath, nil) + if pr.code != http.StatusOK { + t.Fatalf("step 3: %s = %d", prPath, pr.code) + } + var prDoc struct { + AuthorizationServers []string `json:"authorization_servers"` + } + mustJSON(t, pr.body, &prDoc) + + // 4. Authorization-server metadata. + as := a.doAnon("GET", "/.well-known/oauth-authorization-server", nil) + if as.code != http.StatusOK { + t.Fatalf("step 4: status = %d", as.code) + } + var asDoc struct { + AuthorizationEndpoint string `json:"authorization_endpoint"` + TokenEndpoint string `json:"token_endpoint"` + RegistrationEndpoint string `json:"registration_endpoint"` + } + mustJSON(t, as.body, &asDoc) + + // 5. Register, at the advertised endpoint. + reg := a.doAnon("POST", strings.TrimPrefix(asDoc.RegistrationEndpoint, testOAuthIssuer), map[string]any{ + "client_name": "Journey Client", + "redirect_uris": []string{"https://client.example.test/cb"}, + }) + if reg.code != http.StatusCreated { + t.Fatalf("step 5: registration = %d %s", reg.code, reg.body) + } + var regDoc struct { + ClientID string `json:"client_id"` + } + mustJSON(t, reg.body, ®Doc) + + // 6–7. Authorize, SIGNED IN. A cookie is exactly right here: this step is + // a person in a browser. + verifier := "journeyVerifier0123456789abcdefghijklmnopqrs" + q := url.Values{ + "client_id": {regDoc.ClientID}, "redirect_uri": {"https://client.example.test/cb"}, + "response_type": {"code"}, "state": {"journey-state"}, + "code_challenge": {challengeFor(verifier)}, "code_challenge_method": {"S256"}, + "resource": {testMCPResource}, "scope": {"krow.read"}, + } + consent := a.do("GET", "/oauth/authorize?"+q.Encode(), nil) + + // 8. A consent page, not a code. + if consent.code != http.StatusOK { + t.Fatalf("step 8: expected a consent page, got %d %s", consent.code, consent.body) + } + if !strings.Contains(consent.body, "Journey Client") { + t.Error("step 8: the consent page does not name the requesting client") + } + csrf := between(consent.body, `name="csrf" value="`, `"`) + if csrf == "" { + t.Fatal("step 8: no csrf token in the consent form") + } + + // 9–10. Approve; receive a code. + form := url.Values{} + for k, v := range q { + form[k] = v + } + form.Set("decision", "approve") + form.Set("csrf", csrf) + approved := a.doForm("POST", "/oauth/authorize", form) + if approved.code != http.StatusFound { + t.Fatalf("step 10: approve = %d %s", approved.code, approved.body) + } + loc, _ := url.Parse(approved.header.Get("Location")) + code := loc.Query().Get("code") + if code == "" { + t.Fatalf("step 10: no code in %s", loc) + } + if loc.Query().Get("state") != "journey-state" { + t.Errorf("step 10: state = %q", loc.Query().Get("state")) + } + + // 11. Exchange — with NO cookie, as a back-channel call. + tok := a.doAnonForm("POST", strings.TrimPrefix(asDoc.TokenEndpoint, testOAuthIssuer), url.Values{ + "grant_type": {"authorization_code"}, "code": {code}, + "client_id": {regDoc.ClientID}, "redirect_uri": {"https://client.example.test/cb"}, + "code_verifier": {verifier}, + }) + if tok.code != http.StatusOK { + t.Fatalf("step 11: token = %d %s", tok.code, tok.body) + } + var tokens struct { + AccessToken string `json:"access_token"` + TokenType string `json:"token_type"` + } + mustJSON(t, tok.body, &tokens) + if tokens.AccessToken == "" || tokens.TokenType != "Bearer" { + t.Fatalf("step 11: unusable token response: %s", tok.body) + } + + // 12–13. tools/list with the bearer token. + list := a.doAnonWithHeader("POST", "/mcp", map[string]any{ + "jsonrpc": "2.0", "id": 2, "method": "tools/list", + }, "Authorization", "Bearer "+tokens.AccessToken) + if list.code != http.StatusOK { + t.Fatalf("step 13: tools/list = %d %s", list.code, list.body) + } + var listDoc struct { + Result struct { + Tools []struct { + Name string `json:"name"` + } `json:"tools"` + } `json:"result"` + } + mustJSON(t, list.body, &listDoc) + if len(listDoc.Result.Tools) != 16 { + t.Errorf("step 13: %d tools, want 16", len(listDoc.Result.Tools)) + } + for _, tool := range listDoc.Result.Tools { + switch tool.Name { + case "assign_worker", "move_application", "knowledge_search": + t.Errorf("step 13: %q is exposed over the mounted route", tool.Name) + } + } + + // 14. tools/call reaches the existing authorization and real data. + call := a.doAnonWithHeader("POST", "/mcp", map[string]any{ + "jsonrpc": "2.0", "id": 3, "method": "tools/call", + "params": map[string]any{"name": "workspace_summary", "arguments": map[string]any{}}, + }, "Authorization", "Bearer "+tokens.AccessToken) + if call.code != http.StatusOK { + t.Fatalf("step 14: tools/call = %d %s", call.code, call.body) + } + var callDoc struct { + Result struct { + IsError bool `json:"isError"` + Content []struct { + Text string `json:"text"` + } `json:"content"` + } `json:"result"` + } + mustJSON(t, call.body, &callDoc) + if callDoc.Result.IsError { + t.Fatalf("step 14: the tool refused: %s", callDoc.Result.Content[0].Text) + } +} + +// Denial must reach the client correctly and issue nothing. +func TestConsentDenialOverTheMountedRoute(t *testing.T) { + a := newOAuthAPI(t) + + reg := a.doAnon("POST", "/oauth/register", map[string]any{ + "client_name": "Deny Client", "redirect_uris": []string{"https://client.example.test/cb"}, + }) + var regDoc struct { + ClientID string `json:"client_id"` + } + mustJSON(t, reg.body, ®Doc) + + verifier := "denyVerifier0123456789abcdefghijklmnopqrstuv" + q := url.Values{ + "client_id": {regDoc.ClientID}, "redirect_uri": {"https://client.example.test/cb"}, + "response_type": {"code"}, "state": {"deny-state"}, + "code_challenge": {challengeFor(verifier)}, "code_challenge_method": {"S256"}, + "resource": {testMCPResource}, "scope": {"krow.read"}, + } + consent := a.do("GET", "/oauth/authorize?"+q.Encode(), nil) + csrf := between(consent.body, `name="csrf" value="`, `"`) + + form := url.Values{} + for k, v := range q { + form[k] = v + } + form.Set("decision", "deny") + form.Set("csrf", csrf) + + denied := a.doForm("POST", "/oauth/authorize", form) + if denied.code != http.StatusFound { + t.Fatalf("status = %d, want 302", denied.code) + } + loc, _ := url.Parse(denied.header.Get("Location")) + if got := loc.Query().Get("error"); got != "access_denied" { + t.Errorf("error = %q, want access_denied", got) + } + if got := loc.Query().Get("state"); got != "deny-state" { + t.Errorf("state = %q, want deny-state", got) + } + if loc.Query().Get("code") != "" { + t.Error("a denial issued a code") + } +} + +// /oauth/authorize is NOT public: an anonymous visitor must be sent to login. +func TestAuthorizeRequiresASession(t *testing.T) { + a := newOAuthAPI(t) + r := a.doAnon("GET", "/oauth/authorize?client_id=x", nil) + + // Either the middleware refuses it (401) or the handler redirects to + // login. Both are correct; serving a consent page is not. + if r.code == http.StatusOK && strings.Contains(r.body, "Approve") { + t.Fatal("a consent page was served to an anonymous visitor") + } +} + +/* ── Helpers ────────────────────────────────────────────────────────────── */ + +func mustJSON(t *testing.T, body string, dst any) { + t.Helper() + if err := json.Unmarshal([]byte(body), dst); err != nil { + t.Fatalf("response was not JSON: %v\nbody: %s", err, body) + } +} + +// between returns the text between two markers, or "". +func between(s, start, end string) string { + i := strings.Index(s, start) + if i < 0 { + return "" + } + rest := s[i+len(start):] + j := strings.Index(rest, end) + if j < 0 { + return "" + } + return rest[:j] +} + +/* ── Anonymous /oauth/authorize must reach the handler ──────────────────── */ + +// The regression test for the defect a live Claude Web connection exposed. +// +// /oauth/authorize was withheld from publicPaths, so the cookie middleware +// answered a signed-out visitor with its JSON 401 and the handler never ran — +// which meant the handler's redirect-to-login could never execute. A first-time +// connector user is signed out by definition, so OAuth's browser leg was +// unreachable for precisely the people who needed it. +// +// WHY THE EXISTING TESTS MISSED IT, and why this one is shaped differently: +// +// - oauth.TestAuthorizeRedirectsAnonymousToLogin drives AuthorizeHandler +// DIRECTLY, so the middleware is not in the path at all. It passed against +// broken behaviour because it never exercised the thing that was broken. +// - TestAuthorizeRequiresASession (below) asserts only that a consent page is +// not served anonymously — which a 401 satisfies perfectly well. +// +// So this one drives the MOUNTED router and asserts the POSITIVE behaviour: a +// redirect to the login, carrying the original authorization request. +func TestAnonymousAuthorizeReachesTheHandlerAndRedirectsToLogin(t *testing.T) { + a := newOAuthAPI(t) + + // A client to name, so the request is well-formed enough to get past the + // handler's own client/redirect validation and reach the session check. + reg := a.doAnon("POST", "/oauth/register", map[string]any{ + "client_name": "Anonymous Flow", "redirect_uris": []string{"https://client.example.test/cb"}, + }) + var regDoc struct { + ClientID string `json:"client_id"` + } + mustJSON(t, reg.body, ®Doc) + + verifier := "anonVerifier0123456789abcdefghijklmnopqrstu" + q := url.Values{ + "client_id": {regDoc.ClientID}, "redirect_uri": {"https://client.example.test/cb"}, + "response_type": {"code"}, "state": {"anon-state"}, + "code_challenge": {challengeFor(verifier)}, "code_challenge_method": {"S256"}, + "resource": {testMCPResource}, "scope": {"krow.read"}, + } + + // doAnon sends NO session cookie — a first-time connector user. + r := a.doAnon("GET", "/oauth/authorize?"+q.Encode(), nil) + + // The defect: the middleware's JSON 401 instead of the handler's redirect. + if r.code == http.StatusUnauthorized { + t.Fatalf("the middleware refused before the handler ran: %d %s\n"+ + "a signed-out visitor must be sent to sign in, not told 'no'", r.code, r.body) + } + if strings.Contains(r.body, `"code": "unauthorized"`) || + strings.Contains(r.body, `"code":"unauthorized"`) { + t.Fatalf("the response is the middleware's JSON 401, not the handler's: %s", r.body) + } + + if r.code != http.StatusFound { + t.Fatalf("status = %d, want 302 to the login", r.code) + } + location := r.header.Get("Location") + if !strings.HasPrefix(location, "/login?returnTo=") { + t.Fatalf("Location = %q, want a redirect to the configured login path", location) + } + + // The whole authorization request must survive the round trip, or the + // person signs in and lands nowhere. + returnTo, err := url.QueryUnescape(strings.TrimPrefix(location, "/login?returnTo=")) + if err != nil { + t.Fatalf("returnTo is not decodable: %v", err) + } + for name, want := range map[string]string{ + "path": "/oauth/authorize", + "client_id": "client_id=" + regDoc.ClientID, + "state": "state=anon-state", + "code_challenge": "code_challenge=" + challengeFor(verifier), + "code_challenge_method": "code_challenge_method=S256", + "resource": "resource=", + "redirect_uri": "redirect_uri=", + } { + if !strings.Contains(returnTo, want) { + t.Errorf("returnTo has lost the %s: %q", name, returnTo) + } + } +} + +// Listing the path must NOT hand out consent, or a code, to somebody signed +// out. "Public" here means the handler decides — not that the route is open. +func TestAnonymousAuthorizeStillGrantsNothing(t *testing.T) { + a := newOAuthAPI(t) + + reg := a.doAnon("POST", "/oauth/register", map[string]any{ + "client_name": "Nothing Granted", "redirect_uris": []string{"https://client.example.test/cb"}, + }) + var regDoc struct { + ClientID string `json:"client_id"` + } + mustJSON(t, reg.body, ®Doc) + + verifier := "nothingVerifier0123456789abcdefghijklmnopq" + q := url.Values{ + "client_id": {regDoc.ClientID}, "redirect_uri": {"https://client.example.test/cb"}, + "response_type": {"code"}, "state": {"nothing"}, + "code_challenge": {challengeFor(verifier)}, "code_challenge_method": {"S256"}, + "resource": {testMCPResource}, "scope": {"krow.read"}, + } + + // A GET must not render consent. + get := a.doAnon("GET", "/oauth/authorize?"+q.Encode(), nil) + if strings.Contains(get.body, "Approve") || strings.Contains(get.body, "Authorize access to Krow") { + t.Error("a consent page was served to a signed-out visitor") + } + + // And a POST — skipping the page entirely, as an attacker would — must not + // issue a code. The handler's session check refuses before the CSRF check + // is even relevant. + form := url.Values{} + for k, v := range q { + form[k] = v + } + form.Set("decision", "approve") + form.Set("csrf", "forged") + + post := a.doAnonForm("POST", "/oauth/authorize", form) + if loc := post.header.Get("Location"); strings.Contains(loc, "code=") { + t.Fatalf("an anonymous POST obtained an authorization code: %s", loc) + } + if post.code == http.StatusFound && strings.HasPrefix(post.header.Get("Location"), "https://client.example.test") { + t.Fatalf("an anonymous POST reached the client callback: %s", post.header.Get("Location")) + } +} diff --git a/go-api/internal/httpserver/mcplimit.go b/go-api/internal/httpserver/mcplimit.go new file mode 100644 index 0000000..9b9b148 --- /dev/null +++ b/go-api/internal/httpserver/mcplimit.go @@ -0,0 +1,230 @@ +package httpserver + +import ( + "context" + "net/http" + "strconv" + "time" + + "github.com/krow/krow-backend/go-api/internal/ratelimit" +) + +// Rate limiting for the mounted OAuth and MCP routes. +// +// A middleware rather than a change inside either package, for one reason: the +// SUBJECT of a limit is an HTTP concept. Which IP, which bearer token, which +// form field names the client — none of that is knowable from inside +// internal/oauth, and handing those packages a request so they could work it +// out would put transport details in a layer that has none. +// +// NOTHING SENSITIVE BECOMES A BUCKET KEY. Every subject goes through +// ratelimit.Subject, which hashes it. A bucket naming a bearer token would +// write that token into a table and into any log line mentioning the bucket. + +// limited wraps a handler with one rule, keyed by a subject derived per request. +// +// The subject function returns "" to mean "not limitable" — no token on the +// request, say — and the request passes through. That is correct rather than +// lax: a request with no identifiable subject is refused by the handler itself +// a moment later, and inventing a shared bucket for all of them would let one +// caller exhaust a budget that everybody else then queues behind. +func (s *Server) limited(rule ratelimit.Rule, subject func(*http.Request) string, next http.Handler) http.Handler { + if s.limiter == nil { + return next + } + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + raw := subject(r) + if raw == "" { + next.ServeHTTP(w, r) + return + } + + decision, err := s.limiter.Allow(r.Context(), rule, ratelimit.Subject(raw)) + if err != nil { + // The limiter could not count. It has already decided whether that + // permits the request — fail closed by default — and this logs the + // fault without naming the subject, which is a hash of a + // credential. + s.log.Error("rate limiter unavailable", + "rule", rule.Name, "allowed", decision.Allowed, "error", err) + } + + // Headers on every response, not only refusals, so a well-behaved + // client can slow down before it is refused rather than after. + w.Header().Set("RateLimit-Limit", strconv.Itoa(decision.Limit)) + w.Header().Set("RateLimit-Remaining", strconv.Itoa(decision.Remaining)) + + if !decision.Allowed { + // Retry-After in seconds, rounded up and never zero — "Retry-After: + // 0" invites an immediate retry, which is the one thing a limited + // client must not do. The same reasoning as retryAfterSeconds in + // ratelimit.go, and the same rounding. + w.Header().Set("Retry-After", retryAfterSeconds(decision.RetryAfter)) + s.log.Warn("rate limit exceeded", "rule", rule.Name, "path", r.URL.Path) + + w.Header().Set("Content-Type", "application/json; charset=utf-8") + w.WriteHeader(http.StatusTooManyRequests) + _, _ = w.Write([]byte(`{"error":"rate_limited",` + + `"error_description":"too many requests; retry after the interval in the Retry-After header"}`)) + return + } + + next.ServeHTTP(w, r) + }) +} + +/* ── Subjects ───────────────────────────────────────────────────────────── */ + +// byClientAddr keys by the caller's address, for endpoints with no credential. +// +// A method, not a free function, because the address is no longer a property of +// the request alone: resolving it needs the trusted-proxy set, which is wired +// onto the server. See clientip.go. +func (s *Server) byClientAddr(r *http.Request) string { return s.trust.clientAddr(r) } + +// byAddrAndUser keys the authorization endpoint by address AND signed-in user. +// +// Both, because either alone is wrong: by user only, an attacker could exhaust +// somebody else's budget by naming them; by address only, an office behind one +// NAT shares one person's allowance. +func (s *Server) byAddrAndUser(r *http.Request) string { + subject := s.trust.clientAddr(r) + if identity, ok := (sessionResolver{s}).CurrentUser(r); ok { + subject += "|" + identity.UserID + } + return subject +} + +// byFormClientID keys the token endpoint by the client_id it names. +// +// Reading a form value means parsing the body, which the handler then parses +// again — ParseForm caches on the request, so the second call is free. +func (s *Server) byFormClientID(r *http.Request) string { + if err := r.ParseForm(); err != nil { + return "" + } + if id := r.PostFormValue("client_id"); id != "" { + return id + } + // No client_id: the handler will refuse it. Fall back to the address so a + // caller cannot dodge the limit by omitting the field. + return s.trust.clientAddr(r) +} + +// byRefreshFamily keys refresh by the token being presented. +// +// Keyed by the TOKEN's hash, not the family id, because the family is not +// knowable without a database read this middleware has no business doing. The +// effect is very nearly the same: a rotation produces a new token and therefore +// a new bucket, so the practical limit is per-token-per-window rather than +// per-family — which bounds a loop just as well, since a loop presenting the +// SAME token is exactly what the limit is for. +func byRefreshToken(r *http.Request) string { + if err := r.ParseForm(); err != nil { + return "" + } + return r.PostFormValue("refresh_token") +} + +// byBearerToken keys MCP by the presented access token. +// +// The narrowest identity available on an MCP request, and the right one: it is +// one connection from one client for one user. Keying by user would let a +// person's second client eat the first's budget; keying by IP would make +// Claude's shared egress one bucket for every customer. +func byBearerToken(r *http.Request) string { + const prefix = "Bearer " + header := r.Header.Get("Authorization") + if len(header) <= len(prefix) { + return "" + } + // Case-insensitive prefix, matching mcpserver's own parsing. + if !equalFoldASCII(header[:len(prefix)], prefix) { + return "" + } + return header[len(prefix):] +} + +func equalFoldASCII(a, b string) bool { + if len(a) != len(b) { + return false + } + for i := 0; i < len(a); i++ { + ca, cb := a[i], b[i] + if 'A' <= ca && ca <= 'Z' { + ca += 'a' - 'A' + } + if 'A' <= cb && cb <= 'Z' { + cb += 'a' - 'A' + } + if ca != cb { + return false + } + } + return true +} + +// mcpLimited applies BOTH tool-call limits to the MCP endpoint. +// +// Two rules stacked rather than one, because they stop different things: the +// per-minute rule bounds a spike, and the per-hour rule bounds a slow drain +// that would sit under the per-minute rule indefinitely. Checked minute-first +// so the cheaper refusal happens earlier. +func (s *Server) mcpLimited(next http.Handler) http.Handler { + return s.limited(ratelimit.MCPToolCallPerMinute, byBearerToken, + s.limited(ratelimit.MCPToolCallPerHour, byBearerToken, next)) +} + +// tokenLimited applies the right rule for the grant type being requested. +// +// One endpoint, two grant types, two different abuse shapes — so one limit +// keyed one way would be wrong for the other. A code exchange is bounded per +// client; a refresh is bounded per presented token, so one connection looping +// cannot spend a second connection's budget. +func (s *Server) tokenLimited(next http.Handler) http.Handler { + exchange := s.limited(ratelimit.OAuthToken, s.byFormClientID, next) + refresh := s.limited(ratelimit.OAuthRefresh, byRefreshToken, next) + + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + // ParseForm caches on the request, so the handler's own call is free. + if err := r.ParseForm(); err != nil { + next.ServeHTTP(w, r) // let the handler produce the proper error + return + } + if r.PostFormValue("grant_type") == "refresh_token" { + refresh.ServeHTTP(w, r) + return + } + exchange.ServeHTTP(w, r) + }) +} + +/* ── The per-organisation ceiling ───────────────────────────────────────── */ + +// orgLimiter adapts the shared limiter to mcpserver.OrgLimiter. +// +// It is the ONE limit that cannot live in the middleware above, because 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 exactly the +// identity this surface refuses to trust. +// +// So mcpserver calls this from inside dispatch, after it has an +// authctx.Identity, and passes the org id from that identity. This type has no +// access to the request and therefore no way to be handed a different one. +type orgLimiter struct{ s *Server } + +// AllowOrg counts one call against the organisation's hourly ceiling. +// +// The org id is hashed like every other subject. It is not a secret, but the +// bucket format is uniform and a uuid in a table of counters is one more place +// a tenant identifier exists for no reason. +func (o orgLimiter) AllowOrg(ctx context.Context, orgID string) (bool, time.Duration, error) { + if o.s.limiter == nil || orgID == "" { + // No limiter, or no organisation — the latter cannot happen, because + // mcpserver refuses an identity without one before it reaches here. + return true, 0, nil + } + d, err := o.s.limiter.Allow(ctx, ratelimit.MCPPerOrgPerHour, ratelimit.Subject(orgID)) + return d.Allowed, d.RetryAfter, err +} diff --git a/go-api/internal/httpserver/proxybuckets_test.go b/go-api/internal/httpserver/proxybuckets_test.go new file mode 100644 index 0000000..fa738e9 --- /dev/null +++ b/go-api/internal/httpserver/proxybuckets_test.go @@ -0,0 +1,294 @@ +package httpserver_test + +// Multi-client rate-limit isolation, end to end. +// +// WHAT THIS IS FOR +// +// clientip_test.go proves the address RESOLVER picks the right string. It does +// not prove the string reaches Postgres as a distinct bucket, that the mounted +// route uses the resolver at all, or that the OAuth registration limit is +// actually per-client once it does. Those are different failures — a correct +// resolver wired to nothing looks identical from a unit test — and they are +// what broke in production, so they are tested here against the real handler +// and the real limiter. +// +// Every test drives httpserver.Handler() through the full middleware stack with +// a real database behind it. Nothing is stubbed. + +import ( + "fmt" + "io" + "log/slog" + "net/http" + "net/http/httptest" + "net/netip" + "strings" + "testing" + "time" + + "github.com/krow/krow-backend/go-api/internal/config" + "github.com/krow/krow-backend/go-api/internal/db" + "github.com/krow/krow-backend/go-api/internal/httpserver" + "github.com/krow/krow-backend/go-api/internal/ratelimit" + "github.com/krow/krow-backend/go-api/internal/testutil" +) + +// The proxy this fixture's deployment sits behind, and an address inside it. +const ( + proxyNetwork = "10.42.0.0/16" + proxyAddr = "10.42.0.1:9999" +) + +// proxiedAPI is newOAuthAPI with a trusted proxy configured. +// +// Deliberately not a flag on newOAuthAPI: every existing test in this package +// must keep running with an EMPTY trusted set, because that is the default +// posture and a regression in it is the thing worth catching. +type proxiedAPI struct { + handler http.Handler + h *testutil.Harness +} + +func newProxiedAPI(t *testing.T) *proxiedAPI { + t.Helper() + h := testutil.New(t) + + network, err := netip.ParsePrefix(proxyNetwork) + if err != nil { + t.Fatalf("bad test network: %v", err) + } + + cfg := &config.Config{ + AppEnv: "development", + HTTP: config.HTTPConfig{ + Host: "127.0.0.1", Port: 0, ShutdownTimeout: time.Second, + TrustedProxies: []netip.Prefix{network.Masked()}, + }, + DB: config.DBConfig{Schema: "public"}, + OAuth: config.OAuthConfig{ + Issuer: testOAuthIssuer, + Resource: testMCPResource, + LoginPath: "/login", + }, + } + srv, err := httpserver.New(cfg, &db.DB{Pool: h.Pool, Schema: "public"}, + slog.New(slog.NewTextHandler(io.Discard, nil))) + if err != nil { + t.Fatalf("build the server: %v", err) + } + return &proxiedAPI{handler: srv.Handler(), h: h} +} + +// register performs one DCR as a caller arriving via the proxy. +// +// peer is what net/http would report as RemoteAddr; forwarded is the +// X-Forwarded-For the proxy appended. An empty forwarded value sends no header. +func (a *proxiedAPI) register(t *testing.T, peer, forwarded, name string) int { + t.Helper() + body := fmt.Sprintf( + `{"client_name":%q,"redirect_uris":["https://client.example.test/cb"]}`, name) + req, err := http.NewRequest("POST", "/oauth/register", strings.NewReader(body)) + if err != nil { + t.Fatalf("build request: %v", err) + } + req.Header.Set("Content-Type", "application/json") + req.RemoteAddr = peer + if forwarded != "" { + req.Header.Set("X-Forwarded-For", forwarded) + } + rec := httptest.NewRecorder() + a.handler.ServeHTTP(rec, req) + return rec.Code +} + +// exhaust registers until the limit refuses, and returns how many succeeded. +// It stops well past the limit so a failure reports a number rather than hanging. +func (a *proxiedAPI) exhaust(t *testing.T, peer, forwarded, name string) int { + t.Helper() + allowed := 0 + for i := 0; i < ratelimit.OAuthRegister.Limit*3; i++ { + code := a.register(t, peer, forwarded, fmt.Sprintf("%s-%d", name, i)) + if code == http.StatusTooManyRequests { + return allowed + } + if code != http.StatusCreated { + t.Fatalf("%s attempt %d: unexpected status %d", name, i, code) + } + allowed++ + } + t.Fatalf("%s was never refused after %d registrations; the limit is not applied", + name, ratelimit.OAuthRegister.Limit*3) + return allowed +} + +/* ── Ten independent clients ────────────────────────────────────────────── */ + +// The headline requirement: many users behind one proxy each get their own +// budget. Before this change every one of these shared a bucket and the +// eleventh registration on the list would have been refused. +func TestTenClientsBehindOneProxyDoNotShareABucket(t *testing.T) { + a := newProxiedAPI(t) + + const clients = 10 + for i := 0; i < clients; i++ { + client := fmt.Sprintf("203.0.113.%d", i+1) + // Each client registers TWICE — Claude mints a new client per connect, + // so a reconnect must not count against anybody else. + for attempt := 0; attempt < 2; attempt++ { + code := a.register(t, proxyAddr, client+", "+"10.42.0.1", + fmt.Sprintf("client-%d-%d", i, attempt)) + if code != http.StatusCreated { + t.Fatalf("client %s attempt %d: status %d, want 201 — clients are sharing a bucket", + client, attempt, code) + } + } + } + + // Twenty registrations went through on a limit of ten per subject, which is + // only possible if the subject really is the forwarded client. + var buckets int + if err := a.h.Pool.QueryRow(t.Context(), + `SELECT count(*) FROM rate_limits WHERE bucket LIKE 'oauth.register:%'`).Scan(&buckets); err != nil { + t.Fatalf("count buckets: %v", err) + } + if buckets != clients { + t.Errorf("%d distinct oauth.register buckets, want %d — one per client", buckets, clients) + } +} + +/* ── One client's exhaustion is its own ─────────────────────────────────── */ + +// Client A burns its whole budget; client B is unaffected. This is the property +// that failed in production, where A's retries refused B outright. +func TestOneClientExhaustingDoesNotBlockAnother(t *testing.T) { + a := newProxiedAPI(t) + + const clientA, clientB = "203.0.113.50", "203.0.113.51" + + allowed := a.exhaust(t, proxyAddr, clientA, "A") + if allowed != ratelimit.OAuthRegister.Limit { + t.Errorf("client A got %d registrations, want %d", allowed, ratelimit.OAuthRegister.Limit) + } + + // A is now refused. + if code := a.register(t, proxyAddr, clientA, "A-again"); code != http.StatusTooManyRequests { + t.Errorf("client A after exhausting: status %d, want 429", code) + } + // B is not. + if code := a.register(t, proxyAddr, clientB, "B"); code != http.StatusCreated { + t.Errorf("client B: status %d, want 201 — A's exhaustion blocked B", code) + } +} + +// A single client reconnecting repeatedly spends only its own budget, which is +// what Claude actually does: a new DCR client on every connect. +func TestRepeatedReconnectConsumesOnlyThatClientsBudget(t *testing.T) { + a := newProxiedAPI(t) + + const reconnecting = "203.0.113.60" + a.exhaust(t, proxyAddr, reconnecting, "reconnector") + + // Nine other clients are untouched by it. + for i := 0; i < 9; i++ { + other := fmt.Sprintf("198.51.100.%d", i+1) + if code := a.register(t, proxyAddr, other, fmt.Sprintf("other-%d", i)); code != http.StatusCreated { + t.Fatalf("client %s: status %d, want 201", other, code) + } + } +} + +/* ── Spoofing still buys nothing ────────────────────────────────────────── */ + +// A caller reaching the API directly — not through the proxy — cannot escape +// its bucket by varying X-Forwarded-For. It gets ONE budget however many +// different values it sends. +func TestUntrustedCallerCannotEscapeItsBucketByForging(t *testing.T) { + a := newProxiedAPI(t) + + const direct = "198.51.100.200:40000" // outside proxyNetwork + + allowed := 0 + for i := 0; i < ratelimit.OAuthRegister.Limit*2; i++ { + // A different forged client on every single request. + code := a.register(t, direct, fmt.Sprintf("203.0.113.%d", i+100), fmt.Sprintf("forger-%d", i)) + if code == http.StatusTooManyRequests { + break + } + if code != http.StatusCreated { + t.Fatalf("attempt %d: unexpected status %d", i, code) + } + allowed++ + } + + if allowed != ratelimit.OAuthRegister.Limit { + t.Errorf("a forging caller got %d registrations, want %d — the header bought extra budget", + allowed, ratelimit.OAuthRegister.Limit) + } + + var buckets int + if err := a.h.Pool.QueryRow(t.Context(), + `SELECT count(*) FROM rate_limits WHERE bucket LIKE 'oauth.register:%'`).Scan(&buckets); err != nil { + t.Fatalf("count buckets: %v", err) + } + if buckets != 1 { + t.Errorf("a forging caller produced %d buckets, want exactly 1", buckets) + } +} + +// Claiming to be the trusted proxy does not make a caller trusted. +func TestClaimingToBeTheProxyDoesNotWork(t *testing.T) { + a := newProxiedAPI(t) + + const direct = "198.51.100.201:40000" + // The forged chain ends in the proxy's own address, which is what an + // attacker who has read this file would try. + if code := a.register(t, direct, "203.0.113.9, 10.42.0.1", "impostor"); code != http.StatusCreated { + t.Fatalf("setup: status %d", code) + } + + var bucket string + if err := a.h.Pool.QueryRow(t.Context(), + `SELECT bucket FROM rate_limits WHERE bucket LIKE 'oauth.register:%'`).Scan(&bucket); err != nil { + t.Fatalf("read bucket: %v", err) + } + // The bucket must be the DIRECT caller's address, not the forged one. + want := "oauth.register:" + ratelimit.Subject("198.51.100.201") + if bucket != want { + t.Errorf("bucket = %q, want %q — a forged chain was believed", bucket, want) + } +} + +/* ── The default posture is unchanged ───────────────────────────────────── */ + +// With no trusted proxies configured — the default, and how every other test in +// this package runs — the forwarded header is ignored and callers share the +// peer's bucket exactly as before. +func TestWithoutTrustedProxiesCallersShareThePeerBucket(t *testing.T) { + a := newOAuthAPI(t) // no TrustedProxies in its config + + for i := 0; i < 3; i++ { + body := fmt.Sprintf( + `{"client_name":"unproxied-%d","redirect_uris":["https://client.example.test/cb"]}`, i) + req, err := http.NewRequest("POST", "/oauth/register", strings.NewReader(body)) + if err != nil { + t.Fatalf("build request: %v", err) + } + req.Header.Set("Content-Type", "application/json") + req.RemoteAddr = "192.0.2.10:5000" + req.Header.Set("X-Forwarded-For", fmt.Sprintf("203.0.113.%d", i+1)) + rec := httptest.NewRecorder() + a.handler.ServeHTTP(rec, req) + if rec.Code != http.StatusCreated { + t.Fatalf("attempt %d: status %d", i, rec.Code) + } + } + + var buckets int + if err := a.h.Pool.QueryRow(t.Context(), + `SELECT count(*) FROM rate_limits WHERE bucket LIKE 'oauth.register:%'`).Scan(&buckets); err != nil { + t.Fatalf("count buckets: %v", err) + } + if buckets != 1 { + t.Errorf("%d buckets with no trusted proxy configured, want 1 — the header was read", buckets) + } +} diff --git a/go-api/internal/httpserver/ratelimit.go b/go-api/internal/httpserver/ratelimit.go index b25322b..7265621 100644 --- a/go-api/internal/httpserver/ratelimit.go +++ b/go-api/internal/httpserver/ratelimit.go @@ -1,9 +1,6 @@ package httpserver import ( - "net" - "net/http" - "strings" "sync" "time" ) @@ -23,12 +20,13 @@ import ( // one instance needs shared state — Redis, or the database — and this // package is the seam where that goes: attemptLimiter is an implementation // detail behind Allow/Fail/Reset. -// - It trusts net/http's RemoteAddr for the client address. Behind a reverse -// proxy every request appears to come from the proxy, so the per-address -// budget becomes global. Reading X-Forwarded-For instead would be worse, -// not better, until there is a trusted-proxy list to validate it against — -// a client can send that header itself and mint a fresh budget per request. -// Deploying behind a proxy means adding that list first. +// - The per-address budget is only as good as the address. That used to be +// net/http's RemoteAddr, which behind a reverse proxy is the proxy on +// every request and makes this budget global. It is now resolved by +// proxyTrust.clientAddr (clientip.go), which reads a forwarded address +// when — and only when — the immediate peer is a configured trusted +// proxy. An unconfigured deployment still gets RemoteAddr, so a proxied +// deployment must set HTTP_TRUSTED_PROXIES for this limit to be per-user. // - It is memory-bounded by pruning, not by a hard cap, so a flood from many // distinct addresses grows the map until the next prune. // @@ -139,16 +137,3 @@ func (l *attemptLimiter) pruneLocked(now time.Time) { } } } - -// clientAddr is the key for per-address limiting. -// -// The port is stripped: a browser uses a new source port for every connection, -// so keying on host:port would give each attempt its own budget and limit -// nothing at all. -func clientAddr(r *http.Request) string { - host, _, err := net.SplitHostPort(strings.TrimSpace(r.RemoteAddr)) - if err != nil { - return strings.TrimSpace(r.RemoteAddr) - } - return host -} diff --git a/go-api/internal/httpserver/server.go b/go-api/internal/httpserver/server.go index 565217f..66bb0b9 100644 --- a/go-api/internal/httpserver/server.go +++ b/go-api/internal/httpserver/server.go @@ -37,6 +37,7 @@ import ( "github.com/krow/krow-backend/go-api/internal/db" "github.com/krow/krow-backend/go-api/internal/definition" "github.com/krow/krow-backend/go-api/internal/knowledge" + "github.com/krow/krow-backend/go-api/internal/ratelimit" "github.com/krow/krow-backend/go-api/internal/runtime" "github.com/krow/krow-backend/go-api/internal/service" "github.com/krow/krow-backend/go-api/internal/tools" @@ -83,7 +84,18 @@ type Server struct { users auth.UserStore credentials *auth.Credentials loginByEmail *attemptLimiter - loginByAddr *attemptLimiter + + // limiter bounds the OAuth and MCP routes, shared across instances via + // Postgres. Nil when those routes are not registered — see mcplimit.go, + // where a nil limiter means the middleware is not installed at all rather + // than installed and permissive. + limiter *ratelimit.Limiter + loginByAddr *attemptLimiter + + // trust resolves a request to the address its per-address limits are keyed + // by, reading a forwarded address only from a configured proxy. Wired once + // here so no handler can be given a different notion of who called. + trust proxyTrust // now is injectable so tests can drive expiry without sleeping. now func() time.Time @@ -223,6 +235,7 @@ func New(cfg *config.Config, database *db.DB, log *slog.Logger, opts ...Option) credentials: auth.NewCredentials(users), loginByEmail: newAttemptLimiter(o.perEmail, o.loginWindow, o.now), loginByAddr: newAttemptLimiter(o.perAddress, o.loginWindow, o.now), + trust: newProxyTrust(cfg.HTTP.TrustedProxies), now: o.now, } @@ -277,11 +290,23 @@ func New(cfg *config.Config, database *db.DB, log *slog.Logger, opts ...Option) } s.definitions = s.definitions.WithCuratedAgents(curated) + // The shared limiter, built only when the routes that use it exist. The + // existing in-process login limiter is untouched: it guards a different + // thing (failed password attempts) with a different model (count failures, + // reset on success), and replacing it is not this change's business. + if cfg.OAuth.Enabled() { + s.limiter = ratelimit.New(database.Pool) + } + mux := http.NewServeMux() mux.HandleFunc("GET /health", s.handleHealth) s.endpoints = s.routeAuth(mux) + s.routeResources(mux) + s.routeMe(mux) + s.routeDefinitions(mux) + s.routeWorkflows(mux) + s.routeOwliver(mux) + - s.routeRuns(mux) + s.routeVersion(mux) + s.routeTools(mux) + s.routeRuns(mux) + s.routeVersion(mux) + s.routeTools(mux) + + // The MCP surface and the OAuth server behind it. Both return 0 and + // register nothing when OAUTH_ISSUER and MCP_RESOURCE are unset, which + // is every deployment that has not asked for them. + s.routeOAuth(mux) + s.routeMCP(mux) handler := jsonErrors(mux) // Authentication sits where devOrgMiddleware used to, so every route below diff --git a/go-api/internal/mcpserver/auth.go b/go-api/internal/mcpserver/auth.go new file mode 100644 index 0000000..1663c99 --- /dev/null +++ b/go-api/internal/mcpserver/auth.go @@ -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" + } +} diff --git a/go-api/internal/mcpserver/auth_test.go b/go-api/internal/mcpserver/auth_test.go new file mode 100644 index 0000000..ec23e65 --- /dev/null +++ b/go-api/internal/mcpserver/auth_test.go @@ -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 +} diff --git a/go-api/internal/mcpserver/jsonrpc.go b/go-api/internal/mcpserver/jsonrpc.go new file mode 100644 index 0000000..c606729 --- /dev/null +++ b/go-api/internal/mcpserver/jsonrpc.go @@ -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 +} diff --git a/go-api/internal/mcpserver/mcpserver_test.go b/go-api/internal/mcpserver/mcpserver_test.go new file mode 100644 index 0000000..7595368 --- /dev/null +++ b/go-api/internal/mcpserver/mcpserver_test.go @@ -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) + } +} diff --git a/go-api/internal/mcpserver/orglimit_test.go b/go-api/internal/mcpserver/orglimit_test.go new file mode 100644 index 0000000..e1e27df --- /dev/null +++ b/go-api/internal/mcpserver/orglimit_test.go @@ -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):] + } +} diff --git a/go-api/internal/mcpserver/server.go b/go-api/internal/mcpserver/server.go new file mode 100644 index 0000000..7517d86 --- /dev/null +++ b/go-api/internal/mcpserver/server.go @@ -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) +} diff --git a/go-api/internal/mcpserver/tenant_test.go b/go-api/internal/mcpserver/tenant_test.go new file mode 100644 index 0000000..c2ef676 --- /dev/null +++ b/go-api/internal/mcpserver/tenant_test.go @@ -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) +} diff --git a/go-api/internal/mcpserver/tools.go b/go-api/internal/mcpserver/tools.go new file mode 100644 index 0000000..4f79824 --- /dev/null +++ b/go-api/internal/mcpserver/tools.go @@ -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), + }, + } +} diff --git a/go-api/internal/mcpserver/transport.go b/go-api/internal/mcpserver/transport.go new file mode 100644 index 0000000..9db292a --- /dev/null +++ b/go-api/internal/mcpserver/transport.go @@ -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) +} diff --git a/go-api/internal/oauth/abuse_test.go b/go-api/internal/oauth/abuse_test.go new file mode 100644 index 0000000..5139960 --- /dev/null +++ b/go-api/internal/oauth/abuse_test.go @@ -0,0 +1,374 @@ +package oauth + +import ( + "context" + "net/http" + "net/http/httptest" + "net/url" + "strings" + "sync" + "testing" +) + +// Abuse cases and input limits. +// +// The property every test here shares: a refusal must not teach the caller +// anything. Not whether a token existed, not whether an account is suspended +// rather than deleted, not what the database is called, and never the value of +// a credential that was presented. + +/* ── Nothing sensitive reaches a response ───────────────────────────────── */ + +// The broadest check in this file: drive every failure path with known secret +// values, and assert none of them comes back. +func TestNoSecretEverAppearsInAResponse(t *testing.T) { + h := newHarness(t) + clientID := h.register() + + const ( + secretVerifier = "SENTINELverifier0123456789abcdefghijklmnop" + secretCode = "SENTINELcodevalue" + secretToken = "SENTINELtokenvalue" + ) + + bodies := map[string]string{} + + // A failed exchange, with a sentinel code and verifier. + bodies["bad code"] = h.exchange(clientID, secretCode, secretVerifier).Body.String() + + // A real code with the wrong verifier. + real := h.authorizeOK(clientID, verifier43) + bodies["bad verifier"] = h.exchange(clientID, real, secretVerifier).Body.String() + + // A refresh with a sentinel token. + bodies["bad refresh"] = h.token(url.Values{ + "grant_type": {"refresh_token"}, "refresh_token": {secretToken}, "client_id": {clientID}, + }).Body.String() + + // Revocation of an unknown token. + form := url.Values{"token": {secretToken}} + req := httptest.NewRequest(http.MethodPost, "/oauth/revoke", strings.NewReader(form.Encode())) + req.Header.Set("Content-Type", "application/x-www-form-urlencoded") + rec := httptest.NewRecorder() + h.server.RevokeHandler().ServeHTTP(rec, req) + bodies["revoke unknown"] = rec.Body.String() + + for where, body := range bodies { + for what, secret := range map[string]string{ + "code_verifier": secretVerifier, + "authorization code": secretCode, + "token": secretToken, + } { + if strings.Contains(body, secret) { + t.Errorf("%s: the response echoes the presented %s:\n%s", where, what, body) + } + } + // Nor may it leak the shape of the system. + for _, tell := range []string{"SQLSTATE", "pq:", "pgx", "oauth_tokens", "oauth_grants", + "password", "Krow-force", "relation", "column"} { + if strings.Contains(body, tell) { + t.Errorf("%s: the response leaks an internal detail (%q):\n%s", where, tell, body) + } + } + } +} + +/* ── Replay ─────────────────────────────────────────────────────────────── */ + +// Ten attempts to spend one code. Exactly one may succeed. +func TestAuthorizationCodeReplayUnderLoad(t *testing.T) { + h := newHarness(t) + clientID := h.register() + code := h.authorizeOK(clientID, verifier43) + + succeeded := 0 + for i := 0; i < 10; i++ { + if h.exchange(clientID, code, verifier43).Code == http.StatusOK { + succeeded++ + } + } + if succeeded != 1 { + t.Errorf("%d of 10 exchanges of the same code succeeded, want exactly 1", succeeded) + } +} + +// The same, concurrently. A single-use credential redeemed by two racing +// callers must be spent exactly once — this is the property that +// UPDATE … RETURNING buys, and the one a SELECT-then-UPDATE would lose. +// Run with -race. +func TestConcurrentCodeRedemptionSpendsItOnce(t *testing.T) { + h := newHarness(t) + clientID := h.register() + code := h.authorizeOK(clientID, verifier43) + + const racers = 8 + var wg sync.WaitGroup + var mu sync.Mutex + succeeded := 0 + + for i := 0; i < racers; i++ { + wg.Add(1) + go func() { + defer wg.Done() + if h.exchange(clientID, code, verifier43).Code == http.StatusOK { + mu.Lock() + succeeded++ + mu.Unlock() + } + }() + } + wg.Wait() + + if succeeded != 1 { + t.Errorf("%d of %d concurrent redemptions succeeded, want exactly 1", succeeded, racers) + } +} + +// Concurrent refresh of the same token: one rotation, not several. Two +// successes would mean two live families from one credential. +// Run with -race. +func TestConcurrentRefreshRotatesOnce(t *testing.T) { + h := newHarness(t) + clientID := h.register() + code := h.authorizeOK(clientID, verifier43) + tokens := decodeTokens(t, h.exchange(clientID, code, verifier43)) + + const racers = 8 + var wg sync.WaitGroup + var mu sync.Mutex + succeeded := 0 + + for i := 0; i < racers; i++ { + wg.Add(1) + go func() { + defer wg.Done() + rec := h.token(url.Values{ + "grant_type": {"refresh_token"}, "refresh_token": {tokens.RefreshToken}, + "client_id": {clientID}, + }) + if rec.Code == http.StatusOK { + mu.Lock() + succeeded++ + mu.Unlock() + } + }() + } + wg.Wait() + + if succeeded != 1 { + t.Errorf("%d of %d concurrent refreshes succeeded, want exactly 1 — "+ + "a refresh token was double-spent", succeeded, racers) + } +} + +/* ── Open redirect ──────────────────────────────────────────────────────── */ + +// Every shape of redirect tampering, each of which has been a real CVE +// somewhere. None may be honoured, and none may be answered WITH a redirect — +// redirecting an error to an unvalidated URI is the open redirect itself. +func TestOpenRedirectAttempts(t *testing.T) { + h := newHarness(t) + clientID := h.register() + + for name, redirect := range map[string]string{ + "different host": "https://attacker.example/cb", + "prefix extension": testRedirect + ".attacker.example", + "path traversal": testRedirect + "/../../evil", + "userinfo trick": "https://claude.example.test@attacker.example/cb", + "added query": testRedirect + "?next=https://attacker.example", + "protocol swap": strings.Replace(testRedirect, "https", "http", 1), + "case variation": strings.ToUpper(testRedirect), + "trailing slash": testRedirect + "/", + "double slash": "//attacker.example/cb", + "encoded traversal": testRedirect + "/%2e%2e/evil", + "null byte": testRedirect + "\x00.attacker.example", + "newline injection": testRedirect + "\nLocation: https://attacker.example", + "javascript": "javascript:alert(1)", + "completely missing": "", + } { + t.Run(name, func(t *testing.T) { + rec := h.authorize(map[string]string{ + "client_id": clientID, "redirect_uri": redirect, "response_type": "code", + "state": "xyz", "code_challenge": ChallengeFor(verifier43), + "code_challenge_method": "S256", "resource": testResource, + }) + + if rec.Code == http.StatusFound { + location := rec.Header().Get("Location") + t.Fatalf("answered with a redirect to %q — an unregistered target "+ + "must produce a direct error, never a redirect", location) + } + if rec.Code != http.StatusBadRequest { + t.Errorf("status = %d, want 400", rec.Code) + } + if strings.Contains(rec.Body.String(), "attacker.example") { + t.Error("the error echoes the attacker's host back") + } + }) + } +} + +/* ── Malformed and oversized input ──────────────────────────────────────── */ + +func TestRegistrationRejectsMalformedBodies(t *testing.T) { + h := newHarness(t) + + for name, body := range map[string]string{ + "not json": `not json at all`, + "truncated": `{"client_name":`, + "null": `null`, + "array": `[]`, + "deeply nested": `{"client_name":` + strings.Repeat(`[`, 2000) + strings.Repeat(`]`, 2000) + `}`, + "empty": ``, + "wrong types": `{"client_name":123,"redirect_uris":"not-an-array"}`, + "huge name": `{"client_name":"` + strings.Repeat("A", 100_000) + `","redirect_uris":["https://a.test/cb"]}`, + "too many uris": `{"client_name":"x","redirect_uris":[` + strings.TrimSuffix(strings.Repeat(`"https://a.test/cb",`, 50), ",") + `]}`, + "oversized body": `{"client_name":"` + strings.Repeat("A", 20<<10) + `"}`, + } { + t.Run(name, func(t *testing.T) { + req := httptest.NewRequest(http.MethodPost, "/oauth/register", strings.NewReader(body)) + rec := httptest.NewRecorder() + // A panic here fails the test by crashing it, which is the assertion. + h.server.RegisterHandler().ServeHTTP(rec, req) + + if rec.Code == http.StatusCreated { + // Only the "huge name" case may legitimately succeed, truncated. + if name != "huge name" { + t.Errorf("status = %d; a malformed registration was accepted", rec.Code) + } + return + } + if rec.Code < 400 || rec.Code >= 500 { + t.Errorf("status = %d, want a 4xx", rec.Code) + } + }) + } +} + +// A client name is stored, shown on a consent screen, and attacker-controlled. +// It must be bounded, or registration becomes free storage. +func TestClientNameIsBounded(t *testing.T) { + h := newHarness(t) + body := `{"client_name":"` + strings.Repeat("A", 5000) + `","redirect_uris":["https://a.test/cb"]}` + + rec := httptest.NewRecorder() + h.server.RegisterHandler().ServeHTTP(rec, + httptest.NewRequest(http.MethodPost, "/oauth/register", strings.NewReader(body))) + if rec.Code != http.StatusCreated { + t.Fatalf("status = %d", rec.Code) + } + + var stored string + if err := h.h.Pool.QueryRow(context.Background(), + `SELECT client_name FROM oauth_clients ORDER BY created_date DESC LIMIT 1`).Scan(&stored); err != nil { + t.Fatalf("read: %v", err) + } + if len(stored) > 200 { + t.Errorf("stored client_name is %d characters; the column's CHECK allows 200", len(stored)) + } +} + +func TestTokenEndpointRejectsMalformedRequests(t *testing.T) { + h := newHarness(t) + + for name, tc := range map[string]struct { + body string + contentType string + }{ + "no content type": {"grant_type=authorization_code", ""}, + "json body": {`{"grant_type":"authorization_code"}`, "application/json"}, + "empty": {"", "application/x-www-form-urlencoded"}, + "garbage": {"%%%%", "application/x-www-form-urlencoded"}, + "huge": {"grant_type=authorization_code&code=" + strings.Repeat("A", 200_000), "application/x-www-form-urlencoded"}, + "repeated params": {"grant_type=authorization_code&grant_type=password", "application/x-www-form-urlencoded"}, + "null grant": {"grant_type=", "application/x-www-form-urlencoded"}, + "unknown grant": {"grant_type=magic", "application/x-www-form-urlencoded"}, + "injection in code": {"grant_type=authorization_code&code=' OR 1=1 --&client_id=x&redirect_uri=y&code_verifier=" + + strings.Repeat("a", 43), "application/x-www-form-urlencoded"}, + } { + t.Run(name, func(t *testing.T) { + req := httptest.NewRequest(http.MethodPost, "/oauth/token", strings.NewReader(tc.body)) + if tc.contentType != "" { + req.Header.Set("Content-Type", tc.contentType) + } + rec := httptest.NewRecorder() + h.server.TokenHandler().ServeHTTP(rec, req) + + if rec.Code == http.StatusOK { + t.Errorf("a malformed token request succeeded: %s", rec.Body.String()) + } + if rec.Code >= 500 { + t.Errorf("status = %d; a malformed request must not be an internal error: %s", + rec.Code, rec.Body.String()) + } + }) + } +} + +/* ── Scope escalation ───────────────────────────────────────────────────── */ + +// krow.write must be unreachable from every angle: registration, authorization, +// and the consent POST. +func TestWriteScopeIsUnreachable(t *testing.T) { + h := newHarness(t) + + t.Run("at registration", func(t *testing.T) { + body := `{"client_name":"x","redirect_uris":["` + testRedirect + `"],"scope":"` + ScopeWrite + `"}` + rec := httptest.NewRecorder() + h.server.RegisterHandler().ServeHTTP(rec, + httptest.NewRequest(http.MethodPost, "/oauth/register", strings.NewReader(body))) + if rec.Code == http.StatusCreated { + t.Error("a client registered for krow.write") + } + }) + + t.Run("at authorization", func(t *testing.T) { + clientID := h.register() + rec := h.authorize(map[string]string{ + "client_id": clientID, "redirect_uri": testRedirect, "response_type": "code", + "state": "xyz", "code_challenge": ChallengeFor(verifier43), + "code_challenge_method": "S256", "resource": testResource, + "scope": ScopeRead + " " + ScopeWrite, + }) + loc, _ := url.Parse(rec.Header().Get("Location")) + if loc.Query().Get("code") != "" { + t.Error("an authorization requesting krow.write produced a code") + } + }) + + t.Run("no issued token carries it", func(t *testing.T) { + clientID := h.register() + code := h.authorizeOK(clientID, verifier43) + tokens := decodeTokens(t, h.exchange(clientID, code, verifier43)) + if strings.Contains(tokens.Scope, ScopeWrite) { + t.Errorf("an issued token carries %q", tokens.Scope) + } + }) +} + +/* ── Cache and transport headers ────────────────────────────────────────── */ + +// A credential-bearing response must never be cached, and no endpoint may put +// a token in a URL. +func TestSensitiveResponsesAreNotCacheable(t *testing.T) { + h := newHarness(t) + clientID := h.register() + code := h.authorizeOK(clientID, verifier43) + + rec := h.exchange(clientID, code, verifier43) + for header, want := range map[string]string{ + "Cache-Control": "no-store", + "Pragma": "no-cache", + } { + if got := rec.Header().Get(header); !strings.Contains(got, want) { + t.Errorf("%s = %q, want it to contain %q", header, got, want) + } + } + + // The authorization redirect carries a code in its query — that is the + // protocol — but it must never carry a token. + approved := h.authorize(authorizeParamsFor(clientID, verifier43)) + if strings.Contains(approved.Header().Get("Location"), "access_token") { + t.Error("an access token appeared in a redirect URL") + } +} diff --git a/go-api/internal/oauth/authenticator.go b/go-api/internal/oauth/authenticator.go new file mode 100644 index 0000000..fe89366 --- /dev/null +++ b/go-api/internal/oauth/authenticator.go @@ -0,0 +1,172 @@ +package oauth + +import ( + "context" + "crypto/rand" + "encoding/hex" + "errors" + "fmt" + "log/slog" + "strings" + + "github.com/krow/krow-backend/go-api/internal/auth" + "github.com/krow/krow-backend/go-api/internal/authctx" +) + +// Authenticator is the production implementation of +// mcpserver.TokenAuthenticator. +// +// This is where Phase 2's seam is filled in, and the shape of it is the whole +// argument for having defined the interface first: one method, taking a raw +// token, returning the same authctx.Identity a cookie produces. Nothing +// downstream — not tools.Context, not the policy table, not a single handler — +// can tell which path built the identity, so authorization cannot drift between +// them. +// +// THE IDENTITY IS BUILT FROM THE USER ROW, NOT FROM THE TOKEN. +// +// oauth_tokens carries org_id, and it would be cheaper to read it from there. +// It is deliberately not: the token row records the tenant AT ISSUE TIME, and a +// token can outlive the fact. A user moved to another organisation, or +// suspended, would keep working against a stale claim until the token expired. +// Re-reading the user costs one indexed lookup and makes suspension take effect +// on the next call — which is exactly what httpserver/auth.go already does for +// cookies, and the bearer path must not be weaker than the cookie path. +type Authenticator struct { + store *Store + users UserLookup + log *slog.Logger + + // audience is this deployment's canonical MCP resource URI. A token whose + // audience is anything else is refused — see the note in Authenticate. + audience string +} + +// UserLookup is the subset of the existing user store this needs. auth.UserStore +// satisfies it; nothing here builds a second user table or password store. +type UserLookup interface { + FindByID(ctx context.Context, id string) (auth.User, error) +} + +// NewAuthenticator builds the production token authenticator. +func NewAuthenticator(store *Store, users UserLookup, audience string, log *slog.Logger) *Authenticator { + if log == nil { + log = slog.Default() + } + return &Authenticator{store: store, users: users, audience: audience, log: log} +} + +// ErrAudienceMismatch is internal. It never reaches a client — see the single +// return below — but it is distinct so the log can say what happened. +var ErrAudienceMismatch = errors.New("oauth: token audience does not match this resource") + +// Authenticate resolves a bearer token into a KROW identity. +// +// EVERY failure returns the same error. Unknown, expired, revoked, wrong +// audience, suspended user, deleted user — one answer, because a caller who can +// tell them apart learns things they should not: that a token once existed, +// that an account was suspended rather than deleted, that this server is not +// the intended audience for a token they hold. Same discipline as +// tools.Denied() and the session path's identical answer to "not found" and +// "expired". +// +// The reason goes to the log, at warn, where the operator is. +func (a *Authenticator) Authenticate(ctx context.Context, rawToken string) (authctx.Identity, error) { + if strings.TrimSpace(rawToken) == "" { + return authctx.Identity{}, ErrTokenUnusable + } + + // 1. The token must exist, be an access token, be unexpired and unrevoked. + // All four are in the query's predicate. + token, err := a.store.FindAccessToken(ctx, rawToken) + if err != nil { + a.log.Warn("mcp bearer refused", "reason", "token_unusable") + return authctx.Identity{}, ErrTokenUnusable + } + + // 2. Audience. RFC 8707 and the MCP spec both require a server to verify + // that a token was issued FOR IT. Without this check, a token minted by + // this authorization server for some other resource would be spendable + // here — the confused-deputy problem the spec calls out explicitly. The + // comparison is against configuration, never against anything in the + // request: a resource value supplied by the caller would let the caller + // choose their own audience. + if token.Audience != a.audience { + a.log.Warn("mcp bearer refused", + "reason", "audience_mismatch", + "token_id", token.ID, + "expected", a.audience, + "presented", token.Audience) + return authctx.Identity{}, ErrTokenUnusable + } + + // 3. Scope. krow.read is the only scope this phase issues, and the MCP + // surface is read-only, so a token without it has no business here. The + // check is present rather than implied so that adding krow.write later + // is a change in one place. + if !hasScope(token.Scopes, ScopeRead) { + a.log.Warn("mcp bearer refused", "reason", "missing_scope", "token_id", token.ID) + return authctx.Identity{}, ErrTokenUnusable + } + + // 4. The user, re-read live. See the type comment for why this is not taken + // from the token row. + user, err := a.users.FindByID(ctx, token.UserID) + if err != nil { + // The FK cascades, so a missing user should be unreachable. If it + // happens the token is orphaned and worth killing. + a.log.Warn("mcp bearer refused", "reason", "user_missing", "token_id", token.ID) + _ = a.store.RevokeFamily(ctx, token.FamilyID, "user_missing") + return authctx.Identity{}, ErrTokenUnusable + } + + // 5. Suspension revokes on contact, exactly as the cookie path does. Not + // "the token stops working at expiry" — a suspended account must lose + // access on its next request, and leaving the family alive would mean it + // kept a working credential for up to thirty days. + if !user.IsActive() { + a.log.Warn("mcp bearer refused", + "reason", "user_inactive", "user_id", user.ID, "status", user.Status) + _ = a.store.RevokeFamily(ctx, token.FamilyID, "user_suspended") + return authctx.Identity{}, ErrTokenUnusable + } + + // The same construction httpserver/auth.go performs for a cookie. SessionID + // and ExpiresAt are deliberately left zero: there is no session row behind + // this identity, and inventing one would make a token look like something + // logout could end. + return authctx.Identity{ + UserID: user.ID, + OrgID: user.OrgID, + Email: user.Email, + FullName: user.FullName, + Role: user.Role, + AccountType: user.AccountType, + Status: user.Status, + }, nil +} + +// hasScope reports whether a scope was granted. +func hasScope(granted []string, want string) bool { + for _, s := range granted { + if s == want { + return true + } + } + return false +} + +// newUUID returns a random UUID v4 string, for family ids. +// +// Hand-rolled rather than adding a dependency: the module is stdlib plus pgx, +// and one 16-byte read with two bits set is not worth a third-party package. +func newUUID() (string, error) { + var b [16]byte + if _, err := rand.Read(b[:]); err != nil { + return "", fmt.Errorf("oauth: generate uuid: %w", err) + } + b[6] = (b[6] & 0x0f) | 0x40 // version 4 + b[8] = (b[8] & 0x3f) | 0x80 // variant 10 + h := hex.EncodeToString(b[:]) + return h[0:8] + "-" + h[8:12] + "-" + h[12:16] + "-" + h[16:20] + "-" + h[20:32], nil +} diff --git a/go-api/internal/oauth/authserver.go b/go-api/internal/oauth/authserver.go new file mode 100644 index 0000000..450771a --- /dev/null +++ b/go-api/internal/oauth/authserver.go @@ -0,0 +1,857 @@ +package oauth + +import ( + "crypto/hmac" + "crypto/rand" + "crypto/sha256" + "encoding/hex" + "encoding/json" + "errors" + "log/slog" + "net/http" + "net/url" + "strings" + + "github.com/krow/krow-backend/go-api/internal/authctx" +) + +// The authorization server's HTTP surface: register, authorize, token, revoke. +// +// HOW A PERSON IS AUTHENTICATED HERE +// +// They are not, by this package. The authorization endpoint requires a KROW +// user to already be signed in, and it learns who that is from the SessionResolver +// the server was built with — which the HTTP layer implements using the +// existing cookie session. There is no second password store, no second login +// form, and no credential of any kind in this package. +// +// That is also why the authorization endpoint is the only part of OAuth that +// touches cookies: it runs in a browser, as a person, mid-redirect. Everything +// after it — the token endpoint, the MCP endpoint — is a back-channel call from +// the client and uses no cookie at all. + +// SessionResolver reports who is signed in, for the authorization endpoint. +// +// Implemented by the HTTP layer over the existing session manager. An interface +// rather than a direct dependency so this package does not reach into +// httpserver, and so a test can drive the flow without a browser. +type SessionResolver interface { + // CurrentUser returns the signed-in identity, or false when there is none. + CurrentUser(r *http.Request) (authctx.Identity, bool) +} + +// Server is the OAuth authorization server. +type Server struct { + cfg Config + store *Store + sessions SessionResolver + log *slog.Logger + + // loginPath is where an unauthenticated person is sent, with a return + // target, so they can sign in and come back to the consent screen. + loginPath string + + // csrfKey signs consent-form tokens. Per-process and never persisted — + // see csrfFor. + csrfKey []byte +} + +// NewServer builds the authorization server. +func NewServer(cfg Config, store *Store, sessions SessionResolver, loginPath string, log *slog.Logger) *Server { + if log == nil { + log = slog.Default() + } + if loginPath == "" { + loginPath = "/login" + } + key := make([]byte, 32) + if _, err := rand.Read(key); err != nil { + // Unreachable short of the OS entropy source failing. Panicking is + // correct: a server that cannot generate a CSRF key cannot render a + // consent form safely, and starting without one would mean serving a + // form nothing protects. + panic("oauth: could not generate a consent CSRF key: " + err.Error()) + } + return &Server{ + cfg: cfg.Normalise(), + store: store, + sessions: sessions, + log: log, + loginPath: loginPath, + csrfKey: key, + } +} + +/* ── Errors ─────────────────────────────────────────────────────────────── */ + +// oauthError is RFC 6749's error shape. +type oauthError struct { + Code string `json:"error"` + Description string `json:"error_description,omitempty"` +} + +// Standard error codes. Kept to the set RFC 6749 and 7591 define, because a +// client's error handling switches on these strings. +const ( + errInvalidRequest = "invalid_request" + errInvalidClient = "invalid_client" + errInvalidGrant = "invalid_grant" + errUnauthorizedClient = "unauthorized_client" + errUnsupportedGrantType = "unsupported_grant_type" + errInvalidScope = "invalid_scope" + errInvalidRedirectURI = "invalid_redirect_uri" + errInvalidTarget = "invalid_target" // RFC 8707, for a bad resource + errServerError = "server_error" +) + +func writeOAuthError(w http.ResponseWriter, status int, code, description string) { + w.Header().Set("Content-Type", "application/json; charset=utf-8") + // A token or error response must never be cached: it is specific to one + // request and may carry a credential. + w.Header().Set("Cache-Control", "no-store") + w.Header().Set("Pragma", "no-cache") + writeJSONBody(w, status, oauthError{Code: code, Description: description}) +} + +func writeJSONBody(w http.ResponseWriter, status int, payload any) { + encoded, err := json.Marshal(payload) + if err != nil { + http.Error(w, "internal error", http.StatusInternalServerError) + return + } + w.WriteHeader(status) + _, _ = w.Write(encoded) +} + +/* ── RFC 7591: Dynamic Client Registration ──────────────────────────────── */ + +type registrationRequest struct { + ClientName string `json:"client_name"` + RedirectURIs []string `json:"redirect_uris"` + GrantTypes []string `json:"grant_types,omitempty"` + ResponseTypes []string `json:"response_types,omitempty"` + TokenEndpointAuthMethod string `json:"token_endpoint_auth_method,omitempty"` + Scope string `json:"scope,omitempty"` +} + +type registrationResponse struct { + ClientID string `json:"client_id"` + ClientName string `json:"client_name,omitempty"` + RedirectURIs []string `json:"redirect_uris"` + GrantTypes []string `json:"grant_types"` + ResponseTypes []string `json:"response_types"` + TokenEndpointAuthMethod string `json:"token_endpoint_auth_method"` + Scope string `json:"scope"` + ClientIDIssuedAt int64 `json:"client_id_issued_at"` +} + +// maxRegistrationBytes bounds a registration body. A registration is a name and +// a handful of URIs. +const maxRegistrationBytes = 16 << 10 + +// RegisterHandler serves dynamic client registration. +// +// Open by necessity: a client that has never registered has no credential to +// present, which is the entire point of RFC 7591 and what lets Claude connect +// without anyone provisioning anything by hand. +// +// That openness is why redirect URI validation below is strict, and why +// PHASE 5 MUST ADD RATE LIMITING HERE. This endpoint writes a row for any +// caller that can reach it. It is structured for that — one handler, one +// validation pass, nothing that would have to move — but today it has no limit, +// and that is recorded as a known gap rather than quietly left unsaid. +func (s *Server) RegisterHandler() http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodPost { + w.Header().Set("Allow", http.MethodPost) + writeOAuthError(w, http.StatusMethodNotAllowed, errInvalidRequest, "POST only") + return + } + + var req registrationRequest + if err := json.NewDecoder(http.MaxBytesReader(w, r.Body, maxRegistrationBytes)).Decode(&req); err != nil { + writeOAuthError(w, http.StatusBadRequest, errInvalidRequest, "the request body was not valid JSON") + return + } + + if len(req.RedirectURIs) == 0 { + writeOAuthError(w, http.StatusBadRequest, errInvalidRedirectURI, "at least one redirect_uri is required") + return + } + if len(req.RedirectURIs) > 10 { + writeOAuthError(w, http.StatusBadRequest, errInvalidRedirectURI, "too many redirect_uris") + return + } + for _, uri := range req.RedirectURIs { + if err := validateRedirectURI(uri); err != nil { + writeOAuthError(w, http.StatusBadRequest, errInvalidRedirectURI, err.Error()) + return + } + } + + // Only the scopes this server issues. A client asking for krow.write + // is refused rather than quietly downgraded: silently granting less + // than was asked for produces a client that believes it has a + // capability and fails later, somewhere less obvious. + scopes := []string{ScopeRead} + if strings.TrimSpace(req.Scope) != "" { + requested := strings.Fields(req.Scope) + for _, sc := range requested { + if sc != ScopeRead { + writeOAuthError(w, http.StatusBadRequest, errInvalidScope, + "the only scope available is "+ScopeRead) + return + } + } + scopes = requested + } + + clientID, err := newUUID() + if err != nil { + s.log.Error("oauth: client id generation failed", "error", err) + writeOAuthError(w, http.StatusInternalServerError, errServerError, "") + return + } + + name := strings.TrimSpace(req.ClientName) + if len(name) > 200 { + name = name[:200] + } + + client := Client{ + ClientID: clientID, + ClientName: name, + RedirectURIs: req.RedirectURIs, + GrantTypes: []string{"authorization_code", "refresh_token"}, + Scopes: scopes, + } + if err := s.store.CreateClient(r.Context(), client); err != nil { + s.log.Error("oauth: client registration failed", "error", err) + writeOAuthError(w, http.StatusInternalServerError, errServerError, "") + return + } + + s.log.Info("oauth client registered", + "client_id", clientID, "client_name", name, "redirect_uris", len(req.RedirectURIs)) + + w.Header().Set("Content-Type", "application/json; charset=utf-8") + w.Header().Set("Cache-Control", "no-store") + writeJSONBody(w, http.StatusCreated, registrationResponse{ + ClientID: clientID, + ClientName: name, + RedirectURIs: req.RedirectURIs, + GrantTypes: []string{"authorization_code", "refresh_token"}, + // No client_secret. A public client that was issued one would ship + // it to every user's machine, and a secret everybody has is not a + // secret — OAuth 2.1 handles public clients with PKCE instead. + ResponseTypes: []string{"code"}, + TokenEndpointAuthMethod: "none", + Scope: strings.Join(scopes, " "), + ClientIDIssuedAt: s.store.now().Unix(), + }) + }) +} + +// validateRedirectURI refuses a redirect target that cannot be trusted. +// +// The rules, and why each one is here: +// +// - absolute, with a scheme and host — a relative URI has no meaning in a +// redirect and a client sending one is confused about the flow. +// - no fragment — RFC 6749 forbids it, and the authorization response appends +// its own query parameters; a fragment would be silently dropped or would +// mangle them. +// - https, OR http on loopback only. Plain http anywhere else means the +// authorization code travels in clear text. Loopback is the documented +// exception for native clients (RFC 8252) and is safe because the traffic +// never leaves the machine. +// +// Custom schemes (myapp://callback) are NOT accepted. They are legal per RFC +// 8252 and are a real mechanism for native apps, but any application on the +// machine can register the same scheme and steal the code. Claude's connectors +// use https and loopback, so accepting custom schemes would widen the surface +// for no caller that exists. +func validateRedirectURI(raw string) error { + parsed, err := url.Parse(raw) + if err != nil { + return errMsg("redirect_uri is not a valid URI") + } + if parsed.Scheme == "" || parsed.Host == "" { + return errMsg("redirect_uri must be absolute, with a scheme and host") + } + if parsed.Fragment != "" || strings.Contains(raw, "#") { + return errMsg("redirect_uri must not contain a fragment") + } + + switch strings.ToLower(parsed.Scheme) { + case "https": + return nil + case "http": + if isLoopbackHost(parsed.Hostname()) { + return nil + } + return errMsg("http is only permitted for loopback redirect URIs") + default: + return errMsg("redirect_uri must use https, or http on loopback") + } +} + +func isLoopbackHost(host string) bool { + switch host { + case "127.0.0.1", "::1", "localhost": + return true + } + return false +} + +type errString string + +func (e errString) Error() string { return string(e) } +func errMsg(s string) error { return errString(s) } + +/* ── Authorization endpoint ─────────────────────────────────────────────── */ + +// authorizeParams is a validated authorization request. +type authorizeParams struct { + ClientID string + RedirectURI string + ResponseType string + Scopes []string + State string + CodeChallenge string + CodeChallengeMethod string + Resource string +} + +// AuthorizeHandler serves the authorization endpoint. +// +// THE ORDER OF VALIDATION IS A SECURITY PROPERTY, not a style choice. +// +// The client_id and redirect_uri are validated FIRST, against the registration, +// before anything else is looked at. Only once the redirect target is known to +// be one this client registered may an error be delivered BY REDIRECTING to it. +// Getting this backwards — redirecting an error to an unvalidated URI — is an +// open redirect, and it is the most common way this endpoint is got wrong. +// +// So: a bad client_id or a bad redirect_uri is answered as a direct HTTP error +// that the browser displays. Everything after that is delivered as a redirect +// with `error=` and the client's `state`, because by then the target is known +// to be legitimate. +func (s *Server) AuthorizeHandler() http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodGet && r.Method != http.MethodPost { + w.Header().Set("Allow", "GET, POST") + writeOAuthError(w, http.StatusMethodNotAllowed, errInvalidRequest, "GET or POST only") + return + } + // A POST carries the decision and the flow's parameters in its body, + // re-posted from the consent form's hidden fields. Merging them into + // the query is what lets every validation below read from one place + // regardless of method — and means the POST is validated exactly as + // strictly as the GET that produced it, rather than trusting the form. + q := r.URL.Query() + if r.Method == http.MethodPost { + if err := r.ParseForm(); err != nil { + writeOAuthError(w, http.StatusBadRequest, errInvalidRequest, + "the form could not be parsed") + return + } + q = r.PostForm + } + + // ── Stage 1: the client and its redirect target. Errors here are + // direct responses, never redirects. + clientID := strings.TrimSpace(q.Get("client_id")) + if clientID == "" { + writeOAuthError(w, http.StatusBadRequest, errInvalidClient, "client_id is required") + return + } + client, err := s.store.FindClient(r.Context(), clientID) + if err != nil { + s.log.Warn("oauth authorize refused", "reason", "unknown_client", "client_id", clientID) + writeOAuthError(w, http.StatusBadRequest, errInvalidClient, "unknown client") + return + } + + redirectURI := strings.TrimSpace(q.Get("redirect_uri")) + if redirectURI == "" { + writeOAuthError(w, http.StatusBadRequest, errInvalidRedirectURI, "redirect_uri is required") + return + } + if !client.AllowsRedirect(redirectURI) { + // Deliberately NOT redirected. This is the open-redirect guard. + s.log.Warn("oauth authorize refused", + "reason", "redirect_uri_mismatch", "client_id", clientID) + writeOAuthError(w, http.StatusBadRequest, errInvalidRedirectURI, + "redirect_uri does not match a registered URI for this client") + return + } + + // ── Stage 2: everything else. The target is trusted now, so failures + // are delivered to it. + state := strings.TrimSpace(q.Get("state")) + if state == "" { + // Required, not optional. state is the client's CSRF defence for + // the callback; a flow without one can be completed by an attacker + // who injects their own authorization response. + s.redirectError(w, r, redirectURI, "", errInvalidRequest, "state is required") + return + } + + if rt := q.Get("response_type"); rt != "code" { + s.redirectError(w, r, redirectURI, state, "unsupported_response_type", + "only response_type=code is supported") + return + } + + challenge := strings.TrimSpace(q.Get("code_challenge")) + method := strings.TrimSpace(q.Get("code_challenge_method")) + if challenge == "" { + s.redirectError(w, r, redirectURI, state, errInvalidRequest, + "code_challenge is required; this server requires PKCE") + return + } + if method == "" { + // RFC 7636 defaults an absent method to `plain`. This server does + // not accept plain, so an absent method is an error rather than a + // silent downgrade to the weaker mode. + s.redirectError(w, r, redirectURI, state, errInvalidRequest, + "code_challenge_method is required and must be S256") + return + } + if err := ValidateChallenge(challenge, method); err != nil { + s.redirectError(w, r, redirectURI, state, errInvalidRequest, err.Error()) + return + } + + // RFC 8707. The resource must be THIS server's canonical MCP URI. A + // token is bound to it, so accepting an arbitrary value would let a + // client mint a token aimed at something else. + resource := strings.TrimSpace(q.Get("resource")) + if resource == "" { + s.redirectError(w, r, redirectURI, state, errInvalidTarget, + "resource is required") + return + } + if strings.TrimRight(resource, "/") != s.cfg.Resource { + s.log.Warn("oauth authorize refused", + "reason", "resource_mismatch", "client_id", clientID, "presented", resource) + s.redirectError(w, r, redirectURI, state, errInvalidTarget, + "resource is not a resource this server issues tokens for") + return + } + + scopes := []string{ScopeRead} + if raw := strings.TrimSpace(q.Get("scope")); raw != "" { + scopes = strings.Fields(raw) + for _, sc := range scopes { + if sc != ScopeRead { + s.redirectError(w, r, redirectURI, state, errInvalidScope, + "the only scope available is "+ScopeRead) + return + } + } + } + if !client.AllowsScopes(scopes) { + s.redirectError(w, r, redirectURI, state, errInvalidScope, + "this client is not registered for the requested scope") + return + } + + params := authorizeParams{ + ClientID: clientID, RedirectURI: redirectURI, ResponseType: "code", + Scopes: scopes, State: state, CodeChallenge: challenge, + CodeChallengeMethod: method, Resource: resource, + } + + // ── Stage 3: who is this? + identity, signedIn := s.sessions.CurrentUser(r) + if !signedIn { + // Not signed in. Send them to the existing login, with a return + // target that brings them back to this exact authorization request. + // No credential is handled here — the existing cookie login does + // that, unchanged. + s.redirectToLogin(w, r) + return + } + + // ── Stage 4: consent. + // + // A GET renders the question. Only a POST carrying a session-bound + // CSRF token answers it, so a cross-site navigation can show a person + // the form but cannot approve on their behalf. + csrf := s.csrfFor(identity) + + if r.Method != http.MethodPost { + s.renderConsent(w, r, params, identity, csrf) + return + } + + if !s.csrfValid(identity, r.PostFormValue("csrf")) { + // Not an OAuth protocol error — it is a request that did not come + // from the form this server rendered. Answered directly rather + // than redirected, because the client is not the party at fault + // and telling it "access_denied" would be a lie. + s.log.Warn("oauth consent refused", "reason", "csrf_mismatch", + "client_id", params.ClientID, "user_id", identity.UserID) + writeOAuthError(w, http.StatusForbidden, errInvalidRequest, + "this consent form has expired; start the authorization again") + return + } + + switch r.PostFormValue("decision") { + case "approve": + s.log.Info("oauth consent approved", + "client_id", params.ClientID, "user_id", identity.UserID, + "org_id", identity.OrgID, "scopes", params.Scopes) + s.issueCode(w, r, params, identity) + case "deny": + // RFC 6749 section 4.1.2.1: a refusal is `access_denied`, returned + // to the client at its registered redirect with the state intact. + // NO CODE IS ISSUED — the deny path never reaches issueCode. + s.log.Info("oauth consent denied", + "client_id", params.ClientID, "user_id", identity.UserID) + s.redirectError(w, r, params.RedirectURI, params.State, + "access_denied", "the user declined this authorization") + default: + // A POST with neither decision. Re-render rather than guess: the + // one thing that must not happen is inferring approval. + s.renderConsent(w, r, params, identity, csrf) + } + }) +} + +/* ── Consent CSRF ───────────────────────────────────────────────────────── */ + +// csrfFor derives a token binding the consent form to the signed-in user. +// +// An HMAC over the user id under a per-process key, rather than a random value +// in server-side state. The property needed is only "this form was rendered by +// this server for this user", and an HMAC gives that with nothing to store and +// nothing to expire. +// +// The key is generated at startup and never leaves the process, so a token does +// not survive a restart — which ends any consent form open at that moment. That +// is acceptable: the window between rendering and deciding is seconds, and the +// failure mode is a person clicking Approve and being asked to start again. +func (s *Server) csrfFor(identity authctx.Identity) string { + mac := hmac.New(sha256.New, s.csrfKey) + mac.Write([]byte(identity.UserID)) + return hex.EncodeToString(mac.Sum(nil)) +} + +// csrfValid checks a submitted token in constant time. +func (s *Server) csrfValid(identity authctx.Identity, presented string) bool { + if presented == "" { + return false + } + return hmac.Equal([]byte(s.csrfFor(identity)), []byte(presented)) +} + +// issueCode stores an authorization code and redirects it to the client. +func (s *Server) issueCode(w http.ResponseWriter, r *http.Request, p authorizeParams, identity authctx.Identity) { + code, err := s.store.CreateGrant(r.Context(), Grant{ + ClientID: p.ClientID, + UserID: identity.UserID, + OrgID: identity.OrgID, + RedirectURI: p.RedirectURI, + Scopes: p.Scopes, + Resource: p.Resource, + CodeChallenge: p.CodeChallenge, + CodeChallengeMethod: p.CodeChallengeMethod, + }) + if err != nil { + s.log.Error("oauth: could not create grant", "error", err, "client_id", p.ClientID) + s.redirectError(w, r, p.RedirectURI, p.State, errServerError, "") + return + } + + // The code id is not logged, and neither is the code. What is logged is who + // approved what, which is the audit question worth answering. + s.log.Info("oauth code issued", + "client_id", p.ClientID, "user_id", identity.UserID, + "org_id", identity.OrgID, "scopes", p.Scopes, "resource", p.Resource) + + target, err := url.Parse(p.RedirectURI) + if err != nil { + writeOAuthError(w, http.StatusBadRequest, errInvalidRedirectURI, "redirect_uri is not a valid URI") + return + } + q := target.Query() + q.Set("code", code) + q.Set("state", p.State) + target.RawQuery = q.Encode() + + w.Header().Set("Cache-Control", "no-store") + http.Redirect(w, r, target.String(), http.StatusFound) +} + +// redirectError delivers an error to a VALIDATED redirect target. +// +// Only ever called after the redirect_uri has been matched against the client's +// registration. See the note on AuthorizeHandler. +func (s *Server) redirectError(w http.ResponseWriter, r *http.Request, redirectURI, state, code, description string) { + target, err := url.Parse(redirectURI) + if err != nil { + writeOAuthError(w, http.StatusBadRequest, errInvalidRedirectURI, "redirect_uri is not a valid URI") + return + } + q := target.Query() + q.Set("error", code) + if description != "" { + q.Set("error_description", description) + } + if state != "" { + q.Set("state", state) + } + target.RawQuery = q.Encode() + + w.Header().Set("Cache-Control", "no-store") + http.Redirect(w, r, target.String(), http.StatusFound) +} + +// redirectToLogin sends an unauthenticated person to the existing login. +// +// The return target is this server's own path plus the original query, so the +// authorization request survives the round trip. It is built from r.URL rather +// than from anything the caller supplied, so it cannot be pointed elsewhere. +func (s *Server) redirectToLogin(w http.ResponseWriter, r *http.Request) { + returnTo := r.URL.Path + if r.URL.RawQuery != "" { + returnTo += "?" + r.URL.RawQuery + } + target := s.loginPath + "?returnTo=" + url.QueryEscape(returnTo) + w.Header().Set("Cache-Control", "no-store") + http.Redirect(w, r, target, http.StatusFound) +} + +/* ── Token endpoint ─────────────────────────────────────────────────────── */ + +type tokenResponse struct { + AccessToken string `json:"access_token"` + TokenType string `json:"token_type"` + ExpiresIn int `json:"expires_in"` + RefreshToken string `json:"refresh_token"` + Scope string `json:"scope"` +} + +// TokenHandler serves the token endpoint: code exchange and refresh. +func (s *Server) TokenHandler() http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodPost { + w.Header().Set("Allow", http.MethodPost) + writeOAuthError(w, http.StatusMethodNotAllowed, errInvalidRequest, "POST only") + return + } + if err := r.ParseForm(); err != nil { + writeOAuthError(w, http.StatusBadRequest, errInvalidRequest, "the request body could not be parsed") + return + } + + switch r.PostFormValue("grant_type") { + case "authorization_code": + s.exchangeCode(w, r) + case "refresh_token": + s.refresh(w, r) + case "": + writeOAuthError(w, http.StatusBadRequest, errInvalidRequest, "grant_type is required") + default: + // password, client_credentials, implicit and anything else. Named + // explicitly in the metadata as unsupported, and refused here. + writeOAuthError(w, http.StatusBadRequest, errUnsupportedGrantType, + "only authorization_code and refresh_token are supported") + } + }) +} + +// exchangeCode turns an authorization code into a token pair. +// +// Every binding recorded at authorization is re-verified. A code is not a +// bearer credential on its own: it is a credential for one client, one redirect +// target, one resource, and one PKCE verifier, and a mismatch on any of them +// means the code is being spent by someone other than the client it was issued +// to. +func (s *Server) exchangeCode(w http.ResponseWriter, r *http.Request) { + code := r.PostFormValue("code") + clientID := r.PostFormValue("client_id") + redirectURI := r.PostFormValue("redirect_uri") + verifier := r.PostFormValue("code_verifier") + resource := strings.TrimSpace(r.PostFormValue("resource")) + + if code == "" || clientID == "" || redirectURI == "" { + writeOAuthError(w, http.StatusBadRequest, errInvalidRequest, + "code, client_id and redirect_uri are required") + return + } + if verifier == "" { + writeOAuthError(w, http.StatusBadRequest, errInvalidRequest, + "code_verifier is required; this server requires PKCE") + return + } + + // Redeeming CONSUMES the code, whatever happens next. That is deliberate: + // if a later check fails, the code is still spent, so an attacker cannot + // probe the remaining bindings by retrying the same code with different + // values. One code, one attempt. + grant, err := s.store.RedeemGrant(r.Context(), code) + if err != nil { + s.log.Warn("oauth token refused", "reason", "grant_unusable", "client_id", clientID) + writeOAuthError(w, http.StatusBadRequest, errInvalidGrant, + "the authorization code is invalid, expired or already used") + return + } + + if grant.ClientID != clientID { + s.log.Warn("oauth token refused", "reason", "client_mismatch", "client_id", clientID) + writeOAuthError(w, http.StatusBadRequest, errInvalidGrant, "this code was not issued to this client") + return + } + if grant.RedirectURI != redirectURI { + s.log.Warn("oauth token refused", "reason", "redirect_uri_mismatch", "client_id", clientID) + writeOAuthError(w, http.StatusBadRequest, errInvalidGrant, "redirect_uri does not match the authorization request") + return + } + // The resource is optional at the token endpoint when the code already + // carries one, but if it IS supplied it must agree. + if resource != "" && strings.TrimRight(resource, "/") != grant.Resource { + writeOAuthError(w, http.StatusBadRequest, errInvalidTarget, "resource does not match the authorization request") + return + } + if err := VerifyChallenge(verifier, grant.CodeChallenge, grant.CodeChallengeMethod); err != nil { + s.log.Warn("oauth token refused", "reason", "pkce_mismatch", "client_id", clientID) + writeOAuthError(w, http.StatusBadRequest, errInvalidGrant, "code_verifier does not match") + return + } + + pair, err := s.store.IssuePair(r.Context(), Token{ + ClientID: grant.ClientID, + UserID: grant.UserID, + OrgID: grant.OrgID, + Scopes: grant.Scopes, + Audience: grant.Resource, + }, "") + if err != nil { + s.log.Error("oauth: could not issue tokens", "error", err) + writeOAuthError(w, http.StatusInternalServerError, errServerError, "") + return + } + + // The tokens themselves are NOT in this log line and never will be. + s.log.Info("oauth tokens issued", + "grant_type", "authorization_code", "client_id", grant.ClientID, + "user_id", grant.UserID, "org_id", grant.OrgID, "family_id", pair.FamilyID) + + writeTokenResponse(w, pair) +} + +// refresh rotates a refresh token. +func (s *Server) refresh(w http.ResponseWriter, r *http.Request) { + raw := r.PostFormValue("refresh_token") + clientID := r.PostFormValue("client_id") + + if raw == "" || clientID == "" { + writeOAuthError(w, http.StatusBadRequest, errInvalidRequest, + "refresh_token and client_id are required") + return + } + + old, err := s.store.RedeemRefreshToken(r.Context(), raw) + switch { + case err == nil: + // fall through + case errors.Is(err, ErrRefreshReuse): + // The family has already been revoked by the store. Logged at warn + // because it is either a client bug or a stolen token, and both are + // worth seeing. The CLIENT is told the same thing as for any other bad + // token — distinguishing "reused" would confirm the token was once + // real. + s.log.Warn("oauth refresh refused", "reason", "reuse_detected", "client_id", clientID) + writeOAuthError(w, http.StatusBadRequest, errInvalidGrant, "the refresh token is invalid") + return + default: + s.log.Warn("oauth refresh refused", "reason", "token_unusable", "client_id", clientID) + writeOAuthError(w, http.StatusBadRequest, errInvalidGrant, "the refresh token is invalid") + return + } + + if old.ClientID != clientID { + // Not this client's token. Revoke the family: a refresh token that has + // reached the wrong client has leaked. + _ = s.store.RevokeFamily(r.Context(), old.FamilyID, "client_mismatch_on_refresh") + s.log.Warn("oauth refresh refused", "reason", "client_mismatch", "client_id", clientID) + writeOAuthError(w, http.StatusBadRequest, errInvalidGrant, "the refresh token is invalid") + return + } + + // Same family: the rotation continues the lineage, so reuse detection can + // still revoke every descendant if an older token reappears. + pair, err := s.store.IssuePair(r.Context(), Token{ + ClientID: old.ClientID, + UserID: old.UserID, + OrgID: old.OrgID, + Scopes: old.Scopes, + Audience: old.Audience, + }, old.FamilyID) + if err != nil { + s.log.Error("oauth: could not rotate tokens", "error", err) + writeOAuthError(w, http.StatusInternalServerError, errServerError, "") + return + } + + s.log.Info("oauth tokens issued", + "grant_type", "refresh_token", "client_id", old.ClientID, + "user_id", old.UserID, "family_id", pair.FamilyID) + + writeTokenResponse(w, pair) +} + +func writeTokenResponse(w http.ResponseWriter, pair TokenPair) { + w.Header().Set("Content-Type", "application/json; charset=utf-8") + // RFC 6749 section 5.1 requires both of these on a token response. The + // body is a credential; nothing may cache it. + w.Header().Set("Cache-Control", "no-store") + w.Header().Set("Pragma", "no-cache") + writeJSONBody(w, http.StatusOK, tokenResponse{ + AccessToken: pair.AccessToken, + TokenType: "Bearer", + ExpiresIn: pair.ExpiresIn, + RefreshToken: pair.RefreshToken, + Scope: strings.Join(pair.Scopes, " "), + }) +} + +/* ── Revocation (RFC 7009) ──────────────────────────────────────────────── */ + +// RevokeHandler serves token revocation. +// +// RFC 7009 requires 200 for an unknown token: answering 404 would turn this +// into an oracle for whether a token exists. The store already behaves that +// way; this handler just does not undo it. +func (s *Server) RevokeHandler() http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if r.Method != http.MethodPost { + w.Header().Set("Allow", http.MethodPost) + writeOAuthError(w, http.StatusMethodNotAllowed, errInvalidRequest, "POST only") + return + } + if err := r.ParseForm(); err != nil { + writeOAuthError(w, http.StatusBadRequest, errInvalidRequest, "the request body could not be parsed") + return + } + token := r.PostFormValue("token") + if token == "" { + writeOAuthError(w, http.StatusBadRequest, errInvalidRequest, "token is required") + return + } + + if err := s.store.RevokeToken(r.Context(), token, "client_revocation"); err != nil { + s.log.Error("oauth: revocation failed", "error", err) + writeOAuthError(w, http.StatusInternalServerError, errServerError, "") + return + } + s.log.Info("oauth token revoked", "client_id", r.PostFormValue("client_id")) + + w.Header().Set("Cache-Control", "no-store") + w.WriteHeader(http.StatusOK) + }) +} diff --git a/go-api/internal/oauth/cleanup.go b/go-api/internal/oauth/cleanup.go new file mode 100644 index 0000000..991bb3d --- /dev/null +++ b/go-api/internal/oauth/cleanup.go @@ -0,0 +1,130 @@ +package oauth + +import ( + "context" + "fmt" + "time" +) + +// Cleanup of spent and expired OAuth rows. +// +// WHAT IS DELETED, AND WHAT IS DELIBERATELY NOT +// +// Only rows that can no longer authenticate anything. Every predicate below +// requires the row to be past its expiry — not merely consumed, not merely +// revoked — because those two states are evidence, and evidence is worth +// keeping until it stops being relevant. +// +// A consumed refresh token in particular must outlive its usefulness: it is +// what REUSE DETECTION matches against. Delete it the moment it is spent and a +// stolen token replayed a minute later looks like an unknown token rather than +// a theft, and the family is never revoked. So a consumed refresh token is kept +// until its original expiry, by which point replaying it proves nothing anyway. +// +// A revoked token is kept for the same reason plus one more: "this token was +// revoked at 14:02 for refresh_token_reuse" is an answer to a question someone +// will eventually ask. +// +// GRACE. Everything is deleted a grace period AFTER expiry rather than at it, +// so a clock skewed between two instances cannot delete a row another instance +// still considers live. +// +// SAFE TO RUN TWICE, AND SAFE TO RUN CONCURRENTLY. Every statement is a bounded +// DELETE with a predicate that no longer matches once the row is gone. Two +// workers running at once delete disjoint sets and neither errors. + +// CleanupGrace is how long a dead row is kept past its expiry. +// +// An hour is far beyond any plausible clock skew between instances and short +// enough that the tables do not accumulate. It also means a support question +// asked within the hour can still see the row. +const CleanupGrace = time.Hour + +// CleanupBatch bounds one pass. +// +// Bounded because an unbounded DELETE holds locks for as long as it runs, and +// on a table that every MCP request reads that is a latency spike nobody can +// explain afterwards. Five thousand rows is milliseconds; if there is more, the +// next pass takes it. +const CleanupBatch = 5000 + +// CleanupResult reports what one pass removed. +type CleanupResult struct { + Grants int64 + AccessTokens int64 + RefreshTokens int64 + CompletedInOne bool // false when a batch filled, meaning more remains +} + +// Cleanup removes expired authorization codes and tokens. +// +// Returns counts rather than logging them, so the caller decides the level and +// this function stays usable from a test. +func (s *Store) Cleanup(ctx context.Context) (CleanupResult, error) { + cutoff := s.now().Add(-CleanupGrace) + var out CleanupResult + + // Authorization codes. Sixty-second TTL, so almost every row here is + // already dead; this is the highest-volume and cheapest of the three. + // + // ctid rather than id in the subquery because it is the physical row + // address — the planner can go straight to it without a second index + // lookup, which is what keeps a bounded delete genuinely cheap. + tag, err := s.db.Exec(ctx, + `DELETE FROM oauth_grants + WHERE ctid IN ( + SELECT ctid FROM oauth_grants WHERE expires_at < $1 LIMIT $2 + )`, cutoff, CleanupBatch) + if err != nil { + return out, fmt.Errorf("oauth: cleanup grants: %w", err) + } + out.Grants = tag.RowsAffected() + + // Access tokens. Fifteen-minute TTL. An expired one cannot authenticate — + // FindAccessToken's predicate already excludes it — so deleting it removes + // no capability. + tag, err = s.db.Exec(ctx, + `DELETE FROM oauth_tokens + WHERE ctid IN ( + SELECT ctid FROM oauth_tokens + WHERE token_type = 'access' AND expires_at < $1 + LIMIT $2 + )`, cutoff, CleanupBatch) + if err != nil { + return out, fmt.Errorf("oauth: cleanup access tokens: %w", err) + } + out.AccessTokens = tag.RowsAffected() + + // Refresh tokens, and this is the one with a real constraint on it. + // + // EXPIRY ONLY — not `consumed_at IS NOT NULL`, and not `revoked_at IS NOT + // NULL`. A consumed refresh token is what RedeemRefreshToken matches to + // detect reuse; deleting it early turns a detectable theft into an + // unremarkable "unknown token" and the family is never revoked. Thirty-day + // TTL means these are the longest-lived rows in the schema, which is the + // price of that detection and is worth paying. + tag, err = s.db.Exec(ctx, + `DELETE FROM oauth_tokens + WHERE ctid IN ( + SELECT ctid FROM oauth_tokens + WHERE token_type = 'refresh' AND expires_at < $1 + LIMIT $2 + )`, cutoff, CleanupBatch) + if err != nil { + return out, fmt.Errorf("oauth: cleanup refresh tokens: %w", err) + } + out.RefreshTokens = tag.RowsAffected() + + out.CompletedInOne = out.Grants < CleanupBatch && + out.AccessTokens < CleanupBatch && + out.RefreshTokens < CleanupBatch + return out, nil +} + +// RevokeExpiredFamilies is deliberately absent. +// +// It looks like it belongs here — "tidy up families whose tokens have all +// lapsed" — and it would do nothing. Revocation is a state on a row, and a row +// that has been deleted has no state to set. A family whose every token has +// expired and been swept simply ceases to exist, which is the correct outcome +// and requires no work. diff --git a/go-api/internal/oauth/cleanup_test.go b/go-api/internal/oauth/cleanup_test.go new file mode 100644 index 0000000..8cf0f70 --- /dev/null +++ b/go-api/internal/oauth/cleanup_test.go @@ -0,0 +1,240 @@ +package oauth + +import ( + "context" + "sync" + "testing" + "time" +) + +/* ── What cleanup removes ───────────────────────────────────────────────── */ + +func TestCleanupRemovesOnlyDeadRows(t *testing.T) { + h := newHarness(t) + ctx := context.Background() + clientID := h.register() + + // A live pair, and a spent code. + code := h.authorizeOK(clientID, verifier43) + live := decodeTokens(t, h.exchange(clientID, code, verifier43)) + + // A second, which we let expire. + oldCode := h.authorizeOK(clientID, verifier43) + old := decodeTokens(t, h.exchange(clientID, oldCode, verifier43)) + + before := countRows(t, h) + + // Past the access token TTL and the grace, but well inside the refresh + // token's thirty days. + h.advance(AccessTokenTTL + CleanupGrace + time.Minute) + + result, err := h.store.Cleanup(ctx) + if err != nil { + t.Fatalf("Cleanup: %v", err) + } + + // Both access tokens and both codes are dead; both refresh tokens are not. + if result.AccessTokens != 2 { + t.Errorf("removed %d access tokens, want 2", result.AccessTokens) + } + if result.Grants != 2 { + t.Errorf("removed %d grants, want 2", result.Grants) + } + if result.RefreshTokens != 0 { + t.Errorf("removed %d refresh tokens, want 0 — they live thirty days", result.RefreshTokens) + } + if !result.CompletedInOne { + t.Error("a small cleanup reported that more remained") + } + + after := countRows(t, h) + if after.tokens >= before.tokens { + t.Error("cleanup removed nothing") + } + + // The live refresh tokens must still work. This is the property that + // matters: cleanup must not disconnect anybody. + for name, token := range map[string]string{"live": live.RefreshToken, "old": old.RefreshToken} { + if _, err := h.store.RedeemRefreshToken(ctx, token); err != nil { + t.Errorf("the %s refresh token stopped working after cleanup: %v", name, err) + } + } +} + +// The subtle one: a CONSUMED refresh token must survive until its expiry, +// because it is what reuse detection matches against. Delete it early and a +// replayed stolen token looks unknown rather than stolen, and the family is +// never revoked. +func TestCleanupKeepsConsumedRefreshTokensForReuseDetection(t *testing.T) { + h := newHarness(t) + ctx := context.Background() + clientID := h.register() + + code := h.authorizeOK(clientID, verifier43) + first := decodeTokens(t, h.exchange(clientID, code, verifier43)) + + // Rotate: `first` is now consumed. + if _, err := h.store.RedeemRefreshToken(ctx, first.RefreshToken); err != nil { + t.Fatalf("rotate: %v", err) + } + + // Cleanup well past the access token TTL, but inside the refresh TTL. + h.advance(AccessTokenTTL + CleanupGrace + time.Hour) + if _, err := h.store.Cleanup(ctx); err != nil { + t.Fatalf("Cleanup: %v", err) + } + + // Replaying the consumed token must STILL be detected as reuse. + _, err := h.store.RedeemRefreshToken(ctx, first.RefreshToken) + if err != ErrRefreshReuse { + t.Errorf("err = %v, want ErrRefreshReuse — cleanup destroyed the evidence "+ + "that makes theft detectable", err) + } +} + +// Revoked rows are kept until expiry too: "revoked at 14:02 for +// refresh_token_reuse" is an answer somebody will eventually need. +func TestCleanupKeepsRevokedRowsUntilExpiry(t *testing.T) { + h := newHarness(t) + ctx := context.Background() + clientID := h.register() + + code := h.authorizeOK(clientID, verifier43) + tokens := decodeTokens(t, h.exchange(clientID, code, verifier43)) + if err := h.store.RevokeToken(ctx, tokens.AccessToken, "test"); err != nil { + t.Fatalf("revoke: %v", err) + } + + // Just past the access TTL: the access row goes, the refresh row stays. + h.advance(AccessTokenTTL + CleanupGrace + time.Minute) + if _, err := h.store.Cleanup(ctx); err != nil { + t.Fatalf("Cleanup: %v", err) + } + + var revokedRefresh int + if err := h.h.Pool.QueryRow(ctx, + `SELECT count(*) FROM oauth_tokens WHERE token_type='refresh' AND revoked_at IS NOT NULL`). + Scan(&revokedRefresh); err != nil { + t.Fatalf("count: %v", err) + } + if revokedRefresh != 1 { + t.Errorf("%d revoked refresh rows kept, want 1 — the audit trail was swept", revokedRefresh) + } +} + +// Nothing is deleted before the grace period, so clock skew between instances +// cannot destroy a row another instance still considers live. +func TestCleanupHonoursTheGracePeriod(t *testing.T) { + h := newHarness(t) + ctx := context.Background() + clientID := h.register() + code := h.authorizeOK(clientID, verifier43) + decodeTokens(t, h.exchange(clientID, code, verifier43)) + + // Expired, but inside the grace. + h.advance(AccessTokenTTL + time.Minute) + result, err := h.store.Cleanup(ctx) + if err != nil { + t.Fatalf("Cleanup: %v", err) + } + if result.AccessTokens != 0 { + t.Errorf("removed %d access tokens inside the grace period, want 0", result.AccessTokens) + } +} + +/* ── Safety ─────────────────────────────────────────────────────────────── */ + +func TestCleanupIsIdempotent(t *testing.T) { + h := newHarness(t) + ctx := context.Background() + clientID := h.register() + code := h.authorizeOK(clientID, verifier43) + decodeTokens(t, h.exchange(clientID, code, verifier43)) + + h.advance(AccessTokenTTL + CleanupGrace + time.Minute) + + first, err := h.store.Cleanup(ctx) + if err != nil { + t.Fatalf("first: %v", err) + } + second, err := h.store.Cleanup(ctx) + if err != nil { + t.Fatalf("second: %v", err) + } + if second.AccessTokens != 0 || second.Grants != 0 || second.RefreshTokens != 0 { + t.Errorf("a second cleanup removed more rows: %+v (first was %+v)", second, first) + } +} + +// Two workers running cleanup at once must not error and must not +// double-count. Run with -race. +func TestConcurrentCleanupIsSafe(t *testing.T) { + h := newHarness(t) + ctx := context.Background() + clientID := h.register() + + for i := 0; i < 6; i++ { + code := h.authorizeOK(clientID, verifier43) + decodeTokens(t, h.exchange(clientID, code, verifier43)) + } + h.advance(AccessTokenTTL + CleanupGrace + time.Minute) + + const workers = 4 + var wg sync.WaitGroup + var mu sync.Mutex + var total int64 + errs := make([]error, 0, workers) + + for i := 0; i < workers; i++ { + wg.Add(1) + go func() { + defer wg.Done() + r, err := h.store.Cleanup(ctx) + mu.Lock() + defer mu.Unlock() + if err != nil { + errs = append(errs, err) + return + } + total += r.AccessTokens + }() + } + wg.Wait() + + if len(errs) > 0 { + t.Fatalf("concurrent cleanup errored: %v", errs) + } + // Six access tokens existed; between them the workers removed exactly six. + // More would mean a row was counted twice. + if total != 6 { + t.Errorf("workers removed %d access tokens between them, want 6", total) + } +} + +func TestCleanupOnAnEmptyDatabaseIsHarmless(t *testing.T) { + h := newHarness(t) + result, err := h.store.Cleanup(context.Background()) + if err != nil { + t.Fatalf("Cleanup on empty: %v", err) + } + if result.Grants != 0 || result.AccessTokens != 0 || result.RefreshTokens != 0 { + t.Errorf("cleanup on an empty database removed %+v", result) + } +} + +/* ── Helpers ────────────────────────────────────────────────────────────── */ + +type rowCounts struct{ grants, tokens int } + +func countRows(t *testing.T, h *harness) rowCounts { + t.Helper() + var c rowCounts + ctx := context.Background() + if err := h.h.Pool.QueryRow(ctx, `SELECT count(*) FROM oauth_grants`).Scan(&c.grants); err != nil { + t.Fatalf("count grants: %v", err) + } + if err := h.h.Pool.QueryRow(ctx, `SELECT count(*) FROM oauth_tokens`).Scan(&c.tokens); err != nil { + t.Fatalf("count tokens: %v", err) + } + return c +} diff --git a/go-api/internal/oauth/consent.go b/go-api/internal/oauth/consent.go new file mode 100644 index 0000000..35250bd --- /dev/null +++ b/go-api/internal/oauth/consent.go @@ -0,0 +1,346 @@ +package oauth + +import ( + "html/template" + "net/http" + "net/url" + "strings" + + "github.com/krow/krow-backend/go-api/internal/authctx" +) + +// The consent step: the one place a person decides. +// +// Phase 3 approved a signed-in user's authorization immediately. That was +// honest scaffolding and is not a flow anybody should ship: OAuth's entire +// premise is that a RESOURCE OWNER grants access, and an authorization nobody +// was asked about is a token minted on their behalf without their knowledge. +// Any page on the internet could have linked a person to a crafted authorize +// URL and had Claude connected to their workspace before they read anything. +// +// HOW THIS RESISTS THAT +// +// The consent form carries a CSRF token bound to the session, and approval is +// a POST. A cross-site GET to /oauth/authorize can therefore render the form — +// which is harmless, it is a question — but cannot answer it. Without the POST +// and the token, an attacker who can make a browser navigate cannot make it +// consent. +// +// WHAT IT SHOWS +// +// The client's self-declared name, the organisation being granted, the scope in +// plain words, and the resource. The client name is UNTRUSTED — it is whatever +// the registering client sent — so it is escaped by html/template and is never +// the basis of a decision, only of a label. The organisation is read from the +// signed-in identity, so a person can see which tenant they are about to hand +// over even when they belong to more than one. + +// consentTemplate is the approval page. +// +// Deliberately one self-contained page with inline styles: it renders before a +// person is willing to trust anything, it must work with no stylesheet, no +// script and no font available, and a consent screen that depends on assets is +// a consent screen that can fail open into a blank page with two buttons. +// +// Every interpolation is escaped by html/template. The `.ClientName` in +// particular is attacker-controlled — anyone may register a client called +// `` + body, _ := jsonMarshal(registrationRequest{ + ClientName: payload, RedirectURIs: []string{testRedirect}, + }) + rec := httptest.NewRecorder() + h.server.RegisterHandler().ServeHTTP(rec, + httptest.NewRequest(http.MethodPost, "/oauth/register", strings.NewReader(body))) + var reg registrationResponse + _ = jsonUnmarshal(rec.Body.Bytes(), ®) + + page, _ := h.consent(authorizeParamsFor(reg.ClientID, verifier43)) + if strings.Contains(page.Body.String(), "