mcp connection
Some checks failed
CI / fixture (push) Has been cancelled
CI / test (push) Has been cancelled

This commit is contained in:
2026-09-22 10:58:02 +05:30
parent 4e1f746b22
commit f2aa3b3ad8
53 changed files with 12515 additions and 37 deletions

View File

@@ -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
}

View File

@@ -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
}

View File

@@ -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(&reg); 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.

View File

@@ -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.

View File

@@ -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
}

View File

@@ -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)
}
}

View File

@@ -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()
}
}
}

View File

@@ -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")
}
}

View File

@@ -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)

View File

@@ -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, &regDoc)
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, &regDoc)
// 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, &regDoc)
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, &regDoc)
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, &regDoc)
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"))
}
}

View File

@@ -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
}

View File

@@ -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)
}
}

View File

@@ -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
}

View File

@@ -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

View File

@@ -0,0 +1,193 @@
package mcpserver
import (
"context"
"errors"
"net/http"
"strings"
"github.com/krow/krow-backend/go-api/internal/auth"
"github.com/krow/krow-backend/go-api/internal/authctx"
)
// Bearer authentication for the MCP surface.
//
// This file is a SEAM, not an authentication system. It defines the one
// question the MCP transport needs answered — "which KROW user does this token
// belong to" — and leaves answering it to whatever is plugged in. Phase 3 plugs
// in OAuth 2.1 token validation. Nothing here mints, stores, refreshes or
// validates a token's contents, because doing any of that now would be
// inventing a token format that OAuth then has to replace.
//
// WHY MCP AUTHENTICATES SEPARATELY FROM THE REST OF THE API
//
// The cookie middleware in httpserver/auth.go is deliberately not reused, and
// this is the most important decision in this file. Mounting MCP behind that
// middleware would mean a browser session could authenticate an MCP call: the
// middleware puts an Identity in the context, and any handler downstream that
// reads the ambient identity would accept it. That is a real vulnerability
// rather than a theoretical one — a logged-in user's cookie is sent by the
// browser on requests the user did not intend, which is what SameSite exists to
// limit and what an MCP endpoint has no business relying on.
//
// So identity here is PASSED, never ambient. The transport authenticates, and
// hands the result to Handle as a parameter. There is no code path in this
// package that reads authctx.From on an inbound request, which makes "a cookie
// silently authenticated MCP" structurally impossible rather than merely
// unintended. See TestCookieCannotAuthenticateMCP.
/* ── Errors ─────────────────────────────────────────────────────────────── */
var (
// ErrNoAuthenticator is returned when the surface is running without a
// token authenticator. It is a configuration fault, and it fails CLOSED:
// a deployment that forgot to wire one refuses every call rather than
// serving them unauthenticated.
ErrNoAuthenticator = errors.New("mcpserver: no token authenticator configured")
// ErrMissingToken covers an absent or empty Authorization header.
ErrMissingToken = errors.New("mcpserver: no bearer token")
// ErrMalformedToken covers a header this server could not parse as a
// bearer credential — a missing scheme, a wrong scheme, an empty value.
ErrMalformedToken = errors.New("mcpserver: malformed Authorization header")
// ErrInvalidToken covers a well-formed token that does not resolve to a
// user: unknown, expired, revoked, or issued for something else.
//
// ONE error for all of those, deliberately. Telling a caller that a token
// is "expired" rather than "unknown" confirms it once existed, which is an
// oracle over the token space. Same reasoning as tools.Denied().
ErrInvalidToken = errors.New("mcpserver: invalid bearer token")
)
/* ── The seam ───────────────────────────────────────────────────────────── */
// TokenAuthenticator resolves a raw bearer token into a KROW identity.
//
// Deliberately one method taking a string and returning the SAME
// authctx.Identity the cookie path produces. Two things follow from that shape,
// and both are the point:
//
// - There is no second identity model. Everything downstream — the policy
// table, the org pre-filter, tools.Context — consumes authctx.Identity and
// cannot tell which path produced it, so authorization cannot drift between
// the two.
// - An implementation cannot report anything except an identity or a failure.
// It has no way to return "authenticated, but also here is an org" or any
// other channel a caller might trust. The org is inside the identity,
// which comes from the user row.
//
// Phase 3's OAuth implementation of this interface will: hash the presented
// token, look it up, check expiry, revocation and audience, load the user, and
// build the identity from the USER ROW — never from the token's contents. A
// token that carried its own org claim would be a token whose bearer chose
// their own tenant.
type TokenAuthenticator interface {
// Authenticate resolves a raw token, or returns an error.
//
// Implementations must fail closed and must not distinguish unknown from
// expired from revoked in the returned error.
Authenticate(ctx context.Context, rawToken string) (authctx.Identity, error)
}
// UserLookup is the subset of the existing user store this package needs.
//
// Narrowed to one method so an implementation of TokenAuthenticator can re-read
// the user on every call — which is what makes suspension take effect on
// contact rather than whenever a token happens to lapse. httpserver/auth.go
// does exactly this for cookies (see its comment on re-reading the user row),
// and the bearer path must not be weaker.
//
// auth.UserStore already satisfies this.
type UserLookup interface {
FindByID(ctx context.Context, id string) (auth.User, error)
}
/* ── Header parsing ─────────────────────────────────────────────────────── */
// bearerToken extracts the credential from an Authorization header.
//
// Only the Authorization header is consulted. Not a query parameter — the MCP
// spec forbids tokens in the URI, and a URI is logged, cached, and put in a
// Referer. Not a custom header, not a cookie, not the body. One place, so there
// is one thing to reason about.
func bearerToken(r *http.Request) (string, error) {
header := r.Header.Get("Authorization")
if strings.TrimSpace(header) == "" {
return "", ErrMissingToken
}
scheme, value, found := strings.Cut(header, " ")
if !found {
return "", ErrMalformedToken
}
// Case-insensitive per RFC 7235: "Bearer", "bearer" and "BEARER" are the
// same scheme, and rejecting the variants would fail against clients that
// are behaving correctly.
if !strings.EqualFold(strings.TrimSpace(scheme), "bearer") {
return "", ErrMalformedToken
}
token := strings.TrimSpace(value)
if token == "" {
return "", ErrMalformedToken
}
// A second space means a second value — "Bearer a b" is not a token, and
// accepting the first half would silently authenticate something the
// client did not send.
if strings.ContainsAny(token, " \t") {
return "", ErrMalformedToken
}
return token, nil
}
// authenticate resolves the request's bearer credential into an identity.
//
// Every failure returns the same outward answer — 401 with no detail about
// which stage failed. The reason is recorded in the log, where the operator is.
func (s *Server) authenticate(r *http.Request) (authctx.Identity, error) {
token, err := bearerToken(r)
if err != nil {
return authctx.Identity{}, err
}
if s.tokens == nil {
return authctx.Identity{}, ErrNoAuthenticator
}
identity, err := s.tokens.Authenticate(r.Context(), token)
if err != nil {
return authctx.Identity{}, ErrInvalidToken
}
// Defence in depth against an authenticator that returns a partially
// populated identity. Everything downstream assumes these two are present:
// tools/scope.go refuses an empty OrgID, but it should never be asked to,
// and a missing UserID would produce a query scoped to nobody.
if identity.UserID == "" || identity.OrgID == "" {
return authctx.Identity{}, ErrInvalidToken
}
// A suspended account must not hold a working token. The authenticator is
// expected to check this; repeating it here costs nothing and means a
// mistake in one implementation is not a live account bypass.
if identity.Status != "" && identity.Status != auth.StatusActive {
return authctx.Identity{}, ErrInvalidToken
}
return identity, nil
}
// authFailureReason names the stage that refused, for the log only.
func authFailureReason(err error) string {
switch {
case errors.Is(err, ErrMissingToken):
return "missing_token"
case errors.Is(err, ErrMalformedToken):
return "malformed_header"
case errors.Is(err, ErrNoAuthenticator):
return "no_authenticator_configured"
case errors.Is(err, ErrInvalidToken):
return "invalid_token"
default:
return "error"
}
}

View File

@@ -0,0 +1,512 @@
package mcpserver
import (
"context"
"encoding/json"
"io"
"log/slog"
"net/http"
"net/http/httptest"
"strings"
"testing"
"github.com/krow/krow-backend/go-api/internal/authctx"
"github.com/krow/krow-backend/go-api/internal/runtime"
"github.com/krow/krow-backend/go-api/internal/testutil"
"github.com/krow/krow-backend/go-api/internal/tools"
)
/* ── A test-only authenticator ──────────────────────────────────────────────
TEST-ONLY. This type lives in a _test.go file and is therefore not compiled
into the binary at all — there is no build tag to forget and no flag that
could enable it in production. That is deliberate: the one thing worse than
having no authentication is having a development authenticator that ships.
It is a map from token to identity and nothing else. It does not hash, does
not expire, does not check an audience and does not consult a database,
because none of those are what these tests are testing. What they test is
the SEAM: that a token resolves to an identity, that the identity reaches
the registry, and that existing authorization then decides the answer.
Phase 3 replaces this with an OAuth implementation of the same interface.
Every test below keeps working, because what they assert is the behaviour of
the seam rather than the behaviour of any particular token format. That is
the reason for defining the interface before implementing OAuth rather than
after. */
type fakeTokens struct {
byToken map[string]authctx.Identity
err error
}
func (f *fakeTokens) Authenticate(_ context.Context, raw string) (authctx.Identity, error) {
if f.err != nil {
return authctx.Identity{}, f.err
}
id, ok := f.byToken[raw]
if !ok {
return authctx.Identity{}, ErrInvalidToken
}
return id, nil
}
const validToken = "test-token-valid"
func identityFor(orgID, role string) authctx.Identity {
return authctx.Identity{
UserID: "user-" + role,
OrgID: orgID,
Email: role + "@example.test",
FullName: "Test " + role,
Role: role,
AccountType: "employer",
Status: "active",
}
}
// authedServer wires the real registry to a token map.
func authedServer(t *testing.T, reg *tools.Registry, tokens map[string]authctx.Identity) *Server {
t.Helper()
if reg == nil {
reg = runtime.DefaultTools(nil, nil)
}
return New(reg, &fakeTokens{byToken: tokens}, slog.New(slog.NewTextHandler(io.Discard, nil)))
}
// postWith sends a body with an explicit Authorization header value. An empty
// header value means the header is not sent at all.
func postWith(t *testing.T, s *Server, authHeader, body string) *httptest.ResponseRecorder {
t.Helper()
req := httptest.NewRequest(http.MethodPost, "/mcp", strings.NewReader(body))
req.Header.Set("Content-Type", "application/json")
if authHeader != "" {
req.Header.Set("Authorization", authHeader)
}
rec := httptest.NewRecorder()
s.Handler().ServeHTTP(rec, req)
return rec
}
const listBody = `{"jsonrpc":"2.0","id":1,"method":"tools/list"}`
/* ── 1–5. Header and token rejection ────────────────────────────────────── */
func TestBearerRejection(t *testing.T) {
s := authedServer(t, nil, map[string]authctx.Identity{
validToken: identityFor("org-a", "admin"),
})
for name, header := range map[string]string{
"missing header": "",
"no scheme": "abc123",
"wrong scheme basic": "Basic dXNlcjpwYXNz",
"wrong scheme token": "Token abc123",
"empty bearer": "Bearer ",
"bearer with only ws": "Bearer ",
"two values": "Bearer abc def",
"unknown token": "Bearer not-a-real-token",
"token with whitespace": "Bearer abc\tdef",
} {
t.Run(name, func(t *testing.T) {
rec := postWith(t, s, header, listBody)
if rec.Code != http.StatusUnauthorized {
t.Fatalf("status = %d, want 401 (body: %s)", rec.Code, rec.Body.String())
}
// RFC 9728: the client reads the authorization server location from
// this header. Without it a compliant client cannot begin discovery.
if got := rec.Header().Get("WWW-Authenticate"); !strings.HasPrefix(got, "Bearer") {
t.Errorf("WWW-Authenticate = %q, want a Bearer challenge", got)
}
// The refusal must not say WHICH stage failed. "expired" versus
// "unknown" is an oracle over the token space.
body := rec.Body.String()
for _, leak := range []string{"expired", "revoked", "unknown", "malformed", "not found"} {
if strings.Contains(strings.ToLower(body), leak) {
t.Errorf("the 401 body distinguishes failure modes (%q): %s", leak, body)
}
}
})
}
}
// Case-insensitivity is required by RFC 7235; rejecting "bearer" would fail
// against clients that are behaving correctly.
func TestBearerSchemeIsCaseInsensitive(t *testing.T) {
s := authedServer(t, nil, map[string]authctx.Identity{
validToken: identityFor("org-a", "admin"),
})
for _, scheme := range []string{"Bearer", "bearer", "BEARER", "BeArEr"} {
rec := postWith(t, s, scheme+" "+validToken, listBody)
if rec.Code != http.StatusOK {
t.Errorf("scheme %q: status = %d, want 200", scheme, rec.Code)
}
}
}
// A server with no authenticator wired must refuse everything rather than
// serve it unauthenticated. Misconfiguration fails closed.
func TestNoAuthenticatorFailsClosed(t *testing.T) {
s := New(runtime.DefaultTools(nil, nil), nil, slog.New(slog.NewTextHandler(io.Discard, nil)))
rec := postWith(t, s, "Bearer "+validToken, listBody)
if rec.Code != http.StatusUnauthorized {
t.Fatalf("status = %d, want 401 when no authenticator is configured", rec.Code)
}
}
/* ── 6 & 7. Identity and org resolution ─────────────────────────────────── */
func TestValidBearerReachesToolsList(t *testing.T) {
s := authedServer(t, nil, map[string]authctx.Identity{
validToken: identityFor("org-a", "admin"),
})
rec := postWith(t, s, "Bearer "+validToken, listBody)
var out toolsListResult
resultInto(t, rec, &out)
if len(out.Tools) != len(wantExposed) {
t.Fatalf("exposed %d tools, want %d", len(out.Tools), len(wantExposed))
}
}
// An identity missing a tenant must be refused before it reaches a tool.
// tools/scope.go would refuse it too, but it should never be asked to: an
// authenticator that returns a half-built identity is a bug, not a caller.
func TestIncompleteIdentityIsRefused(t *testing.T) {
for name, id := range map[string]authctx.Identity{
"no org": {UserID: "u1", Role: "admin", Status: "active"},
"no user": {OrgID: "org-a", Role: "admin", Status: "active"},
"suspended": {UserID: "u1", OrgID: "org-a", Role: "admin", Status: "suspended"},
"empty entire": {},
} {
t.Run(name, func(t *testing.T) {
s := authedServer(t, nil, map[string]authctx.Identity{validToken: id})
rec := postWith(t, s, "Bearer "+validToken, listBody)
if rec.Code != http.StatusUnauthorized {
t.Errorf("status = %d, want 401", rec.Code)
}
})
}
}
/* ── 8, 9, 10. Identity cannot be overridden ────────────────────────────── */
// The identity must come from the token and from nothing else. This walks the
// channels a client controls and asserts that none of them moves the tenant.
func TestIdentityCannotBeOverridden(t *testing.T) {
h := testutil.New(t)
orgA := freshOrgFor(t, h, "mcp-override-a")
orgB := freshOrgFor(t, h, "mcp-override-b")
seedActivityRows(t, h, orgA, 5, "a@example.test")
seedActivityRows(t, h, orgB, 40, "b@example.test")
reg := runtime.DefaultTools(h.Pool, nil)
s := authedServer(t, reg, map[string]authctx.Identity{
validToken: {UserID: "u-a", OrgID: orgA, Email: "a@example.test",
Role: "admin", AccountType: "employer", Status: "active"},
})
// Every one of these is a client-controlled channel. None may change which
// tenant is read. orgB has 40 rows and orgA has 5, so a successful override
// is visible as a total of 40 or 45.
attempts := map[string]string{
"tool argument": `{"jsonrpc":"2.0","id":1,"method":"tools/call","params":
{"name":"activity_breakdown","arguments":{"org_id":"` + orgB + `"}}}`,
"jsonrpc meta": `{"jsonrpc":"2.0","id":1,"method":"tools/call","_meta":{"org_id":"` + orgB + `"},
"params":{"name":"activity_breakdown","arguments":{}}}`,
"params meta": `{"jsonrpc":"2.0","id":1,"method":"tools/call","params":
{"name":"activity_breakdown","arguments":{},"_meta":{"org_id":"` + orgB + `"}}}`,
"params orgId": `{"jsonrpc":"2.0","id":1,"method":"tools/call","params":
{"name":"activity_breakdown","orgId":"` + orgB + `","arguments":{}}}`,
"params userId": `{"jsonrpc":"2.0","id":1,"method":"tools/call","params":
{"name":"activity_breakdown","userId":"u-b","arguments":{}}}`,
}
for name, body := range attempts {
t.Run(name, func(t *testing.T) {
rec := postWith(t, s, "Bearer "+validToken, body)
total := totalFromActivityBreakdown(t, rec)
if total != 5 {
t.Errorf("total = %d, want 5 (org A only) — %q moved the tenant", total, name)
}
})
}
}
// Custom headers naming another tenant must be ignored outright.
func TestCustomIdentityHeadersAreIgnored(t *testing.T) {
h := testutil.New(t)
orgA := freshOrgFor(t, h, "mcp-hdr-a")
orgB := freshOrgFor(t, h, "mcp-hdr-b")
seedActivityRows(t, h, orgA, 5, "a@example.test")
seedActivityRows(t, h, orgB, 40, "b@example.test")
reg := runtime.DefaultTools(h.Pool, nil)
s := authedServer(t, reg, map[string]authctx.Identity{
validToken: {UserID: "u-a", OrgID: orgA, Email: "a@example.test",
Role: "admin", AccountType: "employer", Status: "active"},
})
body := `{"jsonrpc":"2.0","id":1,"method":"tools/call","params":
{"name":"activity_breakdown","arguments":{}}}`
req := httptest.NewRequest(http.MethodPost, "/mcp", strings.NewReader(body))
req.Header.Set("Content-Type", "application/json")
req.Header.Set("Authorization", "Bearer "+validToken)
for _, h := range []string{"X-Org-Id", "X-Organization-Id", "X-User-Id", "X-Tenant-Id", "X-Krow-Org"} {
req.Header.Set(h, orgB)
}
rec := httptest.NewRecorder()
s.Handler().ServeHTTP(rec, req)
if total := totalFromActivityBreakdown(t, rec); total != 5 {
t.Errorf("total = %d, want 5 — a custom header moved the tenant", total)
}
}
// A cookie must never authenticate MCP.
//
// This is the test for the decision in auth.go: identity is passed, never
// ambient. Even with a valid KROW identity sitting in the request context —
// which is exactly what the cookie middleware would put there if this endpoint
// were mounted behind it — the call must be refused, because no bearer token
// was presented.
func TestCookieCannotAuthenticateMCP(t *testing.T) {
s := authedServer(t, nil, map[string]authctx.Identity{
validToken: identityFor("org-a", "admin"),
})
req := httptest.NewRequest(http.MethodPost, "/mcp", strings.NewReader(listBody))
req.Header.Set("Content-Type", "application/json")
// A session cookie, as a browser would send it.
req.AddCookie(&http.Cookie{Name: "krow_session", Value: "a-perfectly-valid-session-token"})
// AND a fully populated identity in the context, as the cookie middleware
// would have placed there. This is the strongest form of the test: even if
// somebody mounts MCP behind authenticate(), it must still refuse.
ctx := authctx.With(req.Context(), identityFor("org-a", "admin"))
req = req.WithContext(ctx)
rec := httptest.NewRecorder()
s.Handler().ServeHTTP(rec, req)
if rec.Code != http.StatusUnauthorized {
t.Fatalf("status = %d, want 401 — a cookie session authenticated an MCP call", rec.Code)
}
}
/* ── 11, 12, 13. Authorization still runs ───────────────────────────────── */
// The point of the whole design: MCP authenticates, and the EXISTING policy
// table authorizes. A talent caller must be refused a tool that operators own.
func TestExistingAuthorizationStillRuns(t *testing.T) {
h := testutil.New(t)
org := freshOrgFor(t, h, "mcp-authz")
seedActivityRows(t, h, org, 5, "boss@example.test")
reg := runtime.DefaultTools(h.Pool, nil)
s := authedServer(t, reg, map[string]authctx.Identity{
"admin-token": {UserID: "u-admin", OrgID: org, Email: "boss@example.test",
Role: "admin", AccountType: "employer", Status: "active"},
"talent-token": {UserID: "u-talent", OrgID: org, Email: "worker@example.test",
Role: "talent", AccountType: "talent", Status: "active"},
})
// `staff` is operators-only in the policy table, and hires_recent reads
// job-applications which talent may list only in its own scope.
call := `{"jsonrpc":"2.0","id":1,"method":"tools/call","params":
{"name":"activity_breakdown","arguments":{}}}`
adminRec := postWith(t, s, "Bearer admin-token", call)
if got := totalFromActivityBreakdown(t, adminRec); got != 5 {
t.Errorf("admin total = %d, want 5", got)
}
// A talent caller is authenticated but scoped. Whatever comes back, it must
// come back through the policy table rather than around it — the assertion
// is that the two roles do NOT get the same answer.
talentRec := postWith(t, s, "Bearer talent-token", call)
talentTotal := totalFromActivityBreakdownAllowingDenial(t, talentRec)
if talentTotal == 5 {
t.Error("a talent caller saw the admin's total; authorization did not run")
}
}
/* ── 14 & 15. Exclusions hold under authentication ──────────────────────── */
// Authentication must not become a way to reach a withheld tool. A valid token
// is not a key to the write tools.
func TestExclusionsHoldForAuthenticatedCallers(t *testing.T) {
s := authedServer(t, nil, map[string]authctx.Identity{
validToken: identityFor("org-a", "admin"),
})
for _, name := range []string{"assign_worker", "move_application", "knowledge_search"} {
t.Run(name, func(t *testing.T) {
body := `{"jsonrpc":"2.0","id":1,"method":"tools/call","params":{"name":"` + name + `","arguments":{}}}`
rec := postWith(t, s, "Bearer "+validToken, body)
var out toolsCallResult
resultInto(t, rec, &out)
if !out.IsError {
t.Fatalf("%s was reachable by an authenticated caller", name)
}
if !strings.Contains(out.Content[0].Text, "mcp.unknown_tool") {
t.Errorf("expected mcp.unknown_tool, got: %s", out.Content[0].Text)
}
})
}
}
// tools/list must not vary by caller in a way that reveals the withheld set.
func TestToolsListIsTheSameForEveryRole(t *testing.T) {
s := authedServer(t, nil, map[string]authctx.Identity{
"admin-token": identityFor("org-a", "admin"),
"talent-token": identityFor("org-a", "talent"),
})
names := func(token string) []string {
rec := postWith(t, s, "Bearer "+token, listBody)
var out toolsListResult
resultInto(t, rec, &out)
got := make([]string, 0, len(out.Tools))
for _, tool := range out.Tools {
got = append(got, tool.Name)
}
return got
}
admin, talent := names("admin-token"), names("talent-token")
if strings.Join(admin, ",") != strings.Join(talent, ",") {
t.Errorf("tools/list differs by role:\n admin: %v\ntalent: %v", admin, talent)
}
}
/* ── The handshake requires a token too ─────────────────────────────────── */
// REVERSED IN PHASE 4, deliberately.
//
// This test previously asserted the opposite: that initialize and ping were
// reachable without a token, on the reasoning that a client needs somewhere to
// start. That was wrong, and the integration test over the mounted route is
// what caught it — a client's FIRST request is usually initialize, and
// answering it 200 means the client never sees the WWW-Authenticate challenge
// that begins the OAuth flow. It believes it is connected and finds out
// otherwise at the first real call, with no 401 in hand to discover from.
//
// Requiring a token everywhere means the first request, whatever it is,
// produces the challenge. Nothing is lost: the handshake returns only this
// server's name and capabilities, which are of use solely to a client that
// means to authenticate.
func TestHandshakeMethodsAlsoRequireAuth(t *testing.T) {
s := authedServer(t, nil, map[string]authctx.Identity{
validToken: identityFor("org-a", "admin"),
})
for name, body := range map[string]string{
"initialize": `{"jsonrpc":"2.0","id":1,"method":"initialize","params":{}}`,
"ping": `{"jsonrpc":"2.0","id":1,"method":"ping"}`,
"tools/list": `{"jsonrpc":"2.0","id":1,"method":"tools/list"}`,
"unknown verb": `{"jsonrpc":"2.0","id":1,"method":"resources/list"}`,
} {
t.Run(name+" without a token", func(t *testing.T) {
rec := postWith(t, s, "", body)
if rec.Code != http.StatusUnauthorized {
t.Errorf("status = %d, want 401 — every method must produce the "+
"challenge that starts the OAuth flow", rec.Code)
}
if !strings.HasPrefix(rec.Header().Get("WWW-Authenticate"), "Bearer") {
t.Error("the 401 carries no Bearer challenge")
}
})
}
// And with a token, the handshake works — or the test above would pass by
// the endpoint being broken rather than by it being guarded.
for name, body := range map[string]string{
"initialize": `{"jsonrpc":"2.0","id":1,"method":"initialize","params":{}}`,
"ping": `{"jsonrpc":"2.0","id":1,"method":"ping"}`,
} {
t.Run(name+" with a token", func(t *testing.T) {
if rec := postWith(t, s, "Bearer "+validToken, body); rec.Code != http.StatusOK {
t.Errorf("status = %d, want 200", rec.Code)
}
})
}
}
// The resource_metadata pointer must be built from configuration, so a 401
// tells a client where to look without this package knowing any hostname.
func TestChallengeCarriesTheConfiguredResourceMetadata(t *testing.T) {
const metadataURL = "https://configured.example.test/.well-known/oauth-protected-resource"
s := authedServer(t, nil, map[string]authctx.Identity{}).
WithResourceMetadataURL(metadataURL)
rec := postWith(t, s, "", listBody)
challenge := rec.Header().Get("WWW-Authenticate")
if !strings.Contains(challenge, `resource_metadata="`+metadataURL+`"`) {
t.Errorf("WWW-Authenticate = %q, want it to carry %q", challenge, metadataURL)
}
}
/* ── Helpers ────────────────────────────────────────────────────────────── */
func freshOrgFor(t *testing.T, h *testutil.Harness, slug string) string {
t.Helper()
var id string
if err := h.Pool.QueryRow(context.Background(),
`INSERT INTO organizations (name, slug) VALUES ($1, $2) RETURNING id::text`,
slug, slug).Scan(&id); err != nil {
t.Fatalf("create org %s: %v", slug, err)
}
return id
}
func seedActivityRows(t *testing.T, h *testutil.Harness, org string, n int, email string) {
t.Helper()
for i := 0; i < n; i++ {
if _, err := h.Pool.Exec(context.Background(),
`INSERT INTO user_activity (org_id, event_type, user_email, user_name)
VALUES ($1::uuid, 'login', $2, 'Someone')`, org, email); err != nil {
t.Fatalf("seed activity: %v", err)
}
}
}
// totalFromActivityBreakdown reads the `total` out of a successful tool result.
func totalFromActivityBreakdown(t *testing.T, rec *httptest.ResponseRecorder) int {
t.Helper()
var out toolsCallResult
resultInto(t, rec, &out)
if out.IsError {
t.Fatalf("tool call failed: %s", out.Content[0].Text)
}
return parseTotal(t, out.Content[0].Text)
}
// totalFromActivityBreakdownAllowingDenial returns -1 when the tool refused,
// which is a legitimate authorization outcome rather than a test failure.
func totalFromActivityBreakdownAllowingDenial(t *testing.T, rec *httptest.ResponseRecorder) int {
t.Helper()
var out toolsCallResult
resultInto(t, rec, &out)
if out.IsError {
return -1
}
return parseTotal(t, out.Content[0].Text)
}
func parseTotal(t *testing.T, raw string) int {
t.Helper()
// activity_breakdown's own field name, taken from the handler's output
// rather than guessed: a wrong name here reads as zero, which would make a
// cross-tenant leak look like a pass.
var payload struct {
Data struct {
TotalEvents int `json:"totalEvents"`
} `json:"data"`
}
if err := json.Unmarshal([]byte(raw), &payload); err != nil {
t.Fatalf("tool result did not decode: %v\nraw: %s", err, raw)
}
return payload.Data.TotalEvents
}

