mcp connection
This commit is contained in:
@@ -65,10 +65,14 @@ func run() error {
|
||||
return err
|
||||
}
|
||||
|
||||
// The sweeper's context is cancelled by the same signal that stops the
|
||||
// server, so the ticker goes away with the process rather than outliving
|
||||
// the pool it queries.
|
||||
// The sweepers' context is cancelled by the same signal that stops the
|
||||
// server, so the tickers go away with the process rather than outliving
|
||||
// the pool they query.
|
||||
go sweepSessions(ctx, server.Sessions(), log)
|
||||
// OAuth codes and tokens, and the rate-limit counters. Returns immediately
|
||||
// when the deployment does not serve MCP, so this line costs an unconfigured
|
||||
// deployment one nil check at startup and nothing after.
|
||||
go httpserver.SweepMaintenance(ctx, server.Maintenance(), log)
|
||||
|
||||
errCh := make(chan error, 1)
|
||||
go func() { errCh <- server.Start() }()
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
134
go-api/internal/config/trustedproxies_test.go
Normal file
134
go-api/internal/config/trustedproxies_test.go
Normal 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
|
||||
}
|
||||
@@ -724,6 +724,15 @@ func TestMigrationPairsAreComplete(t *testing.T) {
|
||||
"000009_confirmation_replay.up.sql",
|
||||
"000010_definition_versions.up.sql",
|
||||
"000011_employee_roles.up.sql",
|
||||
// Phase 3: the OAuth 2.1 authorization server behind the MCP surface.
|
||||
// Three tables, added together because they are one feature: a client
|
||||
// registers, is issued a code, and exchanges it for tokens.
|
||||
"000012_oauth_clients.up.sql",
|
||||
"000013_oauth_grants.up.sql",
|
||||
"000014_oauth_tokens.up.sql",
|
||||
// Phase 5: shared rate limit counters, so a limit means the same thing
|
||||
// behind one instance and behind ten.
|
||||
"000015_rate_limits.up.sql",
|
||||
}
|
||||
if len(ups) != len(want) {
|
||||
t.Fatalf("%d migrations, want %d — update this list deliberately", len(ups), len(want))
|
||||
@@ -753,11 +762,28 @@ func TestMigrationsAddOnlyTheTablesWeDecidedOn(t *testing.T) {
|
||||
// 17 from 000001, + auth_sessions (000004), + agent_definitions and
|
||||
// skill_definitions (000005), + agent_runs (000006), + agent_confirmations
|
||||
// (000007), + knowledge_documents and knowledge_chunks (000008),
|
||||
// + definition_versions (000010), + employee_roles (000011).
|
||||
// + definition_versions (000010), + employee_roles (000011),
|
||||
// + oauth_clients (000012), + oauth_grants (000013), + oauth_tokens
|
||||
// (000014), + rate_limits (000015).
|
||||
// schema_migrations is golang-migrate's and is absent when the files are
|
||||
// applied directly.
|
||||
if n != 26 {
|
||||
t.Errorf("%d base tables after every migration, want 26", n)
|
||||
if n != 30 {
|
||||
t.Errorf("%d base tables after every migration, want 30", n)
|
||||
}
|
||||
|
||||
// The three OAuth tables, named rather than merely counted. The count
|
||||
// above catches a table arriving without a decision; this catches one of
|
||||
// these three going missing, which the count alone would not if another
|
||||
// arrived in the same change.
|
||||
for _, required := range []string{"oauth_clients", "oauth_grants", "oauth_tokens", "rate_limits"} {
|
||||
var reg *string
|
||||
if err := f.pool.QueryRow(f.ctx,
|
||||
`SELECT to_regclass('public.' || $1)::text`, required).Scan(®); err != nil {
|
||||
t.Fatalf("check %s: %v", required, err)
|
||||
}
|
||||
if reg == nil {
|
||||
t.Errorf("%s is missing; the MCP OAuth surface cannot work without it", required)
|
||||
}
|
||||
}
|
||||
|
||||
// `definition_versions` was on this list, deferred by the Phase 4B decision.
|
||||
|
||||
@@ -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.
|
||||
|
||||
217
go-api/internal/httpserver/clientip.go
Normal file
217
go-api/internal/httpserver/clientip.go
Normal 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
|
||||
}
|
||||
358
go-api/internal/httpserver/clientip_test.go
Normal file
358
go-api/internal/httpserver/clientip_test.go
Normal 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)
|
||||
}
|
||||
}
|
||||
193
go-api/internal/httpserver/maintenance.go
Normal file
193
go-api/internal/httpserver/maintenance.go
Normal 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()
|
||||
}
|
||||
}
|
||||
}
|
||||
275
go-api/internal/httpserver/maintenance_test.go
Normal file
275
go-api/internal/httpserver/maintenance_test.go
Normal 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")
|
||||
}
|
||||
}
|
||||
195
go-api/internal/httpserver/mcp.go
Normal file
195
go-api/internal/httpserver/mcp.go
Normal 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)
|
||||
785
go-api/internal/httpserver/mcp_routes_test.go
Normal file
785
go-api/internal/httpserver/mcp_routes_test.go
Normal 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, ®Doc)
|
||||
|
||||
verifier := "helperVerifier0123456789abcdefghijklmnopqrst"
|
||||
q := url.Values{
|
||||
"client_id": {regDoc.ClientID}, "redirect_uri": {"https://client.example.test/cb"},
|
||||
"response_type": {"code"}, "state": {"helper"},
|
||||
"code_challenge": {challengeFor(verifier)}, "code_challenge_method": {"S256"},
|
||||
"resource": {testMCPResource}, "scope": {"krow.read"},
|
||||
}
|
||||
consent := a.do("GET", "/oauth/authorize?"+q.Encode(), nil)
|
||||
csrf := between(consent.body, `name="csrf" value="`, `"`)
|
||||
|
||||
form := url.Values{}
|
||||
for k, v := range q {
|
||||
form[k] = v
|
||||
}
|
||||
form.Set("decision", "approve")
|
||||
form.Set("csrf", csrf)
|
||||
approved := a.doForm("POST", "/oauth/authorize", form)
|
||||
loc, _ := url.Parse(approved.header.Get("Location"))
|
||||
|
||||
tok := a.doAnonForm("POST", "/oauth/token", url.Values{
|
||||
"grant_type": {"authorization_code"}, "code": {loc.Query().Get("code")},
|
||||
"client_id": {regDoc.ClientID}, "redirect_uri": {"https://client.example.test/cb"},
|
||||
"code_verifier": {verifier},
|
||||
})
|
||||
var tokens struct {
|
||||
AccessToken string `json:"access_token"`
|
||||
}
|
||||
mustJSON(t, tok.body, &tokens)
|
||||
if tokens.AccessToken == "" {
|
||||
t.Fatalf("could not obtain a token: %s", tok.body)
|
||||
}
|
||||
return tokens.AccessToken
|
||||
}
|
||||
|
||||
// challengeFor derives an S256 challenge, so these tests do not depend on the
|
||||
// oauth package's unexported helpers.
|
||||
func challengeFor(verifier string) string {
|
||||
sum := sha256.Sum256([]byte(verifier))
|
||||
return base64.RawURLEncoding.EncodeToString(sum[:])
|
||||
}
|
||||
|
||||
// The mounted surface, end to end.
|
||||
//
|
||||
// Everything below drives the REAL router — the same mux, the same
|
||||
// authenticate() middleware, the same publicPaths allowlist that serves
|
||||
// production. The point is not to re-test the OAuth package (internal/oauth
|
||||
// does that against its own handlers) but to prove the MOUNTING is right: that
|
||||
// discovery is reachable without a cookie, that /mcp is not, that a cookie
|
||||
// cannot substitute for a bearer token, and that the routes appear at all only
|
||||
// when the deployment is configured for them.
|
||||
|
||||
/* ── Route registration is conditional ──────────────────────────────────── */
|
||||
|
||||
// Without OAUTH_ISSUER and MCP_RESOURCE, none of this exists. An upgrade must
|
||||
// not quietly add an authorization server to a deployment that never asked.
|
||||
func TestMCPRoutesAreAbsentWhenUnconfigured(t *testing.T) {
|
||||
a := newAPI(t) // the standard fixture: no OAuth configuration
|
||||
|
||||
for _, path := range []string{
|
||||
"/mcp",
|
||||
"/oauth/register",
|
||||
"/oauth/authorize",
|
||||
"/oauth/token",
|
||||
"/.well-known/oauth-protected-resource",
|
||||
"/.well-known/oauth-authorization-server",
|
||||
} {
|
||||
r := a.doAnon("POST", path, nil)
|
||||
if r.code != http.StatusNotFound && r.code != http.StatusUnauthorized {
|
||||
t.Errorf("%s = %d on an unconfigured deployment; want 404 or 401, never a served response",
|
||||
path, r.code)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/* ── Discovery is public ────────────────────────────────────────────────── */
|
||||
|
||||
// A client with no token must be able to read both documents, or it can never
|
||||
// discover how to get one.
|
||||
func TestDiscoveryIsReachableWithoutASession(t *testing.T) {
|
||||
a := newOAuthAPI(t)
|
||||
|
||||
t.Run("protected resource", func(t *testing.T) {
|
||||
r := a.doAnon("GET", "/.well-known/oauth-protected-resource", nil)
|
||||
if r.code != http.StatusOK {
|
||||
t.Fatalf("status = %d, want 200 without a cookie: %s", r.code, r.body)
|
||||
}
|
||||
var doc struct {
|
||||
Resource string `json:"resource"`
|
||||
AuthorizationServers []string `json:"authorization_servers"`
|
||||
BearerMethods []string `json:"bearer_methods_supported"`
|
||||
}
|
||||
mustJSON(t, r.body, &doc)
|
||||
|
||||
if doc.Resource != testMCPResource {
|
||||
t.Errorf("resource = %q, want %q", doc.Resource, testMCPResource)
|
||||
}
|
||||
if len(doc.AuthorizationServers) != 1 || doc.AuthorizationServers[0] != testOAuthIssuer {
|
||||
t.Errorf("authorization_servers = %v, want [%q]", doc.AuthorizationServers, testOAuthIssuer)
|
||||
}
|
||||
// The MCP spec forbids a token in the query string.
|
||||
if strings.Join(doc.BearerMethods, ",") != "header" {
|
||||
t.Errorf("bearer_methods_supported = %v, want [header]", doc.BearerMethods)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("authorization server", func(t *testing.T) {
|
||||
r := a.doAnon("GET", "/.well-known/oauth-authorization-server", nil)
|
||||
if r.code != http.StatusOK {
|
||||
t.Fatalf("status = %d, want 200 without a cookie: %s", r.code, r.body)
|
||||
}
|
||||
var doc struct {
|
||||
Issuer string `json:"issuer"`
|
||||
AuthorizationEndpoint string `json:"authorization_endpoint"`
|
||||
TokenEndpoint string `json:"token_endpoint"`
|
||||
RegistrationEndpoint string `json:"registration_endpoint"`
|
||||
Scopes []string `json:"scopes_supported"`
|
||||
ResponseTypes []string `json:"response_types_supported"`
|
||||
GrantTypes []string `json:"grant_types_supported"`
|
||||
PKCEMethods []string `json:"code_challenge_methods_supported"`
|
||||
ResourceIndicators bool `json:"resource_indicators_supported"`
|
||||
}
|
||||
mustJSON(t, r.body, &doc)
|
||||
|
||||
// EVERY url must come from configuration. A hardcoded hostname would
|
||||
// be one deployment's identity baked into every other one.
|
||||
if doc.Issuer != testOAuthIssuer {
|
||||
t.Errorf("issuer = %q, want %q", doc.Issuer, testOAuthIssuer)
|
||||
}
|
||||
for name, got := range map[string]string{
|
||||
"authorization_endpoint": doc.AuthorizationEndpoint,
|
||||
"token_endpoint": doc.TokenEndpoint,
|
||||
"registration_endpoint": doc.RegistrationEndpoint,
|
||||
} {
|
||||
if !strings.HasPrefix(got, testOAuthIssuer) {
|
||||
t.Errorf("%s = %q, want it under the configured issuer", name, got)
|
||||
}
|
||||
}
|
||||
if strings.Join(doc.ResponseTypes, ",") != "code" {
|
||||
t.Errorf("response_types_supported = %v; implicit must not be advertised", doc.ResponseTypes)
|
||||
}
|
||||
if strings.Join(doc.PKCEMethods, ",") != "S256" {
|
||||
t.Errorf("code_challenge_methods_supported = %v, want [S256]", doc.PKCEMethods)
|
||||
}
|
||||
for _, forbidden := range []string{"password", "client_credentials", "implicit"} {
|
||||
for _, advertised := range doc.GrantTypes {
|
||||
if advertised == forbidden {
|
||||
t.Errorf("grant_types_supported advertises %q", forbidden)
|
||||
}
|
||||
}
|
||||
}
|
||||
for _, s := range doc.Scopes {
|
||||
if s == "krow.write" {
|
||||
t.Error("scopes_supported advertises krow.write")
|
||||
}
|
||||
}
|
||||
if !doc.ResourceIndicators {
|
||||
t.Error("resource_indicators_supported must be true")
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
/* ── /mcp authentication ────────────────────────────────────────────────── */
|
||||
|
||||
// No bearer → 401 with a challenge that tells the client where to go.
|
||||
func TestMCPWithoutBearerReturns401AndDiscoveryPointer(t *testing.T) {
|
||||
a := newOAuthAPI(t)
|
||||
r := a.doAnon("POST", "/mcp", map[string]any{
|
||||
"jsonrpc": "2.0", "id": 1, "method": "tools/list",
|
||||
})
|
||||
|
||||
if r.code != http.StatusUnauthorized {
|
||||
t.Fatalf("status = %d, want 401", r.code)
|
||||
}
|
||||
challenge := r.header.Get("WWW-Authenticate")
|
||||
if !strings.HasPrefix(challenge, "Bearer") {
|
||||
t.Fatalf("WWW-Authenticate = %q, want a Bearer challenge", challenge)
|
||||
}
|
||||
// RFC 9728: without resource_metadata the client has a 401 and nowhere to
|
||||
// look. This is the difference between "failed" and "here is how".
|
||||
if !strings.Contains(challenge, `resource_metadata="`+testOAuthIssuer) {
|
||||
t.Errorf("WWW-Authenticate = %q, want resource_metadata built from the configured issuer", challenge)
|
||||
}
|
||||
// And it must be built from config, not baked in.
|
||||
if strings.Contains(challenge, "krowforce.com") {
|
||||
t.Errorf("WWW-Authenticate contains a hardcoded production hostname: %q", challenge)
|
||||
}
|
||||
}
|
||||
|
||||
// THE test for this phase's riskiest decision: a perfectly valid KROW session
|
||||
// cookie must not open the MCP endpoint.
|
||||
func TestMCPRejectsACookieSession(t *testing.T) {
|
||||
a := newOAuthAPI(t)
|
||||
|
||||
// `a.do` sends the authenticated session cookie the rest of the suite uses.
|
||||
r := a.do("POST", "/mcp", map[string]any{
|
||||
"jsonrpc": "2.0", "id": 1, "method": "tools/list",
|
||||
})
|
||||
|
||||
if r.code != http.StatusUnauthorized {
|
||||
t.Fatalf("status = %d, want 401 — a browser cookie authenticated an MCP call", r.code)
|
||||
}
|
||||
}
|
||||
|
||||
func TestMCPRejectsAnInvalidBearer(t *testing.T) {
|
||||
a := newOAuthAPI(t)
|
||||
for name, header := range map[string]string{
|
||||
"unknown token": "Bearer not-a-real-token",
|
||||
"empty": "Bearer ",
|
||||
"wrong scheme": "Basic dXNlcjpwYXNz",
|
||||
"no scheme": "abcdef",
|
||||
} {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
r := a.doAnonWithHeader("POST", "/mcp", map[string]any{
|
||||
"jsonrpc": "2.0", "id": 1, "method": "tools/list",
|
||||
}, "Authorization", header)
|
||||
if r.code != http.StatusUnauthorized {
|
||||
t.Errorf("status = %d, want 401", r.code)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// A token must never be accepted from the query string. The MCP spec forbids
|
||||
// it, and a URL is logged, cached and put in a Referer.
|
||||
func TestMCPIgnoresATokenInTheQueryString(t *testing.T) {
|
||||
a := newOAuthAPI(t)
|
||||
token := a.oauthAccessToken(t)
|
||||
|
||||
r := a.doAnon("POST", "/mcp?access_token="+url.QueryEscape(token), map[string]any{
|
||||
"jsonrpc": "2.0", "id": 1, "method": "tools/list",
|
||||
})
|
||||
if r.code != http.StatusUnauthorized {
|
||||
t.Errorf("status = %d, want 401 — a query-string token was accepted", r.code)
|
||||
}
|
||||
}
|
||||
|
||||
// Custom identity headers must be ignored outright.
|
||||
func TestMCPIgnoresCustomIdentityHeaders(t *testing.T) {
|
||||
a := newOAuthAPI(t)
|
||||
for _, header := range []string{"X-Access-Token", "X-Api-Key", "X-Org-Id", "X-User-Id", "X-Krow-Token"} {
|
||||
r := a.doAnonWithHeader("POST", "/mcp", map[string]any{
|
||||
"jsonrpc": "2.0", "id": 1, "method": "tools/list",
|
||||
}, header, a.oauthAccessToken(t))
|
||||
if r.code != http.StatusUnauthorized {
|
||||
t.Errorf("%s was accepted as a credential: %d", header, r.code)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/* ── The full discovery → consent → token → MCP journey ─────────────────── */
|
||||
|
||||
// Every step a Claude client performs, over the real router, in order.
|
||||
func TestFullMCPConnectionJourney(t *testing.T) {
|
||||
a := newOAuthAPI(t)
|
||||
|
||||
// 1–2. Call /mcp with no token; get 401 and a pointer.
|
||||
unauth := a.doAnon("POST", "/mcp", map[string]any{
|
||||
"jsonrpc": "2.0", "id": 1, "method": "initialize",
|
||||
})
|
||||
if unauth.code != http.StatusUnauthorized {
|
||||
t.Fatalf("step 1: status = %d, want 401", unauth.code)
|
||||
}
|
||||
challenge := unauth.header.Get("WWW-Authenticate")
|
||||
|
||||
// 3. Follow resource_metadata to the protected-resource document.
|
||||
metaURL := between(challenge, `resource_metadata="`, `"`)
|
||||
if metaURL == "" {
|
||||
t.Fatal("step 3: the challenge carries no resource_metadata")
|
||||
}
|
||||
prPath := strings.TrimPrefix(metaURL, testOAuthIssuer)
|
||||
pr := a.doAnon("GET", prPath, nil)
|
||||
if pr.code != http.StatusOK {
|
||||
t.Fatalf("step 3: %s = %d", prPath, pr.code)
|
||||
}
|
||||
var prDoc struct {
|
||||
AuthorizationServers []string `json:"authorization_servers"`
|
||||
}
|
||||
mustJSON(t, pr.body, &prDoc)
|
||||
|
||||
// 4. Authorization-server metadata.
|
||||
as := a.doAnon("GET", "/.well-known/oauth-authorization-server", nil)
|
||||
if as.code != http.StatusOK {
|
||||
t.Fatalf("step 4: status = %d", as.code)
|
||||
}
|
||||
var asDoc struct {
|
||||
AuthorizationEndpoint string `json:"authorization_endpoint"`
|
||||
TokenEndpoint string `json:"token_endpoint"`
|
||||
RegistrationEndpoint string `json:"registration_endpoint"`
|
||||
}
|
||||
mustJSON(t, as.body, &asDoc)
|
||||
|
||||
// 5. Register, at the advertised endpoint.
|
||||
reg := a.doAnon("POST", strings.TrimPrefix(asDoc.RegistrationEndpoint, testOAuthIssuer), map[string]any{
|
||||
"client_name": "Journey Client",
|
||||
"redirect_uris": []string{"https://client.example.test/cb"},
|
||||
})
|
||||
if reg.code != http.StatusCreated {
|
||||
t.Fatalf("step 5: registration = %d %s", reg.code, reg.body)
|
||||
}
|
||||
var regDoc struct {
|
||||
ClientID string `json:"client_id"`
|
||||
}
|
||||
mustJSON(t, reg.body, ®Doc)
|
||||
|
||||
// 6–7. Authorize, SIGNED IN. A cookie is exactly right here: this step is
|
||||
// a person in a browser.
|
||||
verifier := "journeyVerifier0123456789abcdefghijklmnopqrs"
|
||||
q := url.Values{
|
||||
"client_id": {regDoc.ClientID}, "redirect_uri": {"https://client.example.test/cb"},
|
||||
"response_type": {"code"}, "state": {"journey-state"},
|
||||
"code_challenge": {challengeFor(verifier)}, "code_challenge_method": {"S256"},
|
||||
"resource": {testMCPResource}, "scope": {"krow.read"},
|
||||
}
|
||||
consent := a.do("GET", "/oauth/authorize?"+q.Encode(), nil)
|
||||
|
||||
// 8. A consent page, not a code.
|
||||
if consent.code != http.StatusOK {
|
||||
t.Fatalf("step 8: expected a consent page, got %d %s", consent.code, consent.body)
|
||||
}
|
||||
if !strings.Contains(consent.body, "Journey Client") {
|
||||
t.Error("step 8: the consent page does not name the requesting client")
|
||||
}
|
||||
csrf := between(consent.body, `name="csrf" value="`, `"`)
|
||||
if csrf == "" {
|
||||
t.Fatal("step 8: no csrf token in the consent form")
|
||||
}
|
||||
|
||||
// 9–10. Approve; receive a code.
|
||||
form := url.Values{}
|
||||
for k, v := range q {
|
||||
form[k] = v
|
||||
}
|
||||
form.Set("decision", "approve")
|
||||
form.Set("csrf", csrf)
|
||||
approved := a.doForm("POST", "/oauth/authorize", form)
|
||||
if approved.code != http.StatusFound {
|
||||
t.Fatalf("step 10: approve = %d %s", approved.code, approved.body)
|
||||
}
|
||||
loc, _ := url.Parse(approved.header.Get("Location"))
|
||||
code := loc.Query().Get("code")
|
||||
if code == "" {
|
||||
t.Fatalf("step 10: no code in %s", loc)
|
||||
}
|
||||
if loc.Query().Get("state") != "journey-state" {
|
||||
t.Errorf("step 10: state = %q", loc.Query().Get("state"))
|
||||
}
|
||||
|
||||
// 11. Exchange — with NO cookie, as a back-channel call.
|
||||
tok := a.doAnonForm("POST", strings.TrimPrefix(asDoc.TokenEndpoint, testOAuthIssuer), url.Values{
|
||||
"grant_type": {"authorization_code"}, "code": {code},
|
||||
"client_id": {regDoc.ClientID}, "redirect_uri": {"https://client.example.test/cb"},
|
||||
"code_verifier": {verifier},
|
||||
})
|
||||
if tok.code != http.StatusOK {
|
||||
t.Fatalf("step 11: token = %d %s", tok.code, tok.body)
|
||||
}
|
||||
var tokens struct {
|
||||
AccessToken string `json:"access_token"`
|
||||
TokenType string `json:"token_type"`
|
||||
}
|
||||
mustJSON(t, tok.body, &tokens)
|
||||
if tokens.AccessToken == "" || tokens.TokenType != "Bearer" {
|
||||
t.Fatalf("step 11: unusable token response: %s", tok.body)
|
||||
}
|
||||
|
||||
// 12–13. tools/list with the bearer token.
|
||||
list := a.doAnonWithHeader("POST", "/mcp", map[string]any{
|
||||
"jsonrpc": "2.0", "id": 2, "method": "tools/list",
|
||||
}, "Authorization", "Bearer "+tokens.AccessToken)
|
||||
if list.code != http.StatusOK {
|
||||
t.Fatalf("step 13: tools/list = %d %s", list.code, list.body)
|
||||
}
|
||||
var listDoc struct {
|
||||
Result struct {
|
||||
Tools []struct {
|
||||
Name string `json:"name"`
|
||||
} `json:"tools"`
|
||||
} `json:"result"`
|
||||
}
|
||||
mustJSON(t, list.body, &listDoc)
|
||||
if len(listDoc.Result.Tools) != 16 {
|
||||
t.Errorf("step 13: %d tools, want 16", len(listDoc.Result.Tools))
|
||||
}
|
||||
for _, tool := range listDoc.Result.Tools {
|
||||
switch tool.Name {
|
||||
case "assign_worker", "move_application", "knowledge_search":
|
||||
t.Errorf("step 13: %q is exposed over the mounted route", tool.Name)
|
||||
}
|
||||
}
|
||||
|
||||
// 14. tools/call reaches the existing authorization and real data.
|
||||
call := a.doAnonWithHeader("POST", "/mcp", map[string]any{
|
||||
"jsonrpc": "2.0", "id": 3, "method": "tools/call",
|
||||
"params": map[string]any{"name": "workspace_summary", "arguments": map[string]any{}},
|
||||
}, "Authorization", "Bearer "+tokens.AccessToken)
|
||||
if call.code != http.StatusOK {
|
||||
t.Fatalf("step 14: tools/call = %d %s", call.code, call.body)
|
||||
}
|
||||
var callDoc struct {
|
||||
Result struct {
|
||||
IsError bool `json:"isError"`
|
||||
Content []struct {
|
||||
Text string `json:"text"`
|
||||
} `json:"content"`
|
||||
} `json:"result"`
|
||||
}
|
||||
mustJSON(t, call.body, &callDoc)
|
||||
if callDoc.Result.IsError {
|
||||
t.Fatalf("step 14: the tool refused: %s", callDoc.Result.Content[0].Text)
|
||||
}
|
||||
}
|
||||
|
||||
// Denial must reach the client correctly and issue nothing.
|
||||
func TestConsentDenialOverTheMountedRoute(t *testing.T) {
|
||||
a := newOAuthAPI(t)
|
||||
|
||||
reg := a.doAnon("POST", "/oauth/register", map[string]any{
|
||||
"client_name": "Deny Client", "redirect_uris": []string{"https://client.example.test/cb"},
|
||||
})
|
||||
var regDoc struct {
|
||||
ClientID string `json:"client_id"`
|
||||
}
|
||||
mustJSON(t, reg.body, ®Doc)
|
||||
|
||||
verifier := "denyVerifier0123456789abcdefghijklmnopqrstuv"
|
||||
q := url.Values{
|
||||
"client_id": {regDoc.ClientID}, "redirect_uri": {"https://client.example.test/cb"},
|
||||
"response_type": {"code"}, "state": {"deny-state"},
|
||||
"code_challenge": {challengeFor(verifier)}, "code_challenge_method": {"S256"},
|
||||
"resource": {testMCPResource}, "scope": {"krow.read"},
|
||||
}
|
||||
consent := a.do("GET", "/oauth/authorize?"+q.Encode(), nil)
|
||||
csrf := between(consent.body, `name="csrf" value="`, `"`)
|
||||
|
||||
form := url.Values{}
|
||||
for k, v := range q {
|
||||
form[k] = v
|
||||
}
|
||||
form.Set("decision", "deny")
|
||||
form.Set("csrf", csrf)
|
||||
|
||||
denied := a.doForm("POST", "/oauth/authorize", form)
|
||||
if denied.code != http.StatusFound {
|
||||
t.Fatalf("status = %d, want 302", denied.code)
|
||||
}
|
||||
loc, _ := url.Parse(denied.header.Get("Location"))
|
||||
if got := loc.Query().Get("error"); got != "access_denied" {
|
||||
t.Errorf("error = %q, want access_denied", got)
|
||||
}
|
||||
if got := loc.Query().Get("state"); got != "deny-state" {
|
||||
t.Errorf("state = %q, want deny-state", got)
|
||||
}
|
||||
if loc.Query().Get("code") != "" {
|
||||
t.Error("a denial issued a code")
|
||||
}
|
||||
}
|
||||
|
||||
// /oauth/authorize is NOT public: an anonymous visitor must be sent to login.
|
||||
func TestAuthorizeRequiresASession(t *testing.T) {
|
||||
a := newOAuthAPI(t)
|
||||
r := a.doAnon("GET", "/oauth/authorize?client_id=x", nil)
|
||||
|
||||
// Either the middleware refuses it (401) or the handler redirects to
|
||||
// login. Both are correct; serving a consent page is not.
|
||||
if r.code == http.StatusOK && strings.Contains(r.body, "Approve") {
|
||||
t.Fatal("a consent page was served to an anonymous visitor")
|
||||
}
|
||||
}
|
||||
|
||||
/* ── Helpers ────────────────────────────────────────────────────────────── */
|
||||
|
||||
func mustJSON(t *testing.T, body string, dst any) {
|
||||
t.Helper()
|
||||
if err := json.Unmarshal([]byte(body), dst); err != nil {
|
||||
t.Fatalf("response was not JSON: %v\nbody: %s", err, body)
|
||||
}
|
||||
}
|
||||
|
||||
// between returns the text between two markers, or "".
|
||||
func between(s, start, end string) string {
|
||||
i := strings.Index(s, start)
|
||||
if i < 0 {
|
||||
return ""
|
||||
}
|
||||
rest := s[i+len(start):]
|
||||
j := strings.Index(rest, end)
|
||||
if j < 0 {
|
||||
return ""
|
||||
}
|
||||
return rest[:j]
|
||||
}
|
||||
|
||||
/* ── Anonymous /oauth/authorize must reach the handler ──────────────────── */
|
||||
|
||||
// The regression test for the defect a live Claude Web connection exposed.
|
||||
//
|
||||
// /oauth/authorize was withheld from publicPaths, so the cookie middleware
|
||||
// answered a signed-out visitor with its JSON 401 and the handler never ran —
|
||||
// which meant the handler's redirect-to-login could never execute. A first-time
|
||||
// connector user is signed out by definition, so OAuth's browser leg was
|
||||
// unreachable for precisely the people who needed it.
|
||||
//
|
||||
// WHY THE EXISTING TESTS MISSED IT, and why this one is shaped differently:
|
||||
//
|
||||
// - oauth.TestAuthorizeRedirectsAnonymousToLogin drives AuthorizeHandler
|
||||
// DIRECTLY, so the middleware is not in the path at all. It passed against
|
||||
// broken behaviour because it never exercised the thing that was broken.
|
||||
// - TestAuthorizeRequiresASession (below) asserts only that a consent page is
|
||||
// not served anonymously — which a 401 satisfies perfectly well.
|
||||
//
|
||||
// So this one drives the MOUNTED router and asserts the POSITIVE behaviour: a
|
||||
// redirect to the login, carrying the original authorization request.
|
||||
func TestAnonymousAuthorizeReachesTheHandlerAndRedirectsToLogin(t *testing.T) {
|
||||
a := newOAuthAPI(t)
|
||||
|
||||
// A client to name, so the request is well-formed enough to get past the
|
||||
// handler's own client/redirect validation and reach the session check.
|
||||
reg := a.doAnon("POST", "/oauth/register", map[string]any{
|
||||
"client_name": "Anonymous Flow", "redirect_uris": []string{"https://client.example.test/cb"},
|
||||
})
|
||||
var regDoc struct {
|
||||
ClientID string `json:"client_id"`
|
||||
}
|
||||
mustJSON(t, reg.body, ®Doc)
|
||||
|
||||
verifier := "anonVerifier0123456789abcdefghijklmnopqrstu"
|
||||
q := url.Values{
|
||||
"client_id": {regDoc.ClientID}, "redirect_uri": {"https://client.example.test/cb"},
|
||||
"response_type": {"code"}, "state": {"anon-state"},
|
||||
"code_challenge": {challengeFor(verifier)}, "code_challenge_method": {"S256"},
|
||||
"resource": {testMCPResource}, "scope": {"krow.read"},
|
||||
}
|
||||
|
||||
// doAnon sends NO session cookie — a first-time connector user.
|
||||
r := a.doAnon("GET", "/oauth/authorize?"+q.Encode(), nil)
|
||||
|
||||
// The defect: the middleware's JSON 401 instead of the handler's redirect.
|
||||
if r.code == http.StatusUnauthorized {
|
||||
t.Fatalf("the middleware refused before the handler ran: %d %s\n"+
|
||||
"a signed-out visitor must be sent to sign in, not told 'no'", r.code, r.body)
|
||||
}
|
||||
if strings.Contains(r.body, `"code": "unauthorized"`) ||
|
||||
strings.Contains(r.body, `"code":"unauthorized"`) {
|
||||
t.Fatalf("the response is the middleware's JSON 401, not the handler's: %s", r.body)
|
||||
}
|
||||
|
||||
if r.code != http.StatusFound {
|
||||
t.Fatalf("status = %d, want 302 to the login", r.code)
|
||||
}
|
||||
location := r.header.Get("Location")
|
||||
if !strings.HasPrefix(location, "/login?returnTo=") {
|
||||
t.Fatalf("Location = %q, want a redirect to the configured login path", location)
|
||||
}
|
||||
|
||||
// The whole authorization request must survive the round trip, or the
|
||||
// person signs in and lands nowhere.
|
||||
returnTo, err := url.QueryUnescape(strings.TrimPrefix(location, "/login?returnTo="))
|
||||
if err != nil {
|
||||
t.Fatalf("returnTo is not decodable: %v", err)
|
||||
}
|
||||
for name, want := range map[string]string{
|
||||
"path": "/oauth/authorize",
|
||||
"client_id": "client_id=" + regDoc.ClientID,
|
||||
"state": "state=anon-state",
|
||||
"code_challenge": "code_challenge=" + challengeFor(verifier),
|
||||
"code_challenge_method": "code_challenge_method=S256",
|
||||
"resource": "resource=",
|
||||
"redirect_uri": "redirect_uri=",
|
||||
} {
|
||||
if !strings.Contains(returnTo, want) {
|
||||
t.Errorf("returnTo has lost the %s: %q", name, returnTo)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Listing the path must NOT hand out consent, or a code, to somebody signed
|
||||
// out. "Public" here means the handler decides — not that the route is open.
|
||||
func TestAnonymousAuthorizeStillGrantsNothing(t *testing.T) {
|
||||
a := newOAuthAPI(t)
|
||||
|
||||
reg := a.doAnon("POST", "/oauth/register", map[string]any{
|
||||
"client_name": "Nothing Granted", "redirect_uris": []string{"https://client.example.test/cb"},
|
||||
})
|
||||
var regDoc struct {
|
||||
ClientID string `json:"client_id"`
|
||||
}
|
||||
mustJSON(t, reg.body, ®Doc)
|
||||
|
||||
verifier := "nothingVerifier0123456789abcdefghijklmnopq"
|
||||
q := url.Values{
|
||||
"client_id": {regDoc.ClientID}, "redirect_uri": {"https://client.example.test/cb"},
|
||||
"response_type": {"code"}, "state": {"nothing"},
|
||||
"code_challenge": {challengeFor(verifier)}, "code_challenge_method": {"S256"},
|
||||
"resource": {testMCPResource}, "scope": {"krow.read"},
|
||||
}
|
||||
|
||||
// A GET must not render consent.
|
||||
get := a.doAnon("GET", "/oauth/authorize?"+q.Encode(), nil)
|
||||
if strings.Contains(get.body, "Approve") || strings.Contains(get.body, "Authorize access to Krow") {
|
||||
t.Error("a consent page was served to a signed-out visitor")
|
||||
}
|
||||
|
||||
// And a POST — skipping the page entirely, as an attacker would — must not
|
||||
// issue a code. The handler's session check refuses before the CSRF check
|
||||
// is even relevant.
|
||||
form := url.Values{}
|
||||
for k, v := range q {
|
||||
form[k] = v
|
||||
}
|
||||
form.Set("decision", "approve")
|
||||
form.Set("csrf", "forged")
|
||||
|
||||
post := a.doAnonForm("POST", "/oauth/authorize", form)
|
||||
if loc := post.header.Get("Location"); strings.Contains(loc, "code=") {
|
||||
t.Fatalf("an anonymous POST obtained an authorization code: %s", loc)
|
||||
}
|
||||
if post.code == http.StatusFound && strings.HasPrefix(post.header.Get("Location"), "https://client.example.test") {
|
||||
t.Fatalf("an anonymous POST reached the client callback: %s", post.header.Get("Location"))
|
||||
}
|
||||
}
|
||||
230
go-api/internal/httpserver/mcplimit.go
Normal file
230
go-api/internal/httpserver/mcplimit.go
Normal 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
|
||||
}
|
||||
294
go-api/internal/httpserver/proxybuckets_test.go
Normal file
294
go-api/internal/httpserver/proxybuckets_test.go
Normal 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)
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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
|
||||
|
||||
193
go-api/internal/mcpserver/auth.go
Normal file
193
go-api/internal/mcpserver/auth.go
Normal file
@@ -0,0 +1,193 @@
|
||||
package mcpserver
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"net/http"
|
||||
"strings"
|
||||
|
||||
"github.com/krow/krow-backend/go-api/internal/auth"
|
||||
"github.com/krow/krow-backend/go-api/internal/authctx"
|
||||
)
|
||||
|
||||
// Bearer authentication for the MCP surface.
|
||||
//
|
||||
// This file is a SEAM, not an authentication system. It defines the one
|
||||
// question the MCP transport needs answered — "which KROW user does this token
|
||||
// belong to" — and leaves answering it to whatever is plugged in. Phase 3 plugs
|
||||
// in OAuth 2.1 token validation. Nothing here mints, stores, refreshes or
|
||||
// validates a token's contents, because doing any of that now would be
|
||||
// inventing a token format that OAuth then has to replace.
|
||||
//
|
||||
// WHY MCP AUTHENTICATES SEPARATELY FROM THE REST OF THE API
|
||||
//
|
||||
// The cookie middleware in httpserver/auth.go is deliberately not reused, and
|
||||
// this is the most important decision in this file. Mounting MCP behind that
|
||||
// middleware would mean a browser session could authenticate an MCP call: the
|
||||
// middleware puts an Identity in the context, and any handler downstream that
|
||||
// reads the ambient identity would accept it. That is a real vulnerability
|
||||
// rather than a theoretical one — a logged-in user's cookie is sent by the
|
||||
// browser on requests the user did not intend, which is what SameSite exists to
|
||||
// limit and what an MCP endpoint has no business relying on.
|
||||
//
|
||||
// So identity here is PASSED, never ambient. The transport authenticates, and
|
||||
// hands the result to Handle as a parameter. There is no code path in this
|
||||
// package that reads authctx.From on an inbound request, which makes "a cookie
|
||||
// silently authenticated MCP" structurally impossible rather than merely
|
||||
// unintended. See TestCookieCannotAuthenticateMCP.
|
||||
|
||||
/* ── Errors ─────────────────────────────────────────────────────────────── */
|
||||
|
||||
var (
|
||||
// ErrNoAuthenticator is returned when the surface is running without a
|
||||
// token authenticator. It is a configuration fault, and it fails CLOSED:
|
||||
// a deployment that forgot to wire one refuses every call rather than
|
||||
// serving them unauthenticated.
|
||||
ErrNoAuthenticator = errors.New("mcpserver: no token authenticator configured")
|
||||
|
||||
// ErrMissingToken covers an absent or empty Authorization header.
|
||||
ErrMissingToken = errors.New("mcpserver: no bearer token")
|
||||
|
||||
// ErrMalformedToken covers a header this server could not parse as a
|
||||
// bearer credential — a missing scheme, a wrong scheme, an empty value.
|
||||
ErrMalformedToken = errors.New("mcpserver: malformed Authorization header")
|
||||
|
||||
// ErrInvalidToken covers a well-formed token that does not resolve to a
|
||||
// user: unknown, expired, revoked, or issued for something else.
|
||||
//
|
||||
// ONE error for all of those, deliberately. Telling a caller that a token
|
||||
// is "expired" rather than "unknown" confirms it once existed, which is an
|
||||
// oracle over the token space. Same reasoning as tools.Denied().
|
||||
ErrInvalidToken = errors.New("mcpserver: invalid bearer token")
|
||||
)
|
||||
|
||||
/* ── The seam ───────────────────────────────────────────────────────────── */
|
||||
|
||||
// TokenAuthenticator resolves a raw bearer token into a KROW identity.
|
||||
//
|
||||
// Deliberately one method taking a string and returning the SAME
|
||||
// authctx.Identity the cookie path produces. Two things follow from that shape,
|
||||
// and both are the point:
|
||||
//
|
||||
// - There is no second identity model. Everything downstream — the policy
|
||||
// table, the org pre-filter, tools.Context — consumes authctx.Identity and
|
||||
// cannot tell which path produced it, so authorization cannot drift between
|
||||
// the two.
|
||||
// - An implementation cannot report anything except an identity or a failure.
|
||||
// It has no way to return "authenticated, but also here is an org" or any
|
||||
// other channel a caller might trust. The org is inside the identity,
|
||||
// which comes from the user row.
|
||||
//
|
||||
// Phase 3's OAuth implementation of this interface will: hash the presented
|
||||
// token, look it up, check expiry, revocation and audience, load the user, and
|
||||
// build the identity from the USER ROW — never from the token's contents. A
|
||||
// token that carried its own org claim would be a token whose bearer chose
|
||||
// their own tenant.
|
||||
type TokenAuthenticator interface {
|
||||
// Authenticate resolves a raw token, or returns an error.
|
||||
//
|
||||
// Implementations must fail closed and must not distinguish unknown from
|
||||
// expired from revoked in the returned error.
|
||||
Authenticate(ctx context.Context, rawToken string) (authctx.Identity, error)
|
||||
}
|
||||
|
||||
// UserLookup is the subset of the existing user store this package needs.
|
||||
//
|
||||
// Narrowed to one method so an implementation of TokenAuthenticator can re-read
|
||||
// the user on every call — which is what makes suspension take effect on
|
||||
// contact rather than whenever a token happens to lapse. httpserver/auth.go
|
||||
// does exactly this for cookies (see its comment on re-reading the user row),
|
||||
// and the bearer path must not be weaker.
|
||||
//
|
||||
// auth.UserStore already satisfies this.
|
||||
type UserLookup interface {
|
||||
FindByID(ctx context.Context, id string) (auth.User, error)
|
||||
}
|
||||
|
||||
/* ── Header parsing ─────────────────────────────────────────────────────── */
|
||||
|
||||
// bearerToken extracts the credential from an Authorization header.
|
||||
//
|
||||
// Only the Authorization header is consulted. Not a query parameter — the MCP
|
||||
// spec forbids tokens in the URI, and a URI is logged, cached, and put in a
|
||||
// Referer. Not a custom header, not a cookie, not the body. One place, so there
|
||||
// is one thing to reason about.
|
||||
func bearerToken(r *http.Request) (string, error) {
|
||||
header := r.Header.Get("Authorization")
|
||||
if strings.TrimSpace(header) == "" {
|
||||
return "", ErrMissingToken
|
||||
}
|
||||
|
||||
scheme, value, found := strings.Cut(header, " ")
|
||||
if !found {
|
||||
return "", ErrMalformedToken
|
||||
}
|
||||
// Case-insensitive per RFC 7235: "Bearer", "bearer" and "BEARER" are the
|
||||
// same scheme, and rejecting the variants would fail against clients that
|
||||
// are behaving correctly.
|
||||
if !strings.EqualFold(strings.TrimSpace(scheme), "bearer") {
|
||||
return "", ErrMalformedToken
|
||||
}
|
||||
|
||||
token := strings.TrimSpace(value)
|
||||
if token == "" {
|
||||
return "", ErrMalformedToken
|
||||
}
|
||||
// A second space means a second value — "Bearer a b" is not a token, and
|
||||
// accepting the first half would silently authenticate something the
|
||||
// client did not send.
|
||||
if strings.ContainsAny(token, " \t") {
|
||||
return "", ErrMalformedToken
|
||||
}
|
||||
return token, nil
|
||||
}
|
||||
|
||||
// authenticate resolves the request's bearer credential into an identity.
|
||||
//
|
||||
// Every failure returns the same outward answer — 401 with no detail about
|
||||
// which stage failed. The reason is recorded in the log, where the operator is.
|
||||
func (s *Server) authenticate(r *http.Request) (authctx.Identity, error) {
|
||||
token, err := bearerToken(r)
|
||||
if err != nil {
|
||||
return authctx.Identity{}, err
|
||||
}
|
||||
if s.tokens == nil {
|
||||
return authctx.Identity{}, ErrNoAuthenticator
|
||||
}
|
||||
|
||||
identity, err := s.tokens.Authenticate(r.Context(), token)
|
||||
if err != nil {
|
||||
return authctx.Identity{}, ErrInvalidToken
|
||||
}
|
||||
|
||||
// Defence in depth against an authenticator that returns a partially
|
||||
// populated identity. Everything downstream assumes these two are present:
|
||||
// tools/scope.go refuses an empty OrgID, but it should never be asked to,
|
||||
// and a missing UserID would produce a query scoped to nobody.
|
||||
if identity.UserID == "" || identity.OrgID == "" {
|
||||
return authctx.Identity{}, ErrInvalidToken
|
||||
}
|
||||
// A suspended account must not hold a working token. The authenticator is
|
||||
// expected to check this; repeating it here costs nothing and means a
|
||||
// mistake in one implementation is not a live account bypass.
|
||||
if identity.Status != "" && identity.Status != auth.StatusActive {
|
||||
return authctx.Identity{}, ErrInvalidToken
|
||||
}
|
||||
return identity, nil
|
||||
}
|
||||
|
||||
// authFailureReason names the stage that refused, for the log only.
|
||||
func authFailureReason(err error) string {
|
||||
switch {
|
||||
case errors.Is(err, ErrMissingToken):
|
||||
return "missing_token"
|
||||
case errors.Is(err, ErrMalformedToken):
|
||||
return "malformed_header"
|
||||
case errors.Is(err, ErrNoAuthenticator):
|
||||
return "no_authenticator_configured"
|
||||
case errors.Is(err, ErrInvalidToken):
|
||||
return "invalid_token"
|
||||
default:
|
||||
return "error"
|
||||
}
|
||||
}
|
||||
512
go-api/internal/mcpserver/auth_test.go
Normal file
512
go-api/internal/mcpserver/auth_test.go
Normal file
@@ -0,0 +1,512 @@
|
||||
package mcpserver
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"io"
|
||||
"log/slog"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/krow/krow-backend/go-api/internal/authctx"
|
||||
"github.com/krow/krow-backend/go-api/internal/runtime"
|
||||
"github.com/krow/krow-backend/go-api/internal/testutil"
|
||||
"github.com/krow/krow-backend/go-api/internal/tools"
|
||||
)
|
||||
|
||||
/* ── A test-only authenticator ──────────────────────────────────────────────
|
||||
|
||||
TEST-ONLY. This type lives in a _test.go file and is therefore not compiled
|
||||
into the binary at all — there is no build tag to forget and no flag that
|
||||
could enable it in production. That is deliberate: the one thing worse than
|
||||
having no authentication is having a development authenticator that ships.
|
||||
|
||||
It is a map from token to identity and nothing else. It does not hash, does
|
||||
not expire, does not check an audience and does not consult a database,
|
||||
because none of those are what these tests are testing. What they test is
|
||||
the SEAM: that a token resolves to an identity, that the identity reaches
|
||||
the registry, and that existing authorization then decides the answer.
|
||||
|
||||
Phase 3 replaces this with an OAuth implementation of the same interface.
|
||||
Every test below keeps working, because what they assert is the behaviour of
|
||||
the seam rather than the behaviour of any particular token format. That is
|
||||
the reason for defining the interface before implementing OAuth rather than
|
||||
after. */
|
||||
|
||||
type fakeTokens struct {
|
||||
byToken map[string]authctx.Identity
|
||||
err error
|
||||
}
|
||||
|
||||
func (f *fakeTokens) Authenticate(_ context.Context, raw string) (authctx.Identity, error) {
|
||||
if f.err != nil {
|
||||
return authctx.Identity{}, f.err
|
||||
}
|
||||
id, ok := f.byToken[raw]
|
||||
if !ok {
|
||||
return authctx.Identity{}, ErrInvalidToken
|
||||
}
|
||||
return id, nil
|
||||
}
|
||||
|
||||
const validToken = "test-token-valid"
|
||||
|
||||
func identityFor(orgID, role string) authctx.Identity {
|
||||
return authctx.Identity{
|
||||
UserID: "user-" + role,
|
||||
OrgID: orgID,
|
||||
Email: role + "@example.test",
|
||||
FullName: "Test " + role,
|
||||
Role: role,
|
||||
AccountType: "employer",
|
||||
Status: "active",
|
||||
}
|
||||
}
|
||||
|
||||
// authedServer wires the real registry to a token map.
|
||||
func authedServer(t *testing.T, reg *tools.Registry, tokens map[string]authctx.Identity) *Server {
|
||||
t.Helper()
|
||||
if reg == nil {
|
||||
reg = runtime.DefaultTools(nil, nil)
|
||||
}
|
||||
return New(reg, &fakeTokens{byToken: tokens}, slog.New(slog.NewTextHandler(io.Discard, nil)))
|
||||
}
|
||||
|
||||
// postWith sends a body with an explicit Authorization header value. An empty
|
||||
// header value means the header is not sent at all.
|
||||
func postWith(t *testing.T, s *Server, authHeader, body string) *httptest.ResponseRecorder {
|
||||
t.Helper()
|
||||
req := httptest.NewRequest(http.MethodPost, "/mcp", strings.NewReader(body))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
if authHeader != "" {
|
||||
req.Header.Set("Authorization", authHeader)
|
||||
}
|
||||
rec := httptest.NewRecorder()
|
||||
s.Handler().ServeHTTP(rec, req)
|
||||
return rec
|
||||
}
|
||||
|
||||
const listBody = `{"jsonrpc":"2.0","id":1,"method":"tools/list"}`
|
||||
|
||||
/* ── 1–5. Header and token rejection ────────────────────────────────────── */
|
||||
|
||||
func TestBearerRejection(t *testing.T) {
|
||||
s := authedServer(t, nil, map[string]authctx.Identity{
|
||||
validToken: identityFor("org-a", "admin"),
|
||||
})
|
||||
|
||||
for name, header := range map[string]string{
|
||||
"missing header": "",
|
||||
"no scheme": "abc123",
|
||||
"wrong scheme basic": "Basic dXNlcjpwYXNz",
|
||||
"wrong scheme token": "Token abc123",
|
||||
"empty bearer": "Bearer ",
|
||||
"bearer with only ws": "Bearer ",
|
||||
"two values": "Bearer abc def",
|
||||
"unknown token": "Bearer not-a-real-token",
|
||||
"token with whitespace": "Bearer abc\tdef",
|
||||
} {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
rec := postWith(t, s, header, listBody)
|
||||
|
||||
if rec.Code != http.StatusUnauthorized {
|
||||
t.Fatalf("status = %d, want 401 (body: %s)", rec.Code, rec.Body.String())
|
||||
}
|
||||
// RFC 9728: the client reads the authorization server location from
|
||||
// this header. Without it a compliant client cannot begin discovery.
|
||||
if got := rec.Header().Get("WWW-Authenticate"); !strings.HasPrefix(got, "Bearer") {
|
||||
t.Errorf("WWW-Authenticate = %q, want a Bearer challenge", got)
|
||||
}
|
||||
// The refusal must not say WHICH stage failed. "expired" versus
|
||||
// "unknown" is an oracle over the token space.
|
||||
body := rec.Body.String()
|
||||
for _, leak := range []string{"expired", "revoked", "unknown", "malformed", "not found"} {
|
||||
if strings.Contains(strings.ToLower(body), leak) {
|
||||
t.Errorf("the 401 body distinguishes failure modes (%q): %s", leak, body)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// Case-insensitivity is required by RFC 7235; rejecting "bearer" would fail
|
||||
// against clients that are behaving correctly.
|
||||
func TestBearerSchemeIsCaseInsensitive(t *testing.T) {
|
||||
s := authedServer(t, nil, map[string]authctx.Identity{
|
||||
validToken: identityFor("org-a", "admin"),
|
||||
})
|
||||
for _, scheme := range []string{"Bearer", "bearer", "BEARER", "BeArEr"} {
|
||||
rec := postWith(t, s, scheme+" "+validToken, listBody)
|
||||
if rec.Code != http.StatusOK {
|
||||
t.Errorf("scheme %q: status = %d, want 200", scheme, rec.Code)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// A server with no authenticator wired must refuse everything rather than
|
||||
// serve it unauthenticated. Misconfiguration fails closed.
|
||||
func TestNoAuthenticatorFailsClosed(t *testing.T) {
|
||||
s := New(runtime.DefaultTools(nil, nil), nil, slog.New(slog.NewTextHandler(io.Discard, nil)))
|
||||
rec := postWith(t, s, "Bearer "+validToken, listBody)
|
||||
if rec.Code != http.StatusUnauthorized {
|
||||
t.Fatalf("status = %d, want 401 when no authenticator is configured", rec.Code)
|
||||
}
|
||||
}
|
||||
|
||||
/* ── 6 & 7. Identity and org resolution ─────────────────────────────────── */
|
||||
|
||||
func TestValidBearerReachesToolsList(t *testing.T) {
|
||||
s := authedServer(t, nil, map[string]authctx.Identity{
|
||||
validToken: identityFor("org-a", "admin"),
|
||||
})
|
||||
rec := postWith(t, s, "Bearer "+validToken, listBody)
|
||||
|
||||
var out toolsListResult
|
||||
resultInto(t, rec, &out)
|
||||
if len(out.Tools) != len(wantExposed) {
|
||||
t.Fatalf("exposed %d tools, want %d", len(out.Tools), len(wantExposed))
|
||||
}
|
||||
}
|
||||
|
||||
// An identity missing a tenant must be refused before it reaches a tool.
|
||||
// tools/scope.go would refuse it too, but it should never be asked to: an
|
||||
// authenticator that returns a half-built identity is a bug, not a caller.
|
||||
func TestIncompleteIdentityIsRefused(t *testing.T) {
|
||||
for name, id := range map[string]authctx.Identity{
|
||||
"no org": {UserID: "u1", Role: "admin", Status: "active"},
|
||||
"no user": {OrgID: "org-a", Role: "admin", Status: "active"},
|
||||
"suspended": {UserID: "u1", OrgID: "org-a", Role: "admin", Status: "suspended"},
|
||||
"empty entire": {},
|
||||
} {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
s := authedServer(t, nil, map[string]authctx.Identity{validToken: id})
|
||||
rec := postWith(t, s, "Bearer "+validToken, listBody)
|
||||
if rec.Code != http.StatusUnauthorized {
|
||||
t.Errorf("status = %d, want 401", rec.Code)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
/* ── 8, 9, 10. Identity cannot be overridden ────────────────────────────── */
|
||||
|
||||
// The identity must come from the token and from nothing else. This walks the
|
||||
// channels a client controls and asserts that none of them moves the tenant.
|
||||
func TestIdentityCannotBeOverridden(t *testing.T) {
|
||||
h := testutil.New(t)
|
||||
orgA := freshOrgFor(t, h, "mcp-override-a")
|
||||
orgB := freshOrgFor(t, h, "mcp-override-b")
|
||||
seedActivityRows(t, h, orgA, 5, "a@example.test")
|
||||
seedActivityRows(t, h, orgB, 40, "b@example.test")
|
||||
|
||||
reg := runtime.DefaultTools(h.Pool, nil)
|
||||
s := authedServer(t, reg, map[string]authctx.Identity{
|
||||
validToken: {UserID: "u-a", OrgID: orgA, Email: "a@example.test",
|
||||
Role: "admin", AccountType: "employer", Status: "active"},
|
||||
})
|
||||
|
||||
// Every one of these is a client-controlled channel. None may change which
|
||||
// tenant is read. orgB has 40 rows and orgA has 5, so a successful override
|
||||
// is visible as a total of 40 or 45.
|
||||
attempts := map[string]string{
|
||||
"tool argument": `{"jsonrpc":"2.0","id":1,"method":"tools/call","params":
|
||||
{"name":"activity_breakdown","arguments":{"org_id":"` + orgB + `"}}}`,
|
||||
"jsonrpc meta": `{"jsonrpc":"2.0","id":1,"method":"tools/call","_meta":{"org_id":"` + orgB + `"},
|
||||
"params":{"name":"activity_breakdown","arguments":{}}}`,
|
||||
"params meta": `{"jsonrpc":"2.0","id":1,"method":"tools/call","params":
|
||||
{"name":"activity_breakdown","arguments":{},"_meta":{"org_id":"` + orgB + `"}}}`,
|
||||
"params orgId": `{"jsonrpc":"2.0","id":1,"method":"tools/call","params":
|
||||
{"name":"activity_breakdown","orgId":"` + orgB + `","arguments":{}}}`,
|
||||
"params userId": `{"jsonrpc":"2.0","id":1,"method":"tools/call","params":
|
||||
{"name":"activity_breakdown","userId":"u-b","arguments":{}}}`,
|
||||
}
|
||||
|
||||
for name, body := range attempts {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
rec := postWith(t, s, "Bearer "+validToken, body)
|
||||
total := totalFromActivityBreakdown(t, rec)
|
||||
if total != 5 {
|
||||
t.Errorf("total = %d, want 5 (org A only) — %q moved the tenant", total, name)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// Custom headers naming another tenant must be ignored outright.
|
||||
func TestCustomIdentityHeadersAreIgnored(t *testing.T) {
|
||||
h := testutil.New(t)
|
||||
orgA := freshOrgFor(t, h, "mcp-hdr-a")
|
||||
orgB := freshOrgFor(t, h, "mcp-hdr-b")
|
||||
seedActivityRows(t, h, orgA, 5, "a@example.test")
|
||||
seedActivityRows(t, h, orgB, 40, "b@example.test")
|
||||
|
||||
reg := runtime.DefaultTools(h.Pool, nil)
|
||||
s := authedServer(t, reg, map[string]authctx.Identity{
|
||||
validToken: {UserID: "u-a", OrgID: orgA, Email: "a@example.test",
|
||||
Role: "admin", AccountType: "employer", Status: "active"},
|
||||
})
|
||||
|
||||
body := `{"jsonrpc":"2.0","id":1,"method":"tools/call","params":
|
||||
{"name":"activity_breakdown","arguments":{}}}`
|
||||
req := httptest.NewRequest(http.MethodPost, "/mcp", strings.NewReader(body))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
req.Header.Set("Authorization", "Bearer "+validToken)
|
||||
for _, h := range []string{"X-Org-Id", "X-Organization-Id", "X-User-Id", "X-Tenant-Id", "X-Krow-Org"} {
|
||||
req.Header.Set(h, orgB)
|
||||
}
|
||||
rec := httptest.NewRecorder()
|
||||
s.Handler().ServeHTTP(rec, req)
|
||||
|
||||
if total := totalFromActivityBreakdown(t, rec); total != 5 {
|
||||
t.Errorf("total = %d, want 5 — a custom header moved the tenant", total)
|
||||
}
|
||||
}
|
||||
|
||||
// A cookie must never authenticate MCP.
|
||||
//
|
||||
// This is the test for the decision in auth.go: identity is passed, never
|
||||
// ambient. Even with a valid KROW identity sitting in the request context —
|
||||
// which is exactly what the cookie middleware would put there if this endpoint
|
||||
// were mounted behind it — the call must be refused, because no bearer token
|
||||
// was presented.
|
||||
func TestCookieCannotAuthenticateMCP(t *testing.T) {
|
||||
s := authedServer(t, nil, map[string]authctx.Identity{
|
||||
validToken: identityFor("org-a", "admin"),
|
||||
})
|
||||
|
||||
req := httptest.NewRequest(http.MethodPost, "/mcp", strings.NewReader(listBody))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
// A session cookie, as a browser would send it.
|
||||
req.AddCookie(&http.Cookie{Name: "krow_session", Value: "a-perfectly-valid-session-token"})
|
||||
// AND a fully populated identity in the context, as the cookie middleware
|
||||
// would have placed there. This is the strongest form of the test: even if
|
||||
// somebody mounts MCP behind authenticate(), it must still refuse.
|
||||
ctx := authctx.With(req.Context(), identityFor("org-a", "admin"))
|
||||
req = req.WithContext(ctx)
|
||||
|
||||
rec := httptest.NewRecorder()
|
||||
s.Handler().ServeHTTP(rec, req)
|
||||
|
||||
if rec.Code != http.StatusUnauthorized {
|
||||
t.Fatalf("status = %d, want 401 — a cookie session authenticated an MCP call", rec.Code)
|
||||
}
|
||||
}
|
||||
|
||||
/* ── 11, 12, 13. Authorization still runs ───────────────────────────────── */
|
||||
|
||||
// The point of the whole design: MCP authenticates, and the EXISTING policy
|
||||
// table authorizes. A talent caller must be refused a tool that operators own.
|
||||
func TestExistingAuthorizationStillRuns(t *testing.T) {
|
||||
h := testutil.New(t)
|
||||
org := freshOrgFor(t, h, "mcp-authz")
|
||||
seedActivityRows(t, h, org, 5, "boss@example.test")
|
||||
|
||||
reg := runtime.DefaultTools(h.Pool, nil)
|
||||
s := authedServer(t, reg, map[string]authctx.Identity{
|
||||
"admin-token": {UserID: "u-admin", OrgID: org, Email: "boss@example.test",
|
||||
Role: "admin", AccountType: "employer", Status: "active"},
|
||||
"talent-token": {UserID: "u-talent", OrgID: org, Email: "worker@example.test",
|
||||
Role: "talent", AccountType: "talent", Status: "active"},
|
||||
})
|
||||
|
||||
// `staff` is operators-only in the policy table, and hires_recent reads
|
||||
// job-applications which talent may list only in its own scope.
|
||||
call := `{"jsonrpc":"2.0","id":1,"method":"tools/call","params":
|
||||
{"name":"activity_breakdown","arguments":{}}}`
|
||||
|
||||
adminRec := postWith(t, s, "Bearer admin-token", call)
|
||||
if got := totalFromActivityBreakdown(t, adminRec); got != 5 {
|
||||
t.Errorf("admin total = %d, want 5", got)
|
||||
}
|
||||
|
||||
// A talent caller is authenticated but scoped. Whatever comes back, it must
|
||||
// come back through the policy table rather than around it — the assertion
|
||||
// is that the two roles do NOT get the same answer.
|
||||
talentRec := postWith(t, s, "Bearer talent-token", call)
|
||||
talentTotal := totalFromActivityBreakdownAllowingDenial(t, talentRec)
|
||||
if talentTotal == 5 {
|
||||
t.Error("a talent caller saw the admin's total; authorization did not run")
|
||||
}
|
||||
}
|
||||
|
||||
/* ── 14 & 15. Exclusions hold under authentication ──────────────────────── */
|
||||
|
||||
// Authentication must not become a way to reach a withheld tool. A valid token
|
||||
// is not a key to the write tools.
|
||||
func TestExclusionsHoldForAuthenticatedCallers(t *testing.T) {
|
||||
s := authedServer(t, nil, map[string]authctx.Identity{
|
||||
validToken: identityFor("org-a", "admin"),
|
||||
})
|
||||
|
||||
for _, name := range []string{"assign_worker", "move_application", "knowledge_search"} {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
body := `{"jsonrpc":"2.0","id":1,"method":"tools/call","params":{"name":"` + name + `","arguments":{}}}`
|
||||
rec := postWith(t, s, "Bearer "+validToken, body)
|
||||
|
||||
var out toolsCallResult
|
||||
resultInto(t, rec, &out)
|
||||
if !out.IsError {
|
||||
t.Fatalf("%s was reachable by an authenticated caller", name)
|
||||
}
|
||||
if !strings.Contains(out.Content[0].Text, "mcp.unknown_tool") {
|
||||
t.Errorf("expected mcp.unknown_tool, got: %s", out.Content[0].Text)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// tools/list must not vary by caller in a way that reveals the withheld set.
|
||||
func TestToolsListIsTheSameForEveryRole(t *testing.T) {
|
||||
s := authedServer(t, nil, map[string]authctx.Identity{
|
||||
"admin-token": identityFor("org-a", "admin"),
|
||||
"talent-token": identityFor("org-a", "talent"),
|
||||
})
|
||||
|
||||
names := func(token string) []string {
|
||||
rec := postWith(t, s, "Bearer "+token, listBody)
|
||||
var out toolsListResult
|
||||
resultInto(t, rec, &out)
|
||||
got := make([]string, 0, len(out.Tools))
|
||||
for _, tool := range out.Tools {
|
||||
got = append(got, tool.Name)
|
||||
}
|
||||
return got
|
||||
}
|
||||
|
||||
admin, talent := names("admin-token"), names("talent-token")
|
||||
if strings.Join(admin, ",") != strings.Join(talent, ",") {
|
||||
t.Errorf("tools/list differs by role:\n admin: %v\ntalent: %v", admin, talent)
|
||||
}
|
||||
}
|
||||
|
||||
/* ── The handshake requires a token too ─────────────────────────────────── */
|
||||
|
||||
// REVERSED IN PHASE 4, deliberately.
|
||||
//
|
||||
// This test previously asserted the opposite: that initialize and ping were
|
||||
// reachable without a token, on the reasoning that a client needs somewhere to
|
||||
// start. That was wrong, and the integration test over the mounted route is
|
||||
// what caught it — a client's FIRST request is usually initialize, and
|
||||
// answering it 200 means the client never sees the WWW-Authenticate challenge
|
||||
// that begins the OAuth flow. It believes it is connected and finds out
|
||||
// otherwise at the first real call, with no 401 in hand to discover from.
|
||||
//
|
||||
// Requiring a token everywhere means the first request, whatever it is,
|
||||
// produces the challenge. Nothing is lost: the handshake returns only this
|
||||
// server's name and capabilities, which are of use solely to a client that
|
||||
// means to authenticate.
|
||||
func TestHandshakeMethodsAlsoRequireAuth(t *testing.T) {
|
||||
s := authedServer(t, nil, map[string]authctx.Identity{
|
||||
validToken: identityFor("org-a", "admin"),
|
||||
})
|
||||
|
||||
for name, body := range map[string]string{
|
||||
"initialize": `{"jsonrpc":"2.0","id":1,"method":"initialize","params":{}}`,
|
||||
"ping": `{"jsonrpc":"2.0","id":1,"method":"ping"}`,
|
||||
"tools/list": `{"jsonrpc":"2.0","id":1,"method":"tools/list"}`,
|
||||
"unknown verb": `{"jsonrpc":"2.0","id":1,"method":"resources/list"}`,
|
||||
} {
|
||||
t.Run(name+" without a token", func(t *testing.T) {
|
||||
rec := postWith(t, s, "", body)
|
||||
if rec.Code != http.StatusUnauthorized {
|
||||
t.Errorf("status = %d, want 401 — every method must produce the "+
|
||||
"challenge that starts the OAuth flow", rec.Code)
|
||||
}
|
||||
if !strings.HasPrefix(rec.Header().Get("WWW-Authenticate"), "Bearer") {
|
||||
t.Error("the 401 carries no Bearer challenge")
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
// And with a token, the handshake works — or the test above would pass by
|
||||
// the endpoint being broken rather than by it being guarded.
|
||||
for name, body := range map[string]string{
|
||||
"initialize": `{"jsonrpc":"2.0","id":1,"method":"initialize","params":{}}`,
|
||||
"ping": `{"jsonrpc":"2.0","id":1,"method":"ping"}`,
|
||||
} {
|
||||
t.Run(name+" with a token", func(t *testing.T) {
|
||||
if rec := postWith(t, s, "Bearer "+validToken, body); rec.Code != http.StatusOK {
|
||||
t.Errorf("status = %d, want 200", rec.Code)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// The resource_metadata pointer must be built from configuration, so a 401
|
||||
// tells a client where to look without this package knowing any hostname.
|
||||
func TestChallengeCarriesTheConfiguredResourceMetadata(t *testing.T) {
|
||||
const metadataURL = "https://configured.example.test/.well-known/oauth-protected-resource"
|
||||
s := authedServer(t, nil, map[string]authctx.Identity{}).
|
||||
WithResourceMetadataURL(metadataURL)
|
||||
|
||||
rec := postWith(t, s, "", listBody)
|
||||
challenge := rec.Header().Get("WWW-Authenticate")
|
||||
if !strings.Contains(challenge, `resource_metadata="`+metadataURL+`"`) {
|
||||
t.Errorf("WWW-Authenticate = %q, want it to carry %q", challenge, metadataURL)
|
||||
}
|
||||
}
|
||||
|
||||
/* ── Helpers ────────────────────────────────────────────────────────────── */
|
||||
|
||||
func freshOrgFor(t *testing.T, h *testutil.Harness, slug string) string {
|
||||
t.Helper()
|
||||
var id string
|
||||
if err := h.Pool.QueryRow(context.Background(),
|
||||
`INSERT INTO organizations (name, slug) VALUES ($1, $2) RETURNING id::text`,
|
||||
slug, slug).Scan(&id); err != nil {
|
||||
t.Fatalf("create org %s: %v", slug, err)
|
||||
}
|
||||
return id
|
||||
}
|
||||
|
||||
func seedActivityRows(t *testing.T, h *testutil.Harness, org string, n int, email string) {
|
||||
t.Helper()
|
||||
for i := 0; i < n; i++ {
|
||||
if _, err := h.Pool.Exec(context.Background(),
|
||||
`INSERT INTO user_activity (org_id, event_type, user_email, user_name)
|
||||
VALUES ($1::uuid, 'login', $2, 'Someone')`, org, email); err != nil {
|
||||
t.Fatalf("seed activity: %v", err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// totalFromActivityBreakdown reads the `total` out of a successful tool result.
|
||||
func totalFromActivityBreakdown(t *testing.T, rec *httptest.ResponseRecorder) int {
|
||||
t.Helper()
|
||||
var out toolsCallResult
|
||||
resultInto(t, rec, &out)
|
||||
if out.IsError {
|
||||
t.Fatalf("tool call failed: %s", out.Content[0].Text)
|
||||
}
|
||||
return parseTotal(t, out.Content[0].Text)
|
||||
}
|
||||
|
||||
// totalFromActivityBreakdownAllowingDenial returns -1 when the tool refused,
|
||||
// which is a legitimate authorization outcome rather than a test failure.
|
||||
func totalFromActivityBreakdownAllowingDenial(t *testing.T, rec *httptest.ResponseRecorder) int {
|
||||
t.Helper()
|
||||
var out toolsCallResult
|
||||
resultInto(t, rec, &out)
|
||||
if out.IsError {
|
||||
return -1
|
||||
}
|
||||
return parseTotal(t, out.Content[0].Text)
|
||||
}
|
||||
|
||||
func parseTotal(t *testing.T, raw string) int {
|
||||
t.Helper()
|
||||
// activity_breakdown's own field name, taken from the handler's output
|
||||
// rather than guessed: a wrong name here reads as zero, which would make a
|
||||
// cross-tenant leak look like a pass.
|
||||
var payload struct {
|
||||
Data struct {
|
||||
TotalEvents int `json:"totalEvents"`
|
||||
} `json:"data"`
|
||||
}
|
||||
if err := json.Unmarshal([]byte(raw), &payload); err != nil {
|
||||
t.Fatalf("tool result did not decode: %v\nraw: %s", err, raw)
|
||||
}
|
||||
return payload.Data.TotalEvents
|
||||
}
|
||||
243
go-api/internal/mcpserver/jsonrpc.go
Normal file
243
go-api/internal/mcpserver/jsonrpc.go
Normal file
@@ -0,0 +1,243 @@
|
||||
// Package mcpserver is the Model Context Protocol surface: a second way into
|
||||
// the tool layer, for clients that speak MCP rather than HTTP+cookie.
|
||||
//
|
||||
// It is an ADDITIONAL interface and nothing else. It owns no business logic, no
|
||||
// SQL and no authorization rules. Every call it serves ends up in
|
||||
// tools.Registry.Dispatch — the same entry point the agent loop uses — so a
|
||||
// question asked through MCP is answered by the same handler, under the same
|
||||
// policy table, behind the same org pre-filter as the same question asked by
|
||||
// Owliver. That is the whole design, and the reason this package is small.
|
||||
//
|
||||
// What lives here:
|
||||
//
|
||||
// - JSON-RPC 2.0 framing (this file)
|
||||
// - the three methods MCP needs to be useful: initialize, tools/list,
|
||||
// tools/call (server.go)
|
||||
// - which tools are published, derived from the registry (tools.go)
|
||||
// - bearer authentication, as a seam an OAuth implementation plugs into
|
||||
// (auth.go)
|
||||
// - the Streamable HTTP binding (transport.go)
|
||||
//
|
||||
// Identity is established by this package's own bearer authentication and is
|
||||
// PASSED to the handlers, never read from the ambient request context. That is
|
||||
// what stops a browser cookie from authenticating an MCP call — see auth.go.
|
||||
package mcpserver
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"fmt"
|
||||
"time"
|
||||
)
|
||||
|
||||
// jsonRPCVersion is the only version this server speaks. A request naming
|
||||
// anything else is malformed rather than merely unsupported: "2.0" is a
|
||||
// constant in the spec, not a negotiation.
|
||||
const jsonRPCVersion = "2.0"
|
||||
|
||||
/* ── Wire types ─────────────────────────────────────────────────────────── */
|
||||
|
||||
// request is one inbound JSON-RPC message.
|
||||
//
|
||||
// ID is json.RawMessage rather than any, because the spec allows a string, a
|
||||
// number or null, and the response MUST echo it back byte-for-byte. Decoding it
|
||||
// into an `any` turns 1 into 1.0 on the way back out, which is a different id to
|
||||
// a client matching responses to requests.
|
||||
type request struct {
|
||||
JSONRPC string `json:"jsonrpc"`
|
||||
ID json.RawMessage `json:"id,omitempty"`
|
||||
Method string `json:"method"`
|
||||
Params json.RawMessage `json:"params,omitempty"`
|
||||
}
|
||||
|
||||
// isNotification reports that no response is expected.
|
||||
//
|
||||
// A notification is a request with no id. The spec is explicit that a server
|
||||
// must not answer one, so the transport drops the response and returns 202.
|
||||
func (r request) isNotification() bool {
|
||||
return len(r.ID) == 0 || string(r.ID) == "null"
|
||||
}
|
||||
|
||||
// response is one outbound JSON-RPC message.
|
||||
//
|
||||
// Result and Error are pointers so exactly one is ever serialised: the spec
|
||||
// forbids both together, and a non-pointer Result would emit `"result":null`
|
||||
// alongside an error.
|
||||
type response struct {
|
||||
JSONRPC string `json:"jsonrpc"`
|
||||
ID json.RawMessage `json:"id,omitempty"`
|
||||
Result any `json:"result,omitempty"`
|
||||
Error *rpcError `json:"error,omitempty"`
|
||||
}
|
||||
|
||||
// rpcError is a JSON-RPC error object.
|
||||
//
|
||||
// retryAfter is NOT serialised. It exists so a rate-limited refusal produced
|
||||
// deep in the handler can reach the transport, which is the only layer that can
|
||||
// set an HTTP status and a Retry-After header. The alternative — returning a
|
||||
// 200 with a JSON-RPC error and no header — would give a client no way to know
|
||||
// how long to wait.
|
||||
type rpcError struct {
|
||||
Code int `json:"code"`
|
||||
Message string `json:"message"`
|
||||
Data any `json:"data,omitempty"`
|
||||
|
||||
retryAfter time.Duration
|
||||
}
|
||||
|
||||
func (e *rpcError) Error() string { return fmt.Sprintf("jsonrpc %d: %s", e.Code, e.Message) }
|
||||
|
||||
/* ── Error codes ────────────────────────────────────────────────────────── */
|
||||
|
||||
// The standard JSON-RPC 2.0 codes. Reserved range is -32768..-32000; anything
|
||||
// this server invents lives outside it.
|
||||
const (
|
||||
codeParseError = -32700
|
||||
codeInvalidRequest = -32600
|
||||
codeMethodNotFound = -32601
|
||||
codeInvalidParams = -32602
|
||||
codeInternalError = -32603
|
||||
)
|
||||
|
||||
// codeUnauthorized is outside the JSON-RPC reserved range (-32768..-32000),
|
||||
// because it is this server's own condition rather than a protocol fault. It
|
||||
// accompanies an HTTP 401: the transport layer carries the authoritative
|
||||
// signal, and this gives a client reading only the JSON-RPC body the same
|
||||
// answer.
|
||||
const codeUnauthorized = -32001
|
||||
|
||||
// codeRateLimited is this server's own condition, outside the reserved range.
|
||||
// It accompanies an HTTP 429 and a Retry-After header.
|
||||
const codeRateLimited = -32002
|
||||
|
||||
// errRateLimited refuses a call that exceeded its organisation's ceiling.
|
||||
//
|
||||
// The message names no number and no organisation. How much quota a tenant has
|
||||
// and how much of it they have spent is not something one caller should learn
|
||||
// from a refusal — it is the same reasoning as the opaque tool denial.
|
||||
func errRateLimited(retryAfter time.Duration) *rpcError {
|
||||
return &rpcError{
|
||||
Code: codeRateLimited,
|
||||
Message: "too many requests for this organisation; retry after the interval in the Retry-After header",
|
||||
retryAfter: retryAfter,
|
||||
}
|
||||
}
|
||||
|
||||
func errParse(detail string) *rpcError {
|
||||
return &rpcError{Code: codeParseError, Message: "invalid JSON", Data: detail}
|
||||
}
|
||||
|
||||
func errInvalidRequest(detail string) *rpcError {
|
||||
return &rpcError{Code: codeInvalidRequest, Message: "invalid JSON-RPC request", Data: detail}
|
||||
}
|
||||
|
||||
func errMethodNotFound(method string) *rpcError {
|
||||
return &rpcError{
|
||||
Code: codeMethodNotFound,
|
||||
Message: "method not found",
|
||||
Data: fmt.Sprintf("this server implements initialize, tools/list and tools/call; it does not implement %q", method),
|
||||
}
|
||||
}
|
||||
|
||||
func errInvalidParams(detail string) *rpcError {
|
||||
return &rpcError{Code: codeInvalidParams, Message: "invalid params", Data: detail}
|
||||
}
|
||||
|
||||
// errInternal deliberately carries no detail.
|
||||
//
|
||||
// An internal failure is the one case where the thing that went wrong is this
|
||||
// server's business and not the caller's: a wrapped database error or a panic
|
||||
// message is reconnaissance. The detail goes to the log, where the operator is.
|
||||
func errInternal() *rpcError {
|
||||
return &rpcError{Code: codeInternalError, Message: "internal error"}
|
||||
}
|
||||
|
||||
/* ── Parsing ────────────────────────────────────────────────────────────── */
|
||||
|
||||
// errBatch marks a batch request, which this server does not accept.
|
||||
//
|
||||
// Rejecting it explicitly rather than failing to parse it is the point: a
|
||||
// client that batches and gets a parse error will retry the same batch, where
|
||||
// one told that batching is unsupported can fall back to sending messages
|
||||
// singly. The current MCP transport binding sends one message per POST, so
|
||||
// nothing a compliant client does requires batching.
|
||||
var errBatch = errors.New("batch requests are not supported")
|
||||
|
||||
// parseRequest decodes one JSON-RPC message and validates its envelope.
|
||||
//
|
||||
// The two are separate returns because they have different fates: a message
|
||||
// that could not be parsed has no id, so its error answers with a null id,
|
||||
// while a message that parsed but is invalid answers with the id it carried.
|
||||
func parseRequest(body []byte) (request, *rpcError) {
|
||||
trimmed := skipSpace(body)
|
||||
if len(trimmed) == 0 {
|
||||
return request{}, errInvalidRequest("the request body was empty")
|
||||
}
|
||||
if trimmed[0] == '[' {
|
||||
return request{}, errInvalidRequest(errBatch.Error())
|
||||
}
|
||||
|
||||
var req request
|
||||
if err := json.Unmarshal(trimmed, &req); err != nil {
|
||||
return request{}, errParse(err.Error())
|
||||
}
|
||||
|
||||
if req.JSONRPC != jsonRPCVersion {
|
||||
return req, errInvalidRequest(fmt.Sprintf(
|
||||
"jsonrpc must be %q, got %q", jsonRPCVersion, req.JSONRPC))
|
||||
}
|
||||
if req.Method == "" {
|
||||
return req, errInvalidRequest("method is required")
|
||||
}
|
||||
// An id, when present, must be a string or a number. Objects and arrays are
|
||||
// forbidden by the spec, and echoing one back would propagate the mistake.
|
||||
if len(req.ID) > 0 && !isValidID(req.ID) {
|
||||
return req, errInvalidRequest("id must be a string, a number or null")
|
||||
}
|
||||
return req, nil
|
||||
}
|
||||
|
||||
// isValidID reports whether a raw id is a string, a number or null.
|
||||
func isValidID(raw json.RawMessage) bool {
|
||||
t := skipSpace(raw)
|
||||
if len(t) == 0 {
|
||||
return false
|
||||
}
|
||||
switch t[0] {
|
||||
case '{', '[':
|
||||
return false
|
||||
}
|
||||
var v any
|
||||
return json.Unmarshal(t, &v) == nil
|
||||
}
|
||||
|
||||
// decodeParams unmarshals params into dst, treating absent params as an empty
|
||||
// object so a method with only optional fields can be called with none.
|
||||
func decodeParams(raw json.RawMessage, dst any) *rpcError {
|
||||
t := skipSpace(raw)
|
||||
if len(t) == 0 || string(t) == "null" {
|
||||
return nil
|
||||
}
|
||||
// Arrays are legal JSON-RPC (positional params) and are not used by MCP,
|
||||
// whose methods all take an object. Saying so beats a confusing type error.
|
||||
if t[0] == '[' {
|
||||
return errInvalidParams("params must be an object; positional params are not supported")
|
||||
}
|
||||
if err := json.Unmarshal(t, dst); err != nil {
|
||||
return errInvalidParams(err.Error())
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func skipSpace(b []byte) []byte {
|
||||
i := 0
|
||||
for i < len(b) {
|
||||
switch b[i] {
|
||||
case ' ', '\t', '\r', '\n':
|
||||
i++
|
||||
default:
|
||||
return b[i:]
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
595
go-api/internal/mcpserver/mcpserver_test.go
Normal file
595
go-api/internal/mcpserver/mcpserver_test.go
Normal file
@@ -0,0 +1,595 @@
|
||||
package mcpserver
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"io"
|
||||
"log/slog"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/krow/krow-backend/go-api/internal/authctx"
|
||||
"github.com/krow/krow-backend/go-api/internal/runtime"
|
||||
"github.com/krow/krow-backend/go-api/internal/tools"
|
||||
)
|
||||
|
||||
// The tool registry these tests run against is the REAL one — the same
|
||||
// construction the service uses — with a nil database and a nil retriever.
|
||||
//
|
||||
// That is sound because nothing here dispatches: Phase 1's tools/call refuses
|
||||
// before reaching a handler, and every other assertion is about metadata, which
|
||||
// is fixed at registration. Testing against a hand-built fixture registry would
|
||||
// be testing a copy of the thing under test: the whole claim of this package is
|
||||
// "what MCP publishes is what the registry holds", and a fixture would let that
|
||||
// claim pass while being false of the real set.
|
||||
// Both dependencies are nil: the handlers capture them in closures and nothing
|
||||
// here reaches a handler, so nothing dereferences them. If a future test does
|
||||
// dispatch, this will nil-panic loudly rather than quietly reading a database
|
||||
// it should not have.
|
||||
func testRegistry(t *testing.T) *tools.Registry {
|
||||
t.Helper()
|
||||
return runtime.DefaultTools(nil, nil)
|
||||
}
|
||||
|
||||
// newServer wires the real registry to a test-only token authenticator.
|
||||
//
|
||||
// PHASE 2 NOTE: tools/list and tools/call now require a bearer token, so this
|
||||
// fixture authenticates. Not one assertion in this file changed — only the
|
||||
// fixture gained a credential. The behaviour change is the point of Phase 2 and
|
||||
// is asserted directly in auth_test.go (TestBearerRejection), rather than being
|
||||
// papered over here.
|
||||
func newServer(t *testing.T) *Server {
|
||||
t.Helper()
|
||||
return New(testRegistry(t), &fakeTokens{byToken: map[string]authctx.Identity{
|
||||
validToken: identityFor("org-test", "admin"),
|
||||
}}, slog.New(slog.NewTextHandler(io.Discard, nil)))
|
||||
}
|
||||
|
||||
// post sends one raw body to the MCP endpoint and returns the recorder.
|
||||
func post(t *testing.T, s *Server, body string) *httptest.ResponseRecorder {
|
||||
t.Helper()
|
||||
req := httptest.NewRequest(http.MethodPost, "/mcp", strings.NewReader(body))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
req.Header.Set("Authorization", "Bearer "+validToken)
|
||||
rec := httptest.NewRecorder()
|
||||
s.Handler().ServeHTTP(rec, req)
|
||||
return rec
|
||||
}
|
||||
|
||||
// decode reads a JSON-RPC response, failing the test if the envelope is wrong.
|
||||
func decode(t *testing.T, rec *httptest.ResponseRecorder) response {
|
||||
t.Helper()
|
||||
var resp response
|
||||
if err := json.Unmarshal(rec.Body.Bytes(), &resp); err != nil {
|
||||
t.Fatalf("response was not JSON: %v\nbody: %s", err, rec.Body.String())
|
||||
}
|
||||
if resp.JSONRPC != jsonRPCVersion {
|
||||
t.Fatalf("jsonrpc = %q, want %q", resp.JSONRPC, jsonRPCVersion)
|
||||
}
|
||||
return resp
|
||||
}
|
||||
|
||||
// resultInto re-decodes a successful result into dst.
|
||||
func resultInto(t *testing.T, rec *httptest.ResponseRecorder, dst any) {
|
||||
t.Helper()
|
||||
var envelope struct {
|
||||
Result json.RawMessage `json:"result"`
|
||||
Error *rpcError `json:"error"`
|
||||
}
|
||||
if err := json.Unmarshal(rec.Body.Bytes(), &envelope); err != nil {
|
||||
t.Fatalf("response was not JSON: %v", err)
|
||||
}
|
||||
if envelope.Error != nil {
|
||||
t.Fatalf("expected a result, got error %d: %s", envelope.Error.Code, envelope.Error.Message)
|
||||
}
|
||||
if err := json.Unmarshal(envelope.Result, dst); err != nil {
|
||||
t.Fatalf("result did not decode: %v\nresult: %s", err, envelope.Result)
|
||||
}
|
||||
}
|
||||
|
||||
/* ── 1. initialize ──────────────────────────────────────────────────────── */
|
||||
|
||||
func TestInitializeSucceeds(t *testing.T) {
|
||||
s := newServer(t)
|
||||
rec := post(t, s, `{"jsonrpc":"2.0","id":1,"method":"initialize","params":{
|
||||
"protocolVersion":"2025-06-18",
|
||||
"capabilities":{},
|
||||
"clientInfo":{"name":"test-client","version":"1.0"}}}`)
|
||||
|
||||
if rec.Code != http.StatusOK {
|
||||
t.Fatalf("status = %d, want 200", rec.Code)
|
||||
}
|
||||
if ct := rec.Header().Get("Content-Type"); !strings.HasPrefix(ct, "application/json") {
|
||||
t.Errorf("Content-Type = %q, want application/json", ct)
|
||||
}
|
||||
|
||||
var out initializeResult
|
||||
resultInto(t, rec, &out)
|
||||
|
||||
if out.ProtocolVersion != ProtocolVersion {
|
||||
t.Errorf("protocolVersion = %q, want %q", out.ProtocolVersion, ProtocolVersion)
|
||||
}
|
||||
if out.ServerInfo["name"] != ServerName {
|
||||
t.Errorf("serverInfo.name = %v, want %q", out.ServerInfo["name"], ServerName)
|
||||
}
|
||||
// Exactly one capability, because exactly one is implemented. Advertising
|
||||
// resources or prompts would promise methods that answer method-not-found.
|
||||
if _, ok := out.Capabilities["tools"]; !ok {
|
||||
t.Error("capabilities.tools is missing")
|
||||
}
|
||||
for _, unimplemented := range []string{"resources", "prompts", "sampling", "logging"} {
|
||||
if _, ok := out.Capabilities[unimplemented]; ok {
|
||||
t.Errorf("capabilities advertises %q, which is not implemented", unimplemented)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestInitializeEchoesTheRequestID(t *testing.T) {
|
||||
s := newServer(t)
|
||||
// A string id, to prove ids are echoed verbatim rather than coerced.
|
||||
rec := post(t, s, `{"jsonrpc":"2.0","id":"abc-123","method":"initialize","params":{}}`)
|
||||
resp := decode(t, rec)
|
||||
if string(resp.ID) != `"abc-123"` {
|
||||
t.Errorf("id = %s, want \"abc-123\"", resp.ID)
|
||||
}
|
||||
}
|
||||
|
||||
/* ── 2 & 3. tools/list and exact exposure ───────────────────────────────── */
|
||||
|
||||
// wantExposed is the Phase 0 read-only MVP set, written out in full.
|
||||
//
|
||||
// Deliberately a literal rather than a filter over the registry: a test that
|
||||
// derived the expectation the same way the code does would pass no matter what
|
||||
// the rule became. This is the list a human agreed to, and it is the thing that
|
||||
// should fail when the rule changes.
|
||||
var wantExposed = []string{
|
||||
"activity_breakdown",
|
||||
"activity_signals",
|
||||
"available_workers",
|
||||
"candidates_awaiting",
|
||||
"candidates_quality",
|
||||
"hires_performance",
|
||||
"hires_recent",
|
||||
"open_positions",
|
||||
"operations_risk",
|
||||
"positions_risk",
|
||||
"talent_pool",
|
||||
"workforce_attendance",
|
||||
"workforce_coverage",
|
||||
"workforce_overtime",
|
||||
"workforce_training",
|
||||
"workspace_summary",
|
||||
}
|
||||
|
||||
func TestToolsListExposesExactlyTheReadOnlyMVPSet(t *testing.T) {
|
||||
s := newServer(t)
|
||||
rec := post(t, s, `{"jsonrpc":"2.0","id":2,"method":"tools/list"}`)
|
||||
|
||||
var out toolsListResult
|
||||
resultInto(t, rec, &out)
|
||||
|
||||
got := make([]string, 0, len(out.Tools))
|
||||
for _, tool := range out.Tools {
|
||||
got = append(got, tool.Name)
|
||||
}
|
||||
|
||||
if len(got) != len(wantExposed) {
|
||||
t.Fatalf("exposed %d tools, want %d\n got: %v\nwant: %v",
|
||||
len(got), len(wantExposed), got, wantExposed)
|
||||
}
|
||||
for i := range wantExposed {
|
||||
if got[i] != wantExposed[i] {
|
||||
t.Errorf("tool[%d] = %q, want %q", i, got[i], wantExposed[i])
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestToolsListWorksWithoutParams(t *testing.T) {
|
||||
s := newServer(t)
|
||||
// No params key at all — every exposed tool's arguments are optional, so a
|
||||
// bare list call must work.
|
||||
rec := post(t, s, `{"jsonrpc":"2.0","id":3,"method":"tools/list"}`)
|
||||
var out toolsListResult
|
||||
resultInto(t, rec, &out)
|
||||
if len(out.Tools) == 0 {
|
||||
t.Fatal("tools/list returned nothing")
|
||||
}
|
||||
}
|
||||
|
||||
/* ── 4 & 5. write tools and knowledge_search are not exposed ────────────── */
|
||||
|
||||
func TestNoWriteToolIsExposed(t *testing.T) {
|
||||
s := newServer(t)
|
||||
rec := post(t, s, `{"jsonrpc":"2.0","id":4,"method":"tools/list"}`)
|
||||
var out toolsListResult
|
||||
resultInto(t, rec, &out)
|
||||
|
||||
reg := testRegistry(t)
|
||||
for _, published := range out.Tools {
|
||||
tool, ok := reg.Get(published.Name)
|
||||
if !ok {
|
||||
t.Fatalf("tools/list published %q, which is not in the registry", published.Name)
|
||||
}
|
||||
if tool.Effect != tools.EffectRead {
|
||||
t.Errorf("%q is exposed with effect %q; only read tools may be exposed",
|
||||
tool.Name, tool.Effect)
|
||||
}
|
||||
if tool.RequiresConfirmation {
|
||||
t.Errorf("%q is exposed and requires confirmation, which this surface cannot obtain",
|
||||
tool.Name)
|
||||
}
|
||||
// The annotation must agree with the registry, or a client shows a
|
||||
// person the wrong thing before approving.
|
||||
if published.Annotations == nil || !published.Annotations.ReadOnlyHint {
|
||||
t.Errorf("%q is not annotated readOnlyHint", tool.Name)
|
||||
}
|
||||
if published.Annotations != nil && published.Annotations.DestructiveHint {
|
||||
t.Errorf("%q is annotated destructiveHint", tool.Name)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestNamedExclusionsAreNotExposed(t *testing.T) {
|
||||
s := newServer(t)
|
||||
rec := post(t, s, `{"jsonrpc":"2.0","id":5,"method":"tools/list"}`)
|
||||
var out toolsListResult
|
||||
resultInto(t, rec, &out)
|
||||
|
||||
published := map[string]bool{}
|
||||
for _, tool := range out.Tools {
|
||||
published[tool.Name] = true
|
||||
}
|
||||
// The two writes and the one deferred read. Named explicitly, because these
|
||||
// three are the ones a future change is most likely to let through.
|
||||
for _, forbidden := range []string{"assign_worker", "move_application", "knowledge_search"} {
|
||||
if published[forbidden] {
|
||||
t.Errorf("%q must not be exposed", forbidden)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// The registry must still hold the excluded tools — they are withheld from this
|
||||
// surface, not removed from KROW. A test that only checked absence would pass
|
||||
// if somebody deleted them.
|
||||
func TestExcludedToolsStillExistInTheRegistry(t *testing.T) {
|
||||
reg := testRegistry(t)
|
||||
for _, name := range []string{"assign_worker", "move_application", "knowledge_search"} {
|
||||
if _, ok := reg.Get(name); !ok {
|
||||
t.Errorf("%q is missing from the registry entirely", name)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/* ── 6. schemas come from the real registry ─────────────────────────────── */
|
||||
|
||||
func TestToolSchemaIsSourcedFromTheRegistry(t *testing.T) {
|
||||
s := newServer(t)
|
||||
rec := post(t, s, `{"jsonrpc":"2.0","id":6,"method":"tools/list"}`)
|
||||
var out toolsListResult
|
||||
resultInto(t, rec, &out)
|
||||
|
||||
reg := testRegistry(t)
|
||||
for _, published := range out.Tools {
|
||||
tool, _ := reg.Get(published.Name)
|
||||
|
||||
if published.Description != tool.Description {
|
||||
t.Errorf("%q description differs from the registry's", tool.Name)
|
||||
}
|
||||
if published.InputSchema == nil {
|
||||
t.Fatalf("%q published a nil inputSchema", tool.Name)
|
||||
}
|
||||
// Every exposed tool declares a schema, so the published one must be
|
||||
// the registry's own map and not the empty-schema fallback.
|
||||
if tool.InputSchema == nil {
|
||||
t.Errorf("%q has no InputSchema in the registry; expected every exposed tool to declare one", tool.Name)
|
||||
continue
|
||||
}
|
||||
wantJSON, _ := json.Marshal(tool.InputSchema)
|
||||
gotJSON, _ := json.Marshal(published.InputSchema)
|
||||
if string(wantJSON) != string(gotJSON) {
|
||||
t.Errorf("%q schema differs from the registry's\n got: %s\nwant: %s",
|
||||
tool.Name, gotJSON, wantJSON)
|
||||
}
|
||||
// A published schema must be a JSON Schema object, or a client may
|
||||
// refuse the whole list.
|
||||
if published.InputSchema["type"] != "object" {
|
||||
t.Errorf("%q schema type = %v, want \"object\"", tool.Name, published.InputSchema["type"])
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// No exposed tool may take a tenant or principal identifier as an argument.
|
||||
// An argument is something a client chooses; identity is not the client's to
|
||||
// choose. Enforced here rather than by review, because the cost of missing it
|
||||
// once is cross-tenant access.
|
||||
func TestNoToolAcceptsATenantOrPrincipalArgument(t *testing.T) {
|
||||
s := newServer(t)
|
||||
rec := post(t, s, `{"jsonrpc":"2.0","id":7,"method":"tools/list"}`)
|
||||
var out toolsListResult
|
||||
resultInto(t, rec, &out)
|
||||
|
||||
forbidden := []string{
|
||||
"org_id", "orgid", "organization_id", "organisation_id",
|
||||
"tenant_id", "tenantid", "tenant",
|
||||
"user_id", "userid", "principal", "principal_id",
|
||||
"account_id", "caller", "caller_id", "on_behalf_of", "impersonate",
|
||||
}
|
||||
|
||||
for _, tool := range out.Tools {
|
||||
props, ok := tool.InputSchema["properties"].(map[string]any)
|
||||
if !ok {
|
||||
continue
|
||||
}
|
||||
for field := range props {
|
||||
lower := strings.ToLower(field)
|
||||
for _, bad := range forbidden {
|
||||
if lower == bad {
|
||||
t.Errorf("%q accepts %q as an argument; identity must come from the "+
|
||||
"authenticated principal, never from the request", tool.Name, field)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/* ── 7 & 8. unknown method, unknown tool ────────────────────────────────── */
|
||||
|
||||
func TestUnknownMethodReturnsMethodNotFound(t *testing.T) {
|
||||
s := newServer(t)
|
||||
rec := post(t, s, `{"jsonrpc":"2.0","id":8,"method":"resources/list"}`)
|
||||
|
||||
if rec.Code != http.StatusOK {
|
||||
t.Errorf("status = %d, want 200 — a JSON-RPC error is still a successful HTTP exchange", rec.Code)
|
||||
}
|
||||
resp := decode(t, rec)
|
||||
if resp.Error == nil {
|
||||
t.Fatal("expected an error")
|
||||
}
|
||||
if resp.Error.Code != codeMethodNotFound {
|
||||
t.Errorf("code = %d, want %d", resp.Error.Code, codeMethodNotFound)
|
||||
}
|
||||
}
|
||||
|
||||
func TestUnknownToolIsReportedAsAToolError(t *testing.T) {
|
||||
s := newServer(t)
|
||||
rec := post(t, s, `{"jsonrpc":"2.0","id":9,"method":"tools/call",
|
||||
"params":{"name":"no_such_tool","arguments":{}}}`)
|
||||
|
||||
var out toolsCallResult
|
||||
resultInto(t, rec, &out)
|
||||
if !out.IsError {
|
||||
t.Error("expected isError on an unknown tool")
|
||||
}
|
||||
if len(out.Content) == 0 || !strings.Contains(out.Content[0].Text, "mcp.unknown_tool") {
|
||||
t.Errorf("expected mcp.unknown_tool, got %+v", out.Content)
|
||||
}
|
||||
}
|
||||
|
||||
// An unexposed tool must be indistinguishable from a non-existent one, or the
|
||||
// error messages become an inventory of what this surface is withholding.
|
||||
func TestUnexposedToolIsIndistinguishableFromUnknown(t *testing.T) {
|
||||
s := newServer(t)
|
||||
|
||||
unknown := post(t, s, `{"jsonrpc":"2.0","id":10,"method":"tools/call","params":{"name":"definitely_not_a_tool"}}`)
|
||||
withheld := post(t, s, `{"jsonrpc":"2.0","id":10,"method":"tools/call","params":{"name":"assign_worker"}}`)
|
||||
|
||||
var a, b toolsCallResult
|
||||
resultInto(t, unknown, &a)
|
||||
resultInto(t, withheld, &b)
|
||||
|
||||
codeOf := func(r toolsCallResult) string {
|
||||
var payload struct {
|
||||
Error struct {
|
||||
Code string `json:"code"`
|
||||
} `json:"error"`
|
||||
}
|
||||
_ = json.Unmarshal([]byte(r.Content[0].Text), &payload)
|
||||
return payload.Error.Code
|
||||
}
|
||||
if codeOf(a) != codeOf(b) {
|
||||
t.Errorf("an unexposed tool is distinguishable from an unknown one: %q vs %q",
|
||||
codeOf(a), codeOf(b))
|
||||
}
|
||||
if a.IsError != b.IsError {
|
||||
t.Error("isError differs between unknown and unexposed tools")
|
||||
}
|
||||
}
|
||||
|
||||
/* ── 9 & 10. malformed JSON and params ──────────────────────────────────── */
|
||||
|
||||
func TestMalformedJSONReturnsParseError(t *testing.T) {
|
||||
s := newServer(t)
|
||||
rec := post(t, s, `{"jsonrpc":"2.0","id":1,"method":`)
|
||||
|
||||
resp := decode(t, rec)
|
||||
if resp.Error == nil || resp.Error.Code != codeParseError {
|
||||
t.Fatalf("want parse error %d, got %+v", codeParseError, resp.Error)
|
||||
}
|
||||
}
|
||||
|
||||
func TestEnvelopeValidation(t *testing.T) {
|
||||
for name, tc := range map[string]struct {
|
||||
body string
|
||||
want int
|
||||
}{
|
||||
"missing jsonrpc": {`{"id":1,"method":"tools/list"}`, codeInvalidRequest},
|
||||
"wrong jsonrpc": {`{"jsonrpc":"1.0","id":1,"method":"tools/list"}`, codeInvalidRequest},
|
||||
"missing method": {`{"jsonrpc":"2.0","id":1}`, codeInvalidRequest},
|
||||
"empty body": {``, codeInvalidRequest},
|
||||
"object id": {`{"jsonrpc":"2.0","id":{"a":1},"method":"tools/list"}`, codeInvalidRequest},
|
||||
"not an object": {`"a string"`, codeParseError},
|
||||
"params as array": {`{"jsonrpc":"2.0","id":1,"method":"tools/call","params":[1,2]}`, codeInvalidParams},
|
||||
"params wrong type": {`{"jsonrpc":"2.0","id":1,"method":"tools/call","params":{"name":42}}`, codeInvalidParams},
|
||||
"missing tool name": {`{"jsonrpc":"2.0","id":1,"method":"tools/call","params":{}}`, codeInvalidParams},
|
||||
} {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
rec := post(t, newServer(t), tc.body)
|
||||
resp := decode(t, rec)
|
||||
if resp.Error == nil {
|
||||
t.Fatalf("expected an error, got %s", rec.Body.String())
|
||||
}
|
||||
if resp.Error.Code != tc.want {
|
||||
t.Errorf("code = %d, want %d (%s)", resp.Error.Code, tc.want, resp.Error.Message)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
/* ── 11. batching is explicitly rejected ────────────────────────────────── */
|
||||
|
||||
func TestBatchRequestsAreExplicitlyRejected(t *testing.T) {
|
||||
s := newServer(t)
|
||||
rec := post(t, s, `[{"jsonrpc":"2.0","id":1,"method":"tools/list"},
|
||||
{"jsonrpc":"2.0","id":2,"method":"tools/list"}]`)
|
||||
|
||||
resp := decode(t, rec)
|
||||
if resp.Error == nil {
|
||||
t.Fatal("a batch must be refused")
|
||||
}
|
||||
if resp.Error.Code != codeInvalidRequest {
|
||||
t.Errorf("code = %d, want %d", resp.Error.Code, codeInvalidRequest)
|
||||
}
|
||||
// The refusal must say WHY, so a client can fall back to sending singly
|
||||
// rather than retrying the same batch forever.
|
||||
detail, _ := resp.Error.Data.(string)
|
||||
if !strings.Contains(detail, "batch") {
|
||||
t.Errorf("the refusal does not mention batching: %q", detail)
|
||||
}
|
||||
}
|
||||
|
||||
/* ── 12. no panic on malformed input ────────────────────────────────────── */
|
||||
|
||||
func TestNoPanicOnHostileInput(t *testing.T) {
|
||||
// Each of these has crashed a hand-written JSON-RPC server somewhere.
|
||||
bodies := []string{
|
||||
``,
|
||||
` `,
|
||||
`null`,
|
||||
`[]`,
|
||||
`[[[[[]]]]]`,
|
||||
`{`,
|
||||
`}`,
|
||||
`{"jsonrpc":"2.0"}`,
|
||||
`{"jsonrpc":null,"id":null,"method":null}`,
|
||||
`{"jsonrpc":"2.0","id":1,"method":"tools/call","params":null}`,
|
||||
`{"jsonrpc":"2.0","id":1,"method":"tools/call","params":{"name":""}}`,
|
||||
`{"jsonrpc":"2.0","id":1,"method":"tools/call","params":{"name":"activity_breakdown","arguments":"not-an-object"}}`,
|
||||
`{"jsonrpc":"2.0","id":1,"method":"tools/call","params":{"name":"activity_breakdown","arguments":[]}}`,
|
||||
`{"jsonrpc":"2.0","id":[1,2,3],"method":"tools/list"}`,
|
||||
`{"jsonrpc":"2.0","id":1,"method":"` + strings.Repeat("A", 10_000) + `"}`,
|
||||
`{"jsonrpc":"2.0","id":1,"method":"tools/list","params":{"cursor":` + strings.Repeat("[", 500) + `}}`,
|
||||
"\x00\x01\x02",
|
||||
}
|
||||
|
||||
s := newServer(t)
|
||||
for i, body := range bodies {
|
||||
// A panic escaping here fails the test by crashing it, which is the
|
||||
// assertion: the handler must answer every one of these.
|
||||
rec := post(t, s, body)
|
||||
if rec.Code < 200 || rec.Code >= 600 {
|
||||
t.Errorf("body %d produced status %d", i, rec.Code)
|
||||
}
|
||||
if rec.Body.Len() > 0 && rec.Code != http.StatusAccepted {
|
||||
var resp response
|
||||
if err := json.Unmarshal(rec.Body.Bytes(), &resp); err != nil {
|
||||
t.Errorf("body %d produced a non-JSON response: %s", i, rec.Body.String())
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/* ── Transport ──────────────────────────────────────────────────────────── */
|
||||
|
||||
func TestOnlyPOSTIsAccepted(t *testing.T) {
|
||||
s := newServer(t)
|
||||
for _, method := range []string{http.MethodGet, http.MethodPut, http.MethodDelete, http.MethodPatch} {
|
||||
req := httptest.NewRequest(method, "/mcp", nil)
|
||||
rec := httptest.NewRecorder()
|
||||
s.Handler().ServeHTTP(rec, req)
|
||||
|
||||
if rec.Code != http.StatusMethodNotAllowed {
|
||||
t.Errorf("%s: status = %d, want 405", method, rec.Code)
|
||||
}
|
||||
if allow := rec.Header().Get("Allow"); allow != http.MethodPost {
|
||||
t.Errorf("%s: Allow = %q, want POST", method, allow)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestNonJSONContentTypeIsRefused(t *testing.T) {
|
||||
s := newServer(t)
|
||||
req := httptest.NewRequest(http.MethodPost, "/mcp",
|
||||
strings.NewReader(`{"jsonrpc":"2.0","id":1,"method":"tools/list"}`))
|
||||
req.Header.Set("Content-Type", "application/x-www-form-urlencoded")
|
||||
rec := httptest.NewRecorder()
|
||||
s.Handler().ServeHTTP(rec, req)
|
||||
|
||||
if rec.Code != http.StatusUnsupportedMediaType {
|
||||
t.Errorf("status = %d, want 415", rec.Code)
|
||||
}
|
||||
}
|
||||
|
||||
func TestOversizedBodyIsRefused(t *testing.T) {
|
||||
s := newServer(t)
|
||||
huge := `{"jsonrpc":"2.0","id":1,"method":"tools/list","params":{"cursor":"` +
|
||||
strings.Repeat("x", MaxRequestBytes+1) + `"}}`
|
||||
rec := post(t, s, huge)
|
||||
|
||||
if rec.Code != http.StatusRequestEntityTooLarge {
|
||||
t.Errorf("status = %d, want 413", rec.Code)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNotificationGetsNoBody(t *testing.T) {
|
||||
s := newServer(t)
|
||||
// No id — a notification. The spec forbids a response.
|
||||
rec := post(t, s, `{"jsonrpc":"2.0","method":"notifications/initialized"}`)
|
||||
|
||||
if rec.Code != http.StatusAccepted {
|
||||
t.Errorf("status = %d, want 202", rec.Code)
|
||||
}
|
||||
if rec.Body.Len() != 0 {
|
||||
t.Errorf("a notification was answered with a body: %s", rec.Body.String())
|
||||
}
|
||||
}
|
||||
|
||||
/* ── Phase 1 boundary ───────────────────────────────────────────────────── */
|
||||
|
||||
// Handle must refuse a nil identity even when called directly.
|
||||
//
|
||||
// The transport answers 401 before this point in the ordinary case, so this is
|
||||
// the SECOND gate: a future caller of Handle that forgets to authenticate must
|
||||
// fail closed rather than dispatch as nobody. It exists to fail loudly if
|
||||
// somebody later supplies a default identity to "make it work".
|
||||
func TestHandleRefusesANilIdentity(t *testing.T) {
|
||||
s := newServer(t)
|
||||
req := request{
|
||||
JSONRPC: jsonRPCVersion,
|
||||
ID: json.RawMessage(`11`),
|
||||
Method: "tools/call",
|
||||
Params: json.RawMessage(`{"name":"activity_breakdown","arguments":{}}`),
|
||||
}
|
||||
result, rpcErr := s.Handle(context.Background(), nil, req)
|
||||
if rpcErr != nil {
|
||||
t.Fatalf("unexpected rpc error: %v", rpcErr)
|
||||
}
|
||||
out, ok := result.(toolsCallResult)
|
||||
if !ok {
|
||||
t.Fatalf("unexpected result type %T", result)
|
||||
}
|
||||
if !out.IsError || !strings.Contains(out.Content[0].Text, "mcp.unauthenticated") {
|
||||
t.Errorf("expected mcp.unauthenticated, got: %+v", out.Content)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPingIsCheap(t *testing.T) {
|
||||
s := newServer(t)
|
||||
rec := post(t, s, `{"jsonrpc":"2.0","id":12,"method":"ping"}`)
|
||||
var out map[string]any
|
||||
resultInto(t, rec, &out)
|
||||
if len(out) != 0 {
|
||||
t.Errorf("ping returned %v, want an empty result", out)
|
||||
}
|
||||
}
|
||||
314
go-api/internal/mcpserver/orglimit_test.go
Normal file
314
go-api/internal/mcpserver/orglimit_test.go
Normal file
@@ -0,0 +1,314 @@
|
||||
package mcpserver
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/http"
|
||||
"strconv"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
// The per-organisation ceiling.
|
||||
//
|
||||
// The property under test is not "a limit exists" but "the limit is keyed by an
|
||||
// organisation the CALLER CANNOT CHOOSE". Every test below therefore checks
|
||||
// which bucket was charged, not merely that something was refused.
|
||||
|
||||
// countingOrgLimiter records which org was charged and refuses past a limit.
|
||||
type countingOrgLimiter struct {
|
||||
mu sync.Mutex
|
||||
counts map[string]int
|
||||
limit int
|
||||
err error
|
||||
}
|
||||
|
||||
func newCountingOrgLimiter(limit int) *countingOrgLimiter {
|
||||
return &countingOrgLimiter{counts: map[string]int{}, limit: limit}
|
||||
}
|
||||
|
||||
func (c *countingOrgLimiter) AllowOrg(_ context.Context, orgID string) (bool, time.Duration, error) {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
if c.err != nil {
|
||||
return false, time.Minute, c.err
|
||||
}
|
||||
c.counts[orgID]++
|
||||
return c.counts[orgID] <= c.limit, 30 * time.Second, nil
|
||||
}
|
||||
|
||||
func (c *countingOrgLimiter) count(orgID string) int {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
return c.counts[orgID]
|
||||
}
|
||||
|
||||
func (c *countingOrgLimiter) buckets() int {
|
||||
c.mu.Lock()
|
||||
defer c.mu.Unlock()
|
||||
return len(c.counts)
|
||||
}
|
||||
|
||||
// orgLimitEnv is the matrix fixture plus a counting limiter.
|
||||
func orgLimitEnv(t *testing.T, limit int) (*matrixEnv, *countingOrgLimiter) {
|
||||
t.Helper()
|
||||
m := newMatrix(t)
|
||||
limiter := newCountingOrgLimiter(limit)
|
||||
m.server = m.server.WithOrgLimiter(limiter)
|
||||
return m, limiter
|
||||
}
|
||||
|
||||
/* ── Each organisation gets its own bucket ──────────────────────────────── */
|
||||
|
||||
func TestEachOrganisationIsCountedSeparately(t *testing.T) {
|
||||
m, limiter := orgLimitEnv(t, 100)
|
||||
|
||||
for i := 0; i < 3; i++ {
|
||||
m.call(t, "tok-a-admin", "activity_breakdown", "{}")
|
||||
}
|
||||
for i := 0; i < 5; i++ {
|
||||
m.call(t, "tok-b-admin", "activity_breakdown", "{}")
|
||||
}
|
||||
|
||||
if got := limiter.count(m.a.orgID); got != 3 {
|
||||
t.Errorf("org A charged %d, want 3", got)
|
||||
}
|
||||
if got := limiter.count(m.b.orgID); got != 5 {
|
||||
t.Errorf("org B charged %d, want 5", got)
|
||||
}
|
||||
if limiter.buckets() != 2 {
|
||||
t.Errorf("%d buckets, want 2 — the two tenants shared a counter", limiter.buckets())
|
||||
}
|
||||
}
|
||||
|
||||
// One organisation exhausting its quota must not affect another's.
|
||||
func TestOneOrganisationCannotConsumeAnothersQuota(t *testing.T) {
|
||||
m, limiter := orgLimitEnv(t, 3)
|
||||
|
||||
// Burn org A's entire budget and then some.
|
||||
for i := 0; i < 10; i++ {
|
||||
m.call(t, "tok-a-admin", "activity_breakdown", "{}")
|
||||
}
|
||||
|
||||
// Org B must be untouched.
|
||||
body := m.call(t, "tok-b-admin", "activity_breakdown", "{}")
|
||||
if isRateLimited(body) {
|
||||
t.Error("org B was refused because org A exhausted its quota")
|
||||
}
|
||||
if got := limiter.count(m.b.orgID); got != 1 {
|
||||
t.Errorf("org B charged %d, want 1", got)
|
||||
}
|
||||
}
|
||||
|
||||
/* ── The bucket cannot be chosen by the caller ──────────────────────────── */
|
||||
|
||||
// Every channel a client controls, against the ORG LIMITER specifically. A
|
||||
// request naming org B must still be charged to org A.
|
||||
func TestTheOrgBucketCannotBeSelectedByTheRequest(t *testing.T) {
|
||||
for name, tc := range map[string]struct {
|
||||
args string
|
||||
mutate func(*http.Request)
|
||||
path string
|
||||
}{
|
||||
"org_id argument": {
|
||||
args: `{"org_id":"OTHER"}`,
|
||||
},
|
||||
"tenant_id argument": {
|
||||
args: `{"tenant_id":"OTHER","organization_id":"OTHER"}`,
|
||||
},
|
||||
"identity headers": {
|
||||
args: `{}`,
|
||||
mutate: func(r *http.Request) {
|
||||
for _, h := range []string{"X-Org-Id", "X-Tenant-Id", "X-Organization-Id"} {
|
||||
r.Header.Set(h, "OTHER")
|
||||
}
|
||||
},
|
||||
},
|
||||
"query string": {
|
||||
args: `{}`,
|
||||
path: "/mcp?org_id=OTHER&tenant_id=OTHER",
|
||||
},
|
||||
} {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
m, limiter := orgLimitEnv(t, 100)
|
||||
args := replaceAll(tc.args, "OTHER", m.b.orgID)
|
||||
path := tc.path
|
||||
if path == "" {
|
||||
path = "/mcp"
|
||||
}
|
||||
|
||||
m.callWith(t, "tok-a-admin", "activity_breakdown", args, path, tc.mutate)
|
||||
|
||||
if got := limiter.count(m.a.orgID); got != 1 {
|
||||
t.Errorf("org A charged %d, want 1 — the caller's own org must be charged", got)
|
||||
}
|
||||
if got := limiter.count(m.b.orgID); got != 0 {
|
||||
t.Errorf("org B charged %d, want 0 — the request selected another tenant's bucket", got)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// _meta at both levels, which is the channel most likely to be trusted by
|
||||
// accident because it is "protocol" rather than "arguments".
|
||||
func TestMetaCannotSelectTheOrgBucket(t *testing.T) {
|
||||
m, limiter := orgLimitEnv(t, 100)
|
||||
|
||||
body := `{"jsonrpc":"2.0","id":1,"method":"tools/call",` +
|
||||
`"_meta":{"org_id":"` + m.b.orgID + `"},` +
|
||||
`"params":{"name":"activity_breakdown","arguments":{},` +
|
||||
`"_meta":{"org_id":"` + m.b.orgID + `","tenant_id":"` + m.b.orgID + `"}}}`
|
||||
|
||||
m.raw(t, "tok-a-admin", body, "/mcp", nil)
|
||||
|
||||
if got := limiter.count(m.a.orgID); got != 1 {
|
||||
t.Errorf("org A charged %d, want 1", got)
|
||||
}
|
||||
if got := limiter.count(m.b.orgID); got != 0 {
|
||||
t.Errorf("org B charged %d, want 0 — _meta selected another tenant's bucket", got)
|
||||
}
|
||||
}
|
||||
|
||||
/* ── Refusal behaviour ──────────────────────────────────────────────────── */
|
||||
|
||||
func TestOverTheOrgLimitReturns429WithRetryAfter(t *testing.T) {
|
||||
m, _ := orgLimitEnv(t, 2)
|
||||
|
||||
// Two are allowed.
|
||||
for i := 0; i < 2; i++ {
|
||||
if rec := m.callRec(t, "tok-a-admin", "activity_breakdown", "{}"); rec.Code != http.StatusOK {
|
||||
t.Fatalf("call %d: status = %d, want 200", i+1, rec.Code)
|
||||
}
|
||||
}
|
||||
|
||||
rec := m.callRec(t, "tok-a-admin", "activity_breakdown", "{}")
|
||||
if rec.Code != http.StatusTooManyRequests {
|
||||
t.Fatalf("status = %d, want 429", rec.Code)
|
||||
}
|
||||
|
||||
retry := rec.Header().Get("Retry-After")
|
||||
if retry == "" {
|
||||
t.Fatal("a 429 carried no Retry-After")
|
||||
}
|
||||
secs, err := strconv.Atoi(retry)
|
||||
if err != nil || secs < 1 {
|
||||
t.Errorf("Retry-After = %q, want a positive whole number of seconds", retry)
|
||||
}
|
||||
|
||||
// The refusal must not describe the quota or name the organisation — how
|
||||
// much a tenant has spent is not something one caller learns from a 429.
|
||||
body := rec.Body.String()
|
||||
if containsAny(body, []string{m.a.orgID, "5000", "quota", "remaining"}) {
|
||||
t.Errorf("the 429 body leaks quota or tenant detail: %s", body)
|
||||
}
|
||||
}
|
||||
|
||||
// The ceiling is checked BEFORE the tool runs, so a refused call costs no
|
||||
// database work.
|
||||
func TestTheOrgLimitIsCheckedBeforeTheToolRuns(t *testing.T) {
|
||||
m, _ := orgLimitEnv(t, 0) // nothing is allowed
|
||||
|
||||
rec := m.callRec(t, "tok-a-admin", "activity_breakdown", "{}")
|
||||
if rec.Code != http.StatusTooManyRequests {
|
||||
t.Fatalf("status = %d, want 429", rec.Code)
|
||||
}
|
||||
// A tool that had run would have produced a result payload.
|
||||
if containsAny(rec.Body.String(), []string{"totalEvents", "distinctKinds"}) {
|
||||
t.Error("the tool ran despite the organisation being over its limit")
|
||||
}
|
||||
}
|
||||
|
||||
// Unauthenticated requests must be refused before the limiter is consulted —
|
||||
// otherwise an anonymous caller could burn a tenant's quota.
|
||||
func TestTheOrgLimiterIsNotConsultedWithoutAuthentication(t *testing.T) {
|
||||
m, limiter := orgLimitEnv(t, 100)
|
||||
|
||||
m.raw(t, "", `{"jsonrpc":"2.0","id":1,"method":"tools/call","params":`+
|
||||
`{"name":"activity_breakdown","arguments":{}}}`, "/mcp", nil)
|
||||
|
||||
if limiter.buckets() != 0 {
|
||||
t.Errorf("%d buckets charged by an unauthenticated request, want 0", limiter.buckets())
|
||||
}
|
||||
}
|
||||
|
||||
/* ── Concurrency and failure ────────────────────────────────────────────── */
|
||||
|
||||
// Concurrent calls must be counted atomically: the limiter's own contract, here
|
||||
// exercised through the full MCP path. Run with -race.
|
||||
func TestConcurrentCallsAreCountedAtomically(t *testing.T) {
|
||||
m, limiter := orgLimitEnv(t, 1000)
|
||||
|
||||
const callers = 30
|
||||
var wg sync.WaitGroup
|
||||
for i := 0; i < callers; i++ {
|
||||
wg.Add(1)
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
m.call(t, "tok-a-admin", "activity_breakdown", "{}")
|
||||
}()
|
||||
}
|
||||
wg.Wait()
|
||||
|
||||
if got := limiter.count(m.a.orgID); got != callers {
|
||||
t.Errorf("org A charged %d, want %d — increments were lost", got, callers)
|
||||
}
|
||||
}
|
||||
|
||||
// A limiter that errors must not let the call through: fail closed.
|
||||
func TestAFailingOrgLimiterRefusesTheCall(t *testing.T) {
|
||||
m, limiter := orgLimitEnv(t, 100)
|
||||
limiter.err = errString("limiter unavailable")
|
||||
|
||||
rec := m.callRec(t, "tok-a-admin", "activity_breakdown", "{}")
|
||||
if rec.Code != http.StatusTooManyRequests {
|
||||
t.Errorf("status = %d, want 429 — a limiter that cannot count must not permit the call", rec.Code)
|
||||
}
|
||||
}
|
||||
|
||||
// With no limiter installed there is no ceiling, and nothing breaks.
|
||||
func TestNoOrgLimiterMeansNoCeiling(t *testing.T) {
|
||||
m := newMatrix(t) // no WithOrgLimiter
|
||||
for i := 0; i < 20; i++ {
|
||||
if rec := m.callRec(t, "tok-a-admin", "activity_breakdown", "{}"); rec.Code != http.StatusOK {
|
||||
t.Fatalf("call %d: status = %d, want 200", i+1, rec.Code)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/* ── Helpers ────────────────────────────────────────────────────────────── */
|
||||
|
||||
type errString string
|
||||
|
||||
func (e errString) Error() string { return string(e) }
|
||||
|
||||
func containsAny(s string, needles []string) bool {
|
||||
for _, n := range needles {
|
||||
if n != "" && contains(s, n) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func contains(s, sub string) bool { return len(sub) > 0 && indexOf(s, sub) >= 0 }
|
||||
|
||||
func indexOf(s, sub string) int {
|
||||
for i := 0; i+len(sub) <= len(s); i++ {
|
||||
if s[i:i+len(sub)] == sub {
|
||||
return i
|
||||
}
|
||||
}
|
||||
return -1
|
||||
}
|
||||
|
||||
func replaceAll(s, old, new string) string {
|
||||
out := ""
|
||||
for {
|
||||
i := indexOf(s, old)
|
||||
if i < 0 {
|
||||
return out + s
|
||||
}
|
||||
out += s[:i] + new
|
||||
s = s[i+len(old):]
|
||||
}
|
||||
}
|
||||
387
go-api/internal/mcpserver/server.go
Normal file
387
go-api/internal/mcpserver/server.go
Normal file
@@ -0,0 +1,387 @@
|
||||
package mcpserver
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"log/slog"
|
||||
"time"
|
||||
|
||||
"github.com/krow/krow-backend/go-api/internal/authctx"
|
||||
"github.com/krow/krow-backend/go-api/internal/tools"
|
||||
)
|
||||
|
||||
// ProtocolVersion is the MCP revision this server implements.
|
||||
//
|
||||
// Echoed back from initialize. A client asking for a different revision is not
|
||||
// refused: the spec's negotiation is that the server states what it speaks and
|
||||
// the client decides whether it can work with that. Refusing would turn a
|
||||
// version skew into an outage where it is usually a compatible difference.
|
||||
const ProtocolVersion = "2025-06-18"
|
||||
|
||||
// ServerName and ServerVersion identify this implementation to a client.
|
||||
const (
|
||||
ServerName = "krow-mcp"
|
||||
ServerVersion = "0.1.0"
|
||||
)
|
||||
|
||||
// maxToolArgumentBytes bounds one tool call's arguments.
|
||||
//
|
||||
// The same reasoning as runs.go's maxRunRequestBytes: arguments are a handful
|
||||
// of scalars against a schema that sets additionalProperties:false, so anything
|
||||
// large is either a mistake or an attempt to push text past the tool layer into
|
||||
// a prompt. The transport bounds the whole body too; this bounds the part that
|
||||
// reaches a handler.
|
||||
const maxToolArgumentBytes = 64 << 10
|
||||
|
||||
// Server answers MCP methods against the existing tool registry.
|
||||
//
|
||||
// It holds a *tools.Registry and nothing else that matters. There is no second
|
||||
// registry, no adapter table and no per-tool code in this package: what MCP
|
||||
// publishes is what the registry holds, filtered by the rule in tools.go.
|
||||
type Server struct {
|
||||
reg *tools.Registry
|
||||
log *slog.Logger
|
||||
|
||||
// tokens resolves a bearer credential into an identity. Nil means this
|
||||
// surface cannot authenticate anyone, and every call is refused — see
|
||||
// ErrNoAuthenticator. Failing closed is the only safe default for a field
|
||||
// whose absence would otherwise mean "let everyone in".
|
||||
tokens TokenAuthenticator
|
||||
|
||||
// resourceMetadataURL is where a 401 points a client so it can begin
|
||||
// discovery. Empty means the challenge carries no pointer, which is a
|
||||
// valid but less useful 401: a client then has nowhere to look.
|
||||
resourceMetadataURL string
|
||||
|
||||
// orgLimiter bounds tool calls per organisation.
|
||||
//
|
||||
// HERE rather than in the HTTP middleware, and that placement is the whole
|
||||
// point: an organisation is not knowable until the bearer token has been
|
||||
// resolved to a user and that user's row read. A middleware running before
|
||||
// authentication could only key by something the CLIENT supplied, which is
|
||||
// precisely the identity this surface refuses to trust.
|
||||
//
|
||||
// Nil means no per-organisation ceiling, which is the correct default for
|
||||
// a deployment that has not configured one.
|
||||
orgLimiter OrgLimiter
|
||||
}
|
||||
|
||||
// OrgLimiter bounds how much one organisation may ask for.
|
||||
//
|
||||
// Takes an org id that the caller has already established from an authenticated
|
||||
// identity. It cannot be handed anything from a request, because the only
|
||||
// caller is dispatch, which has an authctx.Identity and nothing else.
|
||||
type OrgLimiter interface {
|
||||
// AllowOrg reports whether this organisation may make another call, and
|
||||
// how long until its window rolls over.
|
||||
AllowOrg(ctx context.Context, orgID string) (allowed bool, retryAfter time.Duration, err error)
|
||||
}
|
||||
|
||||
// WithOrgLimiter installs the per-organisation ceiling.
|
||||
func (s *Server) WithOrgLimiter(l OrgLimiter) *Server {
|
||||
s.orgLimiter = l
|
||||
return s
|
||||
}
|
||||
|
||||
// WithResourceMetadataURL sets the RFC 9728 document a 401 points at.
|
||||
//
|
||||
// Supplied by the caller rather than derived here, because this package does
|
||||
// not know its own deployment's URLs and must not invent them. A hardcoded
|
||||
// hostname would be one deployment's identity baked into every other one.
|
||||
func (s *Server) WithResourceMetadataURL(u string) *Server {
|
||||
s.resourceMetadataURL = u
|
||||
return s
|
||||
}
|
||||
|
||||
// challenge builds the WWW-Authenticate header for a 401.
|
||||
//
|
||||
// RFC 9728 section 5.1: the client reads `resource_metadata` from here to find
|
||||
// the protected-resource document, and from there the authorization server.
|
||||
// Without the parameter a compliant client has a 401 and nowhere to go, which
|
||||
// is why this is the difference between "authentication failed" and "here is
|
||||
// how to authenticate".
|
||||
func (s *Server) challenge() string {
|
||||
c := `Bearer realm="` + ServerName + `"`
|
||||
if s.resourceMetadataURL != "" {
|
||||
c += `, resource_metadata="` + s.resourceMetadataURL + `"`
|
||||
}
|
||||
return c
|
||||
}
|
||||
|
||||
// New builds a server over an existing registry.
|
||||
//
|
||||
// The registry is the one the rest of the service already built — the caller
|
||||
// passes runtime.DefaultTools(...)'s result, the same value the HTTP server
|
||||
// uses for its author catalogue. Taking it as a parameter rather than building
|
||||
// one here is what guarantees there is only ever one.
|
||||
func New(reg *tools.Registry, tokens TokenAuthenticator, log *slog.Logger) *Server {
|
||||
if log == nil {
|
||||
log = slog.Default()
|
||||
}
|
||||
return &Server{reg: reg, tokens: tokens, log: log}
|
||||
}
|
||||
|
||||
// Handle dispatches one parsed JSON-RPC request on behalf of an identity.
|
||||
//
|
||||
// The identity is a PARAMETER, not something read from the context, and that is
|
||||
// the security property rather than a style choice. If this function resolved
|
||||
// the caller from ctx, then mounting the endpoint behind the cookie middleware
|
||||
// would make a browser session sufficient to call MCP tools — the middleware
|
||||
// puts an Identity in the context, and this code would find it. Taking it as an
|
||||
// argument means only the MCP transport's own bearer authentication can supply
|
||||
// one. See auth.go.
|
||||
//
|
||||
// A nil identity means unauthenticated. The three handshake methods are allowed
|
||||
// without one; tools/call is not.
|
||||
//
|
||||
// Returns a result or an error, never both. Notifications are handled by the
|
||||
// transport, which discards whatever comes back.
|
||||
func (s *Server) Handle(ctx context.Context, ident *authctx.Identity, req request) (any, *rpcError) {
|
||||
switch req.Method {
|
||||
case "initialize":
|
||||
return s.handleInitialize(req.Params)
|
||||
case "notifications/initialized":
|
||||
// The client telling us it is ready. Nothing to do, and answering is
|
||||
// not required — it arrives as a notification.
|
||||
return map[string]any{}, nil
|
||||
case "ping":
|
||||
// Cheap liveness, defined by the spec as an empty result. Costs nothing
|
||||
// and saves a client from using tools/list as a heartbeat.
|
||||
return map[string]any{}, nil
|
||||
case "tools/list":
|
||||
return s.handleToolsList(req.Params)
|
||||
case "tools/call":
|
||||
return s.handleToolsCall(ctx, ident, req.Params)
|
||||
default:
|
||||
return nil, errMethodNotFound(req.Method)
|
||||
}
|
||||
}
|
||||
|
||||
/* ── initialize ─────────────────────────────────────────────────────────── */
|
||||
|
||||
type initializeParams struct {
|
||||
ProtocolVersion string `json:"protocolVersion"`
|
||||
Capabilities json.RawMessage `json:"capabilities"`
|
||||
ClientInfo struct {
|
||||
Name string `json:"name"`
|
||||
Version string `json:"version"`
|
||||
} `json:"clientInfo"`
|
||||
}
|
||||
|
||||
type initializeResult struct {
|
||||
ProtocolVersion string `json:"protocolVersion"`
|
||||
Capabilities map[string]any `json:"capabilities"`
|
||||
ServerInfo map[string]any `json:"serverInfo"`
|
||||
Instructions string `json:"instructions,omitempty"`
|
||||
}
|
||||
|
||||
// handleInitialize answers the opening handshake.
|
||||
//
|
||||
// Declares exactly one capability, because exactly one is implemented. A server
|
||||
// that advertised resources or prompts here would be promising methods that
|
||||
// answer method-not-found, and a client would reasonably call them.
|
||||
//
|
||||
// listChanged is false: the tool set is fixed at process start by the registry,
|
||||
// so there is no change to notify anyone about.
|
||||
func (s *Server) handleInitialize(raw json.RawMessage) (any, *rpcError) {
|
||||
var p initializeParams
|
||||
if err := decodeParams(raw, &p); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
s.log.Info("mcp initialize",
|
||||
"client_name", p.ClientInfo.Name,
|
||||
"client_version", p.ClientInfo.Version,
|
||||
"client_protocol", p.ProtocolVersion,
|
||||
"server_protocol", ProtocolVersion)
|
||||
|
||||
return initializeResult{
|
||||
ProtocolVersion: ProtocolVersion,
|
||||
Capabilities: map[string]any{
|
||||
"tools": map[string]any{"listChanged": false},
|
||||
},
|
||||
ServerInfo: map[string]any{
|
||||
"name": ServerName,
|
||||
"version": ServerVersion,
|
||||
},
|
||||
Instructions: "Read-only access to KROW workforce and hiring data. " +
|
||||
"Every call is scoped to the authenticated user's organisation and role; " +
|
||||
"results are structured data for you to summarise, not prose.",
|
||||
}, nil
|
||||
}
|
||||
|
||||
/* ── tools/list ─────────────────────────────────────────────────────────── */
|
||||
|
||||
type toolsListResult struct {
|
||||
Tools []mcpTool `json:"tools"`
|
||||
}
|
||||
|
||||
// handleToolsList publishes the exposed tools, straight from the registry.
|
||||
func (s *Server) handleToolsList(raw json.RawMessage) (any, *rpcError) {
|
||||
// Params are optional here (cursor, for pagination this server does not
|
||||
// need), but a malformed object is still worth refusing rather than
|
||||
// ignoring — silently accepting nonsense trains a client to send it.
|
||||
var p struct {
|
||||
Cursor string `json:"cursor,omitempty"`
|
||||
}
|
||||
if err := decodeParams(raw, &p); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
infos := exposed(s.reg)
|
||||
out := make([]mcpTool, 0, len(infos))
|
||||
for _, info := range infos {
|
||||
out = append(out, toMCPTool(info))
|
||||
}
|
||||
return toolsListResult{Tools: out}, nil
|
||||
}
|
||||
|
||||
/* ── tools/call ─────────────────────────────────────────────────────────── */
|
||||
|
||||
type toolsCallParams struct {
|
||||
Name string `json:"name"`
|
||||
Arguments json.RawMessage `json:"arguments,omitempty"`
|
||||
}
|
||||
|
||||
// toolsCallResult is MCP's shape for a tool's output.
|
||||
//
|
||||
// IsError is part of the RESULT, not a JSON-RPC error: a tool that refused is
|
||||
// not a protocol fault, and reporting it as one would deny the model the chance
|
||||
// to read the refusal and do something sensible. It is the same distinction
|
||||
// tools.Result already draws, and runs.go draws for terminations.
|
||||
type toolsCallResult struct {
|
||||
Content []contentBlock `json:"content"`
|
||||
IsError bool `json:"isError,omitempty"`
|
||||
}
|
||||
|
||||
type contentBlock struct {
|
||||
Type string `json:"type"`
|
||||
Text string `json:"text"`
|
||||
}
|
||||
|
||||
func textResult(payload any, isError bool) (toolsCallResult, *rpcError) {
|
||||
encoded, err := json.MarshalIndent(payload, "", " ")
|
||||
if err != nil {
|
||||
return toolsCallResult{}, errInternal()
|
||||
}
|
||||
return toolsCallResult{
|
||||
Content: []contentBlock{{Type: "text", Text: string(encoded)}},
|
||||
IsError: isError,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// handleToolsCall validates a call and dispatches it through the registry.
|
||||
//
|
||||
// The order is: exposure, then bounds, then identity, then dispatch.
|
||||
//
|
||||
// Exposure is checked BEFORE identity on purpose. "There is no such tool here"
|
||||
// does not depend on who is asking, and answering it first means the surface's
|
||||
// tool inventory is not something an attacker can probe by comparing an
|
||||
// authenticated 404 against an unauthenticated 401.
|
||||
//
|
||||
// Authentication is the transport's job and has already happened by the time
|
||||
// this runs; `ident` is nil only when it failed or was never attempted. This
|
||||
// function does not read the ambient context for a caller — see Handle.
|
||||
func (s *Server) handleToolsCall(ctx context.Context, ident *authctx.Identity, raw json.RawMessage) (any, *rpcError) {
|
||||
var p toolsCallParams
|
||||
if err := decodeParams(raw, &p); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if p.Name == "" {
|
||||
return nil, errInvalidParams("name is required")
|
||||
}
|
||||
if len(p.Arguments) > maxToolArgumentBytes {
|
||||
return nil, errInvalidParams("arguments are too large")
|
||||
}
|
||||
|
||||
// Unexposed and unknown are the SAME answer, deliberately. See
|
||||
// isExposedName — distinguishing them inventories what this surface is
|
||||
// hiding.
|
||||
if !isExposedName(s.reg, p.Name) {
|
||||
s.log.Warn("mcp tool call refused", "tool", p.Name, "reason", "not_exposed")
|
||||
return textResult(map[string]any{
|
||||
"error": map[string]any{
|
||||
"code": "mcp.unknown_tool",
|
||||
"message": "there is no tool called " + p.Name + " on this surface",
|
||||
},
|
||||
}, true)
|
||||
}
|
||||
|
||||
// Unauthenticated calls never reach a handler. The transport answers 401
|
||||
// before this point in the ordinary case; this is the second gate, so that
|
||||
// a future caller of Handle that forgets to authenticate fails closed
|
||||
// rather than dispatching as nobody.
|
||||
if ident == nil {
|
||||
s.log.Warn("mcp tool call refused", "tool", p.Name, "reason", "no_identity")
|
||||
return textResult(map[string]any{
|
||||
"error": map[string]any{
|
||||
"code": "mcp.unauthenticated",
|
||||
"message": "this call is not authenticated",
|
||||
},
|
||||
}, true)
|
||||
}
|
||||
|
||||
return s.dispatch(ctx, *ident, p)
|
||||
}
|
||||
|
||||
// dispatch runs the tool through the existing registry.
|
||||
//
|
||||
// This is the only place this package touches the tool layer, and it is four
|
||||
// lines on purpose. Everything that decides what comes back — the policy table,
|
||||
// the org pre-filter, the row scopes, the opaque denial, the truncation — is
|
||||
// inside Dispatch and the handler beneath it, unchanged and unreachable from
|
||||
// here.
|
||||
//
|
||||
// The principal is the caller's, from the context. It is never read from
|
||||
// params: an MCP client that could name its own principal could read anything,
|
||||
// which is the bug I1 exists to prevent.
|
||||
func (s *Server) dispatch(ctx context.Context, identity authctx.Identity, p toolsCallParams) (any, *rpcError) {
|
||||
// The per-organisation ceiling, checked after authentication and before
|
||||
// any work. The org comes from `identity`, which came from the token —
|
||||
// there is no path by which a request can name a different bucket, because
|
||||
// this function is never given anything from the request except the tool
|
||||
// name and its arguments.
|
||||
if s.orgLimiter != nil {
|
||||
allowed, retryAfter, err := s.orgLimiter.AllowOrg(ctx, identity.OrgID)
|
||||
if err != nil {
|
||||
// The limiter has already decided whether a failure permits the
|
||||
// call. Logged without the org's usage, which is not the caller's
|
||||
// business.
|
||||
s.log.Error("org rate limiter unavailable", "error", err)
|
||||
}
|
||||
if !allowed {
|
||||
s.log.Warn("mcp org rate limit exceeded",
|
||||
"org_id", identity.OrgID, "tool", p.Name)
|
||||
return nil, errRateLimited(retryAfter)
|
||||
}
|
||||
}
|
||||
|
||||
args := p.Arguments
|
||||
if len(skipSpace(args)) == 0 {
|
||||
args = json.RawMessage(`{}`)
|
||||
}
|
||||
|
||||
tc := tools.Context{
|
||||
Principal: identity,
|
||||
// No RunID: an MCP call is not an agent run and writes no trajectory.
|
||||
// No KnowledgeSources: there is no spec, which is why knowledge_search
|
||||
// is deferred rather than published — see tools.go.
|
||||
}
|
||||
|
||||
res := s.reg.Dispatch(ctx, tc, p.Name, args)
|
||||
|
||||
s.log.Info("mcp tool call",
|
||||
"tool", p.Name,
|
||||
"user_id", identity.UserID,
|
||||
"org_id", identity.OrgID,
|
||||
"ok", res.Error == nil,
|
||||
"truncated", res.Truncated)
|
||||
|
||||
if res.Error != nil {
|
||||
return textResult(map[string]any{"error": res.Error}, true)
|
||||
}
|
||||
return textResult(map[string]any{
|
||||
"data": res.Data,
|
||||
"truncated": res.Truncated,
|
||||
}, false)
|
||||
}
|
||||
632
go-api/internal/mcpserver/tenant_test.go
Normal file
632
go-api/internal/mcpserver/tenant_test.go
Normal file
@@ -0,0 +1,632 @@
|
||||
package mcpserver
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"log/slog"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/jackc/pgx/v5/pgxpool"
|
||||
|
||||
"github.com/krow/krow-backend/go-api/internal/authctx"
|
||||
"github.com/krow/krow-backend/go-api/internal/runtime"
|
||||
"github.com/krow/krow-backend/go-api/internal/testutil"
|
||||
"github.com/krow/krow-backend/go-api/internal/tools"
|
||||
)
|
||||
|
||||
// The tenant isolation matrix: every exposed tool, both organisations.
|
||||
//
|
||||
// THE METHOD, AND WHY IT IS NOT "CHECK THE ROWS"
|
||||
//
|
||||
// Comparing returned rows catches the obvious leak and misses the ones that
|
||||
// matter. A tool that correctly withholds Org B's records while counting them
|
||||
// in a total has leaked; so has one whose "no data" answer differs from its
|
||||
// "no access" answer, or whose error names a record it will not show.
|
||||
//
|
||||
// So each tool is called twice — once as Org A, once as Org B — over identical
|
||||
// but DISTINGUISHABLE data, and the two complete responses are compared as
|
||||
// text. Any value that differs between tenants must be a value that came from
|
||||
// that tenant. A number, a name, an id or a flag that crosses is caught
|
||||
// whatever part of the payload it hides in, including aggregates, counts,
|
||||
// metadata and error text.
|
||||
//
|
||||
// The fixtures are deliberately lopsided — Org B has several times Org A's
|
||||
// volume — so a leak shows up as a wrong NUMBER, not merely a wrong name. A
|
||||
// total of 60 where 7 was correct is unmistakable in a way that a missing name
|
||||
// is not.
|
||||
|
||||
/* ── Fixtures ───────────────────────────────────────────────────────────── */
|
||||
|
||||
// tenant is one seeded organisation and the identities that can act for it.
|
||||
type tenant struct {
|
||||
orgID string
|
||||
label string
|
||||
admin authctx.Identity
|
||||
employer authctx.Identity
|
||||
talent authctx.Identity
|
||||
|
||||
// scale multiplies every seeded row count, so the two tenants' numbers
|
||||
// cannot coincide by accident.
|
||||
scale int
|
||||
}
|
||||
|
||||
func seedTenant(t *testing.T, h *testutil.Harness, label string, scale int) tenant {
|
||||
t.Helper()
|
||||
ctx := context.Background()
|
||||
|
||||
var orgID string
|
||||
if err := h.Pool.QueryRow(ctx,
|
||||
`INSERT INTO organizations (name, slug) VALUES ($1, $2) RETURNING id::text`,
|
||||
"Tenant "+label, "tenant-"+strings.ToLower(label)).Scan(&orgID); err != nil {
|
||||
t.Fatalf("org %s: %v", label, err)
|
||||
}
|
||||
|
||||
mkUser := func(role, email string) authctx.Identity {
|
||||
var id string
|
||||
if err := h.Pool.QueryRow(ctx,
|
||||
`INSERT INTO users (org_id, email, full_name, role, account_type, status)
|
||||
VALUES ($1::uuid, $2, $3, $4, $5, 'active') RETURNING id::text`,
|
||||
orgID, email, "User "+label, role, accountTypeFor(role)).Scan(&id); err != nil {
|
||||
t.Fatalf("user %s: %v", email, err)
|
||||
}
|
||||
return authctx.Identity{
|
||||
UserID: id, OrgID: orgID, Email: email, FullName: "User " + label,
|
||||
Role: role, AccountType: accountTypeFor(role), Status: "active",
|
||||
}
|
||||
}
|
||||
|
||||
tn := tenant{
|
||||
orgID: orgID, label: label, scale: scale,
|
||||
admin: mkUser("admin", "admin-"+label+"@tenant.test"),
|
||||
employer: mkUser("employer", "employer-"+label+"@tenant.test"),
|
||||
talent: mkUser("talent", "talent-"+label+"@tenant.test"),
|
||||
}
|
||||
seedTenantData(t, h, tn)
|
||||
return tn
|
||||
}
|
||||
|
||||
func accountTypeFor(role string) string {
|
||||
if role == "talent" {
|
||||
return "talent"
|
||||
}
|
||||
return "employer"
|
||||
}
|
||||
|
||||
// seedTenantData fills every table the 16 tools read.
|
||||
//
|
||||
// Every value carries the tenant's label, so a leaked string is identifiable on
|
||||
// sight rather than by cross-referencing ids.
|
||||
func seedTenantData(t *testing.T, h *testutil.Harness, tn tenant) {
|
||||
t.Helper()
|
||||
ctx := context.Background()
|
||||
n := tn.scale
|
||||
|
||||
exec := func(sql string, args ...any) {
|
||||
t.Helper()
|
||||
if _, err := h.Pool.Exec(ctx, sql, args...); err != nil {
|
||||
t.Fatalf("seed %s: %v", tn.label, err)
|
||||
}
|
||||
}
|
||||
|
||||
// Activity, SPREAD ACROSS DAYS. activity_signals refuses to call anything
|
||||
// unusual without at least three days of history, so rows all stamped now
|
||||
// would make it answer "not enough history" for both tenants — which would
|
||||
// let the isolation check pass while testing nothing.
|
||||
for i := 0; i < n*3; i++ {
|
||||
exec(`INSERT INTO user_activity (org_id, event_type, user_email, user_name, details, created_date)
|
||||
VALUES ($1::uuid, 'login', $2, $3, $4, now() - ($5::int * interval '1 day'))`,
|
||||
tn.orgID, "actor-"+tn.label+"@tenant.test", "Actor "+tn.label,
|
||||
"detail-"+tn.label, i%7)
|
||||
}
|
||||
for i := 0; i < n; i++ {
|
||||
exec(`INSERT INTO user_activity (org_id, event_type, user_email, user_name, details, created_date)
|
||||
VALUES ($1::uuid, 'create_position', $2, $3, $4, now() - ($5::int * interval '1 day'))`,
|
||||
tn.orgID, "actor-"+tn.label+"@tenant.test", "Actor "+tn.label,
|
||||
"created-"+tn.label, i%5)
|
||||
}
|
||||
|
||||
// Postings, and the applications against them.
|
||||
for i := 0; i < n; i++ {
|
||||
var postingID string
|
||||
if err := h.Pool.QueryRow(ctx,
|
||||
`INSERT INTO job_postings (org_id, title, status, headcount, priority)
|
||||
VALUES ($1::uuid, $2, 'active', 3, 'normal') RETURNING id::text`,
|
||||
tn.orgID, fmt.Sprintf("Role-%s-%d", tn.label, i)).Scan(&postingID); err != nil {
|
||||
t.Fatalf("posting %s: %v", tn.label, err)
|
||||
}
|
||||
for j := 0; j < n; j++ {
|
||||
exec(`INSERT INTO job_applications
|
||||
(org_id, job_posting_id, applicant_name, email, status, ai_score, job_title)
|
||||
VALUES ($1::uuid, $2::uuid, $3, $4, $5, $6, $7)`,
|
||||
tn.orgID, postingID,
|
||||
fmt.Sprintf("Candidate-%s-%d-%d", tn.label, i, j),
|
||||
fmt.Sprintf("cand-%s-%d-%d@tenant.test", tn.label, i, j),
|
||||
[]string{"applied", "ai_screened", "hired"}[j%3],
|
||||
70+j, fmt.Sprintf("Role-%s-%d", tn.label, i))
|
||||
}
|
||||
}
|
||||
|
||||
// Staff and worker profiles.
|
||||
for i := 0; i < n; i++ {
|
||||
exec(`INSERT INTO staff (org_id, name, email, role, status, hire_date)
|
||||
VALUES ($1::uuid, $2, $3, $4, 'active', CURRENT_DATE)`,
|
||||
tn.orgID, fmt.Sprintf("Staff-%s-%d", tn.label, i),
|
||||
fmt.Sprintf("staff-%s-%d@tenant.test", tn.label, i),
|
||||
"Server-"+tn.label)
|
||||
exec(`INSERT INTO worker_profiles
|
||||
(org_id, full_name, email, krow_score, reliability_score, shifts_completed, status)
|
||||
VALUES ($1::uuid, $2, $3, 80, 90, 5, 'active')`,
|
||||
tn.orgID, fmt.Sprintf("Worker-%s-%d", tn.label, i),
|
||||
fmt.Sprintf("worker-%s-%d@tenant.test", tn.label, i))
|
||||
}
|
||||
|
||||
// Shift records — attendance, overtime and coverage all read these.
|
||||
for i := 0; i < n*2; i++ {
|
||||
status := []string{"present", "late", "absent"}[i%3]
|
||||
actual := 8 + i%4
|
||||
// shift_records_absence_has_no_hours: the schema refuses an absence
|
||||
// that logged time, which is the domain rule rather than a quirk —
|
||||
// honoured here so the fixtures are records the product could produce.
|
||||
if status == "absent" {
|
||||
actual = 0
|
||||
}
|
||||
daysAgo := i % 14
|
||||
exec(`INSERT INTO shift_records
|
||||
(org_id, worker_email, worker_name, role, status, scheduled_hours, actual_hours,
|
||||
shift_date, scheduled_start, scheduled_end, overtime_hours, created_date)
|
||||
VALUES ($1::uuid, $2, $3, $4, $5, 8, $6::numeric,
|
||||
CURRENT_DATE - ($7::int * interval '1 day'),
|
||||
now() - ($7::int * interval '1 day'),
|
||||
now() - ($7::int * interval '1 day') + interval '8 hours',
|
||||
$8::numeric, now())`,
|
||||
tn.orgID, fmt.Sprintf("worker-%s-%d@tenant.test", tn.label, i%n),
|
||||
fmt.Sprintf("Worker-%s-%d", tn.label, i%n), "Server-"+tn.label,
|
||||
status, actual, daysAgo, max(actual-8, 0))
|
||||
}
|
||||
|
||||
// Courses, for workforce_training.
|
||||
for i := 0; i < n; i++ {
|
||||
exec(`INSERT INTO courses (org_id, title, status)
|
||||
VALUES ($1::uuid, $2, 'active')`,
|
||||
tn.orgID, fmt.Sprintf("Course-%s-%d", tn.label, i))
|
||||
}
|
||||
}
|
||||
|
||||
// requiredArgs supplies arguments for the tools that cannot be called bare.
|
||||
//
|
||||
// Only two need anything. available_workers takes a shift window — it is a
|
||||
// lookup for "who could work THIS" — and a call without one is an invalid
|
||||
// input rather than an empty result. Everything else answers a bare {}.
|
||||
//
|
||||
// The values are tenant-neutral on purpose: nothing here names an
|
||||
// organisation, so the only thing that can scope the answer is the token.
|
||||
var requiredArgs = map[string]string{
|
||||
"available_workers": `{"starts_at":"2026-09-20T18:00:00Z","ends_at":"2026-09-21T02:00:00Z"}`,
|
||||
}
|
||||
|
||||
/* ── The matrix ─────────────────────────────────────────────────────────── */
|
||||
|
||||
// exposedToolNames is the set under test, taken from the registry rather than
|
||||
// written out, so a tool added to the surface is automatically covered.
|
||||
func exposedToolNames(reg *tools.Registry) []string {
|
||||
infos := exposed(reg)
|
||||
out := make([]string, 0, len(infos))
|
||||
for _, i := range infos {
|
||||
out = append(out, i.Name)
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
type matrixEnv struct {
|
||||
server *Server
|
||||
pool *pgxpool.Pool
|
||||
a, b tenant
|
||||
tokens map[string]authctx.Identity
|
||||
}
|
||||
|
||||
func newMatrix(t *testing.T) *matrixEnv {
|
||||
t.Helper()
|
||||
h := testutil.New(t)
|
||||
|
||||
// Lopsided on purpose: Org B's numbers are several times Org A's, so a
|
||||
// leaked aggregate is a wrong number rather than a plausible one.
|
||||
a := seedTenant(t, h, "A", 2)
|
||||
b := seedTenant(t, h, "B", 5)
|
||||
|
||||
tokens := map[string]authctx.Identity{
|
||||
"tok-a-admin": a.admin,
|
||||
"tok-a-employer": a.employer,
|
||||
"tok-a-talent": a.talent,
|
||||
"tok-b-admin": b.admin,
|
||||
"tok-b-employer": b.employer,
|
||||
"tok-b-talent": b.talent,
|
||||
}
|
||||
srv := New(runtime.DefaultTools(h.Pool, nil), &fakeTokens{byToken: tokens},
|
||||
slog.New(slog.NewTextHandler(io.Discard, nil)))
|
||||
|
||||
return &matrixEnv{server: srv, pool: h.Pool, a: a, b: b, tokens: tokens}
|
||||
}
|
||||
|
||||
// call invokes one tool and returns the whole response body as text.
|
||||
func (m *matrixEnv) call(t *testing.T, token, tool string, args string) string {
|
||||
t.Helper()
|
||||
return m.callRec(t, token, tool, args).Body.String()
|
||||
}
|
||||
|
||||
// callRec is call, returning the whole recorder so a test can read the status
|
||||
// and the headers — which is what a 429 assertion needs.
|
||||
func (m *matrixEnv) callRec(t *testing.T, token, tool, args string) *httptest.ResponseRecorder {
|
||||
t.Helper()
|
||||
return m.callWith(t, token, tool, args, "/mcp", nil)
|
||||
}
|
||||
|
||||
// callWith is callRec with a path and a hook for mutating the request, so the
|
||||
// header- and query-injection tests can drive the same path.
|
||||
func (m *matrixEnv) callWith(t *testing.T, token, tool, args, path string,
|
||||
mutate func(*http.Request)) *httptest.ResponseRecorder {
|
||||
t.Helper()
|
||||
if args == "" {
|
||||
args = "{}"
|
||||
}
|
||||
body := `{"jsonrpc":"2.0","id":1,"method":"tools/call","params":{"name":"` +
|
||||
tool + `","arguments":` + args + `}}`
|
||||
return m.raw(t, token, body, path, mutate)
|
||||
}
|
||||
|
||||
// raw posts an arbitrary JSON-RPC body, for tests that need to shape the
|
||||
// envelope themselves (_meta injection, for one).
|
||||
func (m *matrixEnv) raw(t *testing.T, token, body, path string,
|
||||
mutate func(*http.Request)) *httptest.ResponseRecorder {
|
||||
t.Helper()
|
||||
req := httptest.NewRequest(http.MethodPost, path, strings.NewReader(body))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
if token != "" {
|
||||
req.Header.Set("Authorization", "Bearer "+token)
|
||||
}
|
||||
if mutate != nil {
|
||||
mutate(req)
|
||||
}
|
||||
rec := httptest.NewRecorder()
|
||||
m.server.Handler().ServeHTTP(rec, req)
|
||||
return rec
|
||||
}
|
||||
|
||||
// isRateLimited reports whether a response body is the org-ceiling refusal.
|
||||
func isRateLimited(body string) bool {
|
||||
return strings.Contains(body, "too many requests for this organisation")
|
||||
}
|
||||
|
||||
// TestTenantIsolationMatrix is the heart of Phase 5.
|
||||
//
|
||||
// Sixteen tools × two organisations. For each, the ENTIRE response for Org A is
|
||||
// searched for every marker belonging to Org B, and the reverse. A marker is
|
||||
// any string that identifies the other tenant — its label, its names, its
|
||||
// emails, its org id.
|
||||
func TestTenantIsolationMatrix(t *testing.T) {
|
||||
m := newMatrix(t)
|
||||
names := exposedToolNames(m.server.reg)
|
||||
|
||||
if len(names) != 16 {
|
||||
t.Fatalf("%d exposed tools, want 16 — the matrix must cover all of them", len(names))
|
||||
}
|
||||
|
||||
for _, tool := range names {
|
||||
t.Run(tool, func(t *testing.T) {
|
||||
args := requiredArgs[tool]
|
||||
asA := m.call(t, "tok-a-admin", tool, args)
|
||||
asB := m.call(t, "tok-b-admin", tool, args)
|
||||
|
||||
// Neither response may be an internal failure: a tool that errors
|
||||
// for both tenants would pass a leak check vacuously.
|
||||
for label, body := range map[string]string{"A": asA, "B": asB} {
|
||||
if strings.Contains(body, `"code":-32603`) {
|
||||
t.Fatalf("org %s: the tool failed internally, so isolation is untested: %s",
|
||||
label, truncate(body))
|
||||
}
|
||||
}
|
||||
|
||||
assertNoLeak(t, "A", asA, m.b)
|
||||
assertNoLeak(t, "B", asB, m.a)
|
||||
|
||||
// The two tenants must not produce IDENTICAL payloads. If they do,
|
||||
// either the tool ignores the tenant entirely (a leak) or it
|
||||
// returns nothing for both (in which case this test proves nothing
|
||||
// and should be known to prove nothing).
|
||||
if asA == asB && !strings.Contains(asA, `"data":null`) {
|
||||
t.Errorf("both tenants received a byte-identical response; "+
|
||||
"the tool may not be scoping by organisation at all:\n%s", truncate(asA))
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// assertNoLeak searches one tenant's response for any trace of the other.
|
||||
func assertNoLeak(t *testing.T, whose, body string, other tenant) {
|
||||
t.Helper()
|
||||
|
||||
markers := map[string]string{
|
||||
"organisation id": other.orgID,
|
||||
"actor email": "actor-" + other.label + "@tenant.test",
|
||||
"staff name": "Staff-" + other.label,
|
||||
"worker name": "Worker-" + other.label,
|
||||
"candidate name": "Candidate-" + other.label,
|
||||
"posting title": "Role-" + other.label,
|
||||
"course title": "Course-" + other.label,
|
||||
"role label": "Server-" + other.label,
|
||||
"detail text": "detail-" + other.label,
|
||||
"admin email": other.admin.Email,
|
||||
"user id": other.admin.UserID,
|
||||
}
|
||||
|
||||
for what, marker := range markers {
|
||||
if strings.Contains(body, marker) {
|
||||
t.Errorf("org %s's response contains org %s's %s (%q):\n%s",
|
||||
whose, other.label, what, marker, truncate(body))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func truncate(s string) string {
|
||||
if len(s) > 1200 {
|
||||
return s[:1200] + "… [truncated]"
|
||||
}
|
||||
return s
|
||||
}
|
||||
|
||||
/* ── Aggregates and side channels ───────────────────────────────────────── */
|
||||
|
||||
// A leak through a NUMBER rather than a name.
|
||||
//
|
||||
// Org B has far more of everything. If a tool's totals for Org A are affected
|
||||
// by Org B's rows, the number will be wrong even though no name crosses. This
|
||||
// asserts the arithmetic directly against the database rather than against the
|
||||
// other tenant's response, so it catches a tool that counts everything and
|
||||
// shows only some.
|
||||
func TestAggregatesAreScopedToTheTenant(t *testing.T) {
|
||||
m := newMatrix(t)
|
||||
ctx := context.Background()
|
||||
|
||||
// activity_breakdown reports totalEvents, which must equal exactly this
|
||||
// tenant's rows and not one more.
|
||||
for _, tn := range []tenant{m.a, m.b} {
|
||||
token := "tok-" + strings.ToLower(tn.label) + "-admin"
|
||||
body := m.call(t, token, "activity_breakdown", "{}")
|
||||
|
||||
var want int
|
||||
if err := m.pool.QueryRow(ctx,
|
||||
`SELECT count(*) FROM user_activity WHERE org_id = $1::uuid`, tn.orgID).Scan(&want); err != nil {
|
||||
t.Fatalf("count: %v", err)
|
||||
}
|
||||
|
||||
got := extractInt(t, body, "totalEvents")
|
||||
if got != want {
|
||||
t.Errorf("org %s: totalEvents = %d, want %d (this tenant's rows only)",
|
||||
tn.label, got, want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// An empty tenant must not be able to infer that another tenant is not empty.
|
||||
//
|
||||
// The classic side channel: Org A has no data of some kind, Org B has plenty,
|
||||
// and the "nothing here" answer differs from the "nothing you may see" answer
|
||||
// in a way that reveals the difference.
|
||||
func TestAnEmptyTenantLearnsNothingAboutAFullOne(t *testing.T) {
|
||||
h := testutil.New(t)
|
||||
full := seedTenant(t, h, "Full", 6)
|
||||
empty := seedTenantEmpty(t, h, "Empty")
|
||||
|
||||
srv := New(runtime.DefaultTools(h.Pool, nil), &fakeTokens{byToken: map[string]authctx.Identity{
|
||||
"tok-empty": empty.admin,
|
||||
"tok-full": full.admin,
|
||||
}}, slog.New(slog.NewTextHandler(io.Discard, nil)))
|
||||
m := &matrixEnv{server: srv, pool: h.Pool, a: empty, b: full}
|
||||
|
||||
for _, tool := range exposedToolNames(srv.reg) {
|
||||
t.Run(tool, func(t *testing.T) {
|
||||
body := m.call(t, "tok-empty", tool, requiredArgs[tool])
|
||||
|
||||
// Nothing of the full tenant's may appear.
|
||||
assertNoLeak(t, "Empty", body, full)
|
||||
|
||||
// And no number in the empty tenant's response may match the full
|
||||
// tenant's scale, which would mean a count escaped its filter.
|
||||
for _, n := range []string{`:6`, `:36`, `:12`} {
|
||||
if strings.Contains(strings.ReplaceAll(body, " ", ""), n) &&
|
||||
strings.Contains(body, "total") {
|
||||
t.Logf("note: %s contains %s; verify it is not the other tenant's count", tool, n)
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func seedTenantEmpty(t *testing.T, h *testutil.Harness, label string) tenant {
|
||||
t.Helper()
|
||||
ctx := context.Background()
|
||||
var orgID string
|
||||
if err := h.Pool.QueryRow(ctx,
|
||||
`INSERT INTO organizations (name, slug) VALUES ($1, $2) RETURNING id::text`,
|
||||
"Tenant "+label, "tenant-"+strings.ToLower(label)).Scan(&orgID); err != nil {
|
||||
t.Fatalf("org: %v", err)
|
||||
}
|
||||
var id string
|
||||
if err := h.Pool.QueryRow(ctx,
|
||||
`INSERT INTO users (org_id, email, full_name, role, account_type, status)
|
||||
VALUES ($1::uuid, $2, 'Empty Admin', 'admin', 'employer', 'active') RETURNING id::text`,
|
||||
orgID, "admin-"+label+"@tenant.test").Scan(&id); err != nil {
|
||||
t.Fatalf("user: %v", err)
|
||||
}
|
||||
return tenant{orgID: orgID, label: label, admin: authctx.Identity{
|
||||
UserID: id, OrgID: orgID, Email: "admin-" + label + "@tenant.test",
|
||||
Role: "admin", AccountType: "employer", Status: "active",
|
||||
}}
|
||||
}
|
||||
|
||||
/* ── Role matrix ────────────────────────────────────────────────────────── */
|
||||
|
||||
// The three real KROW roles against all 16 tools.
|
||||
//
|
||||
// This does NOT assert which tools each role may reach — that is the policy
|
||||
// table's business and it is the source of truth, not this test. What it
|
||||
// asserts is the two properties that must hold whatever the policy says:
|
||||
// a refusal must be opaque, and no role may see another tenant.
|
||||
func TestRoleMatrixAcrossBothTenants(t *testing.T) {
|
||||
m := newMatrix(t)
|
||||
names := exposedToolNames(m.server.reg)
|
||||
|
||||
for _, role := range []string{"admin", "employer", "talent"} {
|
||||
for _, tn := range []struct {
|
||||
label string
|
||||
other tenant
|
||||
}{{"a", m.b}, {"b", m.a}} {
|
||||
token := "tok-" + tn.label + "-" + role
|
||||
for _, tool := range names {
|
||||
t.Run(role+"/"+tn.label+"/"+tool, func(t *testing.T) {
|
||||
body := m.call(t, token, tool, requiredArgs[tool])
|
||||
|
||||
// Whatever the policy decides, the other tenant must not
|
||||
// appear in the answer — including in a refusal.
|
||||
assertNoLeak(t, role+"/"+tn.label, body, tn.other)
|
||||
|
||||
// A denial must be the single opaque one. A refusal that
|
||||
// explained itself would describe the shape of what it is
|
||||
// hiding.
|
||||
if strings.Contains(body, "tool.denied") {
|
||||
if !strings.Contains(body, "the caller does not have access to this") {
|
||||
t.Errorf("a denial carried detail beyond the standard message: %s", truncate(body))
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/* ── Injection: no request-supplied identity may influence anything ─────── */
|
||||
|
||||
// Every channel a client controls, against every tool.
|
||||
//
|
||||
// The earlier phases tested this on one tool. Here it is every exposed tool,
|
||||
// because a single handler that read an argument it should not would be enough.
|
||||
func TestNoRequestSuppliedIdentityInfluencesAnyTool(t *testing.T) {
|
||||
m := newMatrix(t)
|
||||
names := exposedToolNames(m.server.reg)
|
||||
|
||||
// Arguments naming the other tenant, in every spelling a caller might try.
|
||||
hostileArgs := `{"org_id":"` + m.b.orgID + `","organization_id":"` + m.b.orgID +
|
||||
`","tenant_id":"` + m.b.orgID + `","user_id":"` + m.b.admin.UserID +
|
||||
`","orgId":"` + m.b.orgID + `","principal":"` + m.b.admin.Email +
|
||||
`","email":"` + m.b.admin.Email + `","on_behalf_of":"` + m.b.admin.Email + `"}`
|
||||
|
||||
for _, tool := range names {
|
||||
t.Run(tool, func(t *testing.T) {
|
||||
clean := m.call(t, "tok-a-admin", tool, requiredArgs[tool])
|
||||
hostile := m.call(t, "tok-a-admin", tool, hostileArgs)
|
||||
|
||||
// Whatever the tool does with unknown arguments — ignore them, or
|
||||
// refuse the call — Org B must not appear.
|
||||
assertNoLeak(t, "A(hostile args)", hostile, m.b)
|
||||
|
||||
// And the answer must not have CHANGED in a way that suggests the
|
||||
// arguments were honoured. A tool that refuses unknown fields is
|
||||
// fine; one that returns different DATA is not.
|
||||
if hostile != clean && !strings.Contains(hostile, "error") {
|
||||
t.Errorf("hostile arguments changed a successful response:\nclean: %s\nhostile: %s",
|
||||
truncate(clean), truncate(hostile))
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// Identity in JSON-RPC metadata, headers and the query string.
|
||||
func TestIdentityChannelsOutsideArgumentsAreIgnored(t *testing.T) {
|
||||
m := newMatrix(t)
|
||||
|
||||
baseline := m.call(t, "tok-a-admin", "activity_breakdown", "{}")
|
||||
baselineTotal := extractInt(t, baseline, "totalEvents")
|
||||
|
||||
send := func(t *testing.T, mutate func(*http.Request), path string) string {
|
||||
t.Helper()
|
||||
body := `{"jsonrpc":"2.0","id":1,"method":"tools/call","_meta":{"org_id":"` + m.b.orgID +
|
||||
`","user_id":"` + m.b.admin.UserID + `"},"params":{"name":"activity_breakdown",` +
|
||||
`"arguments":{},"_meta":{"org_id":"` + m.b.orgID + `"}}}`
|
||||
req := httptest.NewRequest(http.MethodPost, path, strings.NewReader(body))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
req.Header.Set("Authorization", "Bearer tok-a-admin")
|
||||
mutate(req)
|
||||
rec := httptest.NewRecorder()
|
||||
m.server.Handler().ServeHTTP(rec, req)
|
||||
return rec.Body.String()
|
||||
}
|
||||
|
||||
t.Run("jsonrpc _meta at both levels", func(t *testing.T) {
|
||||
got := extractInt(t, send(t, func(*http.Request) {}, "/mcp"), "totalEvents")
|
||||
if got != baselineTotal {
|
||||
t.Errorf("totalEvents = %d, want %d — _meta moved the tenant", got, baselineTotal)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("identity headers", func(t *testing.T) {
|
||||
body := send(t, func(r *http.Request) {
|
||||
for _, h := range []string{
|
||||
"X-Org-Id", "X-Organization-Id", "X-Tenant-Id", "X-User-Id",
|
||||
"X-Krow-Org", "X-Krow-User", "X-On-Behalf-Of", "X-Forwarded-User",
|
||||
} {
|
||||
r.Header.Set(h, m.b.orgID)
|
||||
}
|
||||
}, "/mcp")
|
||||
if got := extractInt(t, body, "totalEvents"); got != baselineTotal {
|
||||
t.Errorf("totalEvents = %d, want %d — a header moved the tenant", got, baselineTotal)
|
||||
}
|
||||
assertNoLeak(t, "A(headers)", body, m.b)
|
||||
})
|
||||
|
||||
t.Run("query string", func(t *testing.T) {
|
||||
body := send(t, func(*http.Request) {},
|
||||
"/mcp?org_id="+m.b.orgID+"&tenant_id="+m.b.orgID+"&user_id="+m.b.admin.UserID)
|
||||
if got := extractInt(t, body, "totalEvents"); got != baselineTotal {
|
||||
t.Errorf("totalEvents = %d, want %d — the query string moved the tenant", got, baselineTotal)
|
||||
}
|
||||
assertNoLeak(t, "A(query)", body, m.b)
|
||||
})
|
||||
}
|
||||
|
||||
/* ── Helpers ────────────────────────────────────────────────────────────── */
|
||||
|
||||
// extractInt pulls a named integer out of a tools/call response.
|
||||
func extractInt(t *testing.T, body, field string) int {
|
||||
t.Helper()
|
||||
var envelope struct {
|
||||
Result struct {
|
||||
Content []struct {
|
||||
Text string `json:"text"`
|
||||
} `json:"content"`
|
||||
IsError bool `json:"isError"`
|
||||
} `json:"result"`
|
||||
}
|
||||
if err := json.Unmarshal([]byte(body), &envelope); err != nil {
|
||||
t.Fatalf("response: %v\n%s", err, truncate(body))
|
||||
}
|
||||
if len(envelope.Result.Content) == 0 {
|
||||
t.Fatalf("no content: %s", truncate(body))
|
||||
}
|
||||
var payload map[string]any
|
||||
if err := json.Unmarshal([]byte(envelope.Result.Content[0].Text), &payload); err != nil {
|
||||
t.Fatalf("payload: %v", err)
|
||||
}
|
||||
data, ok := payload["data"].(map[string]any)
|
||||
if !ok {
|
||||
t.Fatalf("no data object: %s", truncate(body))
|
||||
}
|
||||
value, ok := data[field].(float64)
|
||||
if !ok {
|
||||
t.Fatalf("no %s in %v", field, data)
|
||||
}
|
||||
return int(value)
|
||||
}
|
||||
145
go-api/internal/mcpserver/tools.go
Normal file
145
go-api/internal/mcpserver/tools.go
Normal file
@@ -0,0 +1,145 @@
|
||||
package mcpserver
|
||||
|
||||
import (
|
||||
"sort"
|
||||
|
||||
"github.com/krow/krow-backend/go-api/internal/tools"
|
||||
)
|
||||
|
||||
// Which tools this surface publishes, and why it is a rule rather than a list.
|
||||
//
|
||||
// The set is DERIVED from the registry on every call, not enumerated. A list
|
||||
// would be a promise someone has to keep: register a write tool tomorrow,
|
||||
// forget to update the list, and it ships to every connected client. A
|
||||
// derivation cannot forget. The only hand-maintained part is `deferred` below,
|
||||
// which names tools held back for a reason other than their effect — and
|
||||
// holding something back is the safe direction to be wrong in.
|
||||
//
|
||||
// Three conditions, all required:
|
||||
//
|
||||
// 1. Effect is read. I4's whole point is that a write is gated; a write
|
||||
// reachable over a surface with no confirmation round-trip is I4 defeated.
|
||||
// 2. RequiresConfirmation is false. Belt and braces: Register already forces
|
||||
// it true for a write, so this catches a READ tool that opted in — some
|
||||
// reads are expensive enough to be worth asking about — which this surface
|
||||
// has no way to ask about yet.
|
||||
// 3. Not in `deferred`.
|
||||
|
||||
// deferred names tools held back for a reason that is not their effect.
|
||||
//
|
||||
// knowledge_search is read-only and still cannot ship. Its corpora come from
|
||||
// tools.Context.KnowledgeSources, which the agent loop fills from the running
|
||||
// agent's SPEC — deliberately, so that which documents may be read is not
|
||||
// something a model can choose. An MCP call has no spec, so the field is empty,
|
||||
// and retrieval refuses an empty source list rather than treating it as "all of
|
||||
// them". The tool would therefore fail every call; publishing it would advertise
|
||||
// a capability that cannot work.
|
||||
//
|
||||
// Making it work is a design decision, not an omission: either the connection
|
||||
// binds to an agent spec whose sources it inherits, or sources are derived from
|
||||
// the caller's org ACL (which needs a reindex). Passing them as a tool argument
|
||||
// is the one option that is ruled out, because that is exactly what the field's
|
||||
// placement in the spec exists to prevent.
|
||||
var deferred = map[string]string{
|
||||
"knowledge_search": "corpora come from an agent spec, which an MCP call does not have",
|
||||
}
|
||||
|
||||
// exposed returns the tools this surface publishes, sorted by name.
|
||||
//
|
||||
// Sorted because tools/list is a set, and a stable order makes it diffable in a
|
||||
// test and in a log.
|
||||
func exposed(reg *tools.Registry) []tools.ToolInfo {
|
||||
out := make([]tools.ToolInfo, 0, 16)
|
||||
for _, info := range reg.Catalogue() {
|
||||
if !isExposable(info) {
|
||||
continue
|
||||
}
|
||||
out = append(out, info)
|
||||
}
|
||||
sort.Slice(out, func(i, j int) bool { return out[i].Name < out[j].Name })
|
||||
return out
|
||||
}
|
||||
|
||||
// isExposable is the rule, in one place, so the list and the check cannot
|
||||
// disagree.
|
||||
func isExposable(info tools.ToolInfo) bool {
|
||||
if info.Effect != string(tools.EffectRead) {
|
||||
return false
|
||||
}
|
||||
if info.RequiresConfirmation {
|
||||
return false
|
||||
}
|
||||
if _, held := deferred[info.Name]; held {
|
||||
return false
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
// isExposedName reports whether a tool may be called by name over this surface.
|
||||
//
|
||||
// tools/call consults this BEFORE the registry, so an unexposed tool answers
|
||||
// exactly as an unknown one does. The alternative — dispatching and letting
|
||||
// authorization refuse — would make "this tool exists but you may not reach it
|
||||
// here" distinguishable from "no such tool", which is an inventory of the
|
||||
// surface's own blind spots.
|
||||
func isExposedName(reg *tools.Registry, name string) bool {
|
||||
t, ok := reg.Get(name)
|
||||
if !ok {
|
||||
return false
|
||||
}
|
||||
return isExposable(tools.ToolInfo{
|
||||
Name: t.Name,
|
||||
Effect: string(t.Effect),
|
||||
RequiresConfirmation: t.RequiresConfirmation,
|
||||
})
|
||||
}
|
||||
|
||||
/* ── MCP shapes ─────────────────────────────────────────────────────────── */
|
||||
|
||||
// mcpTool is one entry in a tools/list result.
|
||||
//
|
||||
// Every field is copied from the registry rather than restated. The annotations
|
||||
// are hints a client may show a person before approving a call; they are
|
||||
// derived from Effect so they cannot contradict it.
|
||||
type mcpTool struct {
|
||||
Name string `json:"name"`
|
||||
Description string `json:"description"`
|
||||
InputSchema map[string]any `json:"inputSchema"`
|
||||
Annotations *annotations `json:"annotations,omitempty"`
|
||||
}
|
||||
|
||||
// annotations are the advisory hints from the MCP tool definition.
|
||||
type annotations struct {
|
||||
ReadOnlyHint bool `json:"readOnlyHint"`
|
||||
DestructiveHint bool `json:"destructiveHint"`
|
||||
}
|
||||
|
||||
// emptySchema is what a tool with no declared schema publishes.
|
||||
//
|
||||
// tools/list requires an inputSchema per tool, and a client given `null` may
|
||||
// reasonably refuse the whole list. An object that accepts nothing is the
|
||||
// honest rendering of "this tool takes no arguments".
|
||||
func emptySchema() map[string]any {
|
||||
return map[string]any{
|
||||
"type": "object",
|
||||
"properties": map[string]any{},
|
||||
"additionalProperties": false,
|
||||
}
|
||||
}
|
||||
|
||||
// toMCPTool converts a registry entry into its wire form.
|
||||
func toMCPTool(info tools.ToolInfo) mcpTool {
|
||||
schema := info.InputSchema
|
||||
if schema == nil {
|
||||
schema = emptySchema()
|
||||
}
|
||||
return mcpTool{
|
||||
Name: info.Name,
|
||||
Description: info.Description,
|
||||
InputSchema: schema,
|
||||
Annotations: &annotations{
|
||||
ReadOnlyHint: info.Effect == string(tools.EffectRead),
|
||||
DestructiveHint: info.Effect == string(tools.EffectWrite),
|
||||
},
|
||||
}
|
||||
}
|
||||
216
go-api/internal/mcpserver/transport.go
Normal file
216
go-api/internal/mcpserver/transport.go
Normal file
@@ -0,0 +1,216 @@
|
||||
package mcpserver
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"io"
|
||||
"net/http"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/krow/krow-backend/go-api/internal/authctx"
|
||||
)
|
||||
|
||||
// The Streamable HTTP binding: one endpoint, one message per POST.
|
||||
//
|
||||
// The current MCP spec defines two standard transports — stdio and Streamable
|
||||
// HTTP — and only the second can serve a client that is not launching a
|
||||
// subprocess. Each message is an HTTP POST to a single MCP endpoint, and the
|
||||
// reply is a JSON object or a request-scoped SSE stream.
|
||||
//
|
||||
// This implementation answers with JSON objects and no stream, which is a
|
||||
// complete implementation of the binding for this server's methods rather than
|
||||
// a shortcut: tools/list is a fixed list, and a tool result is structured data
|
||||
// already bounded by MaxResultBytes. There is nothing to deliver incrementally.
|
||||
// Streaming becomes worth adding if a long-running method is ever exposed.
|
||||
//
|
||||
// STATELESS. No session is minted and no Mcp-Session-Id is required, because
|
||||
// nothing here is worth remembering between calls: every request carries its own
|
||||
// bearer credential and every method is independent. Adding session
|
||||
// state now would be state to expire, to bind to a token, to revalidate and to
|
||||
// leak — for no behaviour this server has.
|
||||
|
||||
// MaxRequestBytes bounds an inbound MCP message.
|
||||
//
|
||||
// Well above any legitimate tools/call — arguments are a few scalars — and far
|
||||
// below anything that would be worth sending here. Enforced with
|
||||
// http.MaxBytesReader so the body is refused as it arrives rather than after it
|
||||
// has been buffered.
|
||||
const MaxRequestBytes = 1 << 20 // 1 MiB
|
||||
|
||||
// Handler returns the HTTP handler for the MCP endpoint.
|
||||
//
|
||||
// Deliberately an http.Handler rather than a registered route: this package
|
||||
// does not know its own path, and the server that mounts it decides where it
|
||||
// lives. Note that it does NOT decide what authenticates it — this handler
|
||||
// authenticates its own callers from the Authorization header, and ignores
|
||||
// whatever middleware sits in front. Mounting it behind the cookie middleware
|
||||
// therefore does not make a cookie sufficient to call it.
|
||||
func (s *Server) Handler() http.Handler {
|
||||
return http.HandlerFunc(s.serveHTTP)
|
||||
}
|
||||
|
||||
func (s *Server) serveHTTP(w http.ResponseWriter, r *http.Request) {
|
||||
// One method. GET is what the binding uses for a server-initiated stream,
|
||||
// which this server does not open; saying 405 with an Allow header is more
|
||||
// use to a client than a 404 that suggests the endpoint is absent.
|
||||
if r.Method != http.MethodPost {
|
||||
w.Header().Set("Allow", http.MethodPost)
|
||||
writeRPCError(w, s, http.StatusMethodNotAllowed, nil,
|
||||
errInvalidRequest("this endpoint accepts POST only"))
|
||||
return
|
||||
}
|
||||
|
||||
// Content-Type is checked rather than assumed. A form post or a stray
|
||||
// upload that happened to be valid JSON would otherwise be processed as a
|
||||
// protocol message.
|
||||
if ct := r.Header.Get("Content-Type"); ct != "" && !isJSONContentType(ct) {
|
||||
writeRPCError(w, s, http.StatusUnsupportedMediaType, nil,
|
||||
errInvalidRequest("Content-Type must be application/json"))
|
||||
return
|
||||
}
|
||||
|
||||
body, err := io.ReadAll(http.MaxBytesReader(w, r.Body, MaxRequestBytes))
|
||||
if err != nil {
|
||||
// MaxBytesReader's error is indistinguishable from a truncated upload
|
||||
// without type assertions that buy nothing here: both mean the body is
|
||||
// unusable, and 413 is the more actionable of the two answers.
|
||||
writeRPCError(w, s, http.StatusRequestEntityTooLarge, nil,
|
||||
errInvalidRequest("the request body was too large or could not be read"))
|
||||
return
|
||||
}
|
||||
|
||||
req, rpcErr := parseRequest(body)
|
||||
if rpcErr != nil {
|
||||
// A parse failure has no usable id, so the response carries the id the
|
||||
// message did parse with — null when it parsed with none. HTTP stays
|
||||
// 200: the transport succeeded and the JSON-RPC error IS the answer.
|
||||
writeRPCError(w, s, http.StatusOK, req.ID, rpcErr)
|
||||
return
|
||||
}
|
||||
|
||||
// Authentication, for every method including the handshake — see
|
||||
// methodRequiresAuth. The 401 below is not merely a refusal: its
|
||||
// WWW-Authenticate header is the first step of the MCP authorization flow,
|
||||
// and a client's very first request is what should produce it.
|
||||
var ident *authctx.Identity
|
||||
if resolved, err := s.authenticate(r); err == nil {
|
||||
ident = &resolved
|
||||
} else if methodRequiresAuth(req.Method) {
|
||||
s.log.Warn("mcp request refused",
|
||||
"method", req.Method, "reason", authFailureReason(err))
|
||||
// WWW-Authenticate is not decoration: RFC 9728 has the client read the
|
||||
// resource-metadata URL from this header to find the authorization
|
||||
// server, and from there where to get a token. It is the difference
|
||||
// between "authentication failed" and "here is how to authenticate".
|
||||
w.Header().Set("WWW-Authenticate", s.challenge())
|
||||
writeRPCError(w, s, http.StatusUnauthorized, req.ID,
|
||||
&rpcError{Code: codeUnauthorized, Message: "authentication required"})
|
||||
return
|
||||
}
|
||||
|
||||
// A panic in a handler must not take the process down or leak a stack into
|
||||
// the response. The registry has its own recover around each tool; this is
|
||||
// the outer net for everything else in this package.
|
||||
result, rpcErr := s.handleRecovered(r, ident, req)
|
||||
|
||||
// A notification gets no response body at all — the spec forbids one.
|
||||
if req.isNotification() {
|
||||
w.WriteHeader(http.StatusAccepted)
|
||||
return
|
||||
}
|
||||
if rpcErr != nil {
|
||||
// A rate-limited refusal is the one JSON-RPC error that also carries an
|
||||
// HTTP status, because 429 and Retry-After are how a client knows to
|
||||
// back off. Everything else is 200 with an error body: the transport
|
||||
// succeeded and the error IS the answer.
|
||||
if rpcErr.Code == codeRateLimited {
|
||||
if rpcErr.retryAfter > 0 {
|
||||
w.Header().Set("Retry-After", retryAfterSeconds(rpcErr.retryAfter))
|
||||
}
|
||||
writeRPCError(w, s, http.StatusTooManyRequests, req.ID, rpcErr)
|
||||
return
|
||||
}
|
||||
writeRPCError(w, s, http.StatusOK, req.ID, rpcErr)
|
||||
return
|
||||
}
|
||||
writeJSON(w, s, http.StatusOK, response{
|
||||
JSONRPC: jsonRPCVersion,
|
||||
ID: req.ID,
|
||||
Result: result,
|
||||
})
|
||||
}
|
||||
|
||||
// methodRequiresAuth reports whether a method may only run for a known caller.
|
||||
//
|
||||
// EVERY method does, including the handshake. An earlier revision left
|
||||
// initialize, ping and notifications/initialized open, on the reasoning that a
|
||||
// client needs somewhere to start — and that was wrong in a way worth
|
||||
// recording, because it is the kind of mistake that looks like helpfulness.
|
||||
//
|
||||
// The MCP authorization flow begins with the client making an MCP request
|
||||
// WITHOUT a token and reading the 401's WWW-Authenticate header. A client's
|
||||
// first request is usually initialize. Answering that one with a cheerful 200
|
||||
// means the client never sees the challenge, believes it is connected, and
|
||||
// discovers otherwise only when the first real call fails — by which time it
|
||||
// has no 401 in hand to discover from. Requiring a token everywhere means the
|
||||
// very first request, whatever it is, produces the challenge that starts the
|
||||
// flow.
|
||||
//
|
||||
// Nothing is lost. The handshake is not information a stranger needs: it
|
||||
// returns this server's name and capabilities, which are only useful to a
|
||||
// client that intends to authenticate anyway.
|
||||
//
|
||||
// Kept as a function rather than inlined because it is the single place that
|
||||
// decision lives, and a future method that genuinely must be open should have
|
||||
// to be written down here to become so.
|
||||
func methodRequiresAuth(method string) bool { return true }
|
||||
|
||||
// handleRecovered runs Handle with a recover, converting a panic into an
|
||||
// internal error whose detail goes to the log and not to the caller.
|
||||
func (s *Server) handleRecovered(r *http.Request, ident *authctx.Identity, req request) (result any, rpcErr *rpcError) {
|
||||
defer func() {
|
||||
if p := recover(); p != nil {
|
||||
s.log.Error("mcp handler panicked", "method", req.Method, "panic", p)
|
||||
result, rpcErr = nil, errInternal()
|
||||
}
|
||||
}()
|
||||
return s.Handle(r.Context(), ident, req)
|
||||
}
|
||||
|
||||
// retryAfterSeconds renders a duration for the Retry-After header, rounded up
|
||||
// and never below one second — "Retry-After: 0" invites an immediate retry,
|
||||
// which is the one thing a limited client must not do.
|
||||
func retryAfterSeconds(d time.Duration) string {
|
||||
secs := int(d.Round(time.Second) / time.Second)
|
||||
if secs < 1 {
|
||||
secs = 1
|
||||
}
|
||||
return strconv.Itoa(secs)
|
||||
}
|
||||
|
||||
// isJSONContentType reports whether a Content-Type header names JSON,
|
||||
// tolerating parameters such as "; charset=utf-8".
|
||||
func isJSONContentType(ct string) bool {
|
||||
media := strings.TrimSpace(strings.SplitN(ct, ";", 2)[0])
|
||||
return strings.EqualFold(media, "application/json")
|
||||
}
|
||||
|
||||
func writeRPCError(w http.ResponseWriter, s *Server, status int, id json.RawMessage, e *rpcError) {
|
||||
writeJSON(w, s, status, response{JSONRPC: jsonRPCVersion, ID: id, Error: e})
|
||||
}
|
||||
|
||||
func writeJSON(w http.ResponseWriter, s *Server, status int, payload response) {
|
||||
encoded, err := json.Marshal(payload)
|
||||
if err != nil {
|
||||
// Encoding our own response failed, so there is nothing safe left to
|
||||
// say in JSON. Log it and send a bare 500.
|
||||
s.log.Error("mcp response could not be encoded", "error", err)
|
||||
http.Error(w, "internal error", http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
w.Header().Set("Content-Type", "application/json; charset=utf-8")
|
||||
w.Header().Set("Cache-Control", "no-store")
|
||||
w.WriteHeader(status)
|
||||
_, _ = w.Write(encoded)
|
||||
}
|
||||
374
go-api/internal/oauth/abuse_test.go
Normal file
374
go-api/internal/oauth/abuse_test.go
Normal 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")
|
||||
}
|
||||
}
|
||||
172
go-api/internal/oauth/authenticator.go
Normal file
172
go-api/internal/oauth/authenticator.go
Normal 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
|
||||
}
|
||||
857
go-api/internal/oauth/authserver.go
Normal file
857
go-api/internal/oauth/authserver.go
Normal 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)
|
||||
})
|
||||
}
|
||||
130
go-api/internal/oauth/cleanup.go
Normal file
130
go-api/internal/oauth/cleanup.go
Normal 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.
|
||||
240
go-api/internal/oauth/cleanup_test.go
Normal file
240
go-api/internal/oauth/cleanup_test.go
Normal 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
|
||||
}
|
||||
346
go-api/internal/oauth/consent.go
Normal file
346
go-api/internal/oauth/consent.go
Normal 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 · 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
|
||||
}
|
||||
501
go-api/internal/oauth/consent_test.go
Normal file
501
go-api/internal/oauth/consent_test.go
Normal 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(), ®)
|
||||
|
||||
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(), "<script>") {
|
||||
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(), ®)
|
||||
|
||||
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(), ®)
|
||||
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
|
||||
}
|
||||
337
go-api/internal/oauth/mcp_integration_test.go
Normal file
337
go-api/internal/oauth/mcp_integration_test.go
Normal 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)
|
||||
}
|
||||
}
|
||||
}
|
||||
186
go-api/internal/oauth/metadata.go
Normal file
186
go-api/internal/oauth/metadata.go
Normal 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)
|
||||
}
|
||||
952
go-api/internal/oauth/oauth_test.go
Normal file
952
go-api/internal/oauth/oauth_test.go
Normal 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)
|
||||
}
|
||||
}
|
||||
130
go-api/internal/oauth/pkce.go
Normal file
130
go-api/internal/oauth/pkce.go
Normal 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[:])
|
||||
}
|
||||
107
go-api/internal/oauth/pkce_test.go
Normal file
107
go-api/internal/oauth/pkce_test.go
Normal 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)
|
||||
}
|
||||
}
|
||||
}
|
||||
444
go-api/internal/oauth/store.go
Normal file
444
go-api/internal/oauth/store.go
Normal 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
|
||||
}
|
||||
209
go-api/internal/ratelimit/ratelimit.go
Normal file
209
go-api/internal/ratelimit/ratelimit.go
Normal 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
|
||||
}
|
||||
452
go-api/internal/ratelimit/ratelimit_test.go
Normal file
452
go-api/internal/ratelimit/ratelimit_test.go
Normal 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)
|
||||
}
|
||||
}
|
||||
}
|
||||
163
go-api/internal/ratelimit/rules.go
Normal file
163
go-api/internal/ratelimit/rules.go
Normal 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.
|
||||
*/
|
||||
@@ -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 })
|
||||
|
||||
Reference in New Issue
Block a user