Files
krow_backend/go-api/internal/oauth/oauth_test.go
Aravind f2aa3b3ad8
Some checks failed
CI / fixture (push) Has been cancelled
CI / test (push) Has been cancelled
mcp connection
2026-09-22 10:58:02 +05:30

953 lines
34 KiB
Go

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