View File

@@ -0,0 +1,243 @@
// Package mcpserver is the Model Context Protocol surface: a second way into
// the tool layer, for clients that speak MCP rather than HTTP+cookie.
//
// It is an ADDITIONAL interface and nothing else. It owns no business logic, no
// SQL and no authorization rules. Every call it serves ends up in
// tools.Registry.Dispatch — the same entry point the agent loop uses — so a
// question asked through MCP is answered by the same handler, under the same
// policy table, behind the same org pre-filter as the same question asked by
// Owliver. That is the whole design, and the reason this package is small.
//
// What lives here:
//
// - JSON-RPC 2.0 framing (this file)
// - the three methods MCP needs to be useful: initialize, tools/list,
// tools/call (server.go)
// - which tools are published, derived from the registry (tools.go)
// - bearer authentication, as a seam an OAuth implementation plugs into
// (auth.go)
// - the Streamable HTTP binding (transport.go)
//
// Identity is established by this package's own bearer authentication and is
// PASSED to the handlers, never read from the ambient request context. That is
// what stops a browser cookie from authenticating an MCP call — see auth.go.
package mcpserver
import (
"encoding/json"
"errors"
"fmt"
"time"
)
// jsonRPCVersion is the only version this server speaks. A request naming
// anything else is malformed rather than merely unsupported: "2.0" is a
// constant in the spec, not a negotiation.
const jsonRPCVersion = "2.0"
/* ── Wire types ─────────────────────────────────────────────────────────── */
// request is one inbound JSON-RPC message.
//
// ID is json.RawMessage rather than any, because the spec allows a string, a
// number or null, and the response MUST echo it back byte-for-byte. Decoding it
// into an `any` turns 1 into 1.0 on the way back out, which is a different id to
// a client matching responses to requests.
type request struct {
JSONRPC string `json:"jsonrpc"`
ID json.RawMessage `json:"id,omitempty"`
Method string `json:"method"`
Params json.RawMessage `json:"params,omitempty"`
}
// isNotification reports that no response is expected.
//
// A notification is a request with no id. The spec is explicit that a server
// must not answer one, so the transport drops the response and returns 202.
func (r request) isNotification() bool {
return len(r.ID) == 0 || string(r.ID) == "null"
}
// response is one outbound JSON-RPC message.
//
// Result and Error are pointers so exactly one is ever serialised: the spec
// forbids both together, and a non-pointer Result would emit `"result":null`
// alongside an error.
type response struct {
JSONRPC string `json:"jsonrpc"`
ID json.RawMessage `json:"id,omitempty"`
Result any `json:"result,omitempty"`
Error *rpcError `json:"error,omitempty"`
}
// rpcError is a JSON-RPC error object.
//
// retryAfter is NOT serialised. It exists so a rate-limited refusal produced
// deep in the handler can reach the transport, which is the only layer that can
// set an HTTP status and a Retry-After header. The alternative — returning a
// 200 with a JSON-RPC error and no header — would give a client no way to know
// how long to wait.
type rpcError struct {
Code int `json:"code"`
Message string `json:"message"`
Data any `json:"data,omitempty"`
retryAfter time.Duration
}
func (e *rpcError) Error() string { return fmt.Sprintf("jsonrpc %d: %s", e.Code, e.Message) }
/* ── Error codes ────────────────────────────────────────────────────────── */
// The standard JSON-RPC 2.0 codes. Reserved range is -32768..-32000; anything
// this server invents lives outside it.
const (
codeParseError = -32700
codeInvalidRequest = -32600
codeMethodNotFound = -32601
codeInvalidParams = -32602
codeInternalError = -32603
)
// codeUnauthorized is outside the JSON-RPC reserved range (-32768..-32000),
// because it is this server's own condition rather than a protocol fault. It
// accompanies an HTTP 401: the transport layer carries the authoritative
// signal, and this gives a client reading only the JSON-RPC body the same
// answer.
const codeUnauthorized = -32001
// codeRateLimited is this server's own condition, outside the reserved range.
// It accompanies an HTTP 429 and a Retry-After header.
const codeRateLimited = -32002
// errRateLimited refuses a call that exceeded its organisation's ceiling.
//
// The message names no number and no organisation. How much quota a tenant has
// and how much of it they have spent is not something one caller should learn
// from a refusal — it is the same reasoning as the opaque tool denial.
func errRateLimited(retryAfter time.Duration) *rpcError {
return &rpcError{
Code: codeRateLimited,
Message: "too many requests for this organisation; retry after the interval in the Retry-After header",
retryAfter: retryAfter,
}
}
func errParse(detail string) *rpcError {
return &rpcError{Code: codeParseError, Message: "invalid JSON", Data: detail}
}
func errInvalidRequest(detail string) *rpcError {
return &rpcError{Code: codeInvalidRequest, Message: "invalid JSON-RPC request", Data: detail}
}
func errMethodNotFound(method string) *rpcError {
return &rpcError{
Code: codeMethodNotFound,
Message: "method not found",
Data: fmt.Sprintf("this server implements initialize, tools/list and tools/call; it does not implement %q", method),
}
}
func errInvalidParams(detail string) *rpcError {
return &rpcError{Code: codeInvalidParams, Message: "invalid params", Data: detail}
}
// errInternal deliberately carries no detail.
//
// An internal failure is the one case where the thing that went wrong is this
// server's business and not the caller's: a wrapped database error or a panic
// message is reconnaissance. The detail goes to the log, where the operator is.
func errInternal() *rpcError {
return &rpcError{Code: codeInternalError, Message: "internal error"}
}
/* ── Parsing ────────────────────────────────────────────────────────────── */
// errBatch marks a batch request, which this server does not accept.
//
// Rejecting it explicitly rather than failing to parse it is the point: a
// client that batches and gets a parse error will retry the same batch, where
// one told that batching is unsupported can fall back to sending messages
// singly. The current MCP transport binding sends one message per POST, so
// nothing a compliant client does requires batching.
var errBatch = errors.New("batch requests are not supported")
// parseRequest decodes one JSON-RPC message and validates its envelope.
//
// The two are separate returns because they have different fates: a message
// that could not be parsed has no id, so its error answers with a null id,
// while a message that parsed but is invalid answers with the id it carried.
func parseRequest(body []byte) (request, *rpcError) {
trimmed := skipSpace(body)
if len(trimmed) == 0 {
return request{}, errInvalidRequest("the request body was empty")
}
if trimmed[0] == '[' {
return request{}, errInvalidRequest(errBatch.Error())
}
var req request
if err := json.Unmarshal(trimmed, &req); err != nil {
return request{}, errParse(err.Error())
}
if req.JSONRPC != jsonRPCVersion {
return req, errInvalidRequest(fmt.Sprintf(
"jsonrpc must be %q, got %q", jsonRPCVersion, req.JSONRPC))
}
if req.Method == "" {
return req, errInvalidRequest("method is required")
}
// An id, when present, must be a string or a number. Objects and arrays are
// forbidden by the spec, and echoing one back would propagate the mistake.
if len(req.ID) > 0 && !isValidID(req.ID) {
return req, errInvalidRequest("id must be a string, a number or null")
}
return req, nil
}
// isValidID reports whether a raw id is a string, a number or null.
func isValidID(raw json.RawMessage) bool {
t := skipSpace(raw)
if len(t) == 0 {
return false
}
switch t[0] {
case '{', '[':
return false
}
var v any
return json.Unmarshal(t, &v) == nil
}
// decodeParams unmarshals params into dst, treating absent params as an empty
// object so a method with only optional fields can be called with none.
func decodeParams(raw json.RawMessage, dst any) *rpcError {
t := skipSpace(raw)
if len(t) == 0 || string(t) == "null" {
return nil
}
// Arrays are legal JSON-RPC (positional params) and are not used by MCP,
// whose methods all take an object. Saying so beats a confusing type error.
if t[0] == '[' {
return errInvalidParams("params must be an object; positional params are not supported")
}
if err := json.Unmarshal(t, dst); err != nil {
return errInvalidParams(err.Error())
}
return nil
}
func skipSpace(b []byte) []byte {
i := 0
for i < len(b) {
switch b[i] {
case ' ', '\t', '\r', '\n':
i++
default:
return b[i:]
}
}
return nil
}

View File

@@ -0,0 +1,595 @@
package mcpserver
import (
"context"
"encoding/json"
"io"
"log/slog"
"net/http"
"net/http/httptest"
"strings"
"testing"
"github.com/krow/krow-backend/go-api/internal/authctx"
"github.com/krow/krow-backend/go-api/internal/runtime"
"github.com/krow/krow-backend/go-api/internal/tools"
)
// The tool registry these tests run against is the REAL one — the same
// construction the service uses — with a nil database and a nil retriever.
//
// That is sound because nothing here dispatches: Phase 1's tools/call refuses
// before reaching a handler, and every other assertion is about metadata, which
// is fixed at registration. Testing against a hand-built fixture registry would
// be testing a copy of the thing under test: the whole claim of this package is
// "what MCP publishes is what the registry holds", and a fixture would let that
// claim pass while being false of the real set.
// Both dependencies are nil: the handlers capture them in closures and nothing
// here reaches a handler, so nothing dereferences them. If a future test does
// dispatch, this will nil-panic loudly rather than quietly reading a database
// it should not have.
func testRegistry(t *testing.T) *tools.Registry {
t.Helper()
return runtime.DefaultTools(nil, nil)
}
// newServer wires the real registry to a test-only token authenticator.
//
// PHASE 2 NOTE: tools/list and tools/call now require a bearer token, so this
// fixture authenticates. Not one assertion in this file changed — only the
// fixture gained a credential. The behaviour change is the point of Phase 2 and
// is asserted directly in auth_test.go (TestBearerRejection), rather than being
// papered over here.
func newServer(t *testing.T) *Server {
t.Helper()
return New(testRegistry(t), &fakeTokens{byToken: map[string]authctx.Identity{
validToken: identityFor("org-test", "admin"),
}}, slog.New(slog.NewTextHandler(io.Discard, nil)))
}
// post sends one raw body to the MCP endpoint and returns the recorder.
func post(t *testing.T, s *Server, body string) *httptest.ResponseRecorder {
t.Helper()
req := httptest.NewRequest(http.MethodPost, "/mcp", strings.NewReader(body))
req.Header.Set("Content-Type", "application/json")
req.Header.Set("Authorization", "Bearer "+validToken)
rec := httptest.NewRecorder()
s.Handler().ServeHTTP(rec, req)
return rec
}
// decode reads a JSON-RPC response, failing the test if the envelope is wrong.
func decode(t *testing.T, rec *httptest.ResponseRecorder) response {
t.Helper()
var resp response
if err := json.Unmarshal(rec.Body.Bytes(), &resp); err != nil {
t.Fatalf("response was not JSON: %v\nbody: %s", err, rec.Body.String())
}
if resp.JSONRPC != jsonRPCVersion {
t.Fatalf("jsonrpc = %q, want %q", resp.JSONRPC, jsonRPCVersion)
}
return resp
}
// resultInto re-decodes a successful result into dst.
func resultInto(t *testing.T, rec *httptest.ResponseRecorder, dst any) {
t.Helper()
var envelope struct {
Result json.RawMessage `json:"result"`
Error *rpcError `json:"error"`
}
if err := json.Unmarshal(rec.Body.Bytes(), &envelope); err != nil {
t.Fatalf("response was not JSON: %v", err)
}
if envelope.Error != nil {
t.Fatalf("expected a result, got error %d: %s", envelope.Error.Code, envelope.Error.Message)
}
if err := json.Unmarshal(envelope.Result, dst); err != nil {
t.Fatalf("result did not decode: %v\nresult: %s", err, envelope.Result)
}
}
/* ── 1. initialize ──────────────────────────────────────────────────────── */
func TestInitializeSucceeds(t *testing.T) {
s := newServer(t)
rec := post(t, s, `{"jsonrpc":"2.0","id":1,"method":"initialize","params":{
"protocolVersion":"2025-06-18",
"capabilities":{},
"clientInfo":{"name":"test-client","version":"1.0"}}}`)
if rec.Code != http.StatusOK {
t.Fatalf("status = %d, want 200", rec.Code)
}
if ct := rec.Header().Get("Content-Type"); !strings.HasPrefix(ct, "application/json") {
t.Errorf("Content-Type = %q, want application/json", ct)
}
var out initializeResult
resultInto(t, rec, &out)
if out.ProtocolVersion != ProtocolVersion {
t.Errorf("protocolVersion = %q, want %q", out.ProtocolVersion, ProtocolVersion)
}
if out.ServerInfo["name"] != ServerName {
t.Errorf("serverInfo.name = %v, want %q", out.ServerInfo["name"], ServerName)
}
// Exactly one capability, because exactly one is implemented. Advertising
// resources or prompts would promise methods that answer method-not-found.
if _, ok := out.Capabilities["tools"]; !ok {
t.Error("capabilities.tools is missing")
}
for _, unimplemented := range []string{"resources", "prompts", "sampling", "logging"} {
if _, ok := out.Capabilities[unimplemented]; ok {
t.Errorf("capabilities advertises %q, which is not implemented", unimplemented)
}
}
}
func TestInitializeEchoesTheRequestID(t *testing.T) {
s := newServer(t)
// A string id, to prove ids are echoed verbatim rather than coerced.
rec := post(t, s, `{"jsonrpc":"2.0","id":"abc-123","method":"initialize","params":{}}`)
resp := decode(t, rec)
if string(resp.ID) != `"abc-123"` {
t.Errorf("id = %s, want \"abc-123\"", resp.ID)
}
}
/* ── 2 & 3. tools/list and exact exposure ───────────────────────────────── */
// wantExposed is the Phase 0 read-only MVP set, written out in full.
//
// Deliberately a literal rather than a filter over the registry: a test that
// derived the expectation the same way the code does would pass no matter what
// the rule became. This is the list a human agreed to, and it is the thing that
// should fail when the rule changes.
var wantExposed = []string{
"activity_breakdown",
"activity_signals",
"available_workers",
"candidates_awaiting",
"candidates_quality",
"hires_performance",
"hires_recent",
"open_positions",
"operations_risk",
"positions_risk",
"talent_pool",
"workforce_attendance",
"workforce_coverage",
"workforce_overtime",
"workforce_training",
"workspace_summary",
}
func TestToolsListExposesExactlyTheReadOnlyMVPSet(t *testing.T) {
s := newServer(t)
rec := post(t, s, `{"jsonrpc":"2.0","id":2,"method":"tools/list"}`)
var out toolsListResult
resultInto(t, rec, &out)
got := make([]string, 0, len(out.Tools))
for _, tool := range out.Tools {
got = append(got, tool.Name)
}
if len(got) != len(wantExposed) {
t.Fatalf("exposed %d tools, want %d\n got: %v\nwant: %v",
len(got), len(wantExposed), got, wantExposed)
}
for i := range wantExposed {
if got[i] != wantExposed[i] {
t.Errorf("tool[%d] = %q, want %q", i, got[i], wantExposed[i])
}
}
}
func TestToolsListWorksWithoutParams(t *testing.T) {
s := newServer(t)
// No params key at all — every exposed tool's arguments are optional, so a
// bare list call must work.
rec := post(t, s, `{"jsonrpc":"2.0","id":3,"method":"tools/list"}`)
var out toolsListResult
resultInto(t, rec, &out)
if len(out.Tools) == 0 {
t.Fatal("tools/list returned nothing")
}
}
/* ── 4 & 5. write tools and knowledge_search are not exposed ────────────── */
func TestNoWriteToolIsExposed(t *testing.T) {
s := newServer(t)
rec := post(t, s, `{"jsonrpc":"2.0","id":4,"method":"tools/list"}`)
var out toolsListResult
resultInto(t, rec, &out)
reg := testRegistry(t)
for _, published := range out.Tools {
tool, ok := reg.Get(published.Name)
if !ok {
t.Fatalf("tools/list published %q, which is not in the registry", published.Name)
}
if tool.Effect != tools.EffectRead {
t.Errorf("%q is exposed with effect %q; only read tools may be exposed",
tool.Name, tool.Effect)
}
if tool.RequiresConfirmation {
t.Errorf("%q is exposed and requires confirmation, which this surface cannot obtain",
tool.Name)
}
// The annotation must agree with the registry, or a client shows a
// person the wrong thing before approving.
if published.Annotations == nil || !published.Annotations.ReadOnlyHint {
t.Errorf("%q is not annotated readOnlyHint", tool.Name)
}
if published.Annotations != nil && published.Annotations.DestructiveHint {
t.Errorf("%q is annotated destructiveHint", tool.Name)
}
}
}
func TestNamedExclusionsAreNotExposed(t *testing.T) {
s := newServer(t)
rec := post(t, s, `{"jsonrpc":"2.0","id":5,"method":"tools/list"}`)
var out toolsListResult
resultInto(t, rec, &out)
published := map[string]bool{}
for _, tool := range out.Tools {
published[tool.Name] = true
}
// The two writes and the one deferred read. Named explicitly, because these
// three are the ones a future change is most likely to let through.
for _, forbidden := range []string{"assign_worker", "move_application", "knowledge_search"} {
if published[forbidden] {
t.Errorf("%q must not be exposed", forbidden)
}
}
}
// The registry must still hold the excluded tools — they are withheld from this
// surface, not removed from KROW. A test that only checked absence would pass
// if somebody deleted them.
func TestExcludedToolsStillExistInTheRegistry(t *testing.T) {
reg := testRegistry(t)
for _, name := range []string{"assign_worker", "move_application", "knowledge_search"} {
if _, ok := reg.Get(name); !ok {
t.Errorf("%q is missing from the registry entirely", name)
}
}
}
/* ── 6. schemas come from the real registry ─────────────────────────────── */
func TestToolSchemaIsSourcedFromTheRegistry(t *testing.T) {
s := newServer(t)
rec := post(t, s, `{"jsonrpc":"2.0","id":6,"method":"tools/list"}`)
var out toolsListResult
resultInto(t, rec, &out)
reg := testRegistry(t)
for _, published := range out.Tools {
tool, _ := reg.Get(published.Name)
if published.Description != tool.Description {
t.Errorf("%q description differs from the registry's", tool.Name)
}
if published.InputSchema == nil {
t.Fatalf("%q published a nil inputSchema", tool.Name)
}
// Every exposed tool declares a schema, so the published one must be
// the registry's own map and not the empty-schema fallback.
if tool.InputSchema == nil {
t.Errorf("%q has no InputSchema in the registry; expected every exposed tool to declare one", tool.Name)
continue
}
wantJSON, _ := json.Marshal(tool.InputSchema)
gotJSON, _ := json.Marshal(published.InputSchema)
if string(wantJSON) != string(gotJSON) {
t.Errorf("%q schema differs from the registry's\n got: %s\nwant: %s",
tool.Name, gotJSON, wantJSON)
}
// A published schema must be a JSON Schema object, or a client may
// refuse the whole list.
if published.InputSchema["type"] != "object" {
t.Errorf("%q schema type = %v, want \"object\"", tool.Name, published.InputSchema["type"])
}
}
}
// No exposed tool may take a tenant or principal identifier as an argument.
// An argument is something a client chooses; identity is not the client's to
// choose. Enforced here rather than by review, because the cost of missing it
// once is cross-tenant access.
func TestNoToolAcceptsATenantOrPrincipalArgument(t *testing.T) {
s := newServer(t)
rec := post(t, s, `{"jsonrpc":"2.0","id":7,"method":"tools/list"}`)
var out toolsListResult
resultInto(t, rec, &out)
forbidden := []string{
"org_id", "orgid", "organization_id", "organisation_id",
"tenant_id", "tenantid", "tenant",
"user_id", "userid", "principal", "principal_id",
"account_id", "caller", "caller_id", "on_behalf_of", "impersonate",
}
for _, tool := range out.Tools {
props, ok := tool.InputSchema["properties"].(map[string]any)
if !ok {
continue
}
for field := range props {
lower := strings.ToLower(field)
for _, bad := range forbidden {
if lower == bad {
t.Errorf("%q accepts %q as an argument; identity must come from the "+
"authenticated principal, never from the request", tool.Name, field)
}
}
}
}
}
/* ── 7 & 8. unknown method, unknown tool ────────────────────────────────── */
func TestUnknownMethodReturnsMethodNotFound(t *testing.T) {
s := newServer(t)
rec := post(t, s, `{"jsonrpc":"2.0","id":8,"method":"resources/list"}`)
if rec.Code != http.StatusOK {
t.Errorf("status = %d, want 200 — a JSON-RPC error is still a successful HTTP exchange", rec.Code)
}
resp := decode(t, rec)
if resp.Error == nil {
t.Fatal("expected an error")
}
if resp.Error.Code != codeMethodNotFound {
t.Errorf("code = %d, want %d", resp.Error.Code, codeMethodNotFound)
}
}
func TestUnknownToolIsReportedAsAToolError(t *testing.T) {
s := newServer(t)
rec := post(t, s, `{"jsonrpc":"2.0","id":9,"method":"tools/call",
"params":{"name":"no_such_tool","arguments":{}}}`)
var out toolsCallResult
resultInto(t, rec, &out)
if !out.IsError {
t.Error("expected isError on an unknown tool")
}
if len(out.Content) == 0 || !strings.Contains(out.Content[0].Text, "mcp.unknown_tool") {
t.Errorf("expected mcp.unknown_tool, got %+v", out.Content)
}
}
// An unexposed tool must be indistinguishable from a non-existent one, or the
// error messages become an inventory of what this surface is withholding.
func TestUnexposedToolIsIndistinguishableFromUnknown(t *testing.T) {
s := newServer(t)
unknown := post(t, s, `{"jsonrpc":"2.0","id":10,"method":"tools/call","params":{"name":"definitely_not_a_tool"}}`)
withheld := post(t, s, `{"jsonrpc":"2.0","id":10,"method":"tools/call","params":{"name":"assign_worker"}}`)
var a, b toolsCallResult
resultInto(t, unknown, &a)
resultInto(t, withheld, &b)
codeOf := func(r toolsCallResult) string {
var payload struct {
Error struct {
Code string `json:"code"`
} `json:"error"`
}
_ = json.Unmarshal([]byte(r.Content[0].Text), &payload)
return payload.Error.Code
}
if codeOf(a) != codeOf(b) {
t.Errorf("an unexposed tool is distinguishable from an unknown one: %q vs %q",
codeOf(a), codeOf(b))
}
if a.IsError != b.IsError {
t.Error("isError differs between unknown and unexposed tools")
}
}
/* ── 9 & 10. malformed JSON and params ──────────────────────────────────── */
func TestMalformedJSONReturnsParseError(t *testing.T) {
s := newServer(t)
rec := post(t, s, `{"jsonrpc":"2.0","id":1,"method":`)
resp := decode(t, rec)
if resp.Error == nil || resp.Error.Code != codeParseError {
t.Fatalf("want parse error %d, got %+v", codeParseError, resp.Error)
}
}
func TestEnvelopeValidation(t *testing.T) {
for name, tc := range map[string]struct {
body string
want int
}{
"missing jsonrpc": {`{"id":1,"method":"tools/list"}`, codeInvalidRequest},
"wrong jsonrpc": {`{"jsonrpc":"1.0","id":1,"method":"tools/list"}`, codeInvalidRequest},
"missing method": {`{"jsonrpc":"2.0","id":1}`, codeInvalidRequest},
"empty body": {``, codeInvalidRequest},
"object id": {`{"jsonrpc":"2.0","id":{"a":1},"method":"tools/list"}`, codeInvalidRequest},
"not an object": {`"a string"`, codeParseError},
"params as array": {`{"jsonrpc":"2.0","id":1,"method":"tools/call","params":[1,2]}`, codeInvalidParams},
"params wrong type": {`{"jsonrpc":"2.0","id":1,"method":"tools/call","params":{"name":42}}`, codeInvalidParams},
"missing tool name": {`{"jsonrpc":"2.0","id":1,"method":"tools/call","params":{}}`, codeInvalidParams},
} {
t.Run(name, func(t *testing.T) {
rec := post(t, newServer(t), tc.body)
resp := decode(t, rec)
if resp.Error == nil {
t.Fatalf("expected an error, got %s", rec.Body.String())
}
if resp.Error.Code != tc.want {
t.Errorf("code = %d, want %d (%s)", resp.Error.Code, tc.want, resp.Error.Message)
}
})
}
}
/* ── 11. batching is explicitly rejected ────────────────────────────────── */
func TestBatchRequestsAreExplicitlyRejected(t *testing.T) {
s := newServer(t)
rec := post(t, s, `[{"jsonrpc":"2.0","id":1,"method":"tools/list"},
{"jsonrpc":"2.0","id":2,"method":"tools/list"}]`)
resp := decode(t, rec)
if resp.Error == nil {
t.Fatal("a batch must be refused")
}
if resp.Error.Code != codeInvalidRequest {
t.Errorf("code = %d, want %d", resp.Error.Code, codeInvalidRequest)
}
// The refusal must say WHY, so a client can fall back to sending singly
// rather than retrying the same batch forever.
detail, _ := resp.Error.Data.(string)
if !strings.Contains(detail, "batch") {
t.Errorf("the refusal does not mention batching: %q", detail)
}
}
/* ── 12. no panic on malformed input ────────────────────────────────────── */
func TestNoPanicOnHostileInput(t *testing.T) {
// Each of these has crashed a hand-written JSON-RPC server somewhere.
bodies := []string{
``,
` `,
`null`,
`[]`,
`[[[[[]]]]]`,
`{`,
`}`,
`{"jsonrpc":"2.0"}`,
`{"jsonrpc":null,"id":null,"method":null}`,
`{"jsonrpc":"2.0","id":1,"method":"tools/call","params":null}`,
`{"jsonrpc":"2.0","id":1,"method":"tools/call","params":{"name":""}}`,
`{"jsonrpc":"2.0","id":1,"method":"tools/call","params":{"name":"activity_breakdown","arguments":"not-an-object"}}`,
`{"jsonrpc":"2.0","id":1,"method":"tools/call","params":{"name":"activity_breakdown","arguments":[]}}`,
`{"jsonrpc":"2.0","id":[1,2,3],"method":"tools/list"}`,
`{"jsonrpc":"2.0","id":1,"method":"` + strings.Repeat("A", 10_000) + `"}`,
`{"jsonrpc":"2.0","id":1,"method":"tools/list","params":{"cursor":` + strings.Repeat("[", 500) + `}}`,
"\x00\x01\x02",
}
s := newServer(t)
for i, body := range bodies {
// A panic escaping here fails the test by crashing it, which is the
// assertion: the handler must answer every one of these.
rec := post(t, s, body)
if rec.Code < 200 || rec.Code >= 600 {
t.Errorf("body %d produced status %d", i, rec.Code)
}
if rec.Body.Len() > 0 && rec.Code != http.StatusAccepted {
var resp response
if err := json.Unmarshal(rec.Body.Bytes(), &resp); err != nil {
t.Errorf("body %d produced a non-JSON response: %s", i, rec.Body.String())
}
}
}
}
/* ── Transport ──────────────────────────────────────────────────────────── */
func TestOnlyPOSTIsAccepted(t *testing.T) {
s := newServer(t)
for _, method := range []string{http.MethodGet, http.MethodPut, http.MethodDelete, http.MethodPatch} {
req := httptest.NewRequest(method, "/mcp", nil)
rec := httptest.NewRecorder()
s.Handler().ServeHTTP(rec, req)
if rec.Code != http.StatusMethodNotAllowed {
t.Errorf("%s: status = %d, want 405", method, rec.Code)
}
if allow := rec.Header().Get("Allow"); allow != http.MethodPost {
t.Errorf("%s: Allow = %q, want POST", method, allow)
}
}
}
func TestNonJSONContentTypeIsRefused(t *testing.T) {
s := newServer(t)
req := httptest.NewRequest(http.MethodPost, "/mcp",
strings.NewReader(`{"jsonrpc":"2.0","id":1,"method":"tools/list"}`))
req.Header.Set("Content-Type", "application/x-www-form-urlencoded")
rec := httptest.NewRecorder()
s.Handler().ServeHTTP(rec, req)
if rec.Code != http.StatusUnsupportedMediaType {
t.Errorf("status = %d, want 415", rec.Code)
}
}
func TestOversizedBodyIsRefused(t *testing.T) {
s := newServer(t)
huge := `{"jsonrpc":"2.0","id":1,"method":"tools/list","params":{"cursor":"` +
strings.Repeat("x", MaxRequestBytes+1) + `"}}`
rec := post(t, s, huge)
if rec.Code != http.StatusRequestEntityTooLarge {
t.Errorf("status = %d, want 413", rec.Code)
}
}
func TestNotificationGetsNoBody(t *testing.T) {
s := newServer(t)
// No id — a notification. The spec forbids a response.
rec := post(t, s, `{"jsonrpc":"2.0","method":"notifications/initialized"}`)
if rec.Code != http.StatusAccepted {
t.Errorf("status = %d, want 202", rec.Code)
}
if rec.Body.Len() != 0 {
t.Errorf("a notification was answered with a body: %s", rec.Body.String())
}
}
/* ── Phase 1 boundary ───────────────────────────────────────────────────── */
// Handle must refuse a nil identity even when called directly.
//
// The transport answers 401 before this point in the ordinary case, so this is
// the SECOND gate: a future caller of Handle that forgets to authenticate must
// fail closed rather than dispatch as nobody. It exists to fail loudly if
// somebody later supplies a default identity to "make it work".
func TestHandleRefusesANilIdentity(t *testing.T) {
s := newServer(t)
req := request{
JSONRPC: jsonRPCVersion,
ID: json.RawMessage(`11`),
Method: "tools/call",
Params: json.RawMessage(`{"name":"activity_breakdown","arguments":{}}`),
}
result, rpcErr := s.Handle(context.Background(), nil, req)
if rpcErr != nil {
t.Fatalf("unexpected rpc error: %v", rpcErr)
}
out, ok := result.(toolsCallResult)
if !ok {
t.Fatalf("unexpected result type %T", result)
}
if !out.IsError || !strings.Contains(out.Content[0].Text, "mcp.unauthenticated") {
t.Errorf("expected mcp.unauthenticated, got: %+v", out.Content)
}
}
func TestPingIsCheap(t *testing.T) {
s := newServer(t)
rec := post(t, s, `{"jsonrpc":"2.0","id":12,"method":"ping"}`)
var out map[string]any
resultInto(t, rec, &out)
if len(out) != 0 {
t.Errorf("ping returned %v, want an empty result", out)
}
}

