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