mcp connection
This commit is contained in:
19
.env.example
19
.env.example
@@ -31,6 +31,25 @@ HTTP_SHUTDOWN_TIMEOUT=10s
|
|||||||
# Origins are matched exactly, echoed back one at a time, and "*" is rejected.
|
# Origins are matched exactly, echoed back one at a time, and "*" is rejected.
|
||||||
# HTTP_CORS_ORIGINS=http://localhost:5173,http://127.0.0.1:5173
|
# HTTP_CORS_ORIGINS=http://localhost:5173,http://127.0.0.1:5173
|
||||||
|
|
||||||
|
# Networks whose X-Forwarded-For header may be believed. Comma-separated CIDR
|
||||||
|
# blocks or bare addresses; both IP families accepted.
|
||||||
|
#
|
||||||
|
# Several limits are keyed by the caller's address: failed logins, OAuth client
|
||||||
|
# registration, and OAuth authorization before sign-in. Behind a reverse proxy
|
||||||
|
# every request arrives FROM the proxy, so without this setting those budgets
|
||||||
|
# describe the proxy rather than the caller and every user shares one — one
|
||||||
|
# person retrying a connector exhausts everybody's allowance.
|
||||||
|
#
|
||||||
|
# Unset means no proxy is trusted and the header is ignored entirely, which is
|
||||||
|
# correct for local development: nothing sits in front of the dev server. Leave
|
||||||
|
# it unset here. A misspelt value cannot open a hole — it only restores the
|
||||||
|
# shared bucket — but a malformed entry stops startup rather than being dropped.
|
||||||
|
#
|
||||||
|
# NEVER set this to 0.0.0.0/0. That trusts every caller's own header, which is
|
||||||
|
# not a weaker limit but no limit at all: anyone could mint a fresh budget per
|
||||||
|
# request simply by changing the value they send.
|
||||||
|
# HTTP_TRUSTED_PROXIES=
|
||||||
|
|
||||||
# ── PostgreSQL ──────────────────────────────────────────────────────────────
|
# ── PostgreSQL ──────────────────────────────────────────────────────────────
|
||||||
# The local development database. DATABASE_NAME is mixed-case and hyphenated,
|
# The local development database. DATABASE_NAME is mixed-case and hyphenated,
|
||||||
# so anything that interpolates it into SQL must quote it: "Krow-force".
|
# so anything that interpolates it into SQL must quote it: "Krow-force".
|
||||||
|
|||||||
@@ -65,10 +65,14 @@ func run() error {
|
|||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
|
|
||||||
// The sweeper's context is cancelled by the same signal that stops the
|
// The sweepers' context is cancelled by the same signal that stops the
|
||||||
// server, so the ticker goes away with the process rather than outliving
|
// server, so the tickers go away with the process rather than outliving
|
||||||
// the pool it queries.
|
// the pool they query.
|
||||||
go sweepSessions(ctx, server.Sessions(), log)
|
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)
|
errCh := make(chan error, 1)
|
||||||
go func() { errCh <- server.Start() }()
|
go func() { errCh <- server.Start() }()
|
||||||
|
|||||||
@@ -12,6 +12,7 @@ package config
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"fmt"
|
"fmt"
|
||||||
|
"net/netip"
|
||||||
"net/url"
|
"net/url"
|
||||||
"os"
|
"os"
|
||||||
"strconv"
|
"strconv"
|
||||||
@@ -58,8 +59,44 @@ type Config struct {
|
|||||||
Agents AgentsConfig
|
Agents AgentsConfig
|
||||||
Model ModelConfig
|
Model ModelConfig
|
||||||
Knowledge KnowledgeConfig
|
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.
|
// KnowledgeConfig routes the retrieval layer's embedding provider.
|
||||||
//
|
//
|
||||||
// The chat provider does not serve embeddings, so the dense half of hybrid
|
// 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
|
// this API, and once authentication exists that becomes a real hole rather
|
||||||
// than a theoretical one.
|
// than a theoretical one.
|
||||||
CORSOrigins []string
|
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 {
|
type DBConfig struct {
|
||||||
@@ -278,6 +351,12 @@ func Load() (*Config, error) {
|
|||||||
return v
|
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{
|
cfg := &Config{
|
||||||
AppEnv: withDefault("APP_ENV", "development"),
|
AppEnv: withDefault("APP_ENV", "development"),
|
||||||
Log: LogConfig{Level: withDefault("LOG_LEVEL", "info")},
|
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
|
// the server derive it from the CORS posture", and an explicit value
|
||||||
// overrides that derivation. See Server.sessionSameSite.
|
// overrides that derivation. See Server.sessionSameSite.
|
||||||
CookieSameSite: strings.ToLower(strings.TrimSpace(os.Getenv("HTTP_COOKIE_SAMESITE"))),
|
CookieSameSite: strings.ToLower(strings.TrimSpace(os.Getenv("HTTP_COOKIE_SAMESITE"))),
|
||||||
|
TrustedProxies: trustedProxies,
|
||||||
},
|
},
|
||||||
Seed: SeedConfig{
|
Seed: SeedConfig{
|
||||||
FixturePath: withDefault("SEED_FIXTURE_PATH", "./seed/fixtures/seed.json"),
|
FixturePath: withDefault("SEED_FIXTURE_PATH", "./seed/fixtures/seed.json"),
|
||||||
@@ -300,6 +380,14 @@ func Load() (*Config, error) {
|
|||||||
Agents: AgentsConfig{
|
Agents: AgentsConfig{
|
||||||
CuratedPath: withDefault("CURATED_AGENTS_PATH", "./agents"),
|
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{
|
Knowledge: KnowledgeConfig{
|
||||||
EmbedProvider: strings.ToLower(strings.TrimSpace(os.Getenv("EMBED_PROVIDER"))),
|
EmbedProvider: strings.ToLower(strings.TrimSpace(os.Getenv("EMBED_PROVIDER"))),
|
||||||
EmbedAPIKey: strings.TrimSpace(os.Getenv("VOYAGE_API_KEY")),
|
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 {
|
if len(missing) > 0 {
|
||||||
return nil, fmt.Errorf("missing required environment variables: %s "+
|
return nil, fmt.Errorf("missing required environment variables: %s "+
|
||||||
"(copy .env.example to .env and fill them in)", strings.Join(missing, ", "))
|
"(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 {
|
if c.Knowledge.EmbedDims < 0 {
|
||||||
return fmt.Errorf("EMBED_DIMENSIONS cannot be negative, got %d", c.Knowledge.EmbedDims)
|
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{
|
for name, model := range map[string]string{
|
||||||
"MODEL_FAST": c.Model.Fast, "MODEL_BALANCED": c.Model.Balanced, "MODEL_DEEP": c.Model.Deep,
|
"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
|
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 {
|
func withDefault(key, fallback string) string {
|
||||||
if v := strings.TrimSpace(os.Getenv(key)); v != "" {
|
if v := strings.TrimSpace(os.Getenv(key)); v != "" {
|
||||||
return 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",
|
"000009_confirmation_replay.up.sql",
|
||||||
"000010_definition_versions.up.sql",
|
"000010_definition_versions.up.sql",
|
||||||
"000011_employee_roles.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) {
|
if len(ups) != len(want) {
|
||||||
t.Fatalf("%d migrations, want %d — update this list deliberately", 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
|
// 17 from 000001, + auth_sessions (000004), + agent_definitions and
|
||||||
// skill_definitions (000005), + agent_runs (000006), + agent_confirmations
|
// skill_definitions (000005), + agent_runs (000006), + agent_confirmations
|
||||||
// (000007), + knowledge_documents and knowledge_chunks (000008),
|
// (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
|
// schema_migrations is golang-migrate's and is absent when the files are
|
||||||
// applied directly.
|
// applied directly.
|
||||||
if n != 26 {
|
if n != 30 {
|
||||||
t.Errorf("%d base tables after every migration, want 26", n)
|
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.
|
// `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
|
// which is wider, stops one host working through many accounts. They are
|
||||||
// separate limiters because they are deliberately different sizes — see the
|
// separate limiters because they are deliberately different sizes — see the
|
||||||
// note on Server.
|
// note on Server.
|
||||||
addr := clientAddr(r)
|
addr := s.trust.clientAddr(r)
|
||||||
emailKey := strings.ToLower(email)
|
emailKey := strings.ToLower(email)
|
||||||
for _, check := range []struct {
|
for _, check := range []struct {
|
||||||
limiter *attemptLimiter
|
limiter *attemptLimiter
|
||||||
@@ -318,6 +318,85 @@ var publicPaths = map[string]bool{
|
|||||||
"/health": true,
|
"/health": true,
|
||||||
"/api/v1/auth/login": true,
|
"/api/v1/auth/login": true,
|
||||||
"/api/v1/auth/logout": 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.
|
// 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
|
package httpserver
|
||||||
|
|
||||||
import (
|
import (
|
||||||
"net"
|
|
||||||
"net/http"
|
|
||||||
"strings"
|
|
||||||
"sync"
|
"sync"
|
||||||
"time"
|
"time"
|
||||||
)
|
)
|
||||||
@@ -23,12 +20,13 @@ import (
|
|||||||
// one instance needs shared state — Redis, or the database — and this
|
// one instance needs shared state — Redis, or the database — and this
|
||||||
// package is the seam where that goes: attemptLimiter is an implementation
|
// package is the seam where that goes: attemptLimiter is an implementation
|
||||||
// detail behind Allow/Fail/Reset.
|
// detail behind Allow/Fail/Reset.
|
||||||
// - It trusts net/http's RemoteAddr for the client address. Behind a reverse
|
// - The per-address budget is only as good as the address. That used to be
|
||||||
// proxy every request appears to come from the proxy, so the per-address
|
// net/http's RemoteAddr, which behind a reverse proxy is the proxy on
|
||||||
// budget becomes global. Reading X-Forwarded-For instead would be worse,
|
// every request and makes this budget global. It is now resolved by
|
||||||
// not better, until there is a trusted-proxy list to validate it against —
|
// proxyTrust.clientAddr (clientip.go), which reads a forwarded address
|
||||||
// a client can send that header itself and mint a fresh budget per request.
|
// when — and only when — the immediate peer is a configured trusted
|
||||||
// Deploying behind a proxy means adding that list first.
|
// 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
|
// - 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.
|
// 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/db"
|
||||||
"github.com/krow/krow-backend/go-api/internal/definition"
|
"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/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/runtime"
|
||||||
"github.com/krow/krow-backend/go-api/internal/service"
|
"github.com/krow/krow-backend/go-api/internal/service"
|
||||||
"github.com/krow/krow-backend/go-api/internal/tools"
|
"github.com/krow/krow-backend/go-api/internal/tools"
|
||||||
@@ -83,7 +84,18 @@ type Server struct {
|
|||||||
users auth.UserStore
|
users auth.UserStore
|
||||||
credentials *auth.Credentials
|
credentials *auth.Credentials
|
||||||
loginByEmail *attemptLimiter
|
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 is injectable so tests can drive expiry without sleeping.
|
||||||
now func() time.Time
|
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),
|
credentials: auth.NewCredentials(users),
|
||||||
loginByEmail: newAttemptLimiter(o.perEmail, o.loginWindow, o.now),
|
loginByEmail: newAttemptLimiter(o.perEmail, o.loginWindow, o.now),
|
||||||
loginByAddr: newAttemptLimiter(o.perAddress, o.loginWindow, o.now),
|
loginByAddr: newAttemptLimiter(o.perAddress, o.loginWindow, o.now),
|
||||||
|
trust: newProxyTrust(cfg.HTTP.TrustedProxies),
|
||||||
now: o.now,
|
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)
|
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 := http.NewServeMux()
|
||||||
mux.HandleFunc("GET /health", s.handleHealth)
|
mux.HandleFunc("GET /health", s.handleHealth)
|
||||||
s.endpoints = s.routeAuth(mux) + s.routeResources(mux) + s.routeMe(mux) +
|
s.endpoints = s.routeAuth(mux) + s.routeResources(mux) + s.routeMe(mux) +
|
||||||
s.routeDefinitions(mux) + s.routeWorkflows(mux) + s.routeOwliver(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)
|
handler := jsonErrors(mux)
|
||||||
// Authentication sits where devOrgMiddleware used to, so every route below
|
// 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
|
}, true
|
||||||
}
|
}
|
||||||
|
|
||||||
// ToolInfo is what a tool looks like to somebody choosing one, rather than to
|
// ToolInfo is what a tool looks like to somebody choosing one, and — since the
|
||||||
// the model calling it.
|
// 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
|
// Effect is present because it is the one thing an author must understand: a
|
||||||
// schema is the model's business. Effect is present because it is the one thing
|
// write tool means their agent can propose changes, which a person will then be
|
||||||
// an author must understand — a write tool means their agent can propose
|
// asked to approve.
|
||||||
// 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 {
|
type ToolInfo struct {
|
||||||
Name string `json:"name"`
|
Name string `json:"name"`
|
||||||
Description string `json:"description"`
|
Description string `json:"description"`
|
||||||
Effect string `json:"effect"`
|
Effect string `json:"effect"`
|
||||||
RequiresConfirmation bool `json:"requiresConfirmation"`
|
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.
|
// Catalogue lists every registered tool, sorted, as choosable metadata.
|
||||||
@@ -440,6 +464,8 @@ func (r *Registry) Catalogue() []ToolInfo {
|
|||||||
Description: t.Description,
|
Description: t.Description,
|
||||||
Effect: string(t.Effect),
|
Effect: string(t.Effect),
|
||||||
RequiresConfirmation: t.RequiresConfirmation,
|
RequiresConfirmation: t.RequiresConfirmation,
|
||||||
|
InputSchema: t.InputSchema,
|
||||||
|
MaxResultBytes: t.MaxResultBytes,
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
sort.Slice(out, func(i, j int) bool { return out[i].Name < out[j].Name })
|
sort.Slice(out, func(i, j int) bool { return out[i].Name < out[j].Name })
|
||||||
|
|||||||
@@ -42,6 +42,41 @@ HTTP_CORS_ORIGINS=https://platform.krowforce.com,https://mcp.krowforce.com
|
|||||||
# anonymous.
|
# anonymous.
|
||||||
HTTP_COOKIE_SAMESITE=lax
|
HTTP_COOKIE_SAMESITE=lax
|
||||||
|
|
||||||
|
# Networks whose X-Forwarded-For header may be believed. Comma-separated CIDR
|
||||||
|
# blocks or bare addresses.
|
||||||
|
#
|
||||||
|
# THIS DEPLOYMENT NEEDS IT SET, AND THE VALUE IS NOT KNOWABLE FROM THIS REPO.
|
||||||
|
#
|
||||||
|
# The API binds loopback and something else terminates TLS in front of it, so
|
||||||
|
# Go sees the proxy's address on every request. Three limits are keyed by that
|
||||||
|
# address — failed logins (20 per 15 minutes), OAuth client registration (10
|
||||||
|
# per hour) and OAuth authorization before sign-in (20 per hour) — and while it
|
||||||
|
# is unset, all three are ONE budget for the whole deployment. The observable
|
||||||
|
# symptoms are users rate-limiting each other: a connector that refuses to
|
||||||
|
# register because somebody else already did, and sign-in refused platform-wide
|
||||||
|
# after twenty bad passwords anywhere.
|
||||||
|
#
|
||||||
|
# Set it to the address or network the ingress reaches the API from. Find it
|
||||||
|
# rather than guess it — on the API host, with the stack running:
|
||||||
|
#
|
||||||
|
# docker inspect -f '{{range .NetworkSettings.Networks}}{{.Gateway}}{{end}}' krow-api
|
||||||
|
#
|
||||||
|
# That gateway is the address a proxy on the host arrives as. If the proxy runs
|
||||||
|
# in a container on a shared Docker network, use that network's subnet instead:
|
||||||
|
#
|
||||||
|
# docker network inspect -f '{{range .IPAM.Config}}{{.Subnet}}{{end}}' <network>
|
||||||
|
#
|
||||||
|
# Confirm before committing to it: with LOG_LEVEL=debug, one request's logged
|
||||||
|
# address should be the real client's, not the proxy's.
|
||||||
|
#
|
||||||
|
# NEVER 0.0.0.0/0. That trusts every caller's own header — not a weaker limit
|
||||||
|
# but no limit at all, since anyone could then mint a budget per request.
|
||||||
|
#
|
||||||
|
# Left unset here deliberately. An operator supplying a wrong value gets the
|
||||||
|
# shared bucket back; an operator supplying 0.0.0.0/0 gets no protection at all,
|
||||||
|
# so this file ships no value rather than a plausible-looking one to copy.
|
||||||
|
HTTP_TRUSTED_PROXIES=
|
||||||
|
|
||||||
# ── Database ────────────────────────────────────────────────────────────────
|
# ── Database ────────────────────────────────────────────────────────────────
|
||||||
# Point at a managed PostgreSQL. With docker-compose.local-db.yml layered on
|
# Point at a managed PostgreSQL. With docker-compose.local-db.yml layered on
|
||||||
# top, set DATABASE_HOST=postgres instead.
|
# top, set DATABASE_HOST=postgres instead.
|
||||||
|
|||||||
@@ -152,6 +152,17 @@ services:
|
|||||||
# Unset, the server derives it from HTTP_CORS_ORIGINS. The default below
|
# Unset, the server derives it from HTTP_CORS_ORIGINS. The default below
|
||||||
# is deliberate: a compose deployment keeps Lax unless told otherwise.
|
# is deliberate: a compose deployment keeps Lax unless told otherwise.
|
||||||
HTTP_COOKIE_SAMESITE: ${HTTP_COOKIE_SAMESITE:-lax}
|
HTTP_COOKIE_SAMESITE: ${HTTP_COOKIE_SAMESITE:-lax}
|
||||||
|
# Networks whose X-Forwarded-For may be believed. Empty means none, and
|
||||||
|
# empty is what this file defaults to on purpose: a value invented here
|
||||||
|
# would be trusted by every deployment that copies it.
|
||||||
|
#
|
||||||
|
# It MUST be set for this topology. The api container publishes on
|
||||||
|
# loopback with a reverse proxy in front, so without it Go sees the
|
||||||
|
# proxy's address on every request and the three address-keyed limits —
|
||||||
|
# failed logins, OAuth registration, OAuth authorization before sign-in —
|
||||||
|
# become one budget shared by every user. See .env.docker.example for how
|
||||||
|
# to find the right value.
|
||||||
|
HTTP_TRUSTED_PROXIES: ${HTTP_TRUSTED_PROXIES:-}
|
||||||
DATABASE_MAX_OPEN_CONNS: ${DATABASE_MAX_OPEN_CONNS:-25}
|
DATABASE_MAX_OPEN_CONNS: ${DATABASE_MAX_OPEN_CONNS:-25}
|
||||||
DATABASE_MIN_IDLE_CONNS: ${DATABASE_MIN_IDLE_CONNS:-2}
|
DATABASE_MIN_IDLE_CONNS: ${DATABASE_MIN_IDLE_CONNS:-2}
|
||||||
DATABASE_CONN_MAX_LIFETIME: ${DATABASE_CONN_MAX_LIFETIME:-30m}
|
DATABASE_CONN_MAX_LIFETIME: ${DATABASE_CONN_MAX_LIFETIME:-30m}
|
||||||
|
|||||||
14
migrations/000012_oauth_clients.down.sql
Normal file
14
migrations/000012_oauth_clients.down.sql
Normal file
@@ -0,0 +1,14 @@
|
|||||||
|
-- Reverses 000012.
|
||||||
|
--
|
||||||
|
-- Dropping oauth_clients unregisters every connected MCP client. That is the
|
||||||
|
-- correct meaning of rolling back OAuth client registration, and it destroys no
|
||||||
|
-- application data: every row here is a registration, not a business record.
|
||||||
|
-- A client whose registration is gone simply registers again.
|
||||||
|
--
|
||||||
|
-- 000013 and 000014 reference this table, so they must be rolled back first.
|
||||||
|
-- golang-migrate applies down migrations in descending order, which does that
|
||||||
|
-- by construction.
|
||||||
|
|
||||||
|
SET search_path = public;
|
||||||
|
|
||||||
|
DROP TABLE IF EXISTS public.oauth_clients;
|
||||||
86
migrations/000012_oauth_clients.up.sql
Normal file
86
migrations/000012_oauth_clients.up.sql
Normal file
@@ -0,0 +1,86 @@
|
|||||||
|
-- ============================================================================
|
||||||
|
-- Krow — OAuth 2.1 clients
|
||||||
|
--
|
||||||
|
-- Phase 3, migration 1 of 3. This table holds the clients that may ask for a
|
||||||
|
-- token: in practice, one row per Claude installation that has connected.
|
||||||
|
--
|
||||||
|
-- Registered DYNAMICALLY (RFC 7591), not seeded. An MCP client discovers this
|
||||||
|
-- server, registers itself, and gets a client_id back. There is deliberately no
|
||||||
|
-- pre-provisioned row and no fixture: a seeded client is a credential in the
|
||||||
|
-- repository, and the whole point of dynamic registration is that nobody has to
|
||||||
|
-- put one there.
|
||||||
|
--
|
||||||
|
-- PUBLIC CLIENTS ONLY, and that is why there is no client_secret column.
|
||||||
|
-- Claude Desktop and Claude Web are public clients — they run on a machine the
|
||||||
|
-- user controls, so any secret shipped to them is a secret the user has. OAuth
|
||||||
|
-- 2.1 handles this with PKCE instead, which is why code_challenge is mandatory
|
||||||
|
-- in 000013 rather than optional. A column for a secret that must never be
|
||||||
|
-- trusted is a column somebody will eventually trust.
|
||||||
|
--
|
||||||
|
-- Target schema: public. No system schema is read or written.
|
||||||
|
-- ============================================================================
|
||||||
|
|
||||||
|
SET search_path = public;
|
||||||
|
|
||||||
|
CREATE TABLE oauth_clients (
|
||||||
|
-- The client_id handed back at registration and presented on every
|
||||||
|
-- authorization and token request. Opaque and server-generated: a client
|
||||||
|
-- that could choose its own id could impersonate one already registered.
|
||||||
|
client_id text PRIMARY KEY,
|
||||||
|
|
||||||
|
-- What the client calls itself, for the consent screen. Untrusted display
|
||||||
|
-- text — it is whatever the registering client sent, so it is shown as a
|
||||||
|
-- name and never used for a decision.
|
||||||
|
client_name text NOT NULL DEFAULT '',
|
||||||
|
|
||||||
|
-- Exact-match redirect targets. An array because a client may legitimately
|
||||||
|
-- register more than one (a desktop loopback port and a hosted callback),
|
||||||
|
-- and the authorization endpoint matches the presented redirect_uri against
|
||||||
|
-- these byte-for-byte. No prefix matching, no wildcards, no normalisation:
|
||||||
|
-- every one of those has been an open-redirect CVE somewhere.
|
||||||
|
redirect_uris text[] NOT NULL,
|
||||||
|
|
||||||
|
-- Recorded for auditing which client asked for what. Constrained rather than
|
||||||
|
-- free text so an unexpected value is a failed insert instead of a row
|
||||||
|
-- nobody notices.
|
||||||
|
grant_types text[] NOT NULL DEFAULT ARRAY['authorization_code', 'refresh_token'],
|
||||||
|
|
||||||
|
-- The scopes this client may request. Held per-client so tightening the
|
||||||
|
-- global policy later does not silently widen an existing registration.
|
||||||
|
scopes text[] NOT NULL DEFAULT ARRAY['krow.read'],
|
||||||
|
|
||||||
|
created_date timestamptz NOT NULL DEFAULT now(),
|
||||||
|
last_used_at timestamptz,
|
||||||
|
|
||||||
|
-- A client can be disabled without deleting it, so its tokens can be
|
||||||
|
-- revoked and its history kept.
|
||||||
|
disabled_at timestamptz,
|
||||||
|
|
||||||
|
-- At least one redirect URI, or the client can never complete a flow. Caught
|
||||||
|
-- here so a malformed registration fails at the point of registration rather
|
||||||
|
-- than at the point a person is staring at a broken consent screen.
|
||||||
|
CONSTRAINT oauth_clients_redirect_uris_present
|
||||||
|
CHECK (array_length(redirect_uris, 1) >= 1),
|
||||||
|
|
||||||
|
-- Bound so a registration cannot be used to store bulk data.
|
||||||
|
CONSTRAINT oauth_clients_redirect_uris_bounded
|
||||||
|
CHECK (array_length(redirect_uris, 1) <= 10),
|
||||||
|
|
||||||
|
CONSTRAINT oauth_clients_name_bounded
|
||||||
|
CHECK (length(client_name) <= 200)
|
||||||
|
);
|
||||||
|
|
||||||
|
-- Listing a user's connected clients is not a query this table answers — that
|
||||||
|
-- comes from oauth_tokens, which carries user_id. The only index here is the
|
||||||
|
-- primary key, which is also the lookup path: every request arrives with a
|
||||||
|
-- client_id and probes exactly that column.
|
||||||
|
|
||||||
|
COMMENT ON TABLE oauth_clients IS
|
||||||
|
'OAuth 2.1 public clients, registered dynamically per RFC 7591. No secrets '
|
||||||
|
'are stored: public clients authenticate with PKCE, not with a credential.';
|
||||||
|
|
||||||
|
COMMENT ON COLUMN oauth_clients.redirect_uris IS
|
||||||
|
'Exact-match redirect targets. Never prefix-matched or normalised.';
|
||||||
|
|
||||||
|
COMMENT ON COLUMN oauth_clients.disabled_at IS
|
||||||
|
'Set to disable a client without losing its registration or audit history.';
|
||||||
11
migrations/000013_oauth_grants.down.sql
Normal file
11
migrations/000013_oauth_grants.down.sql
Normal file
@@ -0,0 +1,11 @@
|
|||||||
|
-- Reverses 000013.
|
||||||
|
--
|
||||||
|
-- Drops every outstanding authorization code. Any flow mid-redirect fails and
|
||||||
|
-- the user reauthorizes, which is the correct meaning of rolling this back: a
|
||||||
|
-- row here is a credential in flight, not a record of anything.
|
||||||
|
--
|
||||||
|
-- No touch to oauth_clients, which 000012 owns.
|
||||||
|
|
||||||
|
SET search_path = public;
|
||||||
|
|
||||||
|
DROP TABLE IF EXISTS public.oauth_grants;
|
||||||
117
migrations/000013_oauth_grants.up.sql
Normal file
117
migrations/000013_oauth_grants.up.sql
Normal file
@@ -0,0 +1,117 @@
|
|||||||
|
-- ============================================================================
|
||||||
|
-- Krow — OAuth 2.1 authorization codes
|
||||||
|
--
|
||||||
|
-- Phase 3, migration 2 of 3. One row per authorization code issued: the short
|
||||||
|
-- window between a person clicking Approve and the client exchanging the code
|
||||||
|
-- for a token.
|
||||||
|
--
|
||||||
|
-- A row here is a bearer credential with a fuse. Three properties make it safe,
|
||||||
|
-- and all three are enforced by this schema rather than by the code that uses
|
||||||
|
-- it:
|
||||||
|
--
|
||||||
|
-- SINGLE-USE consumed_at, set by the same UPDATE that reads the row. A
|
||||||
|
-- code redeemed twice is an attacker replaying a code they
|
||||||
|
-- intercepted, and the second attempt must fail.
|
||||||
|
-- SHORT-LIVED expires_at, minutes not hours. The code is in transit through
|
||||||
|
-- a browser redirect, which is the least trustworthy hop in the
|
||||||
|
-- flow.
|
||||||
|
-- BOUND to client, redirect_uri, user, scope, resource and PKCE
|
||||||
|
-- challenge. Every one of those is re-verified at the token
|
||||||
|
-- endpoint, so a code stolen from one context cannot be spent
|
||||||
|
-- in another.
|
||||||
|
--
|
||||||
|
-- THE CODE ITSELF IS NEVER STORED. code_hash holds SHA-256, exactly as
|
||||||
|
-- sessions.token_hash does, so a dump of this table cannot be replayed.
|
||||||
|
--
|
||||||
|
-- Target schema: public. No system schema is read or written.
|
||||||
|
-- ============================================================================
|
||||||
|
|
||||||
|
SET search_path = public;
|
||||||
|
|
||||||
|
CREATE TABLE oauth_grants (
|
||||||
|
id uuid PRIMARY KEY DEFAULT gen_random_uuid(),
|
||||||
|
|
||||||
|
-- SHA-256 of the authorization code, lowercase hex. The CHECK pins the
|
||||||
|
-- format so a caller cannot accidentally store a raw code here: a raw code is
|
||||||
|
-- base64url of random bytes and fails this pattern. Same guard, same
|
||||||
|
-- reasoning as sessions_token_hash_sha256 in 000004.
|
||||||
|
code_hash text NOT NULL,
|
||||||
|
|
||||||
|
client_id text NOT NULL REFERENCES oauth_clients (client_id) ON DELETE CASCADE,
|
||||||
|
|
||||||
|
-- Who approved. ON DELETE CASCADE: a deleted user must not leave a code
|
||||||
|
-- behind that could still be exchanged for a token authenticating as them.
|
||||||
|
user_id uuid NOT NULL REFERENCES users (id) ON DELETE CASCADE,
|
||||||
|
|
||||||
|
-- The tenant, denormalised from the user row at issue time. Carried for
|
||||||
|
-- auditing only. It is NEVER read back as the authority on tenancy —
|
||||||
|
-- identity is rebuilt from the live user row at every token validation, so a
|
||||||
|
-- user who moved organisation does not keep the old one. See
|
||||||
|
-- oauth.Authenticator.
|
||||||
|
org_id uuid NOT NULL REFERENCES organizations (id) ON DELETE CASCADE,
|
||||||
|
|
||||||
|
-- Re-verified at the token endpoint. RFC 6749 requires the redirect_uri
|
||||||
|
-- presented at exchange to match the one presented at authorization; without
|
||||||
|
-- this column there is nothing to match against.
|
||||||
|
redirect_uri text NOT NULL,
|
||||||
|
|
||||||
|
scopes text[] NOT NULL,
|
||||||
|
|
||||||
|
-- RFC 8707. The MCP server this code is being obtained for. Carried into the
|
||||||
|
-- access token's audience, which is what makes a token issued for one
|
||||||
|
-- resource unusable at another.
|
||||||
|
resource text NOT NULL,
|
||||||
|
|
||||||
|
-- PKCE. Mandatory — the column is NOT NULL, so a code without a challenge
|
||||||
|
-- cannot exist. OAuth 2.1 requires PKCE for public clients and this is where
|
||||||
|
-- that requirement stops being advisory.
|
||||||
|
code_challenge text NOT NULL,
|
||||||
|
code_challenge_method text NOT NULL,
|
||||||
|
|
||||||
|
created_date timestamptz NOT NULL DEFAULT now(),
|
||||||
|
expires_at timestamptz NOT NULL,
|
||||||
|
|
||||||
|
-- Set on redemption, in the same statement that reads the row. NULL means
|
||||||
|
-- unspent.
|
||||||
|
consumed_at timestamptz,
|
||||||
|
|
||||||
|
CONSTRAINT oauth_grants_code_hash_key UNIQUE (code_hash),
|
||||||
|
CONSTRAINT oauth_grants_code_hash_sha256 CHECK (code_hash ~ '^[0-9a-f]{64}$'),
|
||||||
|
|
||||||
|
-- S256 only. `plain` is permitted by RFC 7636 and forbidden by OAuth 2.1 for
|
||||||
|
-- public clients, because it makes the verifier recoverable from the
|
||||||
|
-- challenge — which is the entire attack PKCE exists to stop. Refused at the
|
||||||
|
-- schema level so no code path can relax it.
|
||||||
|
CONSTRAINT oauth_grants_pkce_s256_only CHECK (code_challenge_method = 'S256'),
|
||||||
|
|
||||||
|
-- A challenge is base64url of a 32-byte SHA-256 digest: 43 characters, no
|
||||||
|
-- padding. Anything else is malformed.
|
||||||
|
CONSTRAINT oauth_grants_challenge_shape CHECK (code_challenge ~ '^[A-Za-z0-9_-]{43}$'),
|
||||||
|
|
||||||
|
CONSTRAINT oauth_grants_expires_after_created CHECK (expires_at > created_date),
|
||||||
|
CONSTRAINT oauth_grants_scopes_present CHECK (array_length(scopes, 1) >= 1)
|
||||||
|
);
|
||||||
|
|
||||||
|
-- The redemption path: look up by hash, check unspent and unexpired. The UNIQUE
|
||||||
|
-- constraint above already provides this index.
|
||||||
|
|
||||||
|
-- The sweep of dead rows.
|
||||||
|
CREATE INDEX oauth_grants_expires_idx ON oauth_grants (expires_at);
|
||||||
|
|
||||||
|
-- Revoking every outstanding code for a user, and the FK's own cascade check.
|
||||||
|
CREATE INDEX oauth_grants_user_idx ON oauth_grants (user_id);
|
||||||
|
|
||||||
|
COMMENT ON TABLE oauth_grants IS
|
||||||
|
'OAuth authorization codes: single-use, short-lived, and bound to client, '
|
||||||
|
'redirect_uri, user, scope, resource and PKCE challenge. The raw code is '
|
||||||
|
'never stored — only SHA-256 of it.';
|
||||||
|
|
||||||
|
COMMENT ON COLUMN oauth_grants.code_hash IS
|
||||||
|
'Lowercase hex SHA-256 of the authorization code. Never the code.';
|
||||||
|
|
||||||
|
COMMENT ON COLUMN oauth_grants.consumed_at IS
|
||||||
|
'Set by the redemption UPDATE itself, so a code cannot be spent twice.';
|
||||||
|
|
||||||
|
COMMENT ON COLUMN oauth_grants.org_id IS
|
||||||
|
'The tenant at issue time, for audit only. Tenancy is re-read from the live '
|
||||||
|
'user row on every token validation and is never taken from here.';
|
||||||
14
migrations/000014_oauth_tokens.down.sql
Normal file
14
migrations/000014_oauth_tokens.down.sql
Normal file
@@ -0,0 +1,14 @@
|
|||||||
|
-- Reverses 000014.
|
||||||
|
--
|
||||||
|
-- Drops every access and refresh token, disconnecting every connected MCP
|
||||||
|
-- client. Users reconnect through the normal authorization flow.
|
||||||
|
--
|
||||||
|
-- This destroys no application data: every row is a credential. Cookie sessions
|
||||||
|
-- are in `sessions` and are untouched, so the merchant-facing product and the
|
||||||
|
-- existing console keep working exactly as before.
|
||||||
|
--
|
||||||
|
-- No touch to oauth_clients or oauth_grants, which 000012 and 000013 own.
|
||||||
|
|
||||||
|
SET search_path = public;
|
||||||
|
|
||||||
|
DROP TABLE IF EXISTS public.oauth_tokens;
|
||||||
121
migrations/000014_oauth_tokens.up.sql
Normal file
121
migrations/000014_oauth_tokens.up.sql
Normal file
@@ -0,0 +1,121 @@
|
|||||||
|
-- ============================================================================
|
||||||
|
-- Krow — OAuth 2.1 access and refresh tokens
|
||||||
|
--
|
||||||
|
-- Phase 3, migration 3 of 3. One row per issued token, access and refresh
|
||||||
|
-- alike, because they share every lifecycle question worth asking: is it known,
|
||||||
|
-- has it expired, has it been revoked, whose is it, and what may it reach.
|
||||||
|
--
|
||||||
|
-- THE RAW TOKEN IS NEVER STORED. token_hash holds SHA-256, and the CHECK below
|
||||||
|
-- refuses anything that is not 64 hex characters — so a raw token, which is
|
||||||
|
-- base64url of random bytes, cannot physically be written to this column. This
|
||||||
|
-- mirrors sessions.token_hash from 000004 exactly, and for the same reason: a
|
||||||
|
-- dump of this table must not be replayable as a login.
|
||||||
|
--
|
||||||
|
-- TOKEN FAMILIES AND REUSE DETECTION
|
||||||
|
--
|
||||||
|
-- Refresh tokens rotate: spending one issues its replacement and consumes the
|
||||||
|
-- old one. family_id ties a lineage together, which is what makes theft
|
||||||
|
-- detectable. If a consumed refresh token is presented again, either the
|
||||||
|
-- legitimate client is retrying or an attacker is replaying a stolen token, and
|
||||||
|
-- there is no way to tell which. OAuth 2.1's answer is to assume the worse case
|
||||||
|
-- and revoke the entire family — the attacker loses access, and the legitimate
|
||||||
|
-- client is forced through a fresh authorization it can complete. Without
|
||||||
|
-- family_id the best available response is to revoke nothing.
|
||||||
|
--
|
||||||
|
-- Target schema: public. No system schema is read or written.
|
||||||
|
-- ============================================================================
|
||||||
|
|
||||||
|
SET search_path = public;
|
||||||
|
|
||||||
|
CREATE TABLE oauth_tokens (
|
||||||
|
id uuid PRIMARY KEY DEFAULT gen_random_uuid(),
|
||||||
|
|
||||||
|
-- SHA-256 of the token, lowercase hex. Never the token.
|
||||||
|
token_hash text NOT NULL,
|
||||||
|
|
||||||
|
-- 'access' or 'refresh'. One table, because the questions asked of both are
|
||||||
|
-- the same; a type column rather than two tables, because a lookup that had
|
||||||
|
-- to try two tables would eventually try only one.
|
||||||
|
token_type text NOT NULL,
|
||||||
|
|
||||||
|
-- Rotation lineage. Every token minted from the same authorization shares a
|
||||||
|
-- family, so reuse detection can revoke all of them at once.
|
||||||
|
family_id uuid NOT NULL,
|
||||||
|
|
||||||
|
client_id text NOT NULL REFERENCES oauth_clients (client_id) ON DELETE CASCADE,
|
||||||
|
|
||||||
|
-- ON DELETE CASCADE: a deleted user must not leave a live token behind.
|
||||||
|
user_id uuid NOT NULL REFERENCES users (id) ON DELETE CASCADE,
|
||||||
|
|
||||||
|
-- Denormalised at issue time for auditing and for cheap per-tenant queries
|
||||||
|
-- ("what is this organisation's Claude usage"). NEVER the authority on
|
||||||
|
-- tenancy: identity is rebuilt from the live user row at validation, so a
|
||||||
|
-- suspended or moved user is caught on their next call rather than at expiry.
|
||||||
|
org_id uuid NOT NULL REFERENCES organizations (id) ON DELETE CASCADE,
|
||||||
|
|
||||||
|
scopes text[] NOT NULL,
|
||||||
|
|
||||||
|
-- RFC 8707. The MCP server this token is for. Validation requires it to
|
||||||
|
-- match this deployment's canonical resource URI, which is what stops a
|
||||||
|
-- token minted for another service being spent here.
|
||||||
|
audience text NOT NULL,
|
||||||
|
|
||||||
|
created_date timestamptz NOT NULL DEFAULT now(),
|
||||||
|
|
||||||
|
-- Sliding is not a concept here: an access token expires and the client
|
||||||
|
-- refreshes. expires_at is the only deadline an access token has.
|
||||||
|
expires_at timestamptz NOT NULL,
|
||||||
|
|
||||||
|
-- Set when a refresh token is spent. A consumed refresh token presented
|
||||||
|
-- again is the reuse signal that revokes the family.
|
||||||
|
consumed_at timestamptz,
|
||||||
|
|
||||||
|
-- Set by revocation: an explicit disconnect, a family revocation, or a
|
||||||
|
-- suspended account being cleaned up.
|
||||||
|
revoked_at timestamptz,
|
||||||
|
revoked_reason text,
|
||||||
|
|
||||||
|
last_used_at timestamptz,
|
||||||
|
|
||||||
|
CONSTRAINT oauth_tokens_token_hash_key UNIQUE (token_hash),
|
||||||
|
CONSTRAINT oauth_tokens_token_hash_sha256 CHECK (token_hash ~ '^[0-9a-f]{64}$'),
|
||||||
|
CONSTRAINT oauth_tokens_type_check CHECK (token_type IN ('access', 'refresh')),
|
||||||
|
CONSTRAINT oauth_tokens_expires_after_created CHECK (expires_at > created_date),
|
||||||
|
CONSTRAINT oauth_tokens_scopes_present CHECK (array_length(scopes, 1) >= 1),
|
||||||
|
CONSTRAINT oauth_tokens_audience_present CHECK (length(btrim(audience)) > 0)
|
||||||
|
);
|
||||||
|
|
||||||
|
-- Validation reads by hash on every single MCP request; the UNIQUE constraint
|
||||||
|
-- above is that index.
|
||||||
|
|
||||||
|
-- Reuse detection and family revocation: given one token, revoke its lineage.
|
||||||
|
CREATE INDEX oauth_tokens_family_idx ON oauth_tokens (family_id);
|
||||||
|
|
||||||
|
-- "Which apps has this user connected", and revoking everything for a user
|
||||||
|
-- whose account was suspended. Partial, because revoked rows are never the
|
||||||
|
-- answer to either question.
|
||||||
|
CREATE INDEX oauth_tokens_user_active_idx ON oauth_tokens (user_id)
|
||||||
|
WHERE revoked_at IS NULL;
|
||||||
|
|
||||||
|
-- The sweep of expired rows.
|
||||||
|
CREATE INDEX oauth_tokens_expires_idx ON oauth_tokens (expires_at);
|
||||||
|
|
||||||
|
COMMENT ON TABLE oauth_tokens IS
|
||||||
|
'OAuth access and refresh tokens. Stores SHA-256 of each token and never the '
|
||||||
|
'token itself. family_id ties a rotation lineage together so that replay of a '
|
||||||
|
'consumed refresh token can revoke the whole family.';
|
||||||
|
|
||||||
|
COMMENT ON COLUMN oauth_tokens.token_hash IS
|
||||||
|
'Lowercase hex SHA-256 of the token. Never the token.';
|
||||||
|
|
||||||
|
COMMENT ON COLUMN oauth_tokens.family_id IS
|
||||||
|
'Rotation lineage. Presenting a consumed refresh token revokes every row '
|
||||||
|
'sharing this id — the OAuth 2.1 response to a possible stolen token.';
|
||||||
|
|
||||||
|
COMMENT ON COLUMN oauth_tokens.audience IS
|
||||||
|
'RFC 8707 resource indicator. Must match this deployment''s canonical MCP '
|
||||||
|
'resource URI at validation, or the token is refused.';
|
||||||
|
|
||||||
|
COMMENT ON COLUMN oauth_tokens.org_id IS
|
||||||
|
'The tenant at issue time, for audit and reporting only. Tenancy is re-read '
|
||||||
|
'from the live user row on every validation and is never taken from here.';
|
||||||
13
migrations/000015_rate_limits.down.sql
Normal file
13
migrations/000015_rate_limits.down.sql
Normal file
@@ -0,0 +1,13 @@
|
|||||||
|
-- Reverses 000015.
|
||||||
|
--
|
||||||
|
-- Dropping the counters removes every in-flight rate limit window. The effect
|
||||||
|
-- is that limits reset once, which is the correct meaning of rolling back a
|
||||||
|
-- counter table: it destroys no application data, and a caller who was at their
|
||||||
|
-- limit gets a fresh window rather than a permanent refusal.
|
||||||
|
--
|
||||||
|
-- The in-process limiter in httpserver/ratelimit.go is untouched by this
|
||||||
|
-- migration and by its rollback; login limiting keeps working either way.
|
||||||
|
|
||||||
|
SET search_path = public;
|
||||||
|
|
||||||
|
DROP TABLE IF EXISTS public.rate_limits;
|
||||||
81
migrations/000015_rate_limits.up.sql
Normal file
81
migrations/000015_rate_limits.up.sql
Normal file
@@ -0,0 +1,81 @@
|
|||||||
|
-- ============================================================================
|
||||||
|
-- Krow — rate limit counters
|
||||||
|
--
|
||||||
|
-- Phase 5. One row per (bucket, window), counting requests.
|
||||||
|
--
|
||||||
|
-- WHY POSTGRES AND NOT REDIS
|
||||||
|
--
|
||||||
|
-- The existing limiter (httpserver/ratelimit.go) is an in-process map, which
|
||||||
|
-- has a failure mode that is easy to miss: behind N instances the effective
|
||||||
|
-- limit is N times the configured one, because each instance counts only what
|
||||||
|
-- it saw. A limit that silently multiplies by the replica count is not a limit.
|
||||||
|
--
|
||||||
|
-- Redis would work and is the conventional answer. It is not the right answer
|
||||||
|
-- here: this service has exactly one piece of shared infrastructure, and adding
|
||||||
|
-- a second means another thing to run, monitor, secure and fail over — for a
|
||||||
|
-- counter. Postgres already provides the one primitive this needs, an atomic
|
||||||
|
-- read-modify-write, in a single statement:
|
||||||
|
--
|
||||||
|
-- INSERT … ON CONFLICT (bucket, window_start) DO UPDATE
|
||||||
|
-- SET count = rate_limits.count + 1
|
||||||
|
-- RETURNING count
|
||||||
|
--
|
||||||
|
-- That is correct under concurrency without a transaction, without a lock taken
|
||||||
|
-- in application code, and without a round trip to decide anything.
|
||||||
|
--
|
||||||
|
-- FIXED WINDOWS, NOT A SLIDING LOG
|
||||||
|
--
|
||||||
|
-- A sliding window is more accurate and costs a row per request. A fixed window
|
||||||
|
-- costs one row per bucket per window and admits a known burst — up to 2× the
|
||||||
|
-- limit across a window boundary. For abuse prevention that is an acceptable
|
||||||
|
-- trade, and it is the difference between a counter table and an append-only
|
||||||
|
-- log nobody wants to sweep.
|
||||||
|
--
|
||||||
|
-- NO RAW CREDENTIAL IS EVER A BUCKET KEY. Callers hash anything sensitive
|
||||||
|
-- before it reaches this table — see internal/ratelimit. A bucket naming a
|
||||||
|
-- token would put that token in a table, in a log, and in every EXPLAIN a
|
||||||
|
-- developer ever runs.
|
||||||
|
--
|
||||||
|
-- Target schema: public. No system schema is read or written.
|
||||||
|
-- ============================================================================
|
||||||
|
|
||||||
|
SET search_path = public;
|
||||||
|
|
||||||
|
CREATE TABLE rate_limits (
|
||||||
|
-- The thing being limited: a scope prefix and an already-hashed subject,
|
||||||
|
-- e.g. "oauth.register:ip:<sha256>" or "mcp.call:token:<sha256>".
|
||||||
|
bucket text NOT NULL,
|
||||||
|
|
||||||
|
-- The window this count belongs to, truncated to the window size. Part of
|
||||||
|
-- the key rather than a column to compare, so a new window is a new row and
|
||||||
|
-- expiry is "delete old rows" rather than "reset a counter" — which means
|
||||||
|
-- two instances rolling over at once cannot lose each other's increments.
|
||||||
|
window_start timestamptz NOT NULL,
|
||||||
|
|
||||||
|
count integer NOT NULL DEFAULT 0,
|
||||||
|
|
||||||
|
-- When this row may be deleted. Carried explicitly rather than derived from
|
||||||
|
-- window_start plus a duration the cleanup would have to know, so windows of
|
||||||
|
-- different sizes can share one table and one sweep.
|
||||||
|
expires_at timestamptz NOT NULL,
|
||||||
|
|
||||||
|
PRIMARY KEY (bucket, window_start),
|
||||||
|
|
||||||
|
CONSTRAINT rate_limits_count_non_negative CHECK (count >= 0),
|
||||||
|
CONSTRAINT rate_limits_expires_after_window CHECK (expires_at > window_start)
|
||||||
|
);
|
||||||
|
|
||||||
|
-- The sweep. Ordered by the column it filters on, so deleting a batch is a
|
||||||
|
-- range scan rather than a sequential scan of every live counter.
|
||||||
|
CREATE INDEX rate_limits_expires_idx ON rate_limits (expires_at);
|
||||||
|
|
||||||
|
COMMENT ON TABLE rate_limits IS
|
||||||
|
'Fixed-window rate limit counters, shared across API instances. Incremented '
|
||||||
|
'with a single atomic INSERT … ON CONFLICT DO UPDATE … RETURNING.';
|
||||||
|
|
||||||
|
COMMENT ON COLUMN rate_limits.bucket IS
|
||||||
|
'Scope plus an ALREADY-HASHED subject. Never a raw token, code or password.';
|
||||||
|
|
||||||
|
COMMENT ON COLUMN rate_limits.window_start IS
|
||||||
|
'Start of the fixed window, truncated to its size. Part of the key so a new '
|
||||||
|
'window is a new row rather than a reset of an existing counter.';
|
||||||
114
scripts/backfill_worker_profile_id.sql
Normal file
114
scripts/backfill_worker_profile_id.sql
Normal file
@@ -0,0 +1,114 @@
|
|||||||
|
-- ============================================================================
|
||||||
|
-- One-off backfill: job_applications.worker_profile_id
|
||||||
|
--
|
||||||
|
-- NOT A MIGRATION, and deliberately not one. `worker_profile_id` has existed
|
||||||
|
-- since migration 000001 and the column is unchanged; what was missing was the
|
||||||
|
-- frontend setting it when it filed an application from the talent pool. That
|
||||||
|
-- is fixed going forward. This connects the rows written before the fix.
|
||||||
|
--
|
||||||
|
-- A migration would make this part of the schema's history and run once per
|
||||||
|
-- environment whether or not it was wanted. A script runs when somebody decides
|
||||||
|
-- to run it, against the environment they name, after reading what it will do.
|
||||||
|
--
|
||||||
|
-- WHAT IT MATCHES
|
||||||
|
--
|
||||||
|
-- org_id AND case-insensitive email
|
||||||
|
--
|
||||||
|
-- and nothing else. No fuzzy matching, no name comparison, no "probably the
|
||||||
|
-- same person". Email is the identity `worker_profiles` itself already asserts:
|
||||||
|
--
|
||||||
|
-- CREATE UNIQUE INDEX worker_profiles_org_email_key
|
||||||
|
-- ON worker_profiles (org_id, email);
|
||||||
|
--
|
||||||
|
-- That constraint is why an ambiguous match is not merely unlikely but
|
||||||
|
-- IMPOSSIBLE: at most one worker profile can exist for any (org_id, email), so
|
||||||
|
-- the join below can never find two. The HAVING clause asserts it anyway —
|
||||||
|
-- cheap, and it fails loudly rather than silently picking one if that
|
||||||
|
-- constraint is ever relaxed.
|
||||||
|
--
|
||||||
|
-- `email` is `citext`, so the comparison is already case-insensitive; it is
|
||||||
|
-- written with lower() so the intent survives a column type change.
|
||||||
|
--
|
||||||
|
-- WHAT IT TOUCHES
|
||||||
|
--
|
||||||
|
-- job_applications.worker_profile_id ONLY, and only where it is NULL.
|
||||||
|
--
|
||||||
|
-- No status, no score, no timestamp, no other column, no other table. It never
|
||||||
|
-- creates a worker_profiles row — an application with nobody behind it stays
|
||||||
|
-- unlinked, which is the honest answer.
|
||||||
|
--
|
||||||
|
-- SAFE TO RE-RUN. Rows already linked are excluded, so a second run changes
|
||||||
|
-- nothing.
|
||||||
|
--
|
||||||
|
-- USAGE
|
||||||
|
-- psql -d <database> -f scripts/backfill_worker_profile_id.sql
|
||||||
|
--
|
||||||
|
-- Wrapped in a transaction: the report and the update see the same rows, and a
|
||||||
|
-- failure leaves nothing half-applied.
|
||||||
|
-- ============================================================================
|
||||||
|
|
||||||
|
BEGIN;
|
||||||
|
|
||||||
|
\echo ''
|
||||||
|
\echo '── Before ─────────────────────────────────────────────────────────────'
|
||||||
|
|
||||||
|
WITH m AS (
|
||||||
|
SELECT a.worker_profile_id AS current_link,
|
||||||
|
(SELECT count(*) FROM worker_profiles w
|
||||||
|
WHERE w.org_id = a.org_id
|
||||||
|
AND lower(w.email::text) = lower(a.email::text)) AS candidates
|
||||||
|
FROM job_applications a
|
||||||
|
)
|
||||||
|
SELECT count(*) AS applications,
|
||||||
|
count(*) FILTER (WHERE current_link IS NOT NULL) AS already_linked,
|
||||||
|
count(*) FILTER (WHERE current_link IS NULL AND candidates = 1) AS will_link,
|
||||||
|
count(*) FILTER (WHERE current_link IS NULL AND candidates = 0) AS no_worker,
|
||||||
|
count(*) FILTER (WHERE current_link IS NULL AND candidates > 1) AS ambiguous
|
||||||
|
FROM m;
|
||||||
|
|
||||||
|
-- The guard. `worker_profiles_org_email_key` should make this impossible; if it
|
||||||
|
-- ever returns a row the backfill must not run, because picking one of two
|
||||||
|
-- people is exactly the kind of quiet wrong answer this script exists to avoid.
|
||||||
|
\echo ''
|
||||||
|
\echo '── Ambiguous matches (must be empty) ──────────────────────────────────'
|
||||||
|
|
||||||
|
SELECT a.id AS application_id, a.email, count(w.id) AS matching_workers
|
||||||
|
FROM job_applications a
|
||||||
|
JOIN worker_profiles w
|
||||||
|
ON w.org_id = a.org_id
|
||||||
|
AND lower(w.email::text) = lower(a.email::text)
|
||||||
|
WHERE a.worker_profile_id IS NULL
|
||||||
|
GROUP BY a.id, a.email
|
||||||
|
HAVING count(w.id) > 1;
|
||||||
|
|
||||||
|
\echo ''
|
||||||
|
\echo '── Linking ────────────────────────────────────────────────────────────'
|
||||||
|
|
||||||
|
UPDATE job_applications a
|
||||||
|
SET worker_profile_id = w.id
|
||||||
|
FROM worker_profiles w
|
||||||
|
WHERE a.worker_profile_id IS NULL
|
||||||
|
AND w.org_id = a.org_id
|
||||||
|
AND lower(w.email::text) = lower(a.email::text)
|
||||||
|
-- Belt and braces: only where exactly one profile matches.
|
||||||
|
AND (SELECT count(*) FROM worker_profiles w2
|
||||||
|
WHERE w2.org_id = a.org_id
|
||||||
|
AND lower(w2.email::text) = lower(a.email::text)) = 1;
|
||||||
|
|
||||||
|
\echo ''
|
||||||
|
\echo '── After ──────────────────────────────────────────────────────────────'
|
||||||
|
|
||||||
|
WITH m AS (
|
||||||
|
SELECT a.worker_profile_id AS current_link,
|
||||||
|
(SELECT count(*) FROM worker_profiles w
|
||||||
|
WHERE w.org_id = a.org_id
|
||||||
|
AND lower(w.email::text) = lower(a.email::text)) AS candidates
|
||||||
|
FROM job_applications a
|
||||||
|
)
|
||||||
|
SELECT count(*) AS applications,
|
||||||
|
count(*) FILTER (WHERE current_link IS NOT NULL) AS linked,
|
||||||
|
count(*) FILTER (WHERE current_link IS NULL AND candidates = 1) AS still_matchable,
|
||||||
|
count(*) FILTER (WHERE current_link IS NULL AND candidates = 0) AS no_worker
|
||||||
|
FROM m;
|
||||||
|
|
||||||
|
COMMIT;
|
||||||
Reference in New Issue
Block a user