View File

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

View File

@@ -0,0 +1,387 @@
package mcpserver
import (
"context"
"encoding/json"
"log/slog"
"time"
"github.com/krow/krow-backend/go-api/internal/authctx"
"github.com/krow/krow-backend/go-api/internal/tools"
)
// ProtocolVersion is the MCP revision this server implements.
//
// Echoed back from initialize. A client asking for a different revision is not
// refused: the spec's negotiation is that the server states what it speaks and
// the client decides whether it can work with that. Refusing would turn a
// version skew into an outage where it is usually a compatible difference.
const ProtocolVersion = "2025-06-18"
// ServerName and ServerVersion identify this implementation to a client.
const (
ServerName = "krow-mcp"
ServerVersion = "0.1.0"
)
// maxToolArgumentBytes bounds one tool call's arguments.
//
// The same reasoning as runs.go's maxRunRequestBytes: arguments are a handful
// of scalars against a schema that sets additionalProperties:false, so anything
// large is either a mistake or an attempt to push text past the tool layer into
// a prompt. The transport bounds the whole body too; this bounds the part that
// reaches a handler.
const maxToolArgumentBytes = 64 << 10
// Server answers MCP methods against the existing tool registry.
//
// It holds a *tools.Registry and nothing else that matters. There is no second
// registry, no adapter table and no per-tool code in this package: what MCP
// publishes is what the registry holds, filtered by the rule in tools.go.
type Server struct {
reg *tools.Registry
log *slog.Logger
// tokens resolves a bearer credential into an identity. Nil means this
// surface cannot authenticate anyone, and every call is refused — see
// ErrNoAuthenticator. Failing closed is the only safe default for a field
// whose absence would otherwise mean "let everyone in".
tokens TokenAuthenticator
// resourceMetadataURL is where a 401 points a client so it can begin
// discovery. Empty means the challenge carries no pointer, which is a
// valid but less useful 401: a client then has nowhere to look.
resourceMetadataURL string
// orgLimiter bounds tool calls per organisation.
//
// HERE rather than in the HTTP middleware, and that placement is the whole
// point: an organisation is not knowable until the bearer token has been
// resolved to a user and that user's row read. A middleware running before
// authentication could only key by something the CLIENT supplied, which is
// precisely the identity this surface refuses to trust.
//
// Nil means no per-organisation ceiling, which is the correct default for
// a deployment that has not configured one.
orgLimiter OrgLimiter
}
// OrgLimiter bounds how much one organisation may ask for.
//
// Takes an org id that the caller has already established from an authenticated
// identity. It cannot be handed anything from a request, because the only
// caller is dispatch, which has an authctx.Identity and nothing else.
type OrgLimiter interface {
// AllowOrg reports whether this organisation may make another call, and
// how long until its window rolls over.
AllowOrg(ctx context.Context, orgID string) (allowed bool, retryAfter time.Duration, err error)
}
// WithOrgLimiter installs the per-organisation ceiling.
func (s *Server) WithOrgLimiter(l OrgLimiter) *Server {
s.orgLimiter = l
return s
}
// WithResourceMetadataURL sets the RFC 9728 document a 401 points at.
//
// Supplied by the caller rather than derived here, because this package does
// not know its own deployment's URLs and must not invent them. A hardcoded
// hostname would be one deployment's identity baked into every other one.
func (s *Server) WithResourceMetadataURL(u string) *Server {
s.resourceMetadataURL = u
return s
}
// challenge builds the WWW-Authenticate header for a 401.
//
// RFC 9728 section 5.1: the client reads `resource_metadata` from here to find
// the protected-resource document, and from there the authorization server.
// Without the parameter a compliant client has a 401 and nowhere to go, which
// is why this is the difference between "authentication failed" and "here is
// how to authenticate".
func (s *Server) challenge() string {
c := `Bearer realm="` + ServerName + `"`
if s.resourceMetadataURL != "" {
c += `, resource_metadata="` + s.resourceMetadataURL + `"`
}
return c
}
// New builds a server over an existing registry.
//
// The registry is the one the rest of the service already built — the caller
// passes runtime.DefaultTools(...)'s result, the same value the HTTP server
// uses for its author catalogue. Taking it as a parameter rather than building
// one here is what guarantees there is only ever one.
func New(reg *tools.Registry, tokens TokenAuthenticator, log *slog.Logger) *Server {
if log == nil {
log = slog.Default()
}
return &Server{reg: reg, tokens: tokens, log: log}
}
// Handle dispatches one parsed JSON-RPC request on behalf of an identity.
//
// The identity is a PARAMETER, not something read from the context, and that is
// the security property rather than a style choice. If this function resolved
// the caller from ctx, then mounting the endpoint behind the cookie middleware
// would make a browser session sufficient to call MCP tools — the middleware
// puts an Identity in the context, and this code would find it. Taking it as an
// argument means only the MCP transport's own bearer authentication can supply
// one. See auth.go.
//
// A nil identity means unauthenticated. The three handshake methods are allowed
// without one; tools/call is not.
//
// Returns a result or an error, never both. Notifications are handled by the
// transport, which discards whatever comes back.
func (s *Server) Handle(ctx context.Context, ident *authctx.Identity, req request) (any, *rpcError) {
switch req.Method {
case "initialize":
return s.handleInitialize(req.Params)
case "notifications/initialized":
// The client telling us it is ready. Nothing to do, and answering is
// not required — it arrives as a notification.
return map[string]any{}, nil
case "ping":
// Cheap liveness, defined by the spec as an empty result. Costs nothing
// and saves a client from using tools/list as a heartbeat.
return map[string]any{}, nil
case "tools/list":
return s.handleToolsList(req.Params)
case "tools/call":
return s.handleToolsCall(ctx, ident, req.Params)
default:
return nil, errMethodNotFound(req.Method)
}
}
/* ── initialize ─────────────────────────────────────────────────────────── */
type initializeParams struct {
ProtocolVersion string `json:"protocolVersion"`
Capabilities json.RawMessage `json:"capabilities"`
ClientInfo struct {
Name string `json:"name"`
Version string `json:"version"`
} `json:"clientInfo"`
}
type initializeResult struct {
ProtocolVersion string `json:"protocolVersion"`
Capabilities map[string]any `json:"capabilities"`
ServerInfo map[string]any `json:"serverInfo"`
Instructions string `json:"instructions,omitempty"`
}
// handleInitialize answers the opening handshake.
//
// Declares exactly one capability, because exactly one is implemented. A server
// that advertised resources or prompts here would be promising methods that
// answer method-not-found, and a client would reasonably call them.
//
// listChanged is false: the tool set is fixed at process start by the registry,
// so there is no change to notify anyone about.
func (s *Server) handleInitialize(raw json.RawMessage) (any, *rpcError) {
var p initializeParams
if err := decodeParams(raw, &p); err != nil {
return nil, err
}
s.log.Info("mcp initialize",
"client_name", p.ClientInfo.Name,
"client_version", p.ClientInfo.Version,
"client_protocol", p.ProtocolVersion,
"server_protocol", ProtocolVersion)
return initializeResult{
ProtocolVersion: ProtocolVersion,
Capabilities: map[string]any{
"tools": map[string]any{"listChanged": false},
},
ServerInfo: map[string]any{
"name": ServerName,
"version": ServerVersion,
},
Instructions: "Read-only access to KROW workforce and hiring data. " +
"Every call is scoped to the authenticated user's organisation and role; " +
"results are structured data for you to summarise, not prose.",
}, nil
}
/* ── tools/list ─────────────────────────────────────────────────────────── */
type toolsListResult struct {
Tools []mcpTool `json:"tools"`
}
// handleToolsList publishes the exposed tools, straight from the registry.
func (s *Server) handleToolsList(raw json.RawMessage) (any, *rpcError) {
// Params are optional here (cursor, for pagination this server does not
// need), but a malformed object is still worth refusing rather than
// ignoring — silently accepting nonsense trains a client to send it.
var p struct {
Cursor string `json:"cursor,omitempty"`
}
if err := decodeParams(raw, &p); err != nil {
return nil, err
}
infos := exposed(s.reg)
out := make([]mcpTool, 0, len(infos))
for _, info := range infos {
out = append(out, toMCPTool(info))
}
return toolsListResult{Tools: out}, nil
}
/* ── tools/call ─────────────────────────────────────────────────────────── */
type toolsCallParams struct {
Name string `json:"name"`
Arguments json.RawMessage `json:"arguments,omitempty"`
}
// toolsCallResult is MCP's shape for a tool's output.
//
// IsError is part of the RESULT, not a JSON-RPC error: a tool that refused is
// not a protocol fault, and reporting it as one would deny the model the chance
// to read the refusal and do something sensible. It is the same distinction
// tools.Result already draws, and runs.go draws for terminations.
type toolsCallResult struct {
Content []contentBlock `json:"content"`
IsError bool `json:"isError,omitempty"`
}
type contentBlock struct {
Type string `json:"type"`
Text string `json:"text"`
}
func textResult(payload any, isError bool) (toolsCallResult, *rpcError) {
encoded, err := json.MarshalIndent(payload, "", " ")
if err != nil {
return toolsCallResult{}, errInternal()
}
return toolsCallResult{
Content: []contentBlock{{Type: "text", Text: string(encoded)}},
IsError: isError,
}, nil
}
// handleToolsCall validates a call and dispatches it through the registry.
//
// The order is: exposure, then bounds, then identity, then dispatch.
//
// Exposure is checked BEFORE identity on purpose. "There is no such tool here"
// does not depend on who is asking, and answering it first means the surface's
// tool inventory is not something an attacker can probe by comparing an
// authenticated 404 against an unauthenticated 401.
//
// Authentication is the transport's job and has already happened by the time
// this runs; `ident` is nil only when it failed or was never attempted. This
// function does not read the ambient context for a caller — see Handle.
func (s *Server) handleToolsCall(ctx context.Context, ident *authctx.Identity, raw json.RawMessage) (any, *rpcError) {
var p toolsCallParams
if err := decodeParams(raw, &p); err != nil {
return nil, err
}
if p.Name == "" {
return nil, errInvalidParams("name is required")
}
if len(p.Arguments) > maxToolArgumentBytes {
return nil, errInvalidParams("arguments are too large")
}
// Unexposed and unknown are the SAME answer, deliberately. See
// isExposedName — distinguishing them inventories what this surface is
// hiding.
if !isExposedName(s.reg, p.Name) {
s.log.Warn("mcp tool call refused", "tool", p.Name, "reason", "not_exposed")
return textResult(map[string]any{
"error": map[string]any{
"code": "mcp.unknown_tool",
"message": "there is no tool called " + p.Name + " on this surface",
},
}, true)
}
// Unauthenticated calls never reach a handler. The transport answers 401
// before this point in the ordinary case; this is the second gate, so that
// a future caller of Handle that forgets to authenticate fails closed
// rather than dispatching as nobody.
if ident == nil {
s.log.Warn("mcp tool call refused", "tool", p.Name, "reason", "no_identity")
return textResult(map[string]any{
"error": map[string]any{
"code": "mcp.unauthenticated",
"message": "this call is not authenticated",
},
}, true)
}
return s.dispatch(ctx, *ident, p)
}
// dispatch runs the tool through the existing registry.
//
// This is the only place this package touches the tool layer, and it is four
// lines on purpose. Everything that decides what comes back — the policy table,
// the org pre-filter, the row scopes, the opaque denial, the truncation — is
// inside Dispatch and the handler beneath it, unchanged and unreachable from
// here.
//
// The principal is the caller's, from the context. It is never read from
// params: an MCP client that could name its own principal could read anything,
// which is the bug I1 exists to prevent.
func (s *Server) dispatch(ctx context.Context, identity authctx.Identity, p toolsCallParams) (any, *rpcError) {
// The per-organisation ceiling, checked after authentication and before
// any work. The org comes from `identity`, which came from the token —
// there is no path by which a request can name a different bucket, because
// this function is never given anything from the request except the tool
// name and its arguments.
if s.orgLimiter != nil {
allowed, retryAfter, err := s.orgLimiter.AllowOrg(ctx, identity.OrgID)
if err != nil {
// The limiter has already decided whether a failure permits the
// call. Logged without the org's usage, which is not the caller's
// business.
s.log.Error("org rate limiter unavailable", "error", err)
}
if !allowed {
s.log.Warn("mcp org rate limit exceeded",
"org_id", identity.OrgID, "tool", p.Name)
return nil, errRateLimited(retryAfter)
}
}
args := p.Arguments
if len(skipSpace(args)) == 0 {
args = json.RawMessage(`{}`)
}
tc := tools.Context{
Principal: identity,
// No RunID: an MCP call is not an agent run and writes no trajectory.
// No KnowledgeSources: there is no spec, which is why knowledge_search
// is deferred rather than published — see tools.go.
}
res := s.reg.Dispatch(ctx, tc, p.Name, args)
s.log.Info("mcp tool call",
"tool", p.Name,
"user_id", identity.UserID,
"org_id", identity.OrgID,
"ok", res.Error == nil,
"truncated", res.Truncated)
if res.Error != nil {
return textResult(map[string]any{"error": res.Error}, true)
}
return textResult(map[string]any{
"data": res.Data,
"truncated": res.Truncated,
}, false)
}

View File

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

View File

@@ -0,0 +1,145 @@
package mcpserver
import (
"sort"
"github.com/krow/krow-backend/go-api/internal/tools"
)
// Which tools this surface publishes, and why it is a rule rather than a list.
//
// The set is DERIVED from the registry on every call, not enumerated. A list
// would be a promise someone has to keep: register a write tool tomorrow,
// forget to update the list, and it ships to every connected client. A
// derivation cannot forget. The only hand-maintained part is `deferred` below,
// which names tools held back for a reason other than their effect — and
// holding something back is the safe direction to be wrong in.
//
// Three conditions, all required:
//
// 1. Effect is read. I4's whole point is that a write is gated; a write
// reachable over a surface with no confirmation round-trip is I4 defeated.
// 2. RequiresConfirmation is false. Belt and braces: Register already forces
// it true for a write, so this catches a READ tool that opted in — some
// reads are expensive enough to be worth asking about — which this surface
// has no way to ask about yet.
// 3. Not in `deferred`.
// deferred names tools held back for a reason that is not their effect.
//
// knowledge_search is read-only and still cannot ship. Its corpora come from
// tools.Context.KnowledgeSources, which the agent loop fills from the running
// agent's SPEC — deliberately, so that which documents may be read is not
// something a model can choose. An MCP call has no spec, so the field is empty,
// and retrieval refuses an empty source list rather than treating it as "all of
// them". The tool would therefore fail every call; publishing it would advertise
// a capability that cannot work.
//
// Making it work is a design decision, not an omission: either the connection
// binds to an agent spec whose sources it inherits, or sources are derived from
// the caller's org ACL (which needs a reindex). Passing them as a tool argument
// is the one option that is ruled out, because that is exactly what the field's
// placement in the spec exists to prevent.
var deferred = map[string]string{
"knowledge_search": "corpora come from an agent spec, which an MCP call does not have",
}
// exposed returns the tools this surface publishes, sorted by name.
//
// Sorted because tools/list is a set, and a stable order makes it diffable in a
// test and in a log.
func exposed(reg *tools.Registry) []tools.ToolInfo {
out := make([]tools.ToolInfo, 0, 16)
for _, info := range reg.Catalogue() {
if !isExposable(info) {
continue
}
out = append(out, info)
}
sort.Slice(out, func(i, j int) bool { return out[i].Name < out[j].Name })
return out
}
// isExposable is the rule, in one place, so the list and the check cannot
// disagree.
func isExposable(info tools.ToolInfo) bool {
if info.Effect != string(tools.EffectRead) {
return false
}
if info.RequiresConfirmation {
return false
}
if _, held := deferred[info.Name]; held {
return false
}
return true
}
// isExposedName reports whether a tool may be called by name over this surface.
//
// tools/call consults this BEFORE the registry, so an unexposed tool answers
// exactly as an unknown one does. The alternative — dispatching and letting
// authorization refuse — would make "this tool exists but you may not reach it
// here" distinguishable from "no such tool", which is an inventory of the
// surface's own blind spots.
func isExposedName(reg *tools.Registry, name string) bool {
t, ok := reg.Get(name)
if !ok {
return false
}
return isExposable(tools.ToolInfo{
Name: t.Name,
Effect: string(t.Effect),
RequiresConfirmation: t.RequiresConfirmation,
})
}
/* ── MCP shapes ─────────────────────────────────────────────────────────── */
// mcpTool is one entry in a tools/list result.
//
// Every field is copied from the registry rather than restated. The annotations
// are hints a client may show a person before approving a call; they are
// derived from Effect so they cannot contradict it.
type mcpTool struct {
Name string `json:"name"`
Description string `json:"description"`
InputSchema map[string]any `json:"inputSchema"`
Annotations *annotations `json:"annotations,omitempty"`
}
// annotations are the advisory hints from the MCP tool definition.
type annotations struct {
ReadOnlyHint bool `json:"readOnlyHint"`
DestructiveHint bool `json:"destructiveHint"`
}
// emptySchema is what a tool with no declared schema publishes.
//
// tools/list requires an inputSchema per tool, and a client given `null` may
// reasonably refuse the whole list. An object that accepts nothing is the
// honest rendering of "this tool takes no arguments".
func emptySchema() map[string]any {
return map[string]any{
"type": "object",
"properties": map[string]any{},
"additionalProperties": false,
}
}
// toMCPTool converts a registry entry into its wire form.
func toMCPTool(info tools.ToolInfo) mcpTool {
schema := info.InputSchema
if schema == nil {
schema = emptySchema()
}
return mcpTool{
Name: info.Name,
Description: info.Description,
InputSchema: schema,
Annotations: &annotations{
ReadOnlyHint: info.Effect == string(tools.EffectRead),
DestructiveHint: info.Effect == string(tools.EffectWrite),
},
}
}

View File

@@ -0,0 +1,216 @@
package mcpserver
import (
"encoding/json"
"io"
"net/http"
"strconv"
"strings"
"time"
"github.com/krow/krow-backend/go-api/internal/authctx"
)
// The Streamable HTTP binding: one endpoint, one message per POST.
//
// The current MCP spec defines two standard transports — stdio and Streamable
// HTTP — and only the second can serve a client that is not launching a
// subprocess. Each message is an HTTP POST to a single MCP endpoint, and the
// reply is a JSON object or a request-scoped SSE stream.
//
// This implementation answers with JSON objects and no stream, which is a
// complete implementation of the binding for this server's methods rather than
// a shortcut: tools/list is a fixed list, and a tool result is structured data
// already bounded by MaxResultBytes. There is nothing to deliver incrementally.
// Streaming becomes worth adding if a long-running method is ever exposed.
//
// STATELESS. No session is minted and no Mcp-Session-Id is required, because
// nothing here is worth remembering between calls: every request carries its own
// bearer credential and every method is independent. Adding session
// state now would be state to expire, to bind to a token, to revalidate and to
// leak — for no behaviour this server has.
// MaxRequestBytes bounds an inbound MCP message.
//
// Well above any legitimate tools/call — arguments are a few scalars — and far
// below anything that would be worth sending here. Enforced with
// http.MaxBytesReader so the body is refused as it arrives rather than after it
// has been buffered.
const MaxRequestBytes = 1 << 20 // 1 MiB
// Handler returns the HTTP handler for the MCP endpoint.
//
// Deliberately an http.Handler rather than a registered route: this package
// does not know its own path, and the server that mounts it decides where it
// lives. Note that it does NOT decide what authenticates it — this handler
// authenticates its own callers from the Authorization header, and ignores
// whatever middleware sits in front. Mounting it behind the cookie middleware
// therefore does not make a cookie sufficient to call it.
func (s *Server) Handler() http.Handler {
return http.HandlerFunc(s.serveHTTP)
}
func (s *Server) serveHTTP(w http.ResponseWriter, r *http.Request) {
// One method. GET is what the binding uses for a server-initiated stream,
// which this server does not open; saying 405 with an Allow header is more
// use to a client than a 404 that suggests the endpoint is absent.
if r.Method != http.MethodPost {
w.Header().Set("Allow", http.MethodPost)
writeRPCError(w, s, http.StatusMethodNotAllowed, nil,
errInvalidRequest("this endpoint accepts POST only"))
return
}
// Content-Type is checked rather than assumed. A form post or a stray
// upload that happened to be valid JSON would otherwise be processed as a
// protocol message.
if ct := r.Header.Get("Content-Type"); ct != "" && !isJSONContentType(ct) {
writeRPCError(w, s, http.StatusUnsupportedMediaType, nil,
errInvalidRequest("Content-Type must be application/json"))
return
}
body, err := io.ReadAll(http.MaxBytesReader(w, r.Body, MaxRequestBytes))
if err != nil {
// MaxBytesReader's error is indistinguishable from a truncated upload
// without type assertions that buy nothing here: both mean the body is
// unusable, and 413 is the more actionable of the two answers.
writeRPCError(w, s, http.StatusRequestEntityTooLarge, nil,
errInvalidRequest("the request body was too large or could not be read"))
return
}
req, rpcErr := parseRequest(body)
if rpcErr != nil {
// A parse failure has no usable id, so the response carries the id the
// message did parse with — null when it parsed with none. HTTP stays
// 200: the transport succeeded and the JSON-RPC error IS the answer.
writeRPCError(w, s, http.StatusOK, req.ID, rpcErr)
return
}
// Authentication, for every method including the handshake — see
// methodRequiresAuth. The 401 below is not merely a refusal: its
// WWW-Authenticate header is the first step of the MCP authorization flow,
// and a client's very first request is what should produce it.
var ident *authctx.Identity
if resolved, err := s.authenticate(r); err == nil {
ident = &resolved
} else if methodRequiresAuth(req.Method) {
s.log.Warn("mcp request refused",
"method", req.Method, "reason", authFailureReason(err))
// WWW-Authenticate is not decoration: RFC 9728 has the client read the
// resource-metadata URL from this header to find the authorization
// server, and from there where to get a token. It is the difference
// between "authentication failed" and "here is how to authenticate".
w.Header().Set("WWW-Authenticate", s.challenge())
writeRPCError(w, s, http.StatusUnauthorized, req.ID,
&rpcError{Code: codeUnauthorized, Message: "authentication required"})
return
}
// A panic in a handler must not take the process down or leak a stack into
// the response. The registry has its own recover around each tool; this is
// the outer net for everything else in this package.
result, rpcErr := s.handleRecovered(r, ident, req)
// A notification gets no response body at all — the spec forbids one.
if req.isNotification() {
w.WriteHeader(http.StatusAccepted)
return
}
if rpcErr != nil {
// A rate-limited refusal is the one JSON-RPC error that also carries an
// HTTP status, because 429 and Retry-After are how a client knows to
// back off. Everything else is 200 with an error body: the transport
// succeeded and the error IS the answer.
if rpcErr.Code == codeRateLimited {
if rpcErr.retryAfter > 0 {
w.Header().Set("Retry-After", retryAfterSeconds(rpcErr.retryAfter))
}
writeRPCError(w, s, http.StatusTooManyRequests, req.ID, rpcErr)
return
}
writeRPCError(w, s, http.StatusOK, req.ID, rpcErr)
return
}
writeJSON(w, s, http.StatusOK, response{
JSONRPC: jsonRPCVersion,
ID: req.ID,
Result: result,
})
}
// methodRequiresAuth reports whether a method may only run for a known caller.
//
// EVERY method does, including the handshake. An earlier revision left
// initialize, ping and notifications/initialized open, on the reasoning that a
// client needs somewhere to start — and that was wrong in a way worth
// recording, because it is the kind of mistake that looks like helpfulness.
//
// The MCP authorization flow begins with the client making an MCP request
// WITHOUT a token and reading the 401's WWW-Authenticate header. A client's
// first request is usually initialize. Answering that one with a cheerful 200
// means the client never sees the challenge, believes it is connected, and
// discovers otherwise only when the first real call fails — by which time it
// has no 401 in hand to discover from. Requiring a token everywhere means the
// very first request, whatever it is, produces the challenge that starts the
// flow.
//
// Nothing is lost. The handshake is not information a stranger needs: it
// returns this server's name and capabilities, which are only useful to a
// client that intends to authenticate anyway.
//
// Kept as a function rather than inlined because it is the single place that
// decision lives, and a future method that genuinely must be open should have
// to be written down here to become so.
func methodRequiresAuth(method string) bool { return true }
// handleRecovered runs Handle with a recover, converting a panic into an
// internal error whose detail goes to the log and not to the caller.
func (s *Server) handleRecovered(r *http.Request, ident *authctx.Identity, req request) (result any, rpcErr *rpcError) {
defer func() {
if p := recover(); p != nil {
s.log.Error("mcp handler panicked", "method", req.Method, "panic", p)
result, rpcErr = nil, errInternal()
}
}()
return s.Handle(r.Context(), ident, req)
}
// retryAfterSeconds renders a duration for the Retry-After header, rounded up
// and never below one second — "Retry-After: 0" invites an immediate retry,
// which is the one thing a limited client must not do.
func retryAfterSeconds(d time.Duration) string {
secs := int(d.Round(time.Second) / time.Second)
if secs < 1 {
secs = 1
}
return strconv.Itoa(secs)
}
// isJSONContentType reports whether a Content-Type header names JSON,
// tolerating parameters such as "; charset=utf-8".
func isJSONContentType(ct string) bool {
media := strings.TrimSpace(strings.SplitN(ct, ";", 2)[0])
return strings.EqualFold(media, "application/json")
}
func writeRPCError(w http.ResponseWriter, s *Server, status int, id json.RawMessage, e *rpcError) {
writeJSON(w, s, status, response{JSONRPC: jsonRPCVersion, ID: id, Error: e})
}
func writeJSON(w http.ResponseWriter, s *Server, status int, payload response) {
encoded, err := json.Marshal(payload)
if err != nil {
// Encoding our own response failed, so there is nothing safe left to
// say in JSON. Log it and send a bare 500.
s.log.Error("mcp response could not be encoded", "error", err)
http.Error(w, "internal error", http.StatusInternalServerError)
return
}
w.Header().Set("Content-Type", "application/json; charset=utf-8")
w.Header().Set("Cache-Control", "no-store")
w.WriteHeader(status)
_, _ = w.Write(encoded)
}

View File

@@ -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")
}
}

View File

@@ -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
}

View File

@@ -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)
})
}

View File

@@ -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.

View File

@@ -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
}

View File

@@ -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
// `<script>…` — and the escaping is what makes displaying it safe.
var consentTemplate = template.Must(template.New("consent").Parse(`<!DOCTYPE html>
<html lang="en">
<head>
<meta charset="utf-8">
<meta name="viewport" content="width=device-width, initial-scale=1">
<title>Authorize access &middot; Krow</title>
<style>
:root { color-scheme: light dark; }
body { margin:0; min-height:100vh; display:flex; align-items:center;
justify-content:center; background:#f4f5f7;
font-family:-apple-system,BlinkMacSystemFont,"Segoe UI",Roboto,sans-serif;
color:#14161a; padding:16px; box-sizing:border-box; }
.card { background:#fff; border:1px solid #e3e5e8; border-radius:12px;
max-width:440px; width:100%; padding:28px; box-sizing:border-box; }
h1 { font-size:19px; margin:0 0 4px; }
.sub { color:#5c6270; font-size:14px; margin:0 0 20px; }
dl { margin:0 0 20px; border-top:1px solid #eceef0; }
.row { display:flex; justify-content:space-between; gap:16px;
padding:11px 0; border-bottom:1px solid #eceef0; font-size:14px; }
dt { color:#5c6270; margin:0; flex:0 0 auto; }
dd { margin:0; text-align:right; word-break:break-word; font-weight:500; }
.grants { background:#f7f8f9; border-radius:8px; padding:14px 16px;
font-size:14px; margin:0 0 20px; }
.grants strong { display:block; margin-bottom:6px; font-size:13px;
text-transform:uppercase; letter-spacing:.04em; color:#5c6270; }
.grants ul { margin:0; padding-left:18px; }
.grants li { margin:3px 0; }
.actions { display:flex; gap:10px; }
button { flex:1; padding:11px 16px; border-radius:8px; font-size:15px;
font-weight:500; cursor:pointer; border:1px solid transparent; }
.approve { background:#14161a; color:#fff; }
.deny { background:#fff; color:#14161a; border-color:#d4d7dc; }
.note { margin:16px 0 0; font-size:12.5px; color:#787e8a; line-height:1.5; }
@media (prefers-color-scheme: dark) {
body { background:#0e1013; color:#e9eaec; }
.card { background:#16191d; border-color:#282c33; }
dl,.row { border-color:#282c33; }
.grants { background:#1c2026; }
.approve { background:#e9eaec; color:#14161a; }
.deny { background:#16191d; color:#e9eaec; border-color:#3a3f47; }
dt,.sub,.note,.grants strong { color:#9aa1ad; }
}
</style>
</head>
<body>
<main class="card">
<h1>Authorize access to Krow</h1>
<p class="sub"><strong>{{.ClientName}}</strong> is asking to connect to your Krow workspace.</p>
<dl>
<div class="row"><dt>Application</dt><dd>{{.ClientName}}</dd></div>
<div class="row"><dt>Signed in as</dt><dd>{{.UserEmail}}</dd></div>
<div class="row"><dt>Organisation</dt><dd>{{.OrgName}}</dd></div>
<div class="row"><dt>Connecting to</dt><dd>{{.Resource}}</dd></div>
</dl>
<div class="grants">
<strong>This will allow it to</strong>
<ul>{{range .Grants}}<li>{{.}}</li>{{end}}</ul>
</div>
<form method="POST" action="{{.FormAction}}">
{{range $k, $v := .Hidden}}<input type="hidden" name="{{$k}}" value="{{$v}}">{{end}}
<input type="hidden" name="csrf" value="{{.CSRF}}">
<div class="actions">
<button type="submit" name="decision" value="deny" class="deny">Deny</button>
<button type="submit" name="decision" value="approve" class="approve">Approve</button>
</div>
</form>
<p class="note">Approving lets this application read Krow data that you can
already see, as you, in this organisation. It cannot make changes. You can
disconnect it at any time from your Krow settings.</p>
</main>
</body>
</html>`))
// consentView is what the template renders.
type consentView struct {
ClientName string
UserEmail string
OrgName string
Resource string
Grants []string
FormAction string
Hidden map[string]string
CSRF string
}
// grantsFor renders scopes as sentences a person can act on.
//
// "krow.read" means nothing to the person being asked. A consent screen that
// shows a scope identifier is a consent screen that has not obtained informed
// consent — it has obtained a click.
func grantsFor(scopes []string) []string {
out := make([]string, 0, len(scopes))
for _, scope := range scopes {
switch scope {
case ScopeRead:
out = append(out,
"Read workforce activity, staff, candidates and positions",
"See only what your own Krow account can see",
)
case ScopeWrite:
// Unreachable: krow.write is never issued and never registered.
// Present so that if it ever is, it arrives with words attached
// rather than as a bare identifier on a screen.
out = append(out, "Make changes to your Krow data")
default:
out = append(out, scope)
}
}
return out
}
// renderConsent shows the approval form.
func (s *Server) renderConsent(w http.ResponseWriter, r *http.Request, p authorizeParams, identity authctx.Identity, csrf string) {
client, err := s.store.FindClient(r.Context(), p.ClientID)
if err != nil {
writeOAuthError(w, http.StatusBadRequest, errInvalidClient, "unknown client")
return
}
name := strings.TrimSpace(client.ClientName)
if name == "" {
name = "An application"
}
orgName := s.orgNameFor(r, identity.OrgID)
// Everything needed to complete the flow rides in hidden fields, so the
// POST carries its own context and the server keeps no pending-request
// state. State on the server would be state to expire and to clean up, for
// a decision that is made in the next few seconds.
hidden := map[string]string{
"client_id": p.ClientID,
"redirect_uri": p.RedirectURI,
"response_type": p.ResponseType,
"scope": strings.Join(p.Scopes, " "),
"state": p.State,
"code_challenge": p.CodeChallenge,
"code_challenge_method": p.CodeChallengeMethod,
"resource": p.Resource,
}
w.Header().Set("Content-Type", "text/html; charset=utf-8")
// A consent page names a client and an organisation and must never be
// served from a cache to the next person on a shared machine.
w.Header().Set("Cache-Control", "no-store, private")
w.Header().Set("Pragma", "no-cache")
// Defence in depth for a page that renders an attacker-supplied name:
// no framing (so it cannot be clickjacked into an invisible overlay), no
// referrer (so the query string does not leak to the client's site), and a
// CSP that forbids script entirely — this page has none.
w.Header().Set("X-Frame-Options", "DENY")
w.Header().Set("Referrer-Policy", "no-referrer")
w.Header().Set("Content-Security-Policy", consentCSP(client.RedirectURIs))
w.Header().Set("X-Content-Type-Options", "nosniff")
w.WriteHeader(http.StatusOK)
_ = consentTemplate.Execute(w, consentView{
ClientName: name,
UserEmail: identity.Email,
OrgName: orgName,
Resource: p.Resource,
Grants: grantsFor(p.Scopes),
FormAction: s.cfg.AuthorizePath,
Hidden: hidden,
CSRF: csrf,
})
}
// orgNameFor resolves an organisation's display name.
//
// Best effort: a missing name degrades to the id rather than failing the flow.
// A consent screen that will not render because of a display lookup is worse
// than one that shows a uuid.
func (s *Server) orgNameFor(r *http.Request, orgID string) string {
if orgID == "" {
return "your organisation"
}
var name string
if err := s.store.db.QueryRow(r.Context(),
`SELECT name FROM organizations WHERE id = $1::uuid`, orgID).Scan(&name); err != nil {
return orgID
}
if strings.TrimSpace(name) == "" {
return orgID
}
return name
}
/* ── The consent page's Content-Security-Policy ─────────────────────────── */
// consentCSP builds the policy for the consent page.
//
// WHY form-action CANNOT BE 'self' ALONE
//
// It was, and that was a real bug: the consent form is blocked in the browser
// before it can submit. A consent form's successful submission ends, by
// definition, at the OAuth client's registered redirect_uri — a third party's
// callback, always cross-origin. Browsers enforce form-action across the whole
// navigation chain including redirects (MDN carries an explicit warning that
// this is inconsistent between engines; Chrome blocks, older Firefox did not),
// so `form-action 'self'` makes the flow impossible to complete rather than
// merely strict.
//
// The tests did not catch it because httptest executes no CSP. They asserted
// the header's value, which was set exactly as intended; only a real browser
// could show that what was intended was wrong.
//
// # WHAT IS ALLOWED INSTEAD
//
// 'self', plus the ORIGINS OF THIS CLIENT'S OWN REGISTERED REDIRECT URIs, and
// nothing else. That is narrower than it may look:
//
// - The URIs were validated at registration — absolute, https (or http on
// loopback), no fragment. That validation is untouched.
// - The authorization endpoint still matches the presented redirect_uri
// against the registration byte-for-byte. This policy does not widen what
// a flow may redirect to; it only stops the browser blocking the redirect
// the server was already going to permit.
// - Each client gets its own policy, built from its own registration, so one
// client's callback never appears in another's page.
//
// A URI that cannot be reduced to a safe origin is DROPPED rather than
// broadened. The failure mode is a consent page whose form the browser blocks —
// visible, and the safe direction — never a policy that permits more.
func consentCSP(redirectURIs []string) string {
directives := []string{
"default-src 'none'",
"style-src 'unsafe-inline'",
"frame-ancestors 'none'",
}
formAction := "form-action 'self'"
for _, origin := range redirectOrigins(redirectURIs) {
formAction += " " + origin
}
directives = append(directives, formAction)
return strings.Join(directives, "; ")
}
// redirectOrigins reduces registered redirect URIs to CSP source expressions.
//
// A CSP source is an ORIGIN — scheme, host and port — never a path. Emitting
// the full URI would be wrong twice: CSP would match it as a path prefix, and a
// path is not what a form navigation is checked against.
//
// Every value is dropped unless it is unambiguously safe:
//
// unparseable → dropped (never widened to a bare scheme)
// no scheme or no host → dropped
// scheme other than
// http/https → dropped; a custom scheme in a policy is a source
// any app on the machine could claim
// wildcard or separator → dropped; '*', ';' ',' or whitespace in a source
// would either broaden the policy or split the
// header. Registration already refuses these, so
// this is the second lock on the same door.
//
// Duplicates are collapsed so two URIs on one host produce one source, and the
// order registered is preserved so the header is stable and diffable.
func redirectOrigins(redirectURIs []string) []string {
seen := make(map[string]bool, len(redirectURIs))
out := make([]string, 0, len(redirectURIs))
for _, raw := range redirectURIs {
parsed, err := url.Parse(strings.TrimSpace(raw))
if err != nil {
continue
}
scheme := strings.ToLower(parsed.Scheme)
if scheme != "http" && scheme != "https" {
continue
}
// parsed.Host carries host and port together, which is exactly a CSP
// source's host-part. Empty means the URI was relative or malformed.
host := parsed.Host
if host == "" {
continue
}
origin := scheme + "://" + host
// Nothing that could broaden the policy or break the header out of its
// directive. A registered URI cannot contain these — validateRedirectURI
// rejects them — and this refuses to depend on that being true.
if strings.ContainsAny(origin, "*; ,\t\r\n'\"") {
continue
}
if seen[origin] {
continue
}
seen[origin] = true
out = append(out, origin)
}
return out
}

View File

@@ -0,0 +1,501 @@
package oauth
import (
"context"
"encoding/json"
"net/http"
"net/http/httptest"
"net/url"
"strings"
"testing"
)
/* ── Consent is required ────────────────────────────────────────────────── */
// A GET must ASK, not grant. This is the Phase 4 behaviour change, asserted
// directly: before, a signed-in user's authorization was approved on sight.
func TestAuthorizeRendersConsentRatherThanIssuingACode(t *testing.T) {
h := newHarness(t)
clientID := h.register()
rec := h.authorize(authorizeParamsFor(clientID, verifier43))
if rec.Code == http.StatusFound {
t.Fatalf("a GET issued a code without asking: %s", rec.Header().Get("Location"))
}
if rec.Code != http.StatusOK {
t.Fatalf("status = %d, want 200 with a consent page", rec.Code)
}
if ct := rec.Header().Get("Content-Type"); !strings.HasPrefix(ct, "text/html") {
t.Errorf("Content-Type = %q, want text/html", ct)
}
// No grant row may exist yet: rendering a question must not spend anything.
var codes int
if err := h.h.Pool.QueryRow(context.Background(),
`SELECT count(*) FROM oauth_grants`).Scan(&codes); err != nil {
t.Fatalf("count grants: %v", err)
}
if codes != 0 {
t.Errorf("%d authorization codes exist after merely rendering consent", codes)
}
}
// The page must tell a person what they are agreeing to, in their terms.
func TestConsentPageShowsWhatIsBeingGranted(t *testing.T) {
h := newHarness(t)
clientID := h.register()
rec, _ := h.consent(authorizeParamsFor(clientID, verifier43))
body := rec.Body.String()
for name, want := range map[string]string{
"client name": "Test Client",
"signed-in user": "oauth-user@example.test",
"organisation": "OAuth Test",
"resource": testResource,
"approve control": "approve",
"deny control": "deny",
} {
if !strings.Contains(body, want) {
t.Errorf("the consent page does not show the %s (%q)", name, want)
}
}
// A person asked to approve "krow.read" has not been asked anything.
if strings.Contains(body, ScopeRead) && !strings.Contains(body, "Read workforce activity") {
t.Error("the page shows a raw scope identifier without explaining it")
}
// krow.write must never appear on a screen for a flow that cannot grant it.
if strings.Contains(body, ScopeWrite) {
t.Error("the consent page mentions krow.write")
}
}
// The client name is attacker-controlled: anyone may register a client called
// <script>. It must be escaped, not rendered.
func TestConsentPageEscapesTheClientName(t *testing.T) {
h := newHarness(t)
const payload = `<script>alert('xss')</script>`
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(), &reg)
page, _ := h.consent(authorizeParamsFor(reg.ClientID, verifier43))
if strings.Contains(page.Body.String(), "<script>alert") {
t.Fatal("a registered client name was rendered as live HTML")
}
if !strings.Contains(page.Body.String(), "&lt;script&gt;") {
t.Error("the client name does not appear escaped; check it is shown at all")
}
}
/* ── Approve and deny ───────────────────────────────────────────────────── */
func TestConsentApproveIssuesACode(t *testing.T) {
h := newHarness(t)
clientID := h.register()
params := authorizeParamsFor(clientID, verifier43)
_, csrf := h.consent(params)
rec := h.decide(params, "approve", csrf)
if rec.Code != http.StatusFound {
t.Fatalf("status = %d, want 302", rec.Code)
}
loc, _ := url.Parse(rec.Header().Get("Location"))
if loc.Query().Get("code") == "" {
t.Error("approve issued no code")
}
if loc.Query().Get("state") != "xyz" {
t.Errorf("state = %q, want xyz", loc.Query().Get("state"))
}
}
// Denial must reach the client as access_denied, at its registered redirect,
// with state intact and NO code.
func TestConsentDenyReturnsAccessDeniedAndNoCode(t *testing.T) {
h := newHarness(t)
clientID := h.register()
params := authorizeParamsFor(clientID, verifier43)
_, csrf := h.consent(params)
rec := h.decide(params, "deny", csrf)
if rec.Code != http.StatusFound {
t.Fatalf("status = %d, want 302", rec.Code)
}
loc, err := url.Parse(rec.Header().Get("Location"))
if err != nil {
t.Fatalf("Location: %v", err)
}
if !strings.HasPrefix(loc.String(), testRedirect) {
t.Fatalf("denial went to %q, not the registered redirect", loc)
}
if got := loc.Query().Get("error"); got != "access_denied" {
t.Errorf("error = %q, want access_denied", got)
}
if got := loc.Query().Get("state"); got != "xyz" {
t.Errorf("state = %q, want xyz — the client needs it to match the response", got)
}
if loc.Query().Get("code") != "" {
t.Error("a denial returned an authorization code")
}
// And nothing was written.
var codes int
_ = h.h.Pool.QueryRow(context.Background(), `SELECT count(*) FROM oauth_grants`).Scan(&codes)
if codes != 0 {
t.Errorf("%d authorization codes exist after a denial", codes)
}
}
// A POST with no decision must re-ask, never infer approval.
func TestConsentWithNoDecisionDoesNotApprove(t *testing.T) {
h := newHarness(t)
clientID := h.register()
params := authorizeParamsFor(clientID, verifier43)
_, csrf := h.consent(params)
rec := h.decide(params, "", csrf)
if rec.Code == http.StatusFound {
t.Fatalf("a decision-less POST was treated as a decision: %s", rec.Header().Get("Location"))
}
var codes int
_ = h.h.Pool.QueryRow(context.Background(), `SELECT count(*) FROM oauth_grants`).Scan(&codes)
if codes != 0 {
t.Error("a decision-less POST issued a code")
}
}
/* ── CSRF ───────────────────────────────────────────────────────────────── */
// Without the form's token, a cross-site POST must not be able to approve.
// This is what stops a page on the internet connecting a client to somebody's
// workspace while they are signed in.
func TestConsentRequiresTheFormCSRFToken(t *testing.T) {
h := newHarness(t)
clientID := h.register()
params := authorizeParamsFor(clientID, verifier43)
_, valid := h.consent(params)
for name, token := range map[string]string{
"absent": "",
"garbage": "not-the-token",
"flipped": strings.Repeat("0", len(valid)),
} {
t.Run(name, func(t *testing.T) {
rec := h.decide(params, "approve", token)
if rec.Code == http.StatusFound {
t.Fatalf("approval succeeded without a valid CSRF token: %s",
rec.Header().Get("Location"))
}
if rec.Code != http.StatusForbidden {
t.Errorf("status = %d, want 403", rec.Code)
}
})
}
// The real token still works, or the test above would pass vacuously.
if rec := h.decide(params, "approve", valid); rec.Code != http.StatusFound {
t.Errorf("the valid CSRF token was rejected: %d", rec.Code)
}
}
/* ── Validation still applies on the POST ───────────────────────────────── */
// The POST must be validated as strictly as the GET. Trusting the form's
// hidden fields would let a tampered POST change the redirect, the resource or
// the PKCE challenge after the person read the page.
func TestConsentPostRevalidatesEveryParameter(t *testing.T) {
h := newHarness(t)
clientID := h.register()
params := authorizeParamsFor(clientID, verifier43)
_, csrf := h.consent(params)
tamper := func(k, v string) map[string]string {
out := map[string]string{}
for key, val := range params {
out[key] = val
}
out[k] = v
return out
}
t.Run("redirect swapped", func(t *testing.T) {
rec := h.decide(tamper("redirect_uri", "https://attacker.example/steal"), "approve", csrf)
if rec.Code == http.StatusFound {
t.Fatalf("a tampered redirect_uri was honoured: %s", rec.Header().Get("Location"))
}
if rec.Code != http.StatusBadRequest {
t.Errorf("status = %d, want 400", rec.Code)
}
})
t.Run("resource swapped", func(t *testing.T) {
rec := h.decide(tamper("resource", "https://elsewhere.test/mcp"), "approve", csrf)
loc, _ := url.Parse(rec.Header().Get("Location"))
if loc.Query().Get("code") != "" {
t.Error("a tampered resource still produced a code")
}
})
t.Run("pkce removed", func(t *testing.T) {
rec := h.decide(tamper("code_challenge", ""), "approve", csrf)
loc, _ := url.Parse(rec.Header().Get("Location"))
if loc.Query().Get("code") != "" {
t.Error("PKCE was dropped at the consent POST")
}
})
t.Run("scope escalated", func(t *testing.T) {
rec := h.decide(tamper("scope", ScopeWrite), "approve", csrf)
loc, _ := url.Parse(rec.Header().Get("Location"))
if loc.Query().Get("code") != "" {
t.Error("krow.write was granted through the consent POST")
}
if got := loc.Query().Get("error"); got != errInvalidScope {
t.Errorf("error = %q, want %q", got, errInvalidScope)
}
})
}
// An anonymous visitor must be sent to the existing login, not shown a consent
// screen for nobody.
func TestConsentRequiresAuthentication(t *testing.T) {
h := newHarness(t)
clientID := h.register()
h.session.signedIn = false
rec := h.authorize(authorizeParamsFor(clientID, verifier43))
if rec.Code != http.StatusFound {
t.Fatalf("status = %d, want a redirect to login", rec.Code)
}
if !strings.HasPrefix(rec.Header().Get("Location"), "/login?") {
t.Errorf("Location = %q, want the existing login", rec.Header().Get("Location"))
}
}
// The consent page must never be cached: it names a client and an organisation,
// and the next person on a shared machine must not see it.
func TestConsentPageIsNotCacheable(t *testing.T) {
h := newHarness(t)
clientID := h.register()
rec, _ := h.consent(authorizeParamsFor(clientID, verifier43))
if got := rec.Header().Get("Cache-Control"); !strings.Contains(got, "no-store") {
t.Errorf("Cache-Control = %q, want no-store", got)
}
if got := rec.Header().Get("X-Frame-Options"); got != "DENY" {
t.Errorf("X-Frame-Options = %q, want DENY — a consent screen must not be framed", got)
}
if !strings.Contains(rec.Header().Get("Content-Security-Policy"), "frame-ancestors 'none'") {
t.Error("the CSP does not forbid framing")
}
}
// Thin wrappers so the test reads as prose rather than as error handling.
func jsonMarshal(v any) (string, error) {
b, err := json.Marshal(v)
return string(b), err
}
func jsonUnmarshal(b []byte, v any) error { return json.Unmarshal(b, v) }
/* ── The consent page's CSP ─────────────────────────────────────────────── */
// The regression test for the bug a real browser found and httptest could not.
//
// `form-action 'self'` blocked the consent form before it could submit, because
// a consent form's successful submission ends at the client's registered
// callback — always cross-origin. These tests assert the policy admits exactly
// that callback and nothing else.
func TestConsentCSPAllowsTheRegisteredRedirectOrigin(t *testing.T) {
h := newHarness(t)
body, _ := jsonMarshal(registrationRequest{
ClientName: "Claude", RedirectURIs: []string{"https://claude.ai/api/mcp/auth_callback"},
})
rec := httptest.NewRecorder()
h.server.RegisterHandler().ServeHTTP(rec,
httptest.NewRequest(http.MethodPost, "/oauth/register", strings.NewReader(body)))
var reg registrationResponse
_ = jsonUnmarshal(rec.Body.Bytes(), &reg)
page, _ := h.consent(authorizeParamsForClient(reg.ClientID, verifier43, "https://claude.ai/api/mcp/auth_callback"))
csp := page.Header().Get("Content-Security-Policy")
// The registered ORIGIN — not the full URI. A CSP source is an origin; a
// path would be matched as a prefix and is not what a navigation is
// checked against.
if !strings.Contains(csp, "https://claude.ai") {
t.Errorf("CSP does not permit the registered redirect origin:\n %s", csp)
}
if strings.Contains(csp, "/api/mcp/auth_callback") {
t.Errorf("CSP carries a path rather than an origin:\n %s", csp)
}
// Everything that must survive the change.
for _, required := range []string{
"default-src 'none'",
"style-src 'unsafe-inline'",
"frame-ancestors 'none'",
"form-action 'self'",
} {
if !strings.Contains(csp, required) {
t.Errorf("CSP lost %q:\n %s", required, csp)
}
}
// And everything that must never appear.
for _, forbidden := range []string{"form-action *", "'unsafe-eval'", "'unsafe-inline' 'unsafe", "*;", " *"} {
if strings.Contains(csp, forbidden) {
t.Errorf("CSP contains a broad source %q:\n %s", forbidden, csp)
}
}
}
// One client's callback must never appear in another client's policy.
func TestConsentCSPDoesNotLeakBetweenClients(t *testing.T) {
h := newHarness(t)
register := func(name, redirect string) string {
t.Helper()
body, _ := jsonMarshal(registrationRequest{ClientName: name, RedirectURIs: []string{redirect}})
rec := httptest.NewRecorder()
h.server.RegisterHandler().ServeHTTP(rec,
httptest.NewRequest(http.MethodPost, "/oauth/register", strings.NewReader(body)))
var reg registrationResponse
_ = jsonUnmarshal(rec.Body.Bytes(), &reg)
return reg.ClientID
}
first := register("First", "https://first.example.test/cb")
second := register("Second", "https://second.example.test/cb")
firstPage, _ := h.consent(authorizeParamsForClient(first, verifier43, "https://first.example.test/cb"))
firstCSP := firstPage.Header().Get("Content-Security-Policy")
if !strings.Contains(firstCSP, "https://first.example.test") {
t.Errorf("the first client's own origin is missing:\n %s", firstCSP)
}
if strings.Contains(firstCSP, "second.example.test") {
t.Errorf("another client's redirect origin leaked into this policy:\n %s", firstCSP)
}
secondPage, _ := h.consent(authorizeParamsForClient(second, verifier43, "https://second.example.test/cb"))
secondCSP := secondPage.Header().Get("Content-Security-Policy")
if strings.Contains(secondCSP, "first.example.test") {
t.Errorf("the first client's origin leaked into the second's policy:\n %s", secondCSP)
}
}
// An unrelated origin must never be permitted.
func TestConsentCSPExcludesUnrelatedOrigins(t *testing.T) {
h := newHarness(t)
clientID := h.register() // registered for testRedirect only
page, _ := h.consent(authorizeParamsFor(clientID, verifier43))
csp := page.Header().Get("Content-Security-Policy")
for _, unrelated := range []string{"https://evil.test", "https://attacker.example", "https://google.com"} {
if strings.Contains(csp, unrelated) {
t.Errorf("CSP permits an unrelated origin %q:\n %s", unrelated, csp)
}
}
}
/* ── redirectOrigins, directly ──────────────────────────────────────────── */
// The helper carries the whole safety argument, so it is tested on its own
// rather than only through a rendered page.
func TestRedirectOrigins(t *testing.T) {
for name, tc := range map[string]struct {
in []string
want []string
}{
"https with path": {
[]string{"https://claude.ai/api/mcp/auth_callback"},
[]string{"https://claude.ai"},
},
"port preserved": {
[]string{"https://app.example.test:8443/cb"},
[]string{"https://app.example.test:8443"},
},
"loopback http is kept, per registration rules": {
[]string{"http://127.0.0.1:33418/callback"},
[]string{"http://127.0.0.1:33418"},
},
"localhost loopback": {
[]string{"http://localhost:3000/cb"},
[]string{"http://localhost:3000"},
},
"multiple registered URIs": {
[]string{"https://claude.ai/cb", "http://127.0.0.1:33418/callback"},
[]string{"https://claude.ai", "http://127.0.0.1:33418"},
},
"duplicates collapse to one source": {
[]string{"https://claude.ai/one", "https://claude.ai/two", "https://claude.ai/three"},
[]string{"https://claude.ai"},
},
"order is the order registered": {
[]string{"https://b.test/cb", "https://a.test/cb"},
[]string{"https://b.test", "https://a.test"},
},
// Everything below must be DROPPED, never broadened.
"relative uri": {[]string{"/callback"}, nil},
"no host": {[]string{"https://"}, nil},
"custom scheme": {[]string{"myapp://callback"}, nil},
"javascript scheme": {[]string{"javascript:alert(1)"}, nil},
"data scheme": {[]string{"data:text/html,x"}, nil},
"wildcard host": {[]string{"https://*.evil.test/cb"}, nil},
"semicolon injection": {[]string{"https://evil.test;form-action *"}, nil},
"space injection": {[]string{"https://evil.test /cb"}, nil},
"empty": {[]string{""}, nil},
"whitespace only": {[]string{" "}, nil},
} {
t.Run(name, func(t *testing.T) {
got := redirectOrigins(tc.in)
if len(got) != len(tc.want) {
t.Fatalf("redirectOrigins(%q) = %q, want %q", tc.in, got, tc.want)
}
for i := range tc.want {
if got[i] != tc.want[i] {
t.Errorf("origin[%d] = %q, want %q", i, got[i], tc.want[i])
}
}
})
}
}
// A dropped URI must never widen the policy — the page still renders, and the
// form-action list is simply shorter.
func TestABadRedirectURICannotWidenTheCSP(t *testing.T) {
csp := consentCSP([]string{"https://evil.test;form-action *", "myapp://cb", "https://*.evil.test"})
if strings.Contains(csp, "*") {
t.Errorf("a malformed redirect URI introduced a wildcard:\n %s", csp)
}
if strings.Count(csp, ";") != 3 {
t.Errorf("the header has %d separators, want 3 — a URI broke out of its directive:\n %s",
strings.Count(csp, ";"), csp)
}
if !strings.Contains(csp, "form-action 'self'") {
t.Errorf("form-action lost 'self':\n %s", csp)
}
// With every URI dropped, the policy is exactly the strict one — which
// blocks the flow visibly rather than permitting more.
if strings.Contains(csp, "evil.test") {
t.Errorf("a dropped URI still reached the policy:\n %s", csp)
}
}
// authorizeParamsForClient is authorizeParamsFor with an explicit redirect, so
// a test can drive a client registered for something other than testRedirect.
func authorizeParamsForClient(clientID, verifier, redirect string) map[string]string {
p := authorizeParamsFor(clientID, verifier)
p["redirect_uri"] = redirect
return p
}

View File

@@ -0,0 +1,337 @@
package oauth_test
import (
"context"
"encoding/json"
"io"
"log/slog"
"net/http"
"net/http/httptest"
"net/url"
"strings"
"testing"
"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/runtime"
"github.com/krow/krow-backend/go-api/internal/testutil"
)
// The seam, joined.
//
// This is the test that matters most in Phase 3, and it is in an EXTERNAL test
// package (oauth_test) on purpose: it may use only the exported surface, which
// is exactly what the HTTP layer will use when it wires these two packages
// together in a later phase. If this compiles and passes, the wiring is a
// constructor call and nothing else.
//
// What it proves end to end, with a real database and no fakes anywhere:
//
// OAuth authorization code flow
// → access token
// → mcpserver.TokenAuthenticator (the PRODUCTION implementation)
// → authctx.Identity built from the live user row
// → tools.Registry.Dispatch
// → the existing policy table and org pre-filter
// → real rows from Postgres
const (
itIssuer = "https://api.example.test"
itResource = "https://api.example.test/mcp"
itRedirect = "https://claude.example.test/callback"
itVerifier = "abcdefghijklmnopqrstuvwxyz0123456789ABCDEFG"
)
// itSession is the SessionResolver a signed-in browser satisfies with a cookie.
// The real implementation lives in the HTTP layer; this stands in for it so the
// flow can be driven without a browser.
type itSession struct{ id authctx.Identity }
func (s itSession) CurrentUser(*http.Request) (authctx.Identity, bool) { return s.id, true }
func sessionFor(userID, orgID string) itSession {
return itSession{id: authctx.Identity{
UserID: userID, OrgID: orgID, Role: "admin",
Email: "a@example.test", Status: "active", AccountType: "employer",
}}
}
// itRegister performs dynamic client registration over the real handler.
func itRegister(t *testing.T, as *oauth.Server) string {
t.Helper()
body := `{"client_name":"Integration Client","redirect_uris":["` + itRedirect + `"]}`
req := httptest.NewRequest(http.MethodPost, "/oauth/register", strings.NewReader(body))
rec := httptest.NewRecorder()
as.RegisterHandler().ServeHTTP(rec, req)
if rec.Code != http.StatusCreated {
t.Fatalf("register: %d %s", rec.Code, rec.Body.String())
}
var out struct {
ClientID string `json:"client_id"`
}
if err := json.Unmarshal(rec.Body.Bytes(), &out); err != nil {
t.Fatalf("register response: %v", err)
}
return out.ClientID
}
// itAuthorizeParams is a well-formed authorization request.
func itAuthorizeParams(clientID string) url.Values {
return url.Values{
"client_id": {clientID}, "redirect_uri": {itRedirect}, "response_type": {"code"},
"state": {"st8"}, "code_challenge": {oauth.ChallengeFor(itVerifier)},
"code_challenge_method": {"S256"}, "resource": {itResource}, "scope": {oauth.ScopeRead},
}
}
// itCSRF pulls the consent form's token out of the rendered page.
func itCSRF(t *testing.T, body string) string {
t.Helper()
const marker = `name="csrf" value="`
i := strings.Index(body, marker)
if i < 0 {
t.Fatalf("no csrf field in the consent page")
}
rest := body[i+len(marker):]
return rest[:strings.Index(rest, `"`)]
}
// itDecide posts an approve/deny decision.
func itDecide(t *testing.T, as *oauth.Server, params url.Values, decision, csrf string) *httptest.ResponseRecorder {
t.Helper()
form := url.Values{}
for k, v := range params {
form[k] = v
}
form.Set("decision", decision)
form.Set("csrf", csrf)
req := httptest.NewRequest(http.MethodPost, "/oauth/authorize", strings.NewReader(form.Encode()))
req.Header.Set("Content-Type", "application/x-www-form-urlencoded")
rec := httptest.NewRecorder()
as.AuthorizeHandler().ServeHTTP(rec, req)
return rec
}
// itAuthorize drives the authorization endpoint through CONSENT and returns
// the code.
func itAuthorize(t *testing.T, as *oauth.Server, clientID string) string {
t.Helper()
q := itAuthorizeParams(clientID)
// The consent page first — a GET no longer issues a code.
page := httptest.NewRecorder()
as.AuthorizeHandler().ServeHTTP(page,
httptest.NewRequest(http.MethodGet, "/oauth/authorize?"+q.Encode(), nil))
if page.Code != http.StatusOK {
t.Fatalf("consent page: %d %s", page.Code, page.Body.String())
}
rec := itDecide(t, as, q, "approve", itCSRF(t, page.Body.String()))
if rec.Code != http.StatusFound {
t.Fatalf("authorize: %d %s", rec.Code, rec.Body.String())
}
loc, err := url.Parse(rec.Header().Get("Location"))
if err != nil {
t.Fatalf("Location: %v", err)
}
code := loc.Query().Get("code")
if code == "" {
t.Fatalf("no code: %s", loc)
}
return code
}
// itExchange redeems the code for an access token.
func itExchange(t *testing.T, as *oauth.Server, clientID, code string) string {
t.Helper()
form := url.Values{
"grant_type": {"authorization_code"}, "code": {code}, "client_id": {clientID},
"redirect_uri": {itRedirect}, "code_verifier": {itVerifier},
}
req := httptest.NewRequest(http.MethodPost, "/oauth/token", strings.NewReader(form.Encode()))
req.Header.Set("Content-Type", "application/x-www-form-urlencoded")
rec := httptest.NewRecorder()
as.TokenHandler().ServeHTTP(rec, req)
if rec.Code != http.StatusOK {
t.Fatalf("token: %d %s", rec.Code, rec.Body.String())
}
var out struct {
AccessToken string `json:"access_token"`
}
if err := json.Unmarshal(rec.Body.Bytes(), &out); err != nil {
t.Fatalf("token response: %v", err)
}
return out.AccessToken
}
func TestOAuthTokenReachesMCPToolsAndRealAuthorization(t *testing.T) {
db := testutil.New(t)
ctx := context.Background()
log := slog.New(slog.NewTextHandler(io.Discard, nil))
// Two tenants with different volumes, so a leak is visible as a number.
orgA := mustOrg(t, db, "it-org-a")
orgB := mustOrg(t, db, "it-org-b")
userA := mustUser(t, db, orgA, "a@example.test", "admin")
seedActivity(t, db, orgA, 7, "a@example.test")
seedActivity(t, db, orgB, 55, "b@example.test")
store := oauth.NewStore(db.Pool)
// ── Register, authorize, exchange: the real flow, over the real handlers.
as := oauth.NewServer(
oauth.Config{Issuer: itIssuer, Resource: itResource},
store,
sessionFor(userA, orgA),
"/login", log,
)
clientID := itRegister(t, as)
code := itAuthorize(t, as, clientID)
accessToken := itExchange(t, as, clientID, code)
// ── The production authenticator, plugged into the Phase 2 seam.
authenticator := oauth.NewAuthenticator(store, auth.NewPGUserStore(db.Pool), itResource, log)
mcp := mcpserver.New(runtime.DefaultTools(db.Pool, nil), authenticator, log)
// ── A real MCP tool call, carrying a real OAuth token.
rec := itCall(t, mcp, "Bearer "+accessToken,
`{"jsonrpc":"2.0","id":1,"method":"tools/call","params":
{"name":"activity_breakdown","arguments":{}}}`)
if rec.Code != http.StatusOK {
t.Fatalf("MCP call with an OAuth token: %d %s", rec.Code, rec.Body.String())
}
total := itTotalEvents(t, rec)
if total != 7 {
t.Errorf("totalEvents = %d, want 7 (org A only). Org B has 55; a wrong "+
"number here means the OAuth identity did not scope the query", total)
}
// ── Without the token, the same call must be refused.
if rec := itCall(t, mcp, "",
`{"jsonrpc":"2.0","id":1,"method":"tools/list"}`); rec.Code != http.StatusUnauthorized {
t.Errorf("unauthenticated MCP call = %d, want 401", rec.Code)
}
// ── Revoking disconnects: the same token must stop working immediately,
// not at expiry.
if err := store.RevokeToken(ctx, accessToken, "test_disconnect"); err != nil {
t.Fatalf("revoke: %v", err)
}
if rec := itCall(t, mcp, "Bearer "+accessToken,
`{"jsonrpc":"2.0","id":1,"method":"tools/list"}`); rec.Code != http.StatusUnauthorized {
t.Errorf("a revoked token still reached MCP: %d", rec.Code)
}
}
// A token for another resource must not open the MCP surface, even though this
// same server minted it. The confused-deputy case, end to end.
func TestTokenForAnotherResourceCannotReachMCP(t *testing.T) {
db := testutil.New(t)
log := slog.New(slog.NewTextHandler(io.Discard, nil))
org := mustOrg(t, db, "it-aud-org")
user := mustUser(t, db, org, "aud@example.test", "admin")
store := oauth.NewStore(db.Pool)
as := oauth.NewServer(
oauth.Config{Issuer: itIssuer, Resource: itResource},
store, sessionFor(user, org), "/login", log)
clientID := itRegister(t, as)
pair, err := store.IssuePair(context.Background(), oauth.Token{
ClientID: clientID, UserID: user, OrgID: org,
Scopes: []string{oauth.ScopeRead}, Audience: "https://a-different-service.test/mcp",
}, "")
if err != nil {
t.Fatalf("issue: %v", err)
}
mcp := mcpserver.New(
runtime.DefaultTools(db.Pool, nil),
oauth.NewAuthenticator(store, auth.NewPGUserStore(db.Pool), itResource, log),
log)
if rec := itCall(t, mcp, "Bearer "+pair.AccessToken,
`{"jsonrpc":"2.0","id":1,"method":"tools/list"}`); rec.Code != http.StatusUnauthorized {
t.Errorf("a token for another resource reached MCP: %d", rec.Code)
}
}
/* ── Helpers ────────────────────────────────────────────────────────────── */
func itCall(t *testing.T, s *mcpserver.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
}
func itTotalEvents(t *testing.T, rec *httptest.ResponseRecorder) 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(rec.Body.Bytes(), &envelope); err != nil {
t.Fatalf("response: %v", err)
}
if envelope.Result.IsError || len(envelope.Result.Content) == 0 {
t.Fatalf("tool call failed: %s", rec.Body.String())
}
var payload struct {
Data struct {
TotalEvents int `json:"totalEvents"`
} `json:"data"`
}
if err := json.Unmarshal([]byte(envelope.Result.Content[0].Text), &payload); err != nil {
t.Fatalf("tool payload: %v", err)
}
return payload.Data.TotalEvents
}
func mustOrg(t *testing.T, db *testutil.Harness, slug string) string {
t.Helper()
var id string
if err := db.Pool.QueryRow(context.Background(),
`INSERT INTO organizations (name, slug) VALUES ($1, $2) RETURNING id::text`,
slug, slug).Scan(&id); err != nil {
t.Fatalf("org %s: %v", slug, err)
}
return id
}
func mustUser(t *testing.T, db *testutil.Harness, orgID, email, role string) string {
t.Helper()
var id string
if err := db.Pool.QueryRow(context.Background(),
`INSERT INTO users (org_id, email, full_name, role, account_type, status)
VALUES ($1::uuid, $2, 'Test', $3, 'employer', 'active') RETURNING id::text`,
orgID, email, role).Scan(&id); err != nil {
t.Fatalf("user %s: %v", email, err)
}
return id
}
func seedActivity(t *testing.T, db *testutil.Harness, orgID string, n int, email string) {
t.Helper()
for i := 0; i < n; i++ {
if _, err := db.Pool.Exec(context.Background(),
`INSERT INTO user_activity (org_id, event_type, user_email, user_name)
VALUES ($1::uuid, 'login', $2, 'Someone')`, orgID, email); err != nil {
t.Fatalf("seed: %v", err)
}
}
}

View File

@@ -0,0 +1,186 @@
package oauth
import (
"net/http"
"strings"
)
// Discovery: the two documents an MCP client reads before it can authenticate.
//
// The MCP authorization flow starts with the client calling the MCP endpoint
// with no token, getting a 401, and following its way to an authorization
// server. Two RFCs define the path:
//
// RFC 9728 Protected Resource Metadata — served BY THE RESOURCE (the MCP
// server). Answers "which authorization server issues tokens for
// you". The 401's WWW-Authenticate header points here.
// RFC 8414 Authorization Server Metadata — served by the AS. Answers "where
// are your authorize, token and registration endpoints, and what do
// you support".
//
// Both are unauthenticated by necessity: a client that cannot authenticate yet
// has to be able to read them. Neither contains a secret — they are a map of
// public endpoints, which is exactly what discovery means.
//
// NO URL IS GUESSED OR HARDCODED. Every value comes from configuration, so a
// deployment on a different host is a config change and not a code change, and
// so this file contains no production domain.
// Scopes this server issues.
//
// ScopeWrite is DECLARED and never granted. Naming it here means the constant
// exists for a future phase to use deliberately, rather than being invented at
// the point somebody is trying to make a write work. It appears in no
// scopes_supported list and no issued token.
const (
ScopeRead = "krow.read"
ScopeWrite = "krow.write" // reserved; not issued, not advertised
)
// Config is the deployment's OAuth identity.
//
// Issuer and Resource are separate values that will often look similar, and
// conflating them is a real mistake: the ISSUER identifies the authorization
// server, the RESOURCE identifies the thing a token is good for. A token's
// audience is checked against Resource, and its origin against Issuer.
type Config struct {
// Issuer is the authorization server's identity, 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 is what a client puts in its
// `resource` parameter and what an issued token's audience is set to.
Resource string
// The paths, relative to Issuer. Defaults are applied by Normalise.
AuthorizePath string
TokenPath string
RegistrationPath string
RevocationPath string
}
// Normalise fills defaults and trims trailing slashes.
//
// The canonical form of a resource URI has no trailing slash — RFC 8707 says
// implementations SHOULD use that form — and a mismatch here is a token that
// validates everywhere except the one place it was minted for.
func (c Config) Normalise() Config {
c.Issuer = strings.TrimRight(strings.TrimSpace(c.Issuer), "/")
c.Resource = strings.TrimRight(strings.TrimSpace(c.Resource), "/")
if c.AuthorizePath == "" {
c.AuthorizePath = "/oauth/authorize"
}
if c.TokenPath == "" {
c.TokenPath = "/oauth/token"
}
if c.RegistrationPath == "" {
c.RegistrationPath = "/oauth/register"
}
if c.RevocationPath == "" {
c.RevocationPath = "/oauth/revoke"
}
return c
}
// Valid reports whether this configuration can serve discovery at all.
func (c Config) Valid() bool {
return c.Issuer != "" && c.Resource != ""
}
func (c Config) authorizeURL() string { return c.Issuer + c.AuthorizePath }
func (c Config) tokenURL() string { return c.Issuer + c.TokenPath }
func (c Config) registrationURL() string { return c.Issuer + c.RegistrationPath }
func (c Config) revocationURL() string { return c.Issuer + c.RevocationPath }
/* ── RFC 9728: Protected Resource Metadata ──────────────────────────────── */
type protectedResourceMetadata struct {
Resource string `json:"resource"`
AuthorizationServers []string `json:"authorization_servers"`
ScopesSupported []string `json:"scopes_supported"`
BearerMethodsSupported []string `json:"bearer_methods_supported"`
}
// ProtectedResourceHandler serves /.well-known/oauth-protected-resource.
func (c Config) ProtectedResourceHandler() http.Handler {
cfg := c.Normalise()
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodGet {
w.Header().Set("Allow", http.MethodGet)
http.Error(w, "method not allowed", http.StatusMethodNotAllowed)
return
}
writeMetadata(w, protectedResourceMetadata{
Resource: cfg.Resource,
AuthorizationServers: []string{cfg.Issuer},
ScopesSupported: []string{ScopeRead},
// header only. RFC 6750 also defines a form-encoded body parameter
// and a query parameter; the MCP spec forbids the query form and
// this server accepts neither.
BearerMethodsSupported: []string{"header"},
})
})
}
/* ── RFC 8414: Authorization Server Metadata ────────────────────────────── */
type authorizationServerMetadata struct {
Issuer string `json:"issuer"`
AuthorizationEndpoint string `json:"authorization_endpoint"`
TokenEndpoint string `json:"token_endpoint"`
RegistrationEndpoint string `json:"registration_endpoint"`
RevocationEndpoint string `json:"revocation_endpoint"`
ScopesSupported []string `json:"scopes_supported"`
ResponseTypesSupported []string `json:"response_types_supported"`
GrantTypesSupported []string `json:"grant_types_supported"`
CodeChallengeMethodsSupported []string `json:"code_challenge_methods_supported"`
TokenEndpointAuthMethodsSupported []string `json:"token_endpoint_auth_methods_supported"`
ResourceIndicatorsSupported bool `json:"resource_indicators_supported"`
}
// AuthorizationServerHandler serves /.well-known/oauth-authorization-server.
//
// Every list below is a promise, so each one names only what is implemented:
//
// - response_types: `code`. No `token`, because implicit is gone from OAuth
// 2.1 and advertising it would invite a flow this server refuses.
// - grant_types: authorization_code and refresh_token. No password, no
// client_credentials — neither has a caller here, and both would be a way
// to get a token without a person approving anything.
// - code_challenge_methods: S256 only. Listing `plain` would tell a client it
// may use the method this server rejects.
// - token_endpoint_auth_methods: `none`, which is the correct declaration
// for public clients. They authenticate with PKCE, not a secret.
func (c Config) AuthorizationServerHandler() http.Handler {
cfg := c.Normalise()
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodGet {
w.Header().Set("Allow", http.MethodGet)
http.Error(w, "method not allowed", http.StatusMethodNotAllowed)
return
}
writeMetadata(w, authorizationServerMetadata{
Issuer: cfg.Issuer,
AuthorizationEndpoint: cfg.authorizeURL(),
TokenEndpoint: cfg.tokenURL(),
RegistrationEndpoint: cfg.registrationURL(),
RevocationEndpoint: cfg.revocationURL(),
ScopesSupported: []string{ScopeRead},
ResponseTypesSupported: []string{"code"},
GrantTypesSupported: []string{"authorization_code", "refresh_token"},
CodeChallengeMethodsSupported: []string{MethodS256},
TokenEndpointAuthMethodsSupported: []string{"none"},
ResourceIndicatorsSupported: true,
})
})
}
func writeMetadata(w http.ResponseWriter, payload any) {
w.Header().Set("Content-Type", "application/json; charset=utf-8")
// Discovery documents change only with a deployment, and a client that
// re-reads them on every connection costs nothing to serve. Five minutes
// keeps a stale document from outliving a config change by long.
w.Header().Set("Cache-Control", "public, max-age=300")
writeJSONBody(w, http.StatusOK, payload)
}

View File

@@ -0,0 +1,952 @@
package oauth
import (
"context"
"encoding/json"
"errors"
"io"
"log/slog"
"net/http"
"net/http/httptest"
"net/url"
"strings"
"testing"
"time"
"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/testutil"
)
/* ── Fixtures ───────────────────────────────────────────────────────────── */
const (
testIssuer = "https://api.example.test"
testResource = "https://api.example.test/mcp"
testRedirect = "https://claude.example.test/callback"
)
func testConfig() Config {
return Config{Issuer: testIssuer, Resource: testResource}.Normalise()
}
func discard() *slog.Logger { return slog.New(slog.NewTextHandler(io.Discard, nil)) }
// fakeSession is the SessionResolver a browser would satisfy with a cookie.
type fakeSession struct {
identity authctx.Identity
signedIn bool
}
func (f *fakeSession) CurrentUser(*http.Request) (authctx.Identity, bool) {
return f.identity, f.signedIn
}
// harness wires a real database to a real authorization server.
type harness struct {
t *testing.T
h *testutil.Harness
store *Store
server *Server
session *fakeSession
userID string
orgID string
clock time.Time
}
func newHarness(t *testing.T) *harness {
t.Helper()
db := testutil.New(t)
ctx := context.Background()
var orgID string
if err := db.Pool.QueryRow(ctx,
`INSERT INTO organizations (name, slug) VALUES ('OAuth Test', 'oauth-test')
RETURNING id::text`).Scan(&orgID); err != nil {
t.Fatalf("create org: %v", err)
}
var userID string
if err := db.Pool.QueryRow(ctx,
`INSERT INTO users (org_id, email, full_name, role, account_type, status)
VALUES ($1::uuid, 'oauth-user@example.test', 'OAuth User', 'admin', 'employer', 'active')
RETURNING id::text`, orgID).Scan(&userID); err != nil {
t.Fatalf("create user: %v", err)
}
clock := time.Now()
store := NewStore(db.Pool).WithClock(func() time.Time { return clock })
session := &fakeSession{
identity: authctx.Identity{UserID: userID, OrgID: orgID, Role: "admin",
Email: "oauth-user@example.test", Status: "active"},
signedIn: true,
}
hs := &harness{
t: t, h: db, store: store, session: session,
userID: userID, orgID: orgID, clock: clock,
}
hs.server = NewServer(testConfig(), store, session, "/login", discard())
return hs
}
// advance moves the store's clock, so expiry is tested without sleeping.
func (h *harness) advance(d time.Duration) {
h.clock = h.clock.Add(d)
h.store.WithClock(func() time.Time { return h.clock })
}
// register performs dynamic client registration and returns the client_id.
func (h *harness) register(redirectURIs ...string) string {
h.t.Helper()
if len(redirectURIs) == 0 {
redirectURIs = []string{testRedirect}
}
body, _ := json.Marshal(registrationRequest{
ClientName: "Test Client", RedirectURIs: redirectURIs,
})
rec := httptest.NewRecorder()
h.server.RegisterHandler().ServeHTTP(rec,
httptest.NewRequest(http.MethodPost, "/oauth/register", strings.NewReader(string(body))))
if rec.Code != http.StatusCreated {
h.t.Fatalf("registration failed: %d %s", rec.Code, rec.Body.String())
}
var out registrationResponse
if err := json.Unmarshal(rec.Body.Bytes(), &out); err != nil {
h.t.Fatalf("registration response: %v", err)
}
return out.ClientID
}
// authorize drives the authorization endpoint and returns the response recorder.
func (h *harness) authorize(params map[string]string) *httptest.ResponseRecorder {
h.t.Helper()
q := url.Values{}
for k, v := range params {
if v != "" {
q.Set(k, v)
}
}
rec := httptest.NewRecorder()
h.server.AuthorizeHandler().ServeHTTP(rec,
httptest.NewRequest(http.MethodGet, "/oauth/authorize?"+q.Encode(), nil))
return rec
}
// authorizeParamsFor is a well-formed authorization request.
func authorizeParamsFor(clientID, verifier string) map[string]string {
return map[string]string{
"client_id": clientID, "redirect_uri": testRedirect, "response_type": "code",
"state": "xyz", "code_challenge": ChallengeFor(verifier),
"code_challenge_method": "S256", "resource": testResource, "scope": ScopeRead,
}
}
// csrfFromConsentPage pulls the token out of the rendered form.
//
// PHASE 4: a GET now RENDERS a consent page rather than issuing a code. This
// helper and decide() below are how a test plays the part of the person
// clicking a button. Not one assertion in this file changed — the flow gained
// a step, and the tests walk through it.
func csrfFromConsentPage(t *testing.T, body string) string {
t.Helper()
const marker = `name="csrf" value="`
i := strings.Index(body, marker)
if i < 0 {
t.Fatalf("no csrf field in the consent page:\n%s", body)
}
rest := body[i+len(marker):]
j := strings.Index(rest, `"`)
if j < 0 {
t.Fatal("malformed csrf field")
}
return rest[:j]
}
// decide posts an approve/deny decision to the authorization endpoint.
func (h *harness) decide(params map[string]string, decision, csrf string) *httptest.ResponseRecorder {
h.t.Helper()
form := url.Values{}
for k, v := range params {
if v != "" {
form.Set(k, v)
}
}
form.Set("decision", decision)
form.Set("csrf", csrf)
req := httptest.NewRequest(http.MethodPost, "/oauth/authorize", strings.NewReader(form.Encode()))
req.Header.Set("Content-Type", "application/x-www-form-urlencoded")
rec := httptest.NewRecorder()
h.server.AuthorizeHandler().ServeHTTP(rec, req)
return rec
}
// consent renders the form and returns the page plus its CSRF token.
func (h *harness) consent(params map[string]string) (*httptest.ResponseRecorder, string) {
h.t.Helper()
rec := h.authorize(params)
if rec.Code != http.StatusOK {
h.t.Fatalf("consent page: %d %s", rec.Code, rec.Body.String())
}
return rec, csrfFromConsentPage(h.t, rec.Body.String())
}
// authorizeOK runs a well-formed authorization, APPROVES it, and returns the
// code.
func (h *harness) authorizeOK(clientID, verifier string) string {
h.t.Helper()
params := authorizeParamsFor(clientID, verifier)
_, csrf := h.consent(params)
rec := h.decide(params, "approve", csrf)
if rec.Code != http.StatusFound {
h.t.Fatalf("authorize: %d %s", rec.Code, rec.Body.String())
}
loc, err := url.Parse(rec.Header().Get("Location"))
if err != nil {
h.t.Fatalf("bad Location: %v", err)
}
if e := loc.Query().Get("error"); e != "" {
h.t.Fatalf("authorize returned error=%s (%s)", e, loc.Query().Get("error_description"))
}
code := loc.Query().Get("code")
if code == "" {
h.t.Fatalf("no code in %s", loc)
}
if got := loc.Query().Get("state"); got != "xyz" {
h.t.Errorf("state = %q, want xyz — the client's CSRF defence must be echoed", got)
}
return code
}
// token posts to the token endpoint.
func (h *harness) token(form url.Values) *httptest.ResponseRecorder {
h.t.Helper()
req := httptest.NewRequest(http.MethodPost, "/oauth/token", strings.NewReader(form.Encode()))
req.Header.Set("Content-Type", "application/x-www-form-urlencoded")
rec := httptest.NewRecorder()
h.server.TokenHandler().ServeHTTP(rec, req)
return rec
}
func (h *harness) exchange(clientID, code, verifier string) *httptest.ResponseRecorder {
return h.token(url.Values{
"grant_type": {"authorization_code"}, "code": {code}, "client_id": {clientID},
"redirect_uri": {testRedirect}, "code_verifier": {verifier},
})
}
func decodeTokens(t *testing.T, rec *httptest.ResponseRecorder) tokenResponse {
t.Helper()
if rec.Code != http.StatusOK {
t.Fatalf("token endpoint: %d %s", rec.Code, rec.Body.String())
}
var out tokenResponse
if err := json.Unmarshal(rec.Body.Bytes(), &out); err != nil {
t.Fatalf("token response: %v", err)
}
return out
}
func oauthErrorCode(t *testing.T, rec *httptest.ResponseRecorder) string {
t.Helper()
var out oauthError
_ = json.Unmarshal(rec.Body.Bytes(), &out)
return out.Code
}
// verifier43 is a legal code_verifier.
const verifier43 = "abcdefghijklmnopqrstuvwxyz0123456789ABCDEFG"
/* ── The happy path ─────────────────────────────────────────────────────── */
func TestFullAuthorizationCodeFlow(t *testing.T) {
h := newHarness(t)
clientID := h.register()
code := h.authorizeOK(clientID, verifier43)
tokens := decodeTokens(t, h.exchange(clientID, code, verifier43))
if tokens.TokenType != "Bearer" {
t.Errorf("token_type = %q, want Bearer", tokens.TokenType)
}
if tokens.AccessToken == "" || tokens.RefreshToken == "" {
t.Fatal("a token response must carry both tokens")
}
if tokens.AccessToken == tokens.RefreshToken {
t.Error("access and refresh tokens are identical")
}
// 15 minutes, as committed in the plan.
if tokens.ExpiresIn != int(AccessTokenTTL.Seconds()) {
t.Errorf("expires_in = %d, want %d", tokens.ExpiresIn, int(AccessTokenTTL.Seconds()))
}
if tokens.Scope != ScopeRead {
t.Errorf("scope = %q, want %q", tokens.Scope, ScopeRead)
}
}
// The single most important storage property: a dump of these tables must not
// be replayable. Asserted by searching every text column for the raw values.
func TestPlaintextTokensAreNeverStored(t *testing.T) {
h := newHarness(t)
clientID := h.register()
code := h.authorizeOK(clientID, verifier43)
tokens := decodeTokens(t, h.exchange(clientID, code, verifier43))
for name, secret := range map[string]string{
"authorization code": code,
"access token": tokens.AccessToken,
"refresh token": tokens.RefreshToken,
} {
for _, table := range []string{"oauth_grants", "oauth_tokens"} {
var found int
// Cast the whole row to text and search it. Cruder than naming
// columns and much harder to fool: a future column that stored a
// raw token would be caught without anyone updating this test.
if err := h.h.Pool.QueryRow(context.Background(),
`SELECT count(*) FROM `+table+` t WHERE t::text LIKE '%' || $1 || '%'`,
secret).Scan(&found); err != nil {
t.Fatalf("scan %s: %v", table, err)
}
if found != 0 {
t.Errorf("the %s appears in PLAINTEXT in %s (%d rows)", name, table, found)
}
}
}
}
/* ── Authorization endpoint rejection ───────────────────────────────────── */
func TestAuthorizeRejections(t *testing.T) {
h := newHarness(t)
clientID := h.register()
base := map[string]string{
"client_id": clientID, "redirect_uri": testRedirect, "response_type": "code",
"state": "xyz", "code_challenge": ChallengeFor(verifier43),
"code_challenge_method": "S256", "resource": testResource,
}
with := func(changes map[string]string) map[string]string {
out := map[string]string{}
for k, v := range base {
out[k] = v
}
for k, v := range changes {
out[k] = v
}
return out
}
// These are answered DIRECTLY, never by redirecting — redirecting an error
// to an unvalidated URI is an open redirect.
t.Run("direct errors", func(t *testing.T) {
for name, changes := range map[string]map[string]string{
"unknown client": {"client_id": "00000000-0000-4000-8000-000000000000"},
"missing client": {"client_id": ""},
"missing redirect": {"redirect_uri": ""},
"unregistered redirect": {"redirect_uri": "https://attacker.example/steal"},
"redirect near-miss": {"redirect_uri": testRedirect + "/../evil"},
} {
t.Run(name, func(t *testing.T) {
rec := h.authorize(with(changes))
if rec.Code == http.StatusFound {
t.Fatalf("answered with a REDIRECT to %q — this must be a direct error",
rec.Header().Get("Location"))
}
if rec.Code != http.StatusBadRequest {
t.Errorf("status = %d, want 400", rec.Code)
}
})
}
})
// The redirect target is validated by now, so errors go to it.
t.Run("redirected errors", func(t *testing.T) {
for name, tc := range map[string]struct {
changes map[string]string
want string
}{
"missing state": {map[string]string{"state": ""}, errInvalidRequest},
"missing pkce": {map[string]string{"code_challenge": ""}, errInvalidRequest},
"missing method": {map[string]string{"code_challenge_method": ""}, errInvalidRequest},
"plain pkce": {map[string]string{"code_challenge_method": "plain"}, errInvalidRequest},
"bad challenge": {map[string]string{"code_challenge": "too-short"}, errInvalidRequest},
"implicit flow": {map[string]string{"response_type": "token"}, "unsupported_response_type"},
"missing resource": {map[string]string{"resource": ""}, errInvalidTarget},
"wrong resource": {map[string]string{"resource": "https://elsewhere.test/mcp"}, errInvalidTarget},
"write scope denied": {map[string]string{"scope": ScopeWrite}, errInvalidScope},
} {
t.Run(name, func(t *testing.T) {
rec := h.authorize(with(tc.changes))
if rec.Code != http.StatusFound {
t.Fatalf("status = %d, want a 302 carrying the error", rec.Code)
}
loc, _ := url.Parse(rec.Header().Get("Location"))
if !strings.HasPrefix(loc.String(), testRedirect) {
t.Fatalf("error went to %q, not the registered redirect", loc)
}
if got := loc.Query().Get("error"); got != tc.want {
t.Errorf("error = %q, want %q", got, tc.want)
}
if loc.Query().Get("code") != "" {
t.Error("a failed authorization returned a code")
}
})
}
})
}
// An unauthenticated person is sent to the existing login, not refused and not
// asked for a password by this package.
func TestAuthorizeRedirectsAnonymousToLogin(t *testing.T) {
h := newHarness(t)
clientID := h.register()
h.session.signedIn = false
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,
})
if rec.Code != http.StatusFound {
t.Fatalf("status = %d, want 302 to login", rec.Code)
}
loc := rec.Header().Get("Location")
if !strings.HasPrefix(loc, "/login?returnTo=") {
t.Fatalf("Location = %q, want a redirect to /login carrying returnTo", loc)
}
// The authorization request must survive the round trip, or the person
// signs in and lands nowhere.
if !strings.Contains(loc, url.QueryEscape("client_id="+clientID)) {
t.Error("returnTo does not preserve the authorization request")
}
}
/* ── Registration ───────────────────────────────────────────────────────── */
func TestRegistrationRejectsUnsafeRedirectURIs(t *testing.T) {
h := newHarness(t)
for name, uri := range map[string]string{
"plain http": "http://attacker.example/cb",
"relative": "/callback",
"no host": "https://",
"with fragment": "https://ok.example/cb#frag",
"custom scheme": "myapp://callback",
"javascript": "javascript:alert(1)",
"data uri": "data:text/html,hi",
"missing scheme": "ok.example/cb",
} {
t.Run(name, func(t *testing.T) {
body, _ := json.Marshal(registrationRequest{
ClientName: "x", RedirectURIs: []string{uri},
})
rec := httptest.NewRecorder()
h.server.RegisterHandler().ServeHTTP(rec,
httptest.NewRequest(http.MethodPost, "/oauth/register", strings.NewReader(string(body))))
if rec.Code != http.StatusBadRequest {
t.Errorf("status = %d, want 400 for redirect_uri %q", rec.Code, uri)
}
})
}
}
// http on loopback is the documented exception for native clients (RFC 8252):
// the traffic never leaves the machine.
func TestRegistrationAllowsLoopbackHTTP(t *testing.T) {
for _, uri := range []string{
"http://127.0.0.1:8765/callback",
"http://localhost:3000/cb",
"https://claude.example.test/cb",
} {
if err := validateRedirectURI(uri); err != nil {
t.Errorf("%q was rejected: %v", uri, err)
}
}
}
// A public client must not be issued a secret.
func TestRegistrationIssuesNoClientSecret(t *testing.T) {
h := newHarness(t)
body, _ := json.Marshal(registrationRequest{ClientName: "x", RedirectURIs: []string{testRedirect}})
rec := httptest.NewRecorder()
h.server.RegisterHandler().ServeHTTP(rec,
httptest.NewRequest(http.MethodPost, "/oauth/register", strings.NewReader(string(body))))
if strings.Contains(strings.ToLower(rec.Body.String()), "client_secret") {
t.Errorf("a public client was issued a secret: %s", rec.Body.String())
}
var out registrationResponse
_ = json.Unmarshal(rec.Body.Bytes(), &out)
if out.TokenEndpointAuthMethod != "none" {
t.Errorf("token_endpoint_auth_method = %q, want none", out.TokenEndpointAuthMethod)
}
}
/* ── Token endpoint ─────────────────────────────────────────────────────── */
func TestTokenEndpointRejections(t *testing.T) {
h := newHarness(t)
clientID := h.register()
t.Run("wrong verifier", func(t *testing.T) {
code := h.authorizeOK(clientID, verifier43)
rec := h.exchange(clientID, code, strings.Repeat("z", 43))
if rec.Code != http.StatusBadRequest || oauthErrorCode(t, rec) != errInvalidGrant {
t.Errorf("status=%d error=%q, want 400 invalid_grant", rec.Code, oauthErrorCode(t, rec))
}
})
t.Run("missing verifier", func(t *testing.T) {
code := h.authorizeOK(clientID, verifier43)
rec := h.token(url.Values{
"grant_type": {"authorization_code"}, "code": {code},
"client_id": {clientID}, "redirect_uri": {testRedirect},
})
if rec.Code != http.StatusBadRequest {
t.Errorf("status = %d, want 400 — PKCE is mandatory", rec.Code)
}
})
t.Run("wrong client", func(t *testing.T) {
other := h.register("https://other.example.test/cb")
code := h.authorizeOK(clientID, verifier43)
rec := h.exchange(other, code, verifier43)
if rec.Code != http.StatusBadRequest {
t.Errorf("status = %d, want 400 — a code is bound to its client", rec.Code)
}
})
t.Run("wrong redirect_uri", func(t *testing.T) {
code := h.authorizeOK(clientID, verifier43)
rec := h.token(url.Values{
"grant_type": {"authorization_code"}, "code": {code}, "client_id": {clientID},
"redirect_uri": {"https://attacker.example/steal"}, "code_verifier": {verifier43},
})
if rec.Code != http.StatusBadRequest {
t.Errorf("status = %d, want 400", rec.Code)
}
})
t.Run("unknown code", func(t *testing.T) {
rec := h.exchange(clientID, "not-a-real-code", verifier43)
if rec.Code != http.StatusBadRequest {
t.Errorf("status = %d, want 400", rec.Code)
}
})
}
// A code is single-use. The second attempt must fail even with everything else
// correct — this is replay protection.
func TestAuthorizationCodeIsSingleUse(t *testing.T) {
h := newHarness(t)
clientID := h.register()
code := h.authorizeOK(clientID, verifier43)
if rec := h.exchange(clientID, code, verifier43); rec.Code != http.StatusOK {
t.Fatalf("first exchange failed: %d %s", rec.Code, rec.Body.String())
}
rec := h.exchange(clientID, code, verifier43)
if rec.Code != http.StatusBadRequest {
t.Errorf("status = %d, want 400 — a code must not be redeemable twice", rec.Code)
}
}
// A failed exchange still spends the code, so an attacker cannot probe the
// remaining bindings by retrying with different values.
func TestAFailedExchangeStillConsumesTheCode(t *testing.T) {
h := newHarness(t)
clientID := h.register()
code := h.authorizeOK(clientID, verifier43)
if rec := h.exchange(clientID, code, strings.Repeat("z", 43)); rec.Code != http.StatusBadRequest {
t.Fatalf("expected the wrong verifier to fail, got %d", rec.Code)
}
if rec := h.exchange(clientID, code, verifier43); rec.Code != http.StatusBadRequest {
t.Error("the code was still usable after a failed exchange")
}
}
func TestExpiredAuthorizationCodeIsRejected(t *testing.T) {
h := newHarness(t)
clientID := h.register()
code := h.authorizeOK(clientID, verifier43)
h.advance(GrantTTL + time.Second)
if rec := h.exchange(clientID, code, verifier43); rec.Code != http.StatusBadRequest {
t.Errorf("status = %d, want 400 for an expired code", rec.Code)
}
}
func TestUnsupportedGrantTypesAreRejected(t *testing.T) {
h := newHarness(t)
for _, grant := range []string{"password", "client_credentials", "implicit", "device_code", "nonsense"} {
t.Run(grant, func(t *testing.T) {
rec := h.token(url.Values{
"grant_type": {grant}, "username": {"a"}, "password": {"b"},
})
if rec.Code != http.StatusBadRequest {
t.Fatalf("status = %d, want 400", rec.Code)
}
if got := oauthErrorCode(t, rec); got != errUnsupportedGrantType {
t.Errorf("error = %q, want %q", got, errUnsupportedGrantType)
}
})
}
}
/* ── Refresh rotation and reuse detection ───────────────────────────────── */
func TestRefreshRotatesAndInvalidatesTheOldToken(t *testing.T) {
h := newHarness(t)
clientID := h.register()
code := h.authorizeOK(clientID, verifier43)
first := decodeTokens(t, h.exchange(clientID, code, verifier43))
second := decodeTokens(t, h.token(url.Values{
"grant_type": {"refresh_token"}, "refresh_token": {first.RefreshToken}, "client_id": {clientID},
}))
if second.RefreshToken == first.RefreshToken {
t.Error("the refresh token was not rotated")
}
if second.AccessToken == first.AccessToken {
t.Error("refresh returned the same access token")
}
// The rotated-away token must be dead. Presenting it again is also the
// reuse signal — see the next test.
rec := h.token(url.Values{
"grant_type": {"refresh_token"}, "refresh_token": {first.RefreshToken}, "client_id": {clientID},
})
if rec.Code != http.StatusBadRequest {
t.Errorf("the old refresh token still worked: %d", rec.Code)
}
}
// Replaying a consumed refresh token means either a client bug or a stolen
// token, and there is no way to tell. OAuth 2.1's answer is to assume theft and
// revoke the whole family — so the attacker AND the legitimate holder both lose
// access, and the legitimate one reauthorizes.
func TestRefreshReuseRevokesTheWholeFamily(t *testing.T) {
h := newHarness(t)
clientID := h.register()
code := h.authorizeOK(clientID, verifier43)
first := decodeTokens(t, h.exchange(clientID, code, verifier43))
second := decodeTokens(t, h.token(url.Values{
"grant_type": {"refresh_token"}, "refresh_token": {first.RefreshToken}, "client_id": {clientID},
}))
// The attacker replays the stolen (already rotated) token.
if rec := h.token(url.Values{
"grant_type": {"refresh_token"}, "refresh_token": {first.RefreshToken}, "client_id": {clientID},
}); rec.Code != http.StatusBadRequest {
t.Fatalf("reuse was accepted: %d", rec.Code)
}
// Now the LEGITIMATE current token must also be dead.
if rec := h.token(url.Values{
"grant_type": {"refresh_token"}, "refresh_token": {second.RefreshToken}, "client_id": {clientID},
}); rec.Code != http.StatusBadRequest {
t.Error("the family was not revoked after reuse — the thief keeps access")
}
// And so must the access token it minted.
if _, err := h.store.FindAccessToken(context.Background(), second.AccessToken); !errors.Is(err, ErrTokenUnusable) {
t.Error("an access token in the revoked family still validates")
}
}
func TestRefreshWithTheWrongClientIsRejected(t *testing.T) {
h := newHarness(t)
clientID := h.register()
other := h.register("https://other.example.test/cb")
code := h.authorizeOK(clientID, verifier43)
tokens := decodeTokens(t, h.exchange(clientID, code, verifier43))
if rec := h.token(url.Values{
"grant_type": {"refresh_token"}, "refresh_token": {tokens.RefreshToken}, "client_id": {other},
}); rec.Code != http.StatusBadRequest {
t.Errorf("another client refreshed this token: %d", rec.Code)
}
}
func TestExpiredRefreshTokenIsRejected(t *testing.T) {
h := newHarness(t)
clientID := h.register()
code := h.authorizeOK(clientID, verifier43)
tokens := decodeTokens(t, h.exchange(clientID, code, verifier43))
h.advance(RefreshTokenTTL + time.Hour)
if rec := h.token(url.Values{
"grant_type": {"refresh_token"}, "refresh_token": {tokens.RefreshToken}, "client_id": {clientID},
}); rec.Code != http.StatusBadRequest {
t.Errorf("an expired refresh token was accepted: %d", rec.Code)
}
}
/* ── Revocation ─────────────────────────────────────────────────────────── */
// Revoking must disconnect, which means killing the refresh token too.
// Revoking only the access token would leave the client able to mint another
// within seconds — so the button marked "disconnect" would not disconnect.
func TestRevocationKillsTheWholeFamily(t *testing.T) {
h := newHarness(t)
clientID := h.register()
code := h.authorizeOK(clientID, verifier43)
tokens := decodeTokens(t, h.exchange(clientID, code, verifier43))
form := url.Values{"token": {tokens.AccessToken}, "client_id": {clientID}}
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)
if rec.Code != http.StatusOK {
t.Fatalf("revocation: %d %s", rec.Code, rec.Body.String())
}
if _, err := h.store.FindAccessToken(context.Background(), tokens.AccessToken); !errors.Is(err, ErrTokenUnusable) {
t.Error("the access token still validates after revocation")
}
if r := h.token(url.Values{
"grant_type": {"refresh_token"}, "refresh_token": {tokens.RefreshToken}, "client_id": {clientID},
}); r.Code != http.StatusBadRequest {
t.Error("the refresh token survived revocation — this is not a disconnect")
}
}
// RFC 7009: revoking an unknown token is a success, or the endpoint becomes a
// way to test whether a token exists.
func TestRevokingAnUnknownTokenSucceeds(t *testing.T) {
h := newHarness(t)
form := url.Values{"token": {"not-a-real-token"}}
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)
if rec.Code != http.StatusOK {
t.Errorf("status = %d, want 200 per RFC 7009", rec.Code)
}
}
/* ── The authenticator: audience, scope, suspension ─────────────────────── */
func newAuthenticator(h *harness) *Authenticator {
return NewAuthenticator(h.store, auth.NewPGUserStore(h.h.Pool), testResource, discard())
}
func TestAuthenticatorProducesTheExistingIdentity(t *testing.T) {
h := newHarness(t)
clientID := h.register()
code := h.authorizeOK(clientID, verifier43)
tokens := decodeTokens(t, h.exchange(clientID, code, verifier43))
identity, err := newAuthenticator(h).Authenticate(context.Background(), tokens.AccessToken)
if err != nil {
t.Fatalf("a freshly issued token was refused: %v", err)
}
if identity.UserID != h.userID {
t.Errorf("UserID = %q, want %q", identity.UserID, h.userID)
}
// The tenant must come from the USER ROW, which is what makes a moved or
// suspended user take effect immediately.
if identity.OrgID != h.orgID {
t.Errorf("OrgID = %q, want %q", identity.OrgID, h.orgID)
}
if identity.Role != "admin" {
t.Errorf("Role = %q, want admin", identity.Role)
}
// No session behind a bearer identity; inventing one would make a token
// look like something logout could end.
if identity.SessionID != "" {
t.Errorf("SessionID = %q, want empty for a bearer identity", identity.SessionID)
}
}
// Audience confusion: a token minted by THIS server, for a DIFFERENT resource,
// must not be spendable here. This is the confused-deputy case the MCP spec
// calls out explicitly.
func TestAudienceConfusionIsRejected(t *testing.T) {
h := newHarness(t)
pair, err := h.store.IssuePair(context.Background(), Token{
ClientID: h.register(), UserID: h.userID, OrgID: h.orgID,
Scopes: []string{ScopeRead}, Audience: "https://some-other-service.test/mcp",
}, "")
if err != nil {
t.Fatalf("issue: %v", err)
}
if _, err := newAuthenticator(h).Authenticate(context.Background(), pair.AccessToken); err == nil {
t.Fatal("a token for another resource was accepted here")
}
}
func TestMissingScopeIsRejected(t *testing.T) {
h := newHarness(t)
pair, err := h.store.IssuePair(context.Background(), Token{
ClientID: h.register(), UserID: h.userID, OrgID: h.orgID,
Scopes: []string{"some.other.scope"}, Audience: testResource,
}, "")
if err != nil {
t.Fatalf("issue: %v", err)
}
if _, err := newAuthenticator(h).Authenticate(context.Background(), pair.AccessToken); err == nil {
t.Fatal("a token without krow.read was accepted")
}
}
// Suspension must take effect on the NEXT CALL, not at token expiry. Fifteen
// minutes of access for a suspended account is fifteen minutes too many.
func TestSuspendedUserLosesAccessImmediately(t *testing.T) {
h := newHarness(t)
clientID := h.register()
code := h.authorizeOK(clientID, verifier43)
tokens := decodeTokens(t, h.exchange(clientID, code, verifier43))
authr := newAuthenticator(h)
if _, err := authr.Authenticate(context.Background(), tokens.AccessToken); err != nil {
t.Fatalf("token should work while the user is active: %v", err)
}
if _, err := h.h.Pool.Exec(context.Background(),
`UPDATE users SET status = 'suspended' WHERE id = $1::uuid`, h.userID); err != nil {
t.Fatalf("suspend: %v", err)
}
if _, err := authr.Authenticate(context.Background(), tokens.AccessToken); err == nil {
t.Fatal("a suspended user's token still authenticated")
}
// And the family must be revoked, not merely refused once.
if r := h.token(url.Values{
"grant_type": {"refresh_token"}, "refresh_token": {tokens.RefreshToken}, "client_id": {clientID},
}); r.Code != http.StatusBadRequest {
t.Error("a suspended user could still refresh")
}
}
func TestExpiredAccessTokenIsRejected(t *testing.T) {
h := newHarness(t)
clientID := h.register()
code := h.authorizeOK(clientID, verifier43)
tokens := decodeTokens(t, h.exchange(clientID, code, verifier43))
h.advance(AccessTokenTTL + time.Minute)
if _, err := newAuthenticator(h).Authenticate(context.Background(), tokens.AccessToken); err == nil {
t.Fatal("an expired access token authenticated")
}
}
func TestRefreshTokenCannotBeUsedAsAnAccessToken(t *testing.T) {
h := newHarness(t)
clientID := h.register()
code := h.authorizeOK(clientID, verifier43)
tokens := decodeTokens(t, h.exchange(clientID, code, verifier43))
if _, err := newAuthenticator(h).Authenticate(context.Background(), tokens.RefreshToken); err == nil {
t.Fatal("a refresh token authenticated an MCP request")
}
}
func TestGarbageTokensAreRejected(t *testing.T) {
h := newHarness(t)
authr := newAuthenticator(h)
for name, token := range map[string]string{
"empty": "",
"whitespace": " ",
"random": "not-a-token",
"sql-ish": "' OR 1=1 --",
"very long": strings.Repeat("a", 5000),
} {
t.Run(name, func(t *testing.T) {
if _, err := authr.Authenticate(context.Background(), token); err == nil {
t.Errorf("%q authenticated", name)
}
})
}
}
/* ── Discovery metadata ─────────────────────────────────────────────────── */
func TestProtectedResourceMetadata(t *testing.T) {
rec := httptest.NewRecorder()
testConfig().ProtectedResourceHandler().ServeHTTP(rec,
httptest.NewRequest(http.MethodGet, "/.well-known/oauth-protected-resource", nil))
if rec.Code != http.StatusOK {
t.Fatalf("status = %d, want 200", rec.Code)
}
var out protectedResourceMetadata
if err := json.Unmarshal(rec.Body.Bytes(), &out); err != nil {
t.Fatalf("metadata did not decode: %v", err)
}
if out.Resource != testResource {
t.Errorf("resource = %q, want %q", out.Resource, testResource)
}
if len(out.AuthorizationServers) != 1 || out.AuthorizationServers[0] != testIssuer {
t.Errorf("authorization_servers = %v, want [%q]", out.AuthorizationServers, testIssuer)
}
// The MCP spec forbids a token in the query string; advertising anything
// but "header" would tell a client otherwise.
if len(out.BearerMethodsSupported) != 1 || out.BearerMethodsSupported[0] != "header" {
t.Errorf("bearer_methods_supported = %v, want [header]", out.BearerMethodsSupported)
}
}
func TestAuthorizationServerMetadata(t *testing.T) {
rec := httptest.NewRecorder()
testConfig().AuthorizationServerHandler().ServeHTTP(rec,
httptest.NewRequest(http.MethodGet, "/.well-known/oauth-authorization-server", nil))
var out authorizationServerMetadata
if err := json.Unmarshal(rec.Body.Bytes(), &out); err != nil {
t.Fatalf("metadata did not decode: %v", err)
}
if out.Issuer != testIssuer {
t.Errorf("issuer = %q, want %q", out.Issuer, testIssuer)
}
// Every list is a promise. Each must name only what is implemented.
if strings.Join(out.ResponseTypesSupported, ",") != "code" {
t.Errorf("response_types_supported = %v; implicit must not be advertised", out.ResponseTypesSupported)
}
if strings.Join(out.CodeChallengeMethodsSupported, ",") != MethodS256 {
t.Errorf("code_challenge_methods_supported = %v, want [S256]", out.CodeChallengeMethodsSupported)
}
for _, forbidden := range []string{"password", "client_credentials", "implicit"} {
for _, advertised := range out.GrantTypesSupported {
if advertised == forbidden {
t.Errorf("grant_types_supported advertises %q, which is refused", forbidden)
}
}
}
for _, advertised := range out.ScopesSupported {
if advertised == ScopeWrite {
t.Error("scopes_supported advertises krow.write, which is not issued")
}
}
if !out.ResourceIndicatorsSupported {
t.Error("resource_indicators_supported must be true — RFC 8707 is required by MCP")
}
// Every endpoint comes from configuration, never a hardcoded domain.
for name, got := range map[string]string{
"authorization_endpoint": out.AuthorizationEndpoint,
"token_endpoint": out.TokenEndpoint,
"registration_endpoint": out.RegistrationEndpoint,
} {
if !strings.HasPrefix(got, testIssuer) {
t.Errorf("%s = %q, want it under the configured issuer", name, got)
}
}
}
// A token response must never be cached: the body is a credential.
func TestTokenResponsesAreNotCacheable(t *testing.T) {
h := newHarness(t)
clientID := h.register()
code := h.authorizeOK(clientID, verifier43)
rec := h.exchange(clientID, code, verifier43)
if got := rec.Header().Get("Cache-Control"); !strings.Contains(got, "no-store") {
t.Errorf("Cache-Control = %q, want no-store", got)
}
}

View File

@@ -0,0 +1,130 @@
// Package oauth is KROW's OAuth 2.1 authorization server and the token store
// behind it.
//
// It exists for one caller: the MCP surface, which needs a way to authenticate
// a client that cannot hold a cookie. Everything here is in service of turning
// a browser-based approval into an opaque bearer token that
// mcpserver.TokenAuthenticator can resolve back into the SAME
// authctx.Identity the cookie path produces.
//
// # WHAT THIS PACKAGE DOES NOT DO
//
// It does not authorize anything. It establishes WHO is calling; what they may
// then read is decided by the existing policy table, in the existing tool
// layer, exactly as it is for a cookie session. There is no OAuth scope that
// grants access to a row. `krow.read` says "this client may use the read
// tools"; whether this user may see a particular row is a question
// tools/scope.go answers and this package never touches.
//
// It also does not store a password, check one, or keep a second user table.
// The authorization endpoint authenticates the person using the session they
// already have — see authserver.go.
package oauth
import (
"crypto/sha256"
"crypto/subtle"
"encoding/base64"
"errors"
"regexp"
)
// PKCE — Proof Key for Code Exchange, RFC 7636.
//
// The problem it solves: an authorization code travels back through a browser
// redirect, which is the least trustworthy hop in the flow. Anything that can
// observe that redirect — a malicious app registered for the same custom URL
// scheme, a proxy, a shoulder — can steal the code. For a confidential client
// that does not matter, because redeeming the code also requires a client
// secret. A public client has no secret, so the code alone would be enough.
//
// PKCE gives the client a per-request secret instead. It invents a random
// `code_verifier`, sends only SHA-256 of it with the authorization request, and
// presents the verifier itself at the token endpoint. A stolen code is useless
// without the verifier, which never travelled through the browser.
//
// S256 ONLY. RFC 7636 also defines `plain`, where the challenge IS the
// verifier. That protects against nothing — anyone who stole the code from the
// redirect also stole the challenge, and the challenge is the verifier — and
// OAuth 2.1 forbids it for public clients. It is refused here, and refused
// again by a CHECK constraint in migration 000013, so no code path can relax it.
// MethodS256 is the only code_challenge_method this server accepts.
const MethodS256 = "S256"
var (
// ErrUnsupportedChallengeMethod covers `plain` and anything else.
ErrUnsupportedChallengeMethod = errors.New("oauth: code_challenge_method must be S256")
// ErrMalformedChallenge covers a challenge that is not base64url of a
// SHA-256 digest.
ErrMalformedChallenge = errors.New("oauth: malformed code_challenge")
// ErrMalformedVerifier covers a verifier outside RFC 7636's length or
// character set.
ErrMalformedVerifier = errors.New("oauth: malformed code_verifier")
// ErrVerifierMismatch is the one that matters: a verifier that does not
// hash to the stored challenge.
ErrVerifierMismatch = errors.New("oauth: code_verifier does not match code_challenge")
)
// challengePattern is base64url of a 32-byte digest: 43 characters, unpadded.
// The same pattern migration 000013 enforces in oauth_grants_challenge_shape.
var challengePattern = regexp.MustCompile(`^[A-Za-z0-9_-]{43}$`)
// verifierPattern is RFC 7636 section 4.1's `code_verifier` grammar:
// unreserved characters only, 43 to 128 of them.
var verifierPattern = regexp.MustCompile(`^[A-Za-z0-9._~-]{43,128}$`)
// ValidateChallenge checks a code_challenge and its method at the authorization
// endpoint, before any row is written.
//
// Rejecting a malformed challenge here rather than at the token endpoint means
// the failure lands where the client can act on it — on its own authorization
// request — instead of after a person has been walked through a consent screen
// for a flow that was never going to complete.
func ValidateChallenge(challenge, method string) error {
if method != MethodS256 {
return ErrUnsupportedChallengeMethod
}
if !challengePattern.MatchString(challenge) {
return ErrMalformedChallenge
}
return nil
}
// VerifyChallenge reports whether a verifier matches a stored challenge.
//
// The comparison is constant-time. A byte-by-byte comparison that returned
// early would leak, through timing, how much of a guessed verifier was correct
// — which turns an infeasible search into a feasible one, one character at a
// time. The values being compared are both base64url text of the same fixed
// length, so subtle.ConstantTimeCompare is exactly the right tool.
func VerifyChallenge(verifier, challenge, method string) error {
if method != MethodS256 {
return ErrUnsupportedChallengeMethod
}
if !verifierPattern.MatchString(verifier) {
return ErrMalformedVerifier
}
if !challengePattern.MatchString(challenge) {
return ErrMalformedChallenge
}
computed := ChallengeFor(verifier)
if subtle.ConstantTimeCompare([]byte(computed), []byte(challenge)) != 1 {
return ErrVerifierMismatch
}
return nil
}
// ChallengeFor derives the S256 challenge for a verifier.
//
// base64url WITHOUT padding, per RFC 7636 appendix A. Padding would add a '='
// that has to be escaped in a query string, and a client that padded would
// produce a challenge this server did not recognise.
func ChallengeFor(verifier string) string {
sum := sha256.Sum256([]byte(verifier))
return base64.RawURLEncoding.EncodeToString(sum[:])
}

View File

@@ -0,0 +1,107 @@
package oauth
import (
"errors"
"strings"
"testing"
)
// The RFC 7636 appendix B worked example. Using the spec's own vector rather
// than a value this implementation produced means the test would catch an
// encoding mistake that is self-consistent — base64 standard instead of
// base64url, say, or padded instead of raw — which a round-trip test could not.
const (
specVerifier = "dBjftJeZ4CVP-mB92K27uhbUJU1p1r_wW1gFWFOEjXk"
specChallenge = "E9Melhoa2OwvFrEMTJguCHaoeK1t8URWbuGJSstw-cM"
)
func TestChallengeForMatchesTheRFCVector(t *testing.T) {
if got := ChallengeFor(specVerifier); got != specChallenge {
t.Errorf("ChallengeFor(RFC 7636 verifier) = %q, want %q", got, specChallenge)
}
}
func TestVerifyChallengeAcceptsTheCorrectVerifier(t *testing.T) {
if err := VerifyChallenge(specVerifier, specChallenge, MethodS256); err != nil {
t.Errorf("the RFC's own verifier was rejected: %v", err)
}
}
func TestVerifyChallengeRejectsAWrongVerifier(t *testing.T) {
// Same length and character set, one character different. A comparison
// that was accidentally checking length or prefix would let this through.
wrong := "dBjftJeZ4CVP-mB92K27uhbUJU1p1r_wW1gFWFOEjXX"
err := VerifyChallenge(wrong, specChallenge, MethodS256)
if !errors.Is(err, ErrVerifierMismatch) {
t.Errorf("err = %v, want ErrVerifierMismatch", err)
}
}
func TestVerifyChallengeRejectsAMissingVerifier(t *testing.T) {
if err := VerifyChallenge("", specChallenge, MethodS256); !errors.Is(err, ErrMalformedVerifier) {
t.Errorf("err = %v, want ErrMalformedVerifier", err)
}
}
// `plain` must be refused wherever it appears. It is legal in RFC 7636 and
// forbidden by OAuth 2.1 for public clients, because the challenge IS the
// verifier and anyone who stole one stole both.
func TestPlainMethodIsRejected(t *testing.T) {
for _, method := range []string{"plain", "PLAIN", "", "s256", "S512"} {
t.Run("method="+method, func(t *testing.T) {
if err := ValidateChallenge(specChallenge, method); !errors.Is(err, ErrUnsupportedChallengeMethod) {
t.Errorf("ValidateChallenge: err = %v, want ErrUnsupportedChallengeMethod", err)
}
if err := VerifyChallenge(specVerifier, specChallenge, method); !errors.Is(err, ErrUnsupportedChallengeMethod) {
t.Errorf("VerifyChallenge: err = %v, want ErrUnsupportedChallengeMethod", err)
}
})
}
}
func TestMalformedChallengeIsRejected(t *testing.T) {
for name, challenge := range map[string]string{
"empty": "",
"too short": "abc",
"too long": strings.Repeat("a", 44),
"padded base64": "E9Melhoa2OwvFrEMTJguCHaoeK1t8URWbuGJSstw-cM=",
"standard base64": "E9Melhoa2OwvFrEMTJguCHaoeK1t8URWbuGJSstw+cM",
"illegal char": "E9Melhoa2OwvFrEMTJguCHaoeK1t8URWbuGJSstw!cM",
"whitespace": "E9Melhoa2OwvFrEMTJguCHaoeK1t8URWbuGJSstw cM",
"newline injected": "E9Melhoa2OwvFrEMTJguCHaoeK1t8URWbuGJSstw\ncM",
} {
t.Run(name, func(t *testing.T) {
if err := ValidateChallenge(challenge, MethodS256); !errors.Is(err, ErrMalformedChallenge) {
t.Errorf("err = %v, want ErrMalformedChallenge", err)
}
})
}
}
// RFC 7636 section 4.1 constrains the verifier to 43–128 unreserved characters.
// A verifier outside that range is malformed regardless of what it hashes to.
func TestMalformedVerifierIsRejected(t *testing.T) {
for name, verifier := range map[string]string{
"too short (42)": strings.Repeat("a", 42),
"too long (129)": strings.Repeat("a", 129),
"illegal char": strings.Repeat("a", 42) + "!",
"whitespace": strings.Repeat("a", 42) + " ",
} {
t.Run(name, func(t *testing.T) {
if err := VerifyChallenge(verifier, specChallenge, MethodS256); !errors.Is(err, ErrMalformedVerifier) {
t.Errorf("err = %v, want ErrMalformedVerifier", err)
}
})
}
}
// A verifier at each end of the legal range must be accepted, or clients
// generating the maximum length would fail against this server.
func TestVerifierBoundariesAreAccepted(t *testing.T) {
for _, length := range []int{43, 128} {
verifier := strings.Repeat("a", length)
if err := VerifyChallenge(verifier, ChallengeFor(verifier), MethodS256); err != nil {
t.Errorf("a %d-character verifier was rejected: %v", length, err)
}
}
}

View File

@@ -0,0 +1,444 @@
package oauth
import (
"context"
"errors"
"fmt"
"time"
"github.com/krow/krow-backend/go-api/internal/auth"
"github.com/krow/krow-backend/go-api/internal/repo"
)
// The persistence layer for clients, authorization codes and tokens.
//
// Two rules hold throughout this file and are worth stating once:
//
// 1. NO RAW CREDENTIAL IS EVER WRITTEN. Every code and token is hashed with
// auth.HashToken before it reaches a statement. The CHECK constraints in
// migrations 000013 and 000014 refuse anything that is not 64 hex
// characters, so this is enforced twice — in Go where the mistake would be
// made, and in the schema where it would land.
//
// 2. SINGLE-USE IS ENFORCED BY THE UPDATE, NOT BY A READ. Redeeming a code or
// a refresh token is one statement that marks the row consumed and returns
// it in the same breath. A SELECT followed by an UPDATE has a window
// between them where two concurrent requests both see an unspent row and
// both proceed, which is precisely the replay the single-use rule exists to
// prevent.
var (
// ErrNotFound covers a client, code or token that does not exist.
ErrNotFound = errors.New("oauth: not found")
// ErrGrantUnusable covers a code that is expired, already consumed, or
// simply absent. ONE error for all three: distinguishing them tells a
// caller whether a code they hold was ever real, which is an oracle.
ErrGrantUnusable = errors.New("oauth: authorization code is not usable")
// ErrTokenUnusable covers a token that is unknown, expired, revoked or
// consumed. One error, same reasoning.
ErrTokenUnusable = errors.New("oauth: token is not usable")
// ErrRefreshReuse is raised when a CONSUMED refresh token is presented
// again. It is distinct from ErrTokenUnusable internally because it
// triggers family revocation — but the caller must still answer the client
// with an indistinguishable error.
ErrRefreshReuse = errors.New("oauth: refresh token reuse detected")
)
// Store is the database-backed persistence for this package.
type Store struct {
db repo.Querier
now func() time.Time
}
// NewStore builds a store over the existing pool.
func NewStore(db repo.Querier) *Store {
return &Store{db: db, now: time.Now}
}
// WithClock replaces the clock, so expiry can be tested without sleeping.
func (s *Store) WithClock(now func() time.Time) *Store {
s.now = now
return s
}
/* ── Clients ────────────────────────────────────────────────────────────── */
// Client is a registered OAuth client.
type Client struct {
ClientID string
ClientName string
RedirectURIs []string
GrantTypes []string
Scopes []string
DisabledAt *time.Time
}
// AllowsRedirect reports whether a redirect_uri is registered to this client.
//
// EXACT string equality. Not a prefix match, not a normalised comparison, not
// "same host and port". Every relaxation of this check is an open redirect: a
// prefix match lets `https://good.example/cb.attacker.com` through, and
// normalising lets encoding tricks through. RFC 6749 section 3.1.2.3 says
// exact, and exact is what this is.
func (c Client) AllowsRedirect(uri string) bool {
for _, registered := range c.RedirectURIs {
if registered == uri {
return true
}
}
return false
}
// AllowsScopes reports whether every requested scope is registered.
func (c Client) AllowsScopes(requested []string) bool {
for _, want := range requested {
found := false
for _, have := range c.Scopes {
if have == want {
found = true
break
}
}
if !found {
return false
}
}
return true
}
// CreateClient registers a new public client.
func (s *Store) CreateClient(ctx context.Context, c Client) error {
_, err := s.db.Exec(ctx,
`INSERT INTO oauth_clients (client_id, client_name, redirect_uris, grant_types, scopes)
VALUES ($1, $2, $3, $4, $5)`,
c.ClientID, c.ClientName, c.RedirectURIs, c.GrantTypes, c.Scopes)
if err != nil {
return fmt.Errorf("oauth: create client: %w", err)
}
return nil
}
// FindClient resolves a client_id. A disabled client is reported as not found:
// whether it once existed is not the caller's business.
func (s *Store) FindClient(ctx context.Context, clientID string) (Client, error) {
var c Client
err := s.db.QueryRow(ctx,
`SELECT client_id, client_name, redirect_uris, grant_types, scopes, disabled_at
FROM oauth_clients
WHERE client_id = $1 AND disabled_at IS NULL`,
clientID).Scan(&c.ClientID, &c.ClientName, &c.RedirectURIs, &c.GrantTypes, &c.Scopes, &c.DisabledAt)
if err != nil {
return Client{}, ErrNotFound
}
return c, nil
}
/* ── Authorization codes ────────────────────────────────────────────────── */
// Grant is an issued authorization code, as stored.
type Grant struct {
ID string
ClientID string
UserID string
OrgID string
RedirectURI string
Scopes []string
Resource string
CodeChallenge string
CodeChallengeMethod string
ExpiresAt time.Time
}
// GrantTTL is how long an authorization code stays redeemable.
//
// Sixty seconds. RFC 6749 recommends a maximum of ten minutes and "a maximum
// of 60 seconds is RECOMMENDED" for the code's lifetime in OAuth 2.1 guidance.
// The code is in transit through a browser redirect and is exchanged
// immediately by a client that is already waiting for it; a longer window buys
// nothing and widens the replay opportunity.
const GrantTTL = 60 * time.Second
// CreateGrant stores an authorization code, returning the RAW code exactly
// once.
//
// The raw value is returned and never persisted. The caller puts it in a
// redirect and forgets it.
func (s *Store) CreateGrant(ctx context.Context, g Grant) (rawCode string, err error) {
rawCode, err = auth.GenerateToken()
if err != nil {
return "", fmt.Errorf("oauth: generate code: %w", err)
}
_, err = s.db.Exec(ctx,
`INSERT INTO oauth_grants
(code_hash, client_id, user_id, org_id, redirect_uri, scopes, resource,
code_challenge, code_challenge_method, expires_at)
VALUES ($1, $2, $3::uuid, $4::uuid, $5, $6, $7, $8, $9, $10)`,
auth.HashToken(rawCode), g.ClientID, g.UserID, g.OrgID, g.RedirectURI,
g.Scopes, g.Resource, g.CodeChallenge, g.CodeChallengeMethod,
s.now().Add(GrantTTL))
if err != nil {
return "", fmt.Errorf("oauth: create grant: %w", err)
}
return rawCode, nil
}
// RedeemGrant consumes an authorization code and returns what it was bound to.
//
// ONE STATEMENT. The UPDATE marks the row consumed and RETURNS it, so the read
// and the write cannot be interleaved by a concurrent request. The predicate
// carries the whole single-use rule: `consumed_at IS NULL` means a spent code
// matches nothing, and `expires_at > now()` means an old one does too. A second
// redemption of the same code updates zero rows and therefore fails, which is
// what replay protection looks like when the database enforces it.
func (s *Store) RedeemGrant(ctx context.Context, rawCode string) (Grant, error) {
var g Grant
err := s.db.QueryRow(ctx,
`UPDATE oauth_grants
SET consumed_at = now()
WHERE code_hash = $1
AND consumed_at IS NULL
AND expires_at > $2
RETURNING id::text, client_id, user_id::text, org_id::text, redirect_uri,
scopes, resource, code_challenge, code_challenge_method, expires_at`,
auth.HashToken(rawCode), s.now()).
Scan(&g.ID, &g.ClientID, &g.UserID, &g.OrgID, &g.RedirectURI,
&g.Scopes, &g.Resource, &g.CodeChallenge, &g.CodeChallengeMethod, &g.ExpiresAt)
if err != nil {
// No row: unknown, expired or already spent. Indistinguishable on
// purpose — see ErrGrantUnusable.
return Grant{}, ErrGrantUnusable
}
return g, nil
}
/* ── Tokens ─────────────────────────────────────────────────────────────── */
// Token is an issued access or refresh token, as stored.
type Token struct {
ID string
Type string
FamilyID string
ClientID string
UserID string
OrgID string
Scopes []string
Audience string
ExpiresAt time.Time
}
// Token lifetimes.
//
// Fifteen minutes for an access token is the number the MCP plan committed to,
// and the reasoning is that an access token travels on every single request: it
// is the most exposed credential in the system and the one with the least need
// to be long-lived, because a refresh token exists precisely so the client can
// get another without troubling the user.
//
// Thirty days for a refresh token matches the session's own "remember me"
// ceiling, so a connected client and a remembered browser lapse on the same
// schedule rather than on two different ones nobody can remember.
const (
AccessTokenTTL = 15 * time.Minute
RefreshTokenTTL = 30 * 24 * time.Hour
)
// TokenPair is what a successful token request produces.
//
// The raw values are here and nowhere else: they are returned to the client in
// the token response and are never stored, logged or re-derivable.
type TokenPair struct {
AccessToken string
RefreshToken string
ExpiresIn int
Scopes []string
FamilyID string
}
// IssuePair mints an access and refresh token in one family.
//
// familyID empty starts a new lineage; a supplied one continues an existing
// lineage through a rotation, which is what lets reuse detection revoke every
// descendant of a stolen token.
func (s *Store) IssuePair(ctx context.Context, t Token, familyID string) (TokenPair, error) {
if familyID == "" {
generated, err := newUUID()
if err != nil {
return TokenPair{}, err
}
familyID = generated
}
access, err := auth.GenerateToken()
if err != nil {
return TokenPair{}, fmt.Errorf("oauth: generate access token: %w", err)
}
refresh, err := auth.GenerateToken()
if err != nil {
return TokenPair{}, fmt.Errorf("oauth: generate refresh token: %w", err)
}
now := s.now()
for _, row := range []struct {
raw string
kind string
expires time.Time
}{
{access, "access", now.Add(AccessTokenTTL)},
{refresh, "refresh", now.Add(RefreshTokenTTL)},
} {
if _, err := s.db.Exec(ctx,
`INSERT INTO oauth_tokens
(token_hash, token_type, family_id, client_id, user_id, org_id,
scopes, audience, expires_at)
VALUES ($1, $2, $3::uuid, $4, $5::uuid, $6::uuid, $7, $8, $9)`,
auth.HashToken(row.raw), row.kind, familyID, t.ClientID, t.UserID,
t.OrgID, t.Scopes, t.Audience, row.expires); err != nil {
return TokenPair{}, fmt.Errorf("oauth: store %s token: %w", row.kind, err)
}
}
return TokenPair{
AccessToken: access,
RefreshToken: refresh,
ExpiresIn: int(AccessTokenTTL.Seconds()),
Scopes: t.Scopes,
FamilyID: familyID,
}, nil
}
// FindAccessToken resolves a raw access token for validation.
//
// Read-only: validation happens on every MCP request and must not write. The
// predicate does the whole job — unknown, expired, revoked and wrong-type all
// return no row and therefore the same error.
func (s *Store) FindAccessToken(ctx context.Context, raw string) (Token, error) {
var t Token
err := s.db.QueryRow(ctx,
`SELECT id::text, token_type, family_id::text, client_id, user_id::text,
org_id::text, scopes, audience, expires_at
FROM oauth_tokens
WHERE token_hash = $1
AND token_type = 'access'
AND revoked_at IS NULL
AND expires_at > $2`,
auth.HashToken(raw), s.now()).
Scan(&t.ID, &t.Type, &t.FamilyID, &t.ClientID, &t.UserID, &t.OrgID,
&t.Scopes, &t.Audience, &t.ExpiresAt)
if err != nil {
return Token{}, ErrTokenUnusable
}
return t, nil
}
// RedeemRefreshToken consumes a refresh token, or detects its reuse.
//
// The two-step here is deliberate and is the heart of reuse detection:
//
// 1. Try to consume an unspent, unexpired, unrevoked refresh token. One
// statement, same single-use reasoning as RedeemGrant.
// 2. If that matched nothing, look again WITHOUT the `consumed_at IS NULL`
// predicate. A row that exists but was already consumed is not an ordinary
// failure — it means someone presented a token that had already been
// rotated away, and there is no way to tell the legitimate client retrying
// from an attacker replaying a stolen token.
//
// OAuth 2.1's answer to that ambiguity is to assume the worse case and revoke
// the whole family. The attacker loses access; the legitimate client is pushed
// through a fresh authorization it can complete. Doing nothing would leave a
// thief with a working credential.
func (s *Store) RedeemRefreshToken(ctx context.Context, raw string) (Token, error) {
hash := auth.HashToken(raw)
var t Token
err := s.db.QueryRow(ctx,
`UPDATE oauth_tokens
SET consumed_at = now(), last_used_at = now()
WHERE token_hash = $1
AND token_type = 'refresh'
AND consumed_at IS NULL
AND revoked_at IS NULL
AND expires_at > $2
RETURNING id::text, token_type, family_id::text, client_id, user_id::text,
org_id::text, scopes, audience, expires_at`,
hash, s.now()).
Scan(&t.ID, &t.Type, &t.FamilyID, &t.ClientID, &t.UserID, &t.OrgID,
&t.Scopes, &t.Audience, &t.ExpiresAt)
if err == nil {
return t, nil
}
// Step 2: was this a token that HAD been valid and is now spent?
var familyID string
if probeErr := s.db.QueryRow(ctx,
`SELECT family_id::text FROM oauth_tokens
WHERE token_hash = $1 AND token_type = 'refresh' AND consumed_at IS NOT NULL`,
hash).Scan(&familyID); probeErr == nil {
// Reuse. Revoke the lineage and report it, so the caller can log it at
// a level that gets noticed — while still answering the client with an
// indistinguishable error.
_ = s.RevokeFamily(ctx, familyID, "refresh_token_reuse")
return Token{}, ErrRefreshReuse
}
return Token{}, ErrTokenUnusable
}
/* ── Revocation ─────────────────────────────────────────────────────────── */
// RevokeFamily revokes every token in a rotation lineage.
//
// Idempotent, and it does not care whether the rows were already revoked: the
// predicate narrows to unrevoked rows so a second call is a no-op rather than
// an error, which matters because this is called from an error path.
func (s *Store) RevokeFamily(ctx context.Context, familyID, reason string) error {
_, err := s.db.Exec(ctx,
`UPDATE oauth_tokens
SET revoked_at = now(), revoked_reason = $2
WHERE family_id = $1::uuid AND revoked_at IS NULL`,
familyID, reason)
if err != nil {
return fmt.Errorf("oauth: revoke family: %w", err)
}
return nil
}
// RevokeToken revokes one token by its raw value, and its family with it.
//
// Revoking the family rather than the single row is what makes "disconnect"
// mean what a person expects. Revoking one access token would leave the
// refresh token alive to mint another within seconds, so the button that says
// "disconnect Claude" would not disconnect Claude.
func (s *Store) RevokeToken(ctx context.Context, raw, reason string) error {
var familyID string
if err := s.db.QueryRow(ctx,
`SELECT family_id::text FROM oauth_tokens WHERE token_hash = $1`,
auth.HashToken(raw)).Scan(&familyID); err != nil {
// RFC 7009: revoking an unknown token is a success. Saying otherwise
// turns the revocation endpoint into a way to test whether a token
// exists.
return nil
}
return s.RevokeFamily(ctx, familyID, reason)
}
// RevokeAllForUser revokes every token a user holds.
//
// Called when an account is suspended or a person disconnects every app. Token
// validation already re-reads the user and refuses a suspended one, so this is
// belt to that braces: it stops the tokens existing rather than relying on
// every future validation to notice.
func (s *Store) RevokeAllForUser(ctx context.Context, userID, reason string) error {
_, err := s.db.Exec(ctx,
`UPDATE oauth_tokens
SET revoked_at = now(), revoked_reason = $2
WHERE user_id = $1::uuid AND revoked_at IS NULL`,
userID, reason)
if err != nil {
return fmt.Errorf("oauth: revoke user tokens: %w", err)
}
return nil
}

View File

@@ -0,0 +1,209 @@
// Package ratelimit is a fixed-window request limiter shared across API
// instances.
//
// It exists because the limiter this service already had — httpserver's
// attemptLimiter — is an in-process map, and an in-process limiter behind N
// replicas enforces N times the configured limit. That is fine for the thing it
// guards (failed logins, where the real defence is the password hash's cost)
// and not fine for an endpoint that writes a database row for any caller who
// can reach it.
//
// The existing limiter is deliberately left alone. Replacing it is not this
// phase's job, it would change login behaviour, and the two have different
// shapes: attemptLimiter counts FAILURES and resets on success, which is the
// right model for a password and the wrong one for a request budget.
//
// WHAT THIS GUARANTEES
//
// - Correct under concurrency, including across instances: the count comes
// back from the same statement that increments it, so two callers cannot
// both read "9" and both proceed.
// - Bounded memory: the state is a table, swept by the cleanup job.
// - No credential ever becomes a key: callers pass an already-hashed subject,
// and Key refuses to build a bucket from anything that looks raw.
//
// # WHAT IT DOES NOT GUARANTEE
//
// Exactness at a window boundary. A fixed window admits up to 2× the limit
// across the seam — ten requests at 11:59:59 and ten more at 12:00:01. A
// sliding window would fix that and costs a row per request. For abuse
// prevention the burst is acceptable and the trade is deliberate.
package ratelimit
import (
"context"
"crypto/sha256"
"encoding/hex"
"fmt"
"time"
"github.com/krow/krow-backend/go-api/internal/repo"
)
// Limiter counts requests per bucket per window.
type Limiter struct {
db repo.Querier
now func() time.Time
// failOpen decides what happens when the DATABASE fails, which is the one
// judgement call in this package.
//
// Default false — fail CLOSED. A limiter that cannot count is a limiter
// that is not limiting, and these endpoints are the ones worth protecting
// most when things are already going wrong. The alternative, failing open,
// turns a database blip into an unmetered window on an endpoint that
// writes rows for anonymous callers.
//
// Configurable because that is not the right answer everywhere: a
// deployment that would rather serve MCP degraded than refuse it can say
// so deliberately, in one place, rather than by a comment somebody has to
// remember.
failOpen bool
}
// New builds a limiter over the shared pool.
func New(db repo.Querier) *Limiter {
return &Limiter{db: db, now: time.Now}
}
// WithClock replaces the clock, so window rollover can be tested without
// waiting for one.
func (l *Limiter) WithClock(now func() time.Time) *Limiter {
l.now = now
return l
}
// WithFailOpen makes a database failure permit the request rather than refuse
// it. See the field comment: the default is to refuse.
func (l *Limiter) WithFailOpen(open bool) *Limiter {
l.failOpen = open
return l
}
// Rule is one limit: how many requests, over how long.
type Rule struct {
// Name is the scope, and it becomes the bucket's prefix. Keep it stable —
// renaming a scope resets everyone's counter.
Name string
// Limit is the number of requests permitted per window.
Limit int
// Window is the fixed window's length.
Window time.Duration
}
// Decision is the answer for one request.
type Decision struct {
// Allowed is whether the caller may proceed.
Allowed bool
// Remaining is how many requests are left in this window, never negative.
Remaining int
// RetryAfter is how long until the window rolls over. Rendered into the
// Retry-After header on a 429, so a well-behaved client waits exactly long
// enough rather than guessing.
RetryAfter time.Duration
// Limit and Window echo the rule, for the response headers.
Limit int
Window time.Duration
}
// Subject hashes a bucket subject.
//
// EVERY caller must pass identifying material through this. A bucket key built
// from a raw token would write that token to a table, to any log line naming
// the bucket, and to every slow-query report the row ever appears in. Hashing
// costs nothing here — the value is never read back, only compared.
//
// Truncated to 32 hex characters: 128 bits, far beyond collision risk for a
// counter, and it keeps the keys readable in a psql session while still being
// irreversible.
func Subject(raw string) string {
sum := sha256.Sum256([]byte(raw))
return hex.EncodeToString(sum[:])[:32]
}
// Allow records one request against a rule and reports whether it may proceed.
//
// The whole decision is one statement. It is worth reading, because everything
// this package claims about concurrency rests on it:
//
// INSERT INTO rate_limits (bucket, window_start, count, expires_at)
// VALUES ($1, $2, 1, $3)
// ON CONFLICT (bucket, window_start)
// DO UPDATE SET count = rate_limits.count + 1
// RETURNING count
//
// The row is created or incremented, and the resulting count comes back, in one
// round trip under one implicit transaction. Two instances racing on the same
// bucket serialise on the primary key, and each sees a distinct count. There is
// no read-then-write window for them to slip through.
//
// A request is counted even when it is refused. That is deliberate: a caller
// hammering a limit should not be able to keep their own window open by
// spending it, and the alternative — not counting refusals — makes the limit
// cheaper to probe.
func (l *Limiter) Allow(ctx context.Context, rule Rule, subject string) (Decision, error) {
now := l.now()
windowStart := now.Truncate(rule.Window)
expiresAt := windowStart.Add(rule.Window)
bucket := rule.Name + ":" + subject
var count int
err := l.db.QueryRow(ctx,
`INSERT INTO rate_limits (bucket, window_start, count, expires_at)
VALUES ($1, $2, 1, $3)
ON CONFLICT (bucket, window_start)
DO UPDATE SET count = rate_limits.count + 1
RETURNING count`,
bucket, windowStart, expiresAt).Scan(&count)
if err != nil {
if l.failOpen {
return Decision{Allowed: true, Remaining: rule.Limit, Limit: rule.Limit, Window: rule.Window},
fmt.Errorf("ratelimit: %w", err)
}
return Decision{Allowed: false, RetryAfter: rule.Window, Limit: rule.Limit, Window: rule.Window},
fmt.Errorf("ratelimit: %w", err)
}
remaining := rule.Limit - count
if remaining < 0 {
remaining = 0
}
return Decision{
Allowed: count <= rule.Limit,
Remaining: remaining,
RetryAfter: expiresAt.Sub(now),
Limit: rule.Limit,
Window: rule.Window,
}, nil
}
// Sweep deletes expired counters, in bounded batches.
//
// Bounded because an unbounded DELETE on a busy table takes a lock for as long
// as it takes to finish, and "as long as it takes" is not a number anyone can
// predict at 3am. A batch of a few thousand rows completes in milliseconds and
// can simply be run again.
//
// Safe to run concurrently: two sweeps delete disjoint sets because the
// subquery re-reads under each statement's own snapshot, and a row deleted
// twice is not an error.
func (l *Limiter) Sweep(ctx context.Context, batch int) (int64, error) {
if batch <= 0 {
batch = 5000
}
tag, err := l.db.Exec(ctx,
`DELETE FROM rate_limits
WHERE ctid IN (
SELECT ctid FROM rate_limits WHERE expires_at < $1 LIMIT $2
)`,
l.now(), batch)
if err != nil {
return 0, fmt.Errorf("ratelimit: sweep: %w", err)
}
return tag.RowsAffected(), nil
}

View File

@@ -0,0 +1,452 @@
package ratelimit
import (
"context"
"strings"
"sync"
"testing"
"time"
"github.com/jackc/pgx/v5"
"github.com/jackc/pgx/v5/pgconn"
"github.com/krow/krow-backend/go-api/internal/testutil"
)
func newLimiter(t *testing.T) (*Limiter, *testutil.Harness, *time.Time) {
t.Helper()
h := testutil.New(t)
clock := time.Now().Truncate(time.Hour) // a clean window boundary
l := New(h.Pool).WithClock(func() time.Time { return clock })
return l, h, &clock
}
var testRule = Rule{Name: "test.rule", Limit: 3, Window: time.Minute}
func mustAllow(t *testing.T, l *Limiter, subject string) Decision {
t.Helper()
d, err := l.Allow(context.Background(), testRule, Subject(subject))
if err != nil {
t.Fatalf("Allow: %v", err)
}
return d
}
/* ── The basic contract ─────────────────────────────────────────────────── */
func TestUnderAtAndOverTheLimit(t *testing.T) {
l, _, _ := newLimiter(t)
// Under: each of the first three is allowed, and remaining counts down.
for i := 1; i <= testRule.Limit; i++ {
d := mustAllow(t, l, "alice")
if !d.Allowed {
t.Fatalf("request %d of %d was refused", i, testRule.Limit)
}
if want := testRule.Limit - i; d.Remaining != want {
t.Errorf("request %d: remaining = %d, want %d", i, d.Remaining, want)
}
}
// Over: the next one is refused and carries a usable Retry-After.
d := mustAllow(t, l, "alice")
if d.Allowed {
t.Fatal("the request past the limit was allowed")
}
if d.Remaining != 0 {
t.Errorf("remaining = %d, want 0", d.Remaining)
}
if d.RetryAfter <= 0 || d.RetryAfter > testRule.Window {
t.Errorf("RetryAfter = %v, want a positive interval no longer than the window", d.RetryAfter)
}
}
// A refused request is still counted. Otherwise a caller at their limit could
// keep probing for free, and the limit would be cheaper to test than to respect.
func TestRefusedRequestsStillCount(t *testing.T) {
l, h, _ := newLimiter(t)
for i := 0; i < testRule.Limit+5; i++ {
mustAllow(t, l, "bob")
}
var count int
if err := h.Pool.QueryRow(context.Background(),
`SELECT count FROM rate_limits WHERE bucket LIKE $1`, testRule.Name+":%").Scan(&count); err != nil {
t.Fatalf("read counter: %v", err)
}
if count != testRule.Limit+5 {
t.Errorf("count = %d, want %d — refusals must be counted too", count, testRule.Limit+5)
}
}
/* ── Windows ────────────────────────────────────────────────────────────── */
func TestTheWindowResets(t *testing.T) {
l, _, clock := newLimiter(t)
for i := 0; i < testRule.Limit; i++ {
mustAllow(t, l, "carol")
}
if mustAllow(t, l, "carol").Allowed {
t.Fatal("expected to be at the limit")
}
// Roll into the next window.
*clock = clock.Add(testRule.Window)
l.WithClock(func() time.Time { return *clock })
if d := mustAllow(t, l, "carol"); !d.Allowed {
t.Error("the limit did not reset at the window boundary")
} else if d.Remaining != testRule.Limit-1 {
t.Errorf("remaining = %d, want %d after a reset", d.Remaining, testRule.Limit-1)
}
}
// A new window is a new ROW, not a reset of an existing counter. That is what
// makes two instances rolling over simultaneously safe: neither clobbers the
// other's increments.
func TestANewWindowIsANewRow(t *testing.T) {
l, h, clock := newLimiter(t)
mustAllow(t, l, "dave")
*clock = clock.Add(testRule.Window)
l.WithClock(func() time.Time { return *clock })
mustAllow(t, l, "dave")
var rows int
if err := h.Pool.QueryRow(context.Background(),
`SELECT count(*) FROM rate_limits WHERE bucket LIKE $1`, testRule.Name+":%").Scan(&rows); err != nil {
t.Fatalf("count rows: %v", err)
}
if rows != 2 {
t.Errorf("%d rows, want 2 — each window must be its own row", rows)
}
}
/* ── Buckets are independent ────────────────────────────────────────────── */
func TestSubjectsAreIndependent(t *testing.T) {
l, _, _ := newLimiter(t)
// Exhaust one subject entirely.
for i := 0; i < testRule.Limit+2; i++ {
mustAllow(t, l, "user-a")
}
// A different subject must be untouched.
if d := mustAllow(t, l, "user-b"); !d.Allowed {
t.Error("one subject's limit affected another's")
}
if d := mustAllow(t, l, "org-a|user-a"); !d.Allowed {
t.Error("a compound subject collided with a simple one")
}
}
func TestRulesAreIndependent(t *testing.T) {
l, _, _ := newLimiter(t)
other := Rule{Name: "other.rule", Limit: 3, Window: time.Minute}
for i := 0; i < 5; i++ {
mustAllow(t, l, "shared")
}
d, err := l.Allow(context.Background(), other, Subject("shared"))
if err != nil {
t.Fatalf("Allow: %v", err)
}
if !d.Allowed {
t.Error("exhausting one rule exhausted another for the same subject")
}
}
/* ── No credential becomes a key ────────────────────────────────────────── */
// The property that matters most here: a bucket must never contain the thing it
// identifies. A token in this table is a token in every EXPLAIN, every slow
// query log and every backup.
func TestSubjectsAreHashedNotStored(t *testing.T) {
l, h, _ := newLimiter(t)
const secret = "a-very-secret-bearer-token-value"
mustAllow(t, l, secret)
var found int
if err := h.Pool.QueryRow(context.Background(),
`SELECT count(*) FROM rate_limits WHERE bucket LIKE '%' || $1 || '%'`,
secret).Scan(&found); err != nil {
t.Fatalf("scan: %v", err)
}
if found != 0 {
t.Error("the raw subject appears in the rate_limits table")
}
// And the hash is stable, or a caller would get a fresh budget per request.
if Subject(secret) != Subject(secret) {
t.Error("Subject is not deterministic")
}
if Subject(secret) == secret {
t.Error("Subject returned the raw value")
}
if len(Subject(secret)) != 32 {
t.Errorf("Subject length = %d, want 32", len(Subject(secret)))
}
}
/* ── Concurrency ────────────────────────────────────────────────────────── */
// The claim this package rests on: the count comes back from the statement that
// increments it, so concurrent callers cannot both read the same value and both
// proceed. Run with -race.
func TestConcurrentCallersDoNotLoseIncrements(t *testing.T) {
h := testutil.New(t)
clock := time.Now().Truncate(time.Hour)
l := New(h.Pool).WithClock(func() time.Time { return clock })
const callers = 40
rule := Rule{Name: "concurrent.rule", Limit: 10, Window: time.Minute}
var wg sync.WaitGroup
var mu sync.Mutex
allowed := 0
for i := 0; i < callers; i++ {
wg.Add(1)
go func() {
defer wg.Done()
d, err := l.Allow(context.Background(), rule, Subject("hot-subject"))
if err != nil {
return
}
if d.Allowed {
mu.Lock()
allowed++
mu.Unlock()
}
}()
}
wg.Wait()
// EXACTLY the limit. Not "about" — if increments were lost, more would have
// been allowed; if the statement were not atomic, the count would be wrong
// in either direction.
if allowed != rule.Limit {
t.Errorf("%d of %d concurrent callers allowed, want exactly %d",
allowed, callers, rule.Limit)
}
var count int
if err := h.Pool.QueryRow(context.Background(),
`SELECT count FROM rate_limits WHERE bucket LIKE $1`, rule.Name+":%").Scan(&count); err != nil {
t.Fatalf("read counter: %v", err)
}
if count != callers {
t.Errorf("counter = %d, want %d — increments were lost", count, callers)
}
}
/* ── Failure behaviour ──────────────────────────────────────────────────── */
// A limiter that cannot count must refuse by default. Failing open turns a
// database blip into an unmetered window on the endpoints most worth guarding.
func TestFailsClosedByDefault(t *testing.T) {
l := New(brokenQuerier{}).WithClock(time.Now)
d, err := l.Allow(context.Background(), testRule, Subject("x"))
if err == nil {
t.Fatal("expected an error from a broken database")
}
if d.Allowed {
t.Error("the limiter failed OPEN by default; it must fail closed")
}
if d.RetryAfter <= 0 {
t.Error("a fail-closed decision carries no Retry-After")
}
}
func TestFailOpenIsOptIn(t *testing.T) {
l := New(brokenQuerier{}).WithFailOpen(true)
d, err := l.Allow(context.Background(), testRule, Subject("x"))
if err == nil {
t.Fatal("expected an error")
}
if !d.Allowed {
t.Error("WithFailOpen(true) did not permit the request")
}
}
/* ── Sweep ──────────────────────────────────────────────────────────────── */
func TestSweepRemovesOnlyExpiredWindows(t *testing.T) {
l, h, clock := newLimiter(t)
ctx := context.Background()
mustAllow(t, l, "old")
// Move past the old window, and open a new one.
*clock = clock.Add(2 * testRule.Window)
l.WithClock(func() time.Time { return *clock })
mustAllow(t, l, "current")
removed, err := l.Sweep(ctx, 100)
if err != nil {
t.Fatalf("Sweep: %v", err)
}
if removed != 1 {
t.Errorf("swept %d rows, want 1", removed)
}
// The live window must survive.
var remaining int
_ = h.Pool.QueryRow(ctx, `SELECT count(*) FROM rate_limits`).Scan(&remaining)
if remaining != 1 {
t.Errorf("%d rows left, want 1 — the live window was swept", remaining)
}
// Idempotent: a second sweep removes nothing and does not error.
if again, err := l.Sweep(ctx, 100); err != nil || again != 0 {
t.Errorf("second sweep: removed %d, err %v; want 0, nil", again, err)
}
}
func TestSweepIsBounded(t *testing.T) {
l, h, clock := newLimiter(t)
ctx := context.Background()
for i := 0; i < 10; i++ {
mustAllow(t, l, "subject-"+strings.Repeat("x", i))
}
*clock = clock.Add(2 * testRule.Window)
l.WithClock(func() time.Time { return *clock })
removed, err := l.Sweep(ctx, 4)
if err != nil {
t.Fatalf("Sweep: %v", err)
}
if removed != 4 {
t.Errorf("swept %d, want exactly the batch size 4", removed)
}
var left int
_ = h.Pool.QueryRow(ctx, `SELECT count(*) FROM rate_limits`).Scan(&left)
if left != 6 {
t.Errorf("%d rows left, want 6", left)
}
}
/* ── Helpers ────────────────────────────────────────────────────────────── */
// brokenQuerier fails every call, standing in for an unreachable database.
//
// Satisfies repo.Querier with pgx's real types, so this is the same interface
// the production limiter takes — a hand-rolled stand-in would prove the code
// works against a stand-in.
type brokenQuerier struct{}
func (brokenQuerier) Query(context.Context, string, ...any) (pgx.Rows, error) {
return nil, errBroken
}
func (brokenQuerier) QueryRow(context.Context, string, ...any) pgx.Row { return brokenRow{} }
func (brokenQuerier) Exec(context.Context, string, ...any) (pgconn.CommandTag, error) {
return pgconn.CommandTag{}, errBroken
}
type brokenRow struct{}
func (brokenRow) Scan(...any) error { return errBroken }
var errBroken = errString("ratelimit test: database unavailable")
type errString string
func (e errString) Error() string { return string(e) }
/* ── The fixed-window boundary, measured ────────────────────────────────── */
// The known limitation, asserted rather than assumed.
//
// A fixed window admits up to 2× the limit across a boundary: the limit at the
// end of one window and the limit again at the start of the next. This test
// measures that burst exactly, so the number in the documentation is a fact
// rather than a claim, and so a future change to the algorithm has to
// deliberately update it.
func TestFixedWindowBoundaryBurstIsExactlyTwice(t *testing.T) {
h := testutil.New(t)
// Start just inside a window, so "end of window" is reachable.
clock := time.Now().Truncate(time.Minute).Add(59 * time.Second)
l := New(h.Pool).WithClock(func() time.Time { return clock })
rule := Rule{Name: "boundary.rule", Limit: 5, Window: time.Minute}
allowed := 0
// Spend the whole limit at the very end of window 1.
for i := 0; i < rule.Limit; i++ {
d, err := l.Allow(context.Background(), rule, Subject("boundary"))
if err != nil {
t.Fatalf("Allow: %v", err)
}
if d.Allowed {
allowed++
}
}
// One second later, window 2 begins.
clock = clock.Add(time.Second)
l.WithClock(func() time.Time { return clock })
for i := 0; i < rule.Limit; i++ {
d, err := l.Allow(context.Background(), rule, Subject("boundary"))
if err != nil {
t.Fatalf("Allow: %v", err)
}
if d.Allowed {
allowed++
}
}
// Exactly 2× the limit in just over a second. This is the documented
// worst case — not worse, and not better.
if allowed != rule.Limit*2 {
t.Errorf("%d requests allowed across the boundary, want exactly %d (2× the limit)",
allowed, rule.Limit*2)
}
// And the burst does NOT continue: window 2's budget is now spent.
d, err := l.Allow(context.Background(), rule, Subject("boundary"))
if err != nil {
t.Fatalf("Allow: %v", err)
}
if d.Allowed {
t.Error("the burst continued past 2× the limit")
}
}
// Retry-After must point at the END of the current window, not at a fixed
// duration. A client told to wait a whole window when the window is nearly over
// waits twice as long as it needs to; one told to wait too little retries into
// the same refusal.
func TestRetryAfterPointsAtTheWindowBoundary(t *testing.T) {
h := testutil.New(t)
rule := Rule{Name: "retry.rule", Limit: 1, Window: time.Minute}
for _, offset := range []time.Duration{0, 15 * time.Second, 45 * time.Second, 59 * time.Second} {
clock := time.Now().Truncate(time.Minute).Add(offset)
l := New(h.Pool).WithClock(func() time.Time { return clock })
subject := Subject("retry-" + offset.String())
// Spend the budget, then be refused.
_, _ = l.Allow(context.Background(), rule, subject)
d, err := l.Allow(context.Background(), rule, subject)
if err != nil {
t.Fatalf("Allow: %v", err)
}
if d.Allowed {
t.Fatalf("offset %v: expected a refusal", offset)
}
want := rule.Window - offset
if d.RetryAfter != want {
t.Errorf("offset %v: RetryAfter = %v, want %v (the remainder of the window)",
offset, d.RetryAfter, want)
}
}
}

View File

@@ -0,0 +1,163 @@
package ratelimit
import "time"
// The rule set for the MCP and OAuth surface, in one place.
//
// Every number below is a judgement, so each carries the reasoning that
// produced it. They are starting values: the right way to change one is to
// change it here, with the comment updated, rather than to pass a different
// number at a call site.
//
// TWO PRINCIPLES SHAPE ALL OF THEM
//
// 1. Limit the scarce thing, not the request. Registration writes a row for an
// anonymous caller, so it is limited hard. A tool call reads rows the
// caller may already read through the product, so it is limited loosely —
// the cost there is database load, not access.
//
// 2. Key by the narrowest identity available. An IP is a whole office behind
// NAT; a token is one connection. Limiting an authenticated endpoint by IP
// would make one person's loop everyone's outage.
var (
// Registration: 10 per hour per IP.
//
// The tightest limit here, because /oauth/register is the only endpoint
// that WRITES for a caller with no credential at all — RFC 7591 requires
// exactly that. A legitimate client registers once per installation and
// then never again, so ten is already generous by two orders of magnitude;
// it is set there only so a developer retrying a broken integration does
// not lock themselves out.
//
// Keyed by IP because there is nothing else to key by: the caller is
// anonymous by definition at this point.
OAuthRegister = Rule{Name: "oauth.register", Limit: 10, Window: time.Hour}
// Authorization: 20 per hour per IP+user.
//
// A person clicking Approve does it once. Twenty allows for a browser
// reload, a mistyped password, a client retrying a flow, and a developer
// testing — and stops a script walking the authorization endpoint to farm
// consent pages or probe client ids.
//
// IP AND user, not either alone: keying by user only would let one
// attacker burn an innocent person's budget by naming them, and keying by
// IP only would make an office share one person's allowance.
OAuthAuthorize = Rule{Name: "oauth.authorize", Limit: 20, Window: time.Hour}
// Token exchange: 30 per hour per client.
//
// One exchange per authorization, and an authorization is already limited
// above — so this is not the primary defence. It is here to bound
// brute-forcing a code or a verifier: an authorization code lives 60
// seconds and is single-use, and 30 attempts an hour makes guessing one
// hopeless rather than merely improbable.
OAuthToken = Rule{Name: "oauth.token", Limit: 30, Window: time.Hour}
// Refresh: 60 per hour per token family.
//
// An access token lives 15 minutes, so a well-behaved client refreshes
// about 4 times an hour. Sixty leaves room for a client that refreshes
// eagerly, or one running several sessions, while bounding a loop.
//
// Keyed by FAMILY rather than by token, because the token changes on every
// rotation — keying by token would give each rotation a fresh budget,
// which is the same as no budget at all.
OAuthRefresh = Rule{Name: "oauth.refresh", Limit: 60, Window: time.Hour}
// MCP tool calls: 60 a minute, and 1000 an hour, per token.
//
// BOTH, because they stop different things. The minute limit stops a tight
// loop — a model retrying a failing call, or a bug — from becoming a spike.
// The hour limit stops a slow, sustained drain that would sit under the
// minute limit forever: 59 calls a minute is 3,540 an hour, which is a lot
// of queries for one connection.
//
// Sixty a minute is well above interactive use. A person asking questions
// generates a handful of calls per turn, and a model doing several lookups
// for one answer still lands in single figures.
MCPToolCallPerMinute = Rule{Name: "mcp.call.min", Limit: 60, Window: time.Minute}
MCPToolCallPerHour = Rule{Name: "mcp.call.hour", Limit: 1000, Window: time.Hour}
// Per-organisation ceiling: 5000 an hour.
//
// The backstop for the case the per-token limits cannot see: one tenant
// with many connected clients, each individually well-behaved, together
// saturating the database. Set well above the sum of a few active users so
// it is never reached in ordinary use — it exists to bound a runaway, not
// to ration normal work.
MCPPerOrgPerHour = Rule{Name: "mcp.org.hour", Limit: 5000, Window: time.Hour}
)
// A note on what is NOT rate limited here, and why.
//
// CONCURRENT CONNECTIONS. The plan proposed 10 concurrent MCP connections per
// user. That is not implemented, and it is not an oversight: this transport is
// stateless — one POST per message, no session, nothing held open — so there is
// no such thing as a concurrent connection to count. The thing that limit was
// reaching for is request rate, and the two limits above are that, measured
// directly. Implementing a connection counter over a stateless endpoint would
// mean inventing connection state purely so it could be limited.
//
// DISCOVERY. The two .well-known documents are static, cacheable for five
// minutes, and contain public URLs. Limiting them would add a database write to
// the cheapest endpoints on the surface, to protect nothing.
//
// REVOCATION. Deliberately unlimited. Revocation is the thing a person reaches
// for when something has gone wrong, and an attacker gains nothing by calling
// it — the worst they can do is revoke tokens they already hold. Rate limiting
// the emergency brake is the wrong trade.
/*
FAILURE BEHAVIOUR, RULE BY RULE
===============================
The question this section answers: when the database cannot be reached, does a
request get through?
EVERY RULE HERE FAILS CLOSED. Limiter.failOpen defaults to false and nothing in
this service sets it to true. The reasoning is the same for all of them and is
worth stating once rather than per-rule:
- A limiter that cannot count is not limiting. If a database outage lifted
the limits, then the moment the system is least able to absorb load is
exactly the moment its protections switch off — and an attacker who can
cause or wait for a blip gets an unmetered window on the endpoints that
write rows for anonymous callers.
- The cost of failing closed is bounded and visible: MCP returns 429 and
Claude retries. The cost of failing open is unbounded and silent.
- These endpoints are not load-bearing for the product. If the database is
down, /oauth/token cannot mint a token and /mcp cannot read a row anyway;
the limiter refusing first changes the error message, not the outcome.
WHAT IS EXPLICITLY NOT FAIL-OPEN, AND WHY IT MATTERS MOST
oauth.register Writes a row for a caller with no credential. Failing open
here is an unauthenticated write endpoint with no ceiling.
oauth.token Bounds brute-forcing a code or a verifier. Failing open
turns a 60-second, single-use code into one an attacker may
guess at without limit for the duration of the outage.
oauth.refresh Failing open removes the bound on a loop against a
long-lived credential.
THE ONE PLACE FAIL-OPEN WOULD BE DEFENSIBLE
A deployment that would rather serve MCP degraded than refuse it can call
WithFailOpen(true) on the limiter used for the mcp.* rules only — those guard
database load rather than access, and every call behind them is already
authenticated and already authorized by the policy table. That is a deliberate
operational trade, it is one line, and it is deliberately not the default.
It must NOT be applied to the oauth.* rules. Those guard the credential issuance
path, where the thing being limited is an attacker's number of attempts.
OBSERVABILITY
A limiter failure is logged at ERROR by the middleware (httpserver/mcplimit.go)
with the rule name and the decision, never the subject — the subject is a hash
of a credential. A sustained run of those log lines means the limiter is not
limiting, and is worth an alert.
*/

View File

@@ -409,18 +409,42 @@ func (r *Registry) DispatchApproved(ctx context.Context, tc Context, token strin
}, true
}
// ToolInfo is what a tool looks like to somebody choosing one, rather than to
// the model calling it.
// ToolInfo is what a tool looks like to somebody choosing one, and — since the
// MCP surface — to a client that must publish the schema before calling it.
//
// The InputSchema is deliberately absent: an author picks a capability, and the
// schema is the model's business. Effect is present because it is the one thing
// an author must understand — a write tool means their agent can propose
// changes, which a person will then be asked to approve.
// Effect is present because it is the one thing an author must understand: a
// write tool means their agent can propose changes, which a person will then be
// asked to approve.
//
// InputSchema used to be deliberately absent, on the reasoning that an author
// picks a capability and the schema is the model's business. That is still true
// of the authoring UI, which simply ignores the field. It stopped being true of
// the catalogue as a whole once a second consumer appeared: MCP's tools/list
// must publish a JSON Schema per tool, and the only alternative to carrying
// this one is maintaining a copy. A copy is a second definition of the same
// thing, and the way it fails is silent — a field renamed on the handler side
// leaves a schema that still validates and no longer matches, so the call
// succeeds and the argument is quietly ignored.
//
// Both new fields carry `omitempty`, but be clear about what that does and does
// not buy: every currently registered tool declares a schema and a byte cap, so
// GET /api/v1/tools genuinely does get larger — roughly 700 bytes per tool. The
// change is additive rather than invisible. It is compatible because JSON
// consumers ignore keys they do not read, and the one consumer that exists (the
// agent editor's tool picker) reads name, description and effect; `omitempty`
// covers the remaining case of a tool registered with neither field.
type ToolInfo struct {
Name string `json:"name"`
Description string `json:"description"`
Effect string `json:"effect"`
RequiresConfirmation bool `json:"requiresConfirmation"`
// InputSchema is the tool's JSON Schema, exactly as registered.
InputSchema map[string]any `json:"inputSchema,omitempty"`
// MaxResultBytes is the cap Dispatch truncates at, so a caller can state
// the limit rather than discover it by hitting it.
MaxResultBytes int `json:"maxResultBytes,omitempty"`
}
// Catalogue lists every registered tool, sorted, as choosable metadata.
@@ -440,6 +464,8 @@ func (r *Registry) Catalogue() []ToolInfo {
Description: t.Description,
Effect: string(t.Effect),
RequiresConfirmation: t.RequiresConfirmation,
InputSchema: t.InputSchema,
MaxResultBytes: t.MaxResultBytes,
})
}
sort.Slice(out, func(i, j int) bool { return out[i].Name < out[j].Name })