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

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

View File

@@ -0,0 +1,374 @@
package oauth
import (
"context"
"net/http"
"net/http/httptest"
"net/url"
"strings"
"sync"
"testing"
)
// Abuse cases and input limits.
//
// The property every test here shares: a refusal must not teach the caller
// anything. Not whether a token existed, not whether an account is suspended
// rather than deleted, not what the database is called, and never the value of
// a credential that was presented.
/* ── Nothing sensitive reaches a response ───────────────────────────────── */
// The broadest check in this file: drive every failure path with known secret
// values, and assert none of them comes back.
func TestNoSecretEverAppearsInAResponse(t *testing.T) {
h := newHarness(t)
clientID := h.register()
const (
secretVerifier = "SENTINELverifier0123456789abcdefghijklmnop"
secretCode = "SENTINELcodevalue"
secretToken = "SENTINELtokenvalue"
)
bodies := map[string]string{}
// A failed exchange, with a sentinel code and verifier.
bodies["bad code"] = h.exchange(clientID, secretCode, secretVerifier).Body.String()
// A real code with the wrong verifier.
real := h.authorizeOK(clientID, verifier43)
bodies["bad verifier"] = h.exchange(clientID, real, secretVerifier).Body.String()
// A refresh with a sentinel token.
bodies["bad refresh"] = h.token(url.Values{
"grant_type": {"refresh_token"}, "refresh_token": {secretToken}, "client_id": {clientID},
}).Body.String()
// Revocation of an unknown token.
form := url.Values{"token": {secretToken}}
req := httptest.NewRequest(http.MethodPost, "/oauth/revoke", strings.NewReader(form.Encode()))
req.Header.Set("Content-Type", "application/x-www-form-urlencoded")
rec := httptest.NewRecorder()
h.server.RevokeHandler().ServeHTTP(rec, req)
bodies["revoke unknown"] = rec.Body.String()
for where, body := range bodies {
for what, secret := range map[string]string{
"code_verifier": secretVerifier,
"authorization code": secretCode,
"token": secretToken,
} {
if strings.Contains(body, secret) {
t.Errorf("%s: the response echoes the presented %s:\n%s", where, what, body)
}
}
// Nor may it leak the shape of the system.
for _, tell := range []string{"SQLSTATE", "pq:", "pgx", "oauth_tokens", "oauth_grants",
"password", "Krow-force", "relation", "column"} {
if strings.Contains(body, tell) {
t.Errorf("%s: the response leaks an internal detail (%q):\n%s", where, tell, body)
}
}
}
}
/* ── Replay ─────────────────────────────────────────────────────────────── */
// Ten attempts to spend one code. Exactly one may succeed.
func TestAuthorizationCodeReplayUnderLoad(t *testing.T) {
h := newHarness(t)
clientID := h.register()
code := h.authorizeOK(clientID, verifier43)
succeeded := 0
for i := 0; i < 10; i++ {
if h.exchange(clientID, code, verifier43).Code == http.StatusOK {
succeeded++
}
}
if succeeded != 1 {
t.Errorf("%d of 10 exchanges of the same code succeeded, want exactly 1", succeeded)
}
}
// The same, concurrently. A single-use credential redeemed by two racing
// callers must be spent exactly once — this is the property that
// UPDATE … RETURNING buys, and the one a SELECT-then-UPDATE would lose.
// Run with -race.
func TestConcurrentCodeRedemptionSpendsItOnce(t *testing.T) {
h := newHarness(t)
clientID := h.register()
code := h.authorizeOK(clientID, verifier43)
const racers = 8
var wg sync.WaitGroup
var mu sync.Mutex
succeeded := 0
for i := 0; i < racers; i++ {
wg.Add(1)
go func() {
defer wg.Done()
if h.exchange(clientID, code, verifier43).Code == http.StatusOK {
mu.Lock()
succeeded++
mu.Unlock()
}
}()
}
wg.Wait()
if succeeded != 1 {
t.Errorf("%d of %d concurrent redemptions succeeded, want exactly 1", succeeded, racers)
}
}
// Concurrent refresh of the same token: one rotation, not several. Two
// successes would mean two live families from one credential.
// Run with -race.
func TestConcurrentRefreshRotatesOnce(t *testing.T) {
h := newHarness(t)
clientID := h.register()
code := h.authorizeOK(clientID, verifier43)
tokens := decodeTokens(t, h.exchange(clientID, code, verifier43))
const racers = 8
var wg sync.WaitGroup
var mu sync.Mutex
succeeded := 0
for i := 0; i < racers; i++ {
wg.Add(1)
go func() {
defer wg.Done()
rec := h.token(url.Values{
"grant_type": {"refresh_token"}, "refresh_token": {tokens.RefreshToken},
"client_id": {clientID},
})
if rec.Code == http.StatusOK {
mu.Lock()
succeeded++
mu.Unlock()
}
}()
}
wg.Wait()
if succeeded != 1 {
t.Errorf("%d of %d concurrent refreshes succeeded, want exactly 1 — "+
"a refresh token was double-spent", succeeded, racers)
}
}
/* ── Open redirect ──────────────────────────────────────────────────────── */
// Every shape of redirect tampering, each of which has been a real CVE
// somewhere. None may be honoured, and none may be answered WITH a redirect —
// redirecting an error to an unvalidated URI is the open redirect itself.
func TestOpenRedirectAttempts(t *testing.T) {
h := newHarness(t)
clientID := h.register()
for name, redirect := range map[string]string{
"different host": "https://attacker.example/cb",
"prefix extension": testRedirect + ".attacker.example",
"path traversal": testRedirect + "/../../evil",
"userinfo trick": "https://claude.example.test@attacker.example/cb",
"added query": testRedirect + "?next=https://attacker.example",
"protocol swap": strings.Replace(testRedirect, "https", "http", 1),
"case variation": strings.ToUpper(testRedirect),
"trailing slash": testRedirect + "/",
"double slash": "//attacker.example/cb",
"encoded traversal": testRedirect + "/%2e%2e/evil",
"null byte": testRedirect + "\x00.attacker.example",
"newline injection": testRedirect + "\nLocation: https://attacker.example",
"javascript": "javascript:alert(1)",
"completely missing": "",
} {
t.Run(name, func(t *testing.T) {
rec := h.authorize(map[string]string{
"client_id": clientID, "redirect_uri": redirect, "response_type": "code",
"state": "xyz", "code_challenge": ChallengeFor(verifier43),
"code_challenge_method": "S256", "resource": testResource,
})
if rec.Code == http.StatusFound {
location := rec.Header().Get("Location")
t.Fatalf("answered with a redirect to %q — an unregistered target "+
"must produce a direct error, never a redirect", location)
}
if rec.Code != http.StatusBadRequest {
t.Errorf("status = %d, want 400", rec.Code)
}
if strings.Contains(rec.Body.String(), "attacker.example") {
t.Error("the error echoes the attacker's host back")
}
})
}
}
/* ── Malformed and oversized input ──────────────────────────────────────── */
func TestRegistrationRejectsMalformedBodies(t *testing.T) {
h := newHarness(t)
for name, body := range map[string]string{
"not json": `not json at all`,
"truncated": `{"client_name":`,
"null": `null`,
"array": `[]`,
"deeply nested": `{"client_name":` + strings.Repeat(`[`, 2000) + strings.Repeat(`]`, 2000) + `}`,
"empty": ``,
"wrong types": `{"client_name":123,"redirect_uris":"not-an-array"}`,
"huge name": `{"client_name":"` + strings.Repeat("A", 100_000) + `","redirect_uris":["https://a.test/cb"]}`,
"too many uris": `{"client_name":"x","redirect_uris":[` + strings.TrimSuffix(strings.Repeat(`"https://a.test/cb",`, 50), ",") + `]}`,
"oversized body": `{"client_name":"` + strings.Repeat("A", 20<<10) + `"}`,
} {
t.Run(name, func(t *testing.T) {
req := httptest.NewRequest(http.MethodPost, "/oauth/register", strings.NewReader(body))
rec := httptest.NewRecorder()
// A panic here fails the test by crashing it, which is the assertion.
h.server.RegisterHandler().ServeHTTP(rec, req)
if rec.Code == http.StatusCreated {
// Only the "huge name" case may legitimately succeed, truncated.
if name != "huge name" {
t.Errorf("status = %d; a malformed registration was accepted", rec.Code)
}
return
}
if rec.Code < 400 || rec.Code >= 500 {
t.Errorf("status = %d, want a 4xx", rec.Code)
}
})
}
}
// A client name is stored, shown on a consent screen, and attacker-controlled.
// It must be bounded, or registration becomes free storage.
func TestClientNameIsBounded(t *testing.T) {
h := newHarness(t)
body := `{"client_name":"` + strings.Repeat("A", 5000) + `","redirect_uris":["https://a.test/cb"]}`
rec := httptest.NewRecorder()
h.server.RegisterHandler().ServeHTTP(rec,
httptest.NewRequest(http.MethodPost, "/oauth/register", strings.NewReader(body)))
if rec.Code != http.StatusCreated {
t.Fatalf("status = %d", rec.Code)
}
var stored string
if err := h.h.Pool.QueryRow(context.Background(),
`SELECT client_name FROM oauth_clients ORDER BY created_date DESC LIMIT 1`).Scan(&stored); err != nil {
t.Fatalf("read: %v", err)
}
if len(stored) > 200 {
t.Errorf("stored client_name is %d characters; the column's CHECK allows 200", len(stored))
}
}
func TestTokenEndpointRejectsMalformedRequests(t *testing.T) {
h := newHarness(t)
for name, tc := range map[string]struct {
body string
contentType string
}{
"no content type": {"grant_type=authorization_code", ""},
"json body": {`{"grant_type":"authorization_code"}`, "application/json"},
"empty": {"", "application/x-www-form-urlencoded"},
"garbage": {"%%%%", "application/x-www-form-urlencoded"},
"huge": {"grant_type=authorization_code&code=" + strings.Repeat("A", 200_000), "application/x-www-form-urlencoded"},
"repeated params": {"grant_type=authorization_code&grant_type=password", "application/x-www-form-urlencoded"},
"null grant": {"grant_type=", "application/x-www-form-urlencoded"},
"unknown grant": {"grant_type=magic", "application/x-www-form-urlencoded"},
"injection in code": {"grant_type=authorization_code&code=' OR 1=1 --&client_id=x&redirect_uri=y&code_verifier=" +
strings.Repeat("a", 43), "application/x-www-form-urlencoded"},
} {
t.Run(name, func(t *testing.T) {
req := httptest.NewRequest(http.MethodPost, "/oauth/token", strings.NewReader(tc.body))
if tc.contentType != "" {
req.Header.Set("Content-Type", tc.contentType)
}
rec := httptest.NewRecorder()
h.server.TokenHandler().ServeHTTP(rec, req)
if rec.Code == http.StatusOK {
t.Errorf("a malformed token request succeeded: %s", rec.Body.String())
}
if rec.Code >= 500 {
t.Errorf("status = %d; a malformed request must not be an internal error: %s",
rec.Code, rec.Body.String())
}
})
}
}
/* ── Scope escalation ───────────────────────────────────────────────────── */
// krow.write must be unreachable from every angle: registration, authorization,
// and the consent POST.
func TestWriteScopeIsUnreachable(t *testing.T) {
h := newHarness(t)
t.Run("at registration", func(t *testing.T) {
body := `{"client_name":"x","redirect_uris":["` + testRedirect + `"],"scope":"` + ScopeWrite + `"}`
rec := httptest.NewRecorder()
h.server.RegisterHandler().ServeHTTP(rec,
httptest.NewRequest(http.MethodPost, "/oauth/register", strings.NewReader(body)))
if rec.Code == http.StatusCreated {
t.Error("a client registered for krow.write")
}
})
t.Run("at authorization", func(t *testing.T) {
clientID := h.register()
rec := h.authorize(map[string]string{
"client_id": clientID, "redirect_uri": testRedirect, "response_type": "code",
"state": "xyz", "code_challenge": ChallengeFor(verifier43),
"code_challenge_method": "S256", "resource": testResource,
"scope": ScopeRead + " " + ScopeWrite,
})
loc, _ := url.Parse(rec.Header().Get("Location"))
if loc.Query().Get("code") != "" {
t.Error("an authorization requesting krow.write produced a code")
}
})
t.Run("no issued token carries it", func(t *testing.T) {
clientID := h.register()
code := h.authorizeOK(clientID, verifier43)
tokens := decodeTokens(t, h.exchange(clientID, code, verifier43))
if strings.Contains(tokens.Scope, ScopeWrite) {
t.Errorf("an issued token carries %q", tokens.Scope)
}
})
}
/* ── Cache and transport headers ────────────────────────────────────────── */
// A credential-bearing response must never be cached, and no endpoint may put
// a token in a URL.
func TestSensitiveResponsesAreNotCacheable(t *testing.T) {
h := newHarness(t)
clientID := h.register()
code := h.authorizeOK(clientID, verifier43)
rec := h.exchange(clientID, code, verifier43)
for header, want := range map[string]string{
"Cache-Control": "no-store",
"Pragma": "no-cache",
} {
if got := rec.Header().Get(header); !strings.Contains(got, want) {
t.Errorf("%s = %q, want it to contain %q", header, got, want)
}
}
// The authorization redirect carries a code in its query — that is the
// protocol — but it must never carry a token.
approved := h.authorize(authorizeParamsFor(clientID, verifier43))
if strings.Contains(approved.Header().Get("Location"), "access_token") {
t.Error("an access token appeared in a redirect URL")
}
}

View File

@@ -0,0 +1,172 @@
package oauth
import (
"context"
"crypto/rand"
"encoding/hex"
"errors"
"fmt"
"log/slog"
"strings"
"github.com/krow/krow-backend/go-api/internal/auth"
"github.com/krow/krow-backend/go-api/internal/authctx"
)
// Authenticator is the production implementation of
// mcpserver.TokenAuthenticator.
//
// This is where Phase 2's seam is filled in, and the shape of it is the whole
// argument for having defined the interface first: one method, taking a raw
// token, returning the same authctx.Identity a cookie produces. Nothing
// downstream — not tools.Context, not the policy table, not a single handler —
// can tell which path built the identity, so authorization cannot drift between
// them.
//
// THE IDENTITY IS BUILT FROM THE USER ROW, NOT FROM THE TOKEN.
//
// oauth_tokens carries org_id, and it would be cheaper to read it from there.
// It is deliberately not: the token row records the tenant AT ISSUE TIME, and a
// token can outlive the fact. A user moved to another organisation, or
// suspended, would keep working against a stale claim until the token expired.
// Re-reading the user costs one indexed lookup and makes suspension take effect
// on the next call — which is exactly what httpserver/auth.go already does for
// cookies, and the bearer path must not be weaker than the cookie path.
type Authenticator struct {
store *Store
users UserLookup
log *slog.Logger
// audience is this deployment's canonical MCP resource URI. A token whose
// audience is anything else is refused — see the note in Authenticate.
audience string
}
// UserLookup is the subset of the existing user store this needs. auth.UserStore
// satisfies it; nothing here builds a second user table or password store.
type UserLookup interface {
FindByID(ctx context.Context, id string) (auth.User, error)
}
// NewAuthenticator builds the production token authenticator.
func NewAuthenticator(store *Store, users UserLookup, audience string, log *slog.Logger) *Authenticator {
if log == nil {
log = slog.Default()
}
return &Authenticator{store: store, users: users, audience: audience, log: log}
}
// ErrAudienceMismatch is internal. It never reaches a client — see the single
// return below — but it is distinct so the log can say what happened.
var ErrAudienceMismatch = errors.New("oauth: token audience does not match this resource")
// Authenticate resolves a bearer token into a KROW identity.
//
// EVERY failure returns the same error. Unknown, expired, revoked, wrong
// audience, suspended user, deleted user — one answer, because a caller who can
// tell them apart learns things they should not: that a token once existed,
// that an account was suspended rather than deleted, that this server is not
// the intended audience for a token they hold. Same discipline as
// tools.Denied() and the session path's identical answer to "not found" and
// "expired".
//
// The reason goes to the log, at warn, where the operator is.
func (a *Authenticator) Authenticate(ctx context.Context, rawToken string) (authctx.Identity, error) {
if strings.TrimSpace(rawToken) == "" {
return authctx.Identity{}, ErrTokenUnusable
}
// 1. The token must exist, be an access token, be unexpired and unrevoked.
// All four are in the query's predicate.
token, err := a.store.FindAccessToken(ctx, rawToken)
if err != nil {
a.log.Warn("mcp bearer refused", "reason", "token_unusable")
return authctx.Identity{}, ErrTokenUnusable
}
// 2. Audience. RFC 8707 and the MCP spec both require a server to verify
// that a token was issued FOR IT. Without this check, a token minted by
// this authorization server for some other resource would be spendable
// here — the confused-deputy problem the spec calls out explicitly. The
// comparison is against configuration, never against anything in the
// request: a resource value supplied by the caller would let the caller
// choose their own audience.
if token.Audience != a.audience {
a.log.Warn("mcp bearer refused",
"reason", "audience_mismatch",
"token_id", token.ID,
"expected", a.audience,
"presented", token.Audience)
return authctx.Identity{}, ErrTokenUnusable
}
// 3. Scope. krow.read is the only scope this phase issues, and the MCP
// surface is read-only, so a token without it has no business here. The
// check is present rather than implied so that adding krow.write later
// is a change in one place.
if !hasScope(token.Scopes, ScopeRead) {
a.log.Warn("mcp bearer refused", "reason", "missing_scope", "token_id", token.ID)
return authctx.Identity{}, ErrTokenUnusable
}
// 4. The user, re-read live. See the type comment for why this is not taken
// from the token row.
user, err := a.users.FindByID(ctx, token.UserID)
if err != nil {
// The FK cascades, so a missing user should be unreachable. If it
// happens the token is orphaned and worth killing.
a.log.Warn("mcp bearer refused", "reason", "user_missing", "token_id", token.ID)
_ = a.store.RevokeFamily(ctx, token.FamilyID, "user_missing")
return authctx.Identity{}, ErrTokenUnusable
}
// 5. Suspension revokes on contact, exactly as the cookie path does. Not
// "the token stops working at expiry" — a suspended account must lose
// access on its next request, and leaving the family alive would mean it
// kept a working credential for up to thirty days.
if !user.IsActive() {
a.log.Warn("mcp bearer refused",
"reason", "user_inactive", "user_id", user.ID, "status", user.Status)
_ = a.store.RevokeFamily(ctx, token.FamilyID, "user_suspended")
return authctx.Identity{}, ErrTokenUnusable
}
// The same construction httpserver/auth.go performs for a cookie. SessionID
// and ExpiresAt are deliberately left zero: there is no session row behind
// this identity, and inventing one would make a token look like something
// logout could end.
return authctx.Identity{
UserID: user.ID,
OrgID: user.OrgID,
Email: user.Email,
FullName: user.FullName,
Role: user.Role,
AccountType: user.AccountType,
Status: user.Status,
}, nil
}
// hasScope reports whether a scope was granted.
func hasScope(granted []string, want string) bool {
for _, s := range granted {
if s == want {
return true
}
}
return false
}
// newUUID returns a random UUID v4 string, for family ids.
//
// Hand-rolled rather than adding a dependency: the module is stdlib plus pgx,
// and one 16-byte read with two bits set is not worth a third-party package.
func newUUID() (string, error) {
var b [16]byte
if _, err := rand.Read(b[:]); err != nil {
return "", fmt.Errorf("oauth: generate uuid: %w", err)
}
b[6] = (b[6] & 0x0f) | 0x40 // version 4
b[8] = (b[8] & 0x3f) | 0x80 // variant 10
h := hex.EncodeToString(b[:])
return h[0:8] + "-" + h[8:12] + "-" + h[12:16] + "-" + h[16:20] + "-" + h[20:32], nil
}

View File

@@ -0,0 +1,857 @@
package oauth
import (
"crypto/hmac"
"crypto/rand"
"crypto/sha256"
"encoding/hex"
"encoding/json"
"errors"
"log/slog"
"net/http"
"net/url"
"strings"
"github.com/krow/krow-backend/go-api/internal/authctx"
)
// The authorization server's HTTP surface: register, authorize, token, revoke.
//
// HOW A PERSON IS AUTHENTICATED HERE
//
// They are not, by this package. The authorization endpoint requires a KROW
// user to already be signed in, and it learns who that is from the SessionResolver
// the server was built with — which the HTTP layer implements using the
// existing cookie session. There is no second password store, no second login
// form, and no credential of any kind in this package.
//
// That is also why the authorization endpoint is the only part of OAuth that
// touches cookies: it runs in a browser, as a person, mid-redirect. Everything
// after it — the token endpoint, the MCP endpoint — is a back-channel call from
// the client and uses no cookie at all.
// SessionResolver reports who is signed in, for the authorization endpoint.
//
// Implemented by the HTTP layer over the existing session manager. An interface
// rather than a direct dependency so this package does not reach into
// httpserver, and so a test can drive the flow without a browser.
type SessionResolver interface {
// CurrentUser returns the signed-in identity, or false when there is none.
CurrentUser(r *http.Request) (authctx.Identity, bool)
}
// Server is the OAuth authorization server.
type Server struct {
cfg Config
store *Store
sessions SessionResolver
log *slog.Logger
// loginPath is where an unauthenticated person is sent, with a return
// target, so they can sign in and come back to the consent screen.
loginPath string
// csrfKey signs consent-form tokens. Per-process and never persisted —
// see csrfFor.
csrfKey []byte
}
// NewServer builds the authorization server.
func NewServer(cfg Config, store *Store, sessions SessionResolver, loginPath string, log *slog.Logger) *Server {
if log == nil {
log = slog.Default()
}
if loginPath == "" {
loginPath = "/login"
}
key := make([]byte, 32)
if _, err := rand.Read(key); err != nil {
// Unreachable short of the OS entropy source failing. Panicking is
// correct: a server that cannot generate a CSRF key cannot render a
// consent form safely, and starting without one would mean serving a
// form nothing protects.
panic("oauth: could not generate a consent CSRF key: " + err.Error())
}
return &Server{
cfg: cfg.Normalise(),
store: store,
sessions: sessions,
log: log,
loginPath: loginPath,
csrfKey: key,
}
}
/* ── Errors ─────────────────────────────────────────────────────────────── */
// oauthError is RFC 6749's error shape.
type oauthError struct {
Code string `json:"error"`
Description string `json:"error_description,omitempty"`
}
// Standard error codes. Kept to the set RFC 6749 and 7591 define, because a
// client's error handling switches on these strings.
const (
errInvalidRequest = "invalid_request"
errInvalidClient = "invalid_client"
errInvalidGrant = "invalid_grant"
errUnauthorizedClient = "unauthorized_client"
errUnsupportedGrantType = "unsupported_grant_type"
errInvalidScope = "invalid_scope"
errInvalidRedirectURI = "invalid_redirect_uri"
errInvalidTarget = "invalid_target" // RFC 8707, for a bad resource
errServerError = "server_error"
)
func writeOAuthError(w http.ResponseWriter, status int, code, description string) {
w.Header().Set("Content-Type", "application/json; charset=utf-8")
// A token or error response must never be cached: it is specific to one
// request and may carry a credential.
w.Header().Set("Cache-Control", "no-store")
w.Header().Set("Pragma", "no-cache")
writeJSONBody(w, status, oauthError{Code: code, Description: description})
}
func writeJSONBody(w http.ResponseWriter, status int, payload any) {
encoded, err := json.Marshal(payload)
if err != nil {
http.Error(w, "internal error", http.StatusInternalServerError)
return
}
w.WriteHeader(status)
_, _ = w.Write(encoded)
}
/* ── RFC 7591: Dynamic Client Registration ──────────────────────────────── */
type registrationRequest struct {
ClientName string `json:"client_name"`
RedirectURIs []string `json:"redirect_uris"`
GrantTypes []string `json:"grant_types,omitempty"`
ResponseTypes []string `json:"response_types,omitempty"`
TokenEndpointAuthMethod string `json:"token_endpoint_auth_method,omitempty"`
Scope string `json:"scope,omitempty"`
}
type registrationResponse struct {
ClientID string `json:"client_id"`
ClientName string `json:"client_name,omitempty"`
RedirectURIs []string `json:"redirect_uris"`
GrantTypes []string `json:"grant_types"`
ResponseTypes []string `json:"response_types"`
TokenEndpointAuthMethod string `json:"token_endpoint_auth_method"`
Scope string `json:"scope"`
ClientIDIssuedAt int64 `json:"client_id_issued_at"`
}
// maxRegistrationBytes bounds a registration body. A registration is a name and
// a handful of URIs.
const maxRegistrationBytes = 16 << 10
// RegisterHandler serves dynamic client registration.
//
// Open by necessity: a client that has never registered has no credential to
// present, which is the entire point of RFC 7591 and what lets Claude connect
// without anyone provisioning anything by hand.
//
// That openness is why redirect URI validation below is strict, and why
// PHASE 5 MUST ADD RATE LIMITING HERE. This endpoint writes a row for any
// caller that can reach it. It is structured for that — one handler, one
// validation pass, nothing that would have to move — but today it has no limit,
// and that is recorded as a known gap rather than quietly left unsaid.
func (s *Server) RegisterHandler() http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
w.Header().Set("Allow", http.MethodPost)
writeOAuthError(w, http.StatusMethodNotAllowed, errInvalidRequest, "POST only")
return
}
var req registrationRequest
if err := json.NewDecoder(http.MaxBytesReader(w, r.Body, maxRegistrationBytes)).Decode(&req); err != nil {
writeOAuthError(w, http.StatusBadRequest, errInvalidRequest, "the request body was not valid JSON")
return
}
if len(req.RedirectURIs) == 0 {
writeOAuthError(w, http.StatusBadRequest, errInvalidRedirectURI, "at least one redirect_uri is required")
return
}
if len(req.RedirectURIs) > 10 {
writeOAuthError(w, http.StatusBadRequest, errInvalidRedirectURI, "too many redirect_uris")
return
}
for _, uri := range req.RedirectURIs {
if err := validateRedirectURI(uri); err != nil {
writeOAuthError(w, http.StatusBadRequest, errInvalidRedirectURI, err.Error())
return
}
}
// Only the scopes this server issues. A client asking for krow.write
// is refused rather than quietly downgraded: silently granting less
// than was asked for produces a client that believes it has a
// capability and fails later, somewhere less obvious.
scopes := []string{ScopeRead}
if strings.TrimSpace(req.Scope) != "" {
requested := strings.Fields(req.Scope)
for _, sc := range requested {
if sc != ScopeRead {
writeOAuthError(w, http.StatusBadRequest, errInvalidScope,
"the only scope available is "+ScopeRead)
return
}
}
scopes = requested
}
clientID, err := newUUID()
if err != nil {
s.log.Error("oauth: client id generation failed", "error", err)
writeOAuthError(w, http.StatusInternalServerError, errServerError, "")
return
}
name := strings.TrimSpace(req.ClientName)
if len(name) > 200 {
name = name[:200]
}
client := Client{
ClientID: clientID,
ClientName: name,
RedirectURIs: req.RedirectURIs,
GrantTypes: []string{"authorization_code", "refresh_token"},
Scopes: scopes,
}
if err := s.store.CreateClient(r.Context(), client); err != nil {
s.log.Error("oauth: client registration failed", "error", err)
writeOAuthError(w, http.StatusInternalServerError, errServerError, "")
return
}
s.log.Info("oauth client registered",
"client_id", clientID, "client_name", name, "redirect_uris", len(req.RedirectURIs))
w.Header().Set("Content-Type", "application/json; charset=utf-8")
w.Header().Set("Cache-Control", "no-store")
writeJSONBody(w, http.StatusCreated, registrationResponse{
ClientID: clientID,
ClientName: name,
RedirectURIs: req.RedirectURIs,
GrantTypes: []string{"authorization_code", "refresh_token"},
// No client_secret. A public client that was issued one would ship
// it to every user's machine, and a secret everybody has is not a
// secret — OAuth 2.1 handles public clients with PKCE instead.
ResponseTypes: []string{"code"},
TokenEndpointAuthMethod: "none",
Scope: strings.Join(scopes, " "),
ClientIDIssuedAt: s.store.now().Unix(),
})
})
}
// validateRedirectURI refuses a redirect target that cannot be trusted.
//
// The rules, and why each one is here:
//
// - absolute, with a scheme and host — a relative URI has no meaning in a
// redirect and a client sending one is confused about the flow.
// - no fragment — RFC 6749 forbids it, and the authorization response appends
// its own query parameters; a fragment would be silently dropped or would
// mangle them.
// - https, OR http on loopback only. Plain http anywhere else means the
// authorization code travels in clear text. Loopback is the documented
// exception for native clients (RFC 8252) and is safe because the traffic
// never leaves the machine.
//
// Custom schemes (myapp://callback) are NOT accepted. They are legal per RFC
// 8252 and are a real mechanism for native apps, but any application on the
// machine can register the same scheme and steal the code. Claude's connectors
// use https and loopback, so accepting custom schemes would widen the surface
// for no caller that exists.
func validateRedirectURI(raw string) error {
parsed, err := url.Parse(raw)
if err != nil {
return errMsg("redirect_uri is not a valid URI")
}
if parsed.Scheme == "" || parsed.Host == "" {
return errMsg("redirect_uri must be absolute, with a scheme and host")
}
if parsed.Fragment != "" || strings.Contains(raw, "#") {
return errMsg("redirect_uri must not contain a fragment")
}
switch strings.ToLower(parsed.Scheme) {
case "https":
return nil
case "http":
if isLoopbackHost(parsed.Hostname()) {
return nil
}
return errMsg("http is only permitted for loopback redirect URIs")
default:
return errMsg("redirect_uri must use https, or http on loopback")
}
}
func isLoopbackHost(host string) bool {
switch host {
case "127.0.0.1", "::1", "localhost":
return true
}
return false
}
type errString string
func (e errString) Error() string { return string(e) }
func errMsg(s string) error { return errString(s) }
/* ── Authorization endpoint ─────────────────────────────────────────────── */
// authorizeParams is a validated authorization request.
type authorizeParams struct {
ClientID string
RedirectURI string
ResponseType string
Scopes []string
State string
CodeChallenge string
CodeChallengeMethod string
Resource string
}
// AuthorizeHandler serves the authorization endpoint.
//
// THE ORDER OF VALIDATION IS A SECURITY PROPERTY, not a style choice.
//
// The client_id and redirect_uri are validated FIRST, against the registration,
// before anything else is looked at. Only once the redirect target is known to
// be one this client registered may an error be delivered BY REDIRECTING to it.
// Getting this backwards — redirecting an error to an unvalidated URI — is an
// open redirect, and it is the most common way this endpoint is got wrong.
//
// So: a bad client_id or a bad redirect_uri is answered as a direct HTTP error
// that the browser displays. Everything after that is delivered as a redirect
// with `error=` and the client's `state`, because by then the target is known
// to be legitimate.
func (s *Server) AuthorizeHandler() http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodGet && r.Method != http.MethodPost {
w.Header().Set("Allow", "GET, POST")
writeOAuthError(w, http.StatusMethodNotAllowed, errInvalidRequest, "GET or POST only")
return
}
// A POST carries the decision and the flow's parameters in its body,
// re-posted from the consent form's hidden fields. Merging them into
// the query is what lets every validation below read from one place
// regardless of method — and means the POST is validated exactly as
// strictly as the GET that produced it, rather than trusting the form.
q := r.URL.Query()
if r.Method == http.MethodPost {
if err := r.ParseForm(); err != nil {
writeOAuthError(w, http.StatusBadRequest, errInvalidRequest,
"the form could not be parsed")
return
}
q = r.PostForm
}
// ── Stage 1: the client and its redirect target. Errors here are
// direct responses, never redirects.
clientID := strings.TrimSpace(q.Get("client_id"))
if clientID == "" {
writeOAuthError(w, http.StatusBadRequest, errInvalidClient, "client_id is required")
return
}
client, err := s.store.FindClient(r.Context(), clientID)
if err != nil {
s.log.Warn("oauth authorize refused", "reason", "unknown_client", "client_id", clientID)
writeOAuthError(w, http.StatusBadRequest, errInvalidClient, "unknown client")
return
}
redirectURI := strings.TrimSpace(q.Get("redirect_uri"))
if redirectURI == "" {
writeOAuthError(w, http.StatusBadRequest, errInvalidRedirectURI, "redirect_uri is required")
return
}
if !client.AllowsRedirect(redirectURI) {
// Deliberately NOT redirected. This is the open-redirect guard.
s.log.Warn("oauth authorize refused",
"reason", "redirect_uri_mismatch", "client_id", clientID)
writeOAuthError(w, http.StatusBadRequest, errInvalidRedirectURI,
"redirect_uri does not match a registered URI for this client")
return
}
// ── Stage 2: everything else. The target is trusted now, so failures
// are delivered to it.
state := strings.TrimSpace(q.Get("state"))
if state == "" {
// Required, not optional. state is the client's CSRF defence for
// the callback; a flow without one can be completed by an attacker
// who injects their own authorization response.
s.redirectError(w, r, redirectURI, "", errInvalidRequest, "state is required")
return
}
if rt := q.Get("response_type"); rt != "code" {
s.redirectError(w, r, redirectURI, state, "unsupported_response_type",
"only response_type=code is supported")
return
}
challenge := strings.TrimSpace(q.Get("code_challenge"))
method := strings.TrimSpace(q.Get("code_challenge_method"))
if challenge == "" {
s.redirectError(w, r, redirectURI, state, errInvalidRequest,
"code_challenge is required; this server requires PKCE")
return
}
if method == "" {
// RFC 7636 defaults an absent method to `plain`. This server does
// not accept plain, so an absent method is an error rather than a
// silent downgrade to the weaker mode.
s.redirectError(w, r, redirectURI, state, errInvalidRequest,
"code_challenge_method is required and must be S256")
return
}
if err := ValidateChallenge(challenge, method); err != nil {
s.redirectError(w, r, redirectURI, state, errInvalidRequest, err.Error())
return
}
// RFC 8707. The resource must be THIS server's canonical MCP URI. A
// token is bound to it, so accepting an arbitrary value would let a
// client mint a token aimed at something else.
resource := strings.TrimSpace(q.Get("resource"))
if resource == "" {
s.redirectError(w, r, redirectURI, state, errInvalidTarget,
"resource is required")
return
}
if strings.TrimRight(resource, "/") != s.cfg.Resource {
s.log.Warn("oauth authorize refused",
"reason", "resource_mismatch", "client_id", clientID, "presented", resource)
s.redirectError(w, r, redirectURI, state, errInvalidTarget,
"resource is not a resource this server issues tokens for")
return
}
scopes := []string{ScopeRead}
if raw := strings.TrimSpace(q.Get("scope")); raw != "" {
scopes = strings.Fields(raw)
for _, sc := range scopes {
if sc != ScopeRead {
s.redirectError(w, r, redirectURI, state, errInvalidScope,
"the only scope available is "+ScopeRead)
return
}
}
}
if !client.AllowsScopes(scopes) {
s.redirectError(w, r, redirectURI, state, errInvalidScope,
"this client is not registered for the requested scope")
return
}
params := authorizeParams{
ClientID: clientID, RedirectURI: redirectURI, ResponseType: "code",
Scopes: scopes, State: state, CodeChallenge: challenge,
CodeChallengeMethod: method, Resource: resource,
}
// ── Stage 3: who is this?
identity, signedIn := s.sessions.CurrentUser(r)
if !signedIn {
// Not signed in. Send them to the existing login, with a return
// target that brings them back to this exact authorization request.
// No credential is handled here — the existing cookie login does
// that, unchanged.
s.redirectToLogin(w, r)
return
}
// ── Stage 4: consent.
//
// A GET renders the question. Only a POST carrying a session-bound
// CSRF token answers it, so a cross-site navigation can show a person
// the form but cannot approve on their behalf.
csrf := s.csrfFor(identity)
if r.Method != http.MethodPost {
s.renderConsent(w, r, params, identity, csrf)
return
}
if !s.csrfValid(identity, r.PostFormValue("csrf")) {
// Not an OAuth protocol error — it is a request that did not come
// from the form this server rendered. Answered directly rather
// than redirected, because the client is not the party at fault
// and telling it "access_denied" would be a lie.
s.log.Warn("oauth consent refused", "reason", "csrf_mismatch",
"client_id", params.ClientID, "user_id", identity.UserID)
writeOAuthError(w, http.StatusForbidden, errInvalidRequest,
"this consent form has expired; start the authorization again")
return
}
switch r.PostFormValue("decision") {
case "approve":
s.log.Info("oauth consent approved",
"client_id", params.ClientID, "user_id", identity.UserID,
"org_id", identity.OrgID, "scopes", params.Scopes)
s.issueCode(w, r, params, identity)
case "deny":
// RFC 6749 section 4.1.2.1: a refusal is `access_denied`, returned
// to the client at its registered redirect with the state intact.
// NO CODE IS ISSUED — the deny path never reaches issueCode.
s.log.Info("oauth consent denied",
"client_id", params.ClientID, "user_id", identity.UserID)
s.redirectError(w, r, params.RedirectURI, params.State,
"access_denied", "the user declined this authorization")
default:
// A POST with neither decision. Re-render rather than guess: the
// one thing that must not happen is inferring approval.
s.renderConsent(w, r, params, identity, csrf)
}
})
}
/* ── Consent CSRF ───────────────────────────────────────────────────────── */
// csrfFor derives a token binding the consent form to the signed-in user.
//
// An HMAC over the user id under a per-process key, rather than a random value
// in server-side state. The property needed is only "this form was rendered by
// this server for this user", and an HMAC gives that with nothing to store and
// nothing to expire.
//
// The key is generated at startup and never leaves the process, so a token does
// not survive a restart — which ends any consent form open at that moment. That
// is acceptable: the window between rendering and deciding is seconds, and the
// failure mode is a person clicking Approve and being asked to start again.
func (s *Server) csrfFor(identity authctx.Identity) string {
mac := hmac.New(sha256.New, s.csrfKey)
mac.Write([]byte(identity.UserID))
return hex.EncodeToString(mac.Sum(nil))
}
// csrfValid checks a submitted token in constant time.
func (s *Server) csrfValid(identity authctx.Identity, presented string) bool {
if presented == "" {
return false
}
return hmac.Equal([]byte(s.csrfFor(identity)), []byte(presented))
}
// issueCode stores an authorization code and redirects it to the client.
func (s *Server) issueCode(w http.ResponseWriter, r *http.Request, p authorizeParams, identity authctx.Identity) {
code, err := s.store.CreateGrant(r.Context(), Grant{
ClientID: p.ClientID,
UserID: identity.UserID,
OrgID: identity.OrgID,
RedirectURI: p.RedirectURI,
Scopes: p.Scopes,
Resource: p.Resource,
CodeChallenge: p.CodeChallenge,
CodeChallengeMethod: p.CodeChallengeMethod,
})
if err != nil {
s.log.Error("oauth: could not create grant", "error", err, "client_id", p.ClientID)
s.redirectError(w, r, p.RedirectURI, p.State, errServerError, "")
return
}
// The code id is not logged, and neither is the code. What is logged is who
// approved what, which is the audit question worth answering.
s.log.Info("oauth code issued",
"client_id", p.ClientID, "user_id", identity.UserID,
"org_id", identity.OrgID, "scopes", p.Scopes, "resource", p.Resource)
target, err := url.Parse(p.RedirectURI)
if err != nil {
writeOAuthError(w, http.StatusBadRequest, errInvalidRedirectURI, "redirect_uri is not a valid URI")
return
}
q := target.Query()
q.Set("code", code)
q.Set("state", p.State)
target.RawQuery = q.Encode()
w.Header().Set("Cache-Control", "no-store")
http.Redirect(w, r, target.String(), http.StatusFound)
}
// redirectError delivers an error to a VALIDATED redirect target.
//
// Only ever called after the redirect_uri has been matched against the client's
// registration. See the note on AuthorizeHandler.
func (s *Server) redirectError(w http.ResponseWriter, r *http.Request, redirectURI, state, code, description string) {
target, err := url.Parse(redirectURI)
if err != nil {
writeOAuthError(w, http.StatusBadRequest, errInvalidRedirectURI, "redirect_uri is not a valid URI")
return
}
q := target.Query()
q.Set("error", code)
if description != "" {
q.Set("error_description", description)
}
if state != "" {
q.Set("state", state)
}
target.RawQuery = q.Encode()
w.Header().Set("Cache-Control", "no-store")
http.Redirect(w, r, target.String(), http.StatusFound)
}
// redirectToLogin sends an unauthenticated person to the existing login.
//
// The return target is this server's own path plus the original query, so the
// authorization request survives the round trip. It is built from r.URL rather
// than from anything the caller supplied, so it cannot be pointed elsewhere.
func (s *Server) redirectToLogin(w http.ResponseWriter, r *http.Request) {
returnTo := r.URL.Path
if r.URL.RawQuery != "" {
returnTo += "?" + r.URL.RawQuery
}
target := s.loginPath + "?returnTo=" + url.QueryEscape(returnTo)
w.Header().Set("Cache-Control", "no-store")
http.Redirect(w, r, target, http.StatusFound)
}
/* ── Token endpoint ─────────────────────────────────────────────────────── */
type tokenResponse struct {
AccessToken string `json:"access_token"`
TokenType string `json:"token_type"`
ExpiresIn int `json:"expires_in"`
RefreshToken string `json:"refresh_token"`
Scope string `json:"scope"`
}
// TokenHandler serves the token endpoint: code exchange and refresh.
func (s *Server) TokenHandler() http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
w.Header().Set("Allow", http.MethodPost)
writeOAuthError(w, http.StatusMethodNotAllowed, errInvalidRequest, "POST only")
return
}
if err := r.ParseForm(); err != nil {
writeOAuthError(w, http.StatusBadRequest, errInvalidRequest, "the request body could not be parsed")
return
}
switch r.PostFormValue("grant_type") {
case "authorization_code":
s.exchangeCode(w, r)
case "refresh_token":
s.refresh(w, r)
case "":
writeOAuthError(w, http.StatusBadRequest, errInvalidRequest, "grant_type is required")
default:
// password, client_credentials, implicit and anything else. Named
// explicitly in the metadata as unsupported, and refused here.
writeOAuthError(w, http.StatusBadRequest, errUnsupportedGrantType,
"only authorization_code and refresh_token are supported")
}
})
}
// exchangeCode turns an authorization code into a token pair.
//
// Every binding recorded at authorization is re-verified. A code is not a
// bearer credential on its own: it is a credential for one client, one redirect
// target, one resource, and one PKCE verifier, and a mismatch on any of them
// means the code is being spent by someone other than the client it was issued
// to.
func (s *Server) exchangeCode(w http.ResponseWriter, r *http.Request) {
code := r.PostFormValue("code")
clientID := r.PostFormValue("client_id")
redirectURI := r.PostFormValue("redirect_uri")
verifier := r.PostFormValue("code_verifier")
resource := strings.TrimSpace(r.PostFormValue("resource"))
if code == "" || clientID == "" || redirectURI == "" {
writeOAuthError(w, http.StatusBadRequest, errInvalidRequest,
"code, client_id and redirect_uri are required")
return
}
if verifier == "" {
writeOAuthError(w, http.StatusBadRequest, errInvalidRequest,
"code_verifier is required; this server requires PKCE")
return
}
// Redeeming CONSUMES the code, whatever happens next. That is deliberate:
// if a later check fails, the code is still spent, so an attacker cannot
// probe the remaining bindings by retrying the same code with different
// values. One code, one attempt.
grant, err := s.store.RedeemGrant(r.Context(), code)
if err != nil {
s.log.Warn("oauth token refused", "reason", "grant_unusable", "client_id", clientID)
writeOAuthError(w, http.StatusBadRequest, errInvalidGrant,
"the authorization code is invalid, expired or already used")
return
}
if grant.ClientID != clientID {
s.log.Warn("oauth token refused", "reason", "client_mismatch", "client_id", clientID)
writeOAuthError(w, http.StatusBadRequest, errInvalidGrant, "this code was not issued to this client")
return
}
if grant.RedirectURI != redirectURI {
s.log.Warn("oauth token refused", "reason", "redirect_uri_mismatch", "client_id", clientID)
writeOAuthError(w, http.StatusBadRequest, errInvalidGrant, "redirect_uri does not match the authorization request")
return
}
// The resource is optional at the token endpoint when the code already
// carries one, but if it IS supplied it must agree.
if resource != "" && strings.TrimRight(resource, "/") != grant.Resource {
writeOAuthError(w, http.StatusBadRequest, errInvalidTarget, "resource does not match the authorization request")
return
}
if err := VerifyChallenge(verifier, grant.CodeChallenge, grant.CodeChallengeMethod); err != nil {
s.log.Warn("oauth token refused", "reason", "pkce_mismatch", "client_id", clientID)
writeOAuthError(w, http.StatusBadRequest, errInvalidGrant, "code_verifier does not match")
return
}
pair, err := s.store.IssuePair(r.Context(), Token{
ClientID: grant.ClientID,
UserID: grant.UserID,
OrgID: grant.OrgID,
Scopes: grant.Scopes,
Audience: grant.Resource,
}, "")
if err != nil {
s.log.Error("oauth: could not issue tokens", "error", err)
writeOAuthError(w, http.StatusInternalServerError, errServerError, "")
return
}
// The tokens themselves are NOT in this log line and never will be.
s.log.Info("oauth tokens issued",
"grant_type", "authorization_code", "client_id", grant.ClientID,
"user_id", grant.UserID, "org_id", grant.OrgID, "family_id", pair.FamilyID)
writeTokenResponse(w, pair)
}
// refresh rotates a refresh token.
func (s *Server) refresh(w http.ResponseWriter, r *http.Request) {
raw := r.PostFormValue("refresh_token")
clientID := r.PostFormValue("client_id")
if raw == "" || clientID == "" {
writeOAuthError(w, http.StatusBadRequest, errInvalidRequest,
"refresh_token and client_id are required")
return
}
old, err := s.store.RedeemRefreshToken(r.Context(), raw)
switch {
case err == nil:
// fall through
case errors.Is(err, ErrRefreshReuse):
// The family has already been revoked by the store. Logged at warn
// because it is either a client bug or a stolen token, and both are
// worth seeing. The CLIENT is told the same thing as for any other bad
// token — distinguishing "reused" would confirm the token was once
// real.
s.log.Warn("oauth refresh refused", "reason", "reuse_detected", "client_id", clientID)
writeOAuthError(w, http.StatusBadRequest, errInvalidGrant, "the refresh token is invalid")
return
default:
s.log.Warn("oauth refresh refused", "reason", "token_unusable", "client_id", clientID)
writeOAuthError(w, http.StatusBadRequest, errInvalidGrant, "the refresh token is invalid")
return
}
if old.ClientID != clientID {
// Not this client's token. Revoke the family: a refresh token that has
// reached the wrong client has leaked.
_ = s.store.RevokeFamily(r.Context(), old.FamilyID, "client_mismatch_on_refresh")
s.log.Warn("oauth refresh refused", "reason", "client_mismatch", "client_id", clientID)
writeOAuthError(w, http.StatusBadRequest, errInvalidGrant, "the refresh token is invalid")
return
}
// Same family: the rotation continues the lineage, so reuse detection can
// still revoke every descendant if an older token reappears.
pair, err := s.store.IssuePair(r.Context(), Token{
ClientID: old.ClientID,
UserID: old.UserID,
OrgID: old.OrgID,
Scopes: old.Scopes,
Audience: old.Audience,
}, old.FamilyID)
if err != nil {
s.log.Error("oauth: could not rotate tokens", "error", err)
writeOAuthError(w, http.StatusInternalServerError, errServerError, "")
return
}
s.log.Info("oauth tokens issued",
"grant_type", "refresh_token", "client_id", old.ClientID,
"user_id", old.UserID, "family_id", pair.FamilyID)
writeTokenResponse(w, pair)
}
func writeTokenResponse(w http.ResponseWriter, pair TokenPair) {
w.Header().Set("Content-Type", "application/json; charset=utf-8")
// RFC 6749 section 5.1 requires both of these on a token response. The
// body is a credential; nothing may cache it.
w.Header().Set("Cache-Control", "no-store")
w.Header().Set("Pragma", "no-cache")
writeJSONBody(w, http.StatusOK, tokenResponse{
AccessToken: pair.AccessToken,
TokenType: "Bearer",
ExpiresIn: pair.ExpiresIn,
RefreshToken: pair.RefreshToken,
Scope: strings.Join(pair.Scopes, " "),
})
}
/* ── Revocation (RFC 7009) ──────────────────────────────────────────────── */
// RevokeHandler serves token revocation.
//
// RFC 7009 requires 200 for an unknown token: answering 404 would turn this
// into an oracle for whether a token exists. The store already behaves that
// way; this handler just does not undo it.
func (s *Server) RevokeHandler() http.Handler {
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.Method != http.MethodPost {
w.Header().Set("Allow", http.MethodPost)
writeOAuthError(w, http.StatusMethodNotAllowed, errInvalidRequest, "POST only")
return
}
if err := r.ParseForm(); err != nil {
writeOAuthError(w, http.StatusBadRequest, errInvalidRequest, "the request body could not be parsed")
return
}
token := r.PostFormValue("token")
if token == "" {
writeOAuthError(w, http.StatusBadRequest, errInvalidRequest, "token is required")
return
}
if err := s.store.RevokeToken(r.Context(), token, "client_revocation"); err != nil {
s.log.Error("oauth: revocation failed", "error", err)
writeOAuthError(w, http.StatusInternalServerError, errServerError, "")
return
}
s.log.Info("oauth token revoked", "client_id", r.PostFormValue("client_id"))
w.Header().Set("Cache-Control", "no-store")
w.WriteHeader(http.StatusOK)
})
}

View File

@@ -0,0 +1,130 @@
package oauth
import (
"context"
"fmt"
"time"
)
// Cleanup of spent and expired OAuth rows.
//
// WHAT IS DELETED, AND WHAT IS DELIBERATELY NOT
//
// Only rows that can no longer authenticate anything. Every predicate below
// requires the row to be past its expiry — not merely consumed, not merely
// revoked — because those two states are evidence, and evidence is worth
// keeping until it stops being relevant.
//
// A consumed refresh token in particular must outlive its usefulness: it is
// what REUSE DETECTION matches against. Delete it the moment it is spent and a
// stolen token replayed a minute later looks like an unknown token rather than
// a theft, and the family is never revoked. So a consumed refresh token is kept
// until its original expiry, by which point replaying it proves nothing anyway.
//
// A revoked token is kept for the same reason plus one more: "this token was
// revoked at 14:02 for refresh_token_reuse" is an answer to a question someone
// will eventually ask.
//
// GRACE. Everything is deleted a grace period AFTER expiry rather than at it,
// so a clock skewed between two instances cannot delete a row another instance
// still considers live.
//
// SAFE TO RUN TWICE, AND SAFE TO RUN CONCURRENTLY. Every statement is a bounded
// DELETE with a predicate that no longer matches once the row is gone. Two
// workers running at once delete disjoint sets and neither errors.
// CleanupGrace is how long a dead row is kept past its expiry.
//
// An hour is far beyond any plausible clock skew between instances and short
// enough that the tables do not accumulate. It also means a support question
// asked within the hour can still see the row.
const CleanupGrace = time.Hour
// CleanupBatch bounds one pass.
//
// Bounded because an unbounded DELETE holds locks for as long as it runs, and
// on a table that every MCP request reads that is a latency spike nobody can
// explain afterwards. Five thousand rows is milliseconds; if there is more, the
// next pass takes it.
const CleanupBatch = 5000
// CleanupResult reports what one pass removed.
type CleanupResult struct {
Grants int64
AccessTokens int64
RefreshTokens int64
CompletedInOne bool // false when a batch filled, meaning more remains
}
// Cleanup removes expired authorization codes and tokens.
//
// Returns counts rather than logging them, so the caller decides the level and
// this function stays usable from a test.
func (s *Store) Cleanup(ctx context.Context) (CleanupResult, error) {
cutoff := s.now().Add(-CleanupGrace)
var out CleanupResult
// Authorization codes. Sixty-second TTL, so almost every row here is
// already dead; this is the highest-volume and cheapest of the three.
//
// ctid rather than id in the subquery because it is the physical row
// address — the planner can go straight to it without a second index
// lookup, which is what keeps a bounded delete genuinely cheap.
tag, err := s.db.Exec(ctx,
`DELETE FROM oauth_grants
WHERE ctid IN (
SELECT ctid FROM oauth_grants WHERE expires_at < $1 LIMIT $2
)`, cutoff, CleanupBatch)
if err != nil {
return out, fmt.Errorf("oauth: cleanup grants: %w", err)
}
out.Grants = tag.RowsAffected()
// Access tokens. Fifteen-minute TTL. An expired one cannot authenticate —
// FindAccessToken's predicate already excludes it — so deleting it removes
// no capability.
tag, err = s.db.Exec(ctx,
`DELETE FROM oauth_tokens
WHERE ctid IN (
SELECT ctid FROM oauth_tokens
WHERE token_type = 'access' AND expires_at < $1
LIMIT $2
)`, cutoff, CleanupBatch)
if err != nil {
return out, fmt.Errorf("oauth: cleanup access tokens: %w", err)
}
out.AccessTokens = tag.RowsAffected()
// Refresh tokens, and this is the one with a real constraint on it.
//
// EXPIRY ONLY — not `consumed_at IS NOT NULL`, and not `revoked_at IS NOT
// NULL`. A consumed refresh token is what RedeemRefreshToken matches to
// detect reuse; deleting it early turns a detectable theft into an
// unremarkable "unknown token" and the family is never revoked. Thirty-day
// TTL means these are the longest-lived rows in the schema, which is the
// price of that detection and is worth paying.
tag, err = s.db.Exec(ctx,
`DELETE FROM oauth_tokens
WHERE ctid IN (
SELECT ctid FROM oauth_tokens
WHERE token_type = 'refresh' AND expires_at < $1
LIMIT $2
)`, cutoff, CleanupBatch)
if err != nil {
return out, fmt.Errorf("oauth: cleanup refresh tokens: %w", err)
}
out.RefreshTokens = tag.RowsAffected()
out.CompletedInOne = out.Grants < CleanupBatch &&
out.AccessTokens < CleanupBatch &&
out.RefreshTokens < CleanupBatch
return out, nil
}
// RevokeExpiredFamilies is deliberately absent.
//
// It looks like it belongs here — "tidy up families whose tokens have all
// lapsed" — and it would do nothing. Revocation is a state on a row, and a row
// that has been deleted has no state to set. A family whose every token has
// expired and been swept simply ceases to exist, which is the correct outcome
// and requires no work.

View File

@@ -0,0 +1,240 @@
package oauth
import (
"context"
"sync"
"testing"
"time"
)
/* ── What cleanup removes ───────────────────────────────────────────────── */
func TestCleanupRemovesOnlyDeadRows(t *testing.T) {
h := newHarness(t)
ctx := context.Background()
clientID := h.register()
// A live pair, and a spent code.
code := h.authorizeOK(clientID, verifier43)
live := decodeTokens(t, h.exchange(clientID, code, verifier43))
// A second, which we let expire.
oldCode := h.authorizeOK(clientID, verifier43)
old := decodeTokens(t, h.exchange(clientID, oldCode, verifier43))
before := countRows(t, h)
// Past the access token TTL and the grace, but well inside the refresh
// token's thirty days.
h.advance(AccessTokenTTL + CleanupGrace + time.Minute)
result, err := h.store.Cleanup(ctx)
if err != nil {
t.Fatalf("Cleanup: %v", err)
}
// Both access tokens and both codes are dead; both refresh tokens are not.
if result.AccessTokens != 2 {
t.Errorf("removed %d access tokens, want 2", result.AccessTokens)
}
if result.Grants != 2 {
t.Errorf("removed %d grants, want 2", result.Grants)
}
if result.RefreshTokens != 0 {
t.Errorf("removed %d refresh tokens, want 0 — they live thirty days", result.RefreshTokens)
}
if !result.CompletedInOne {
t.Error("a small cleanup reported that more remained")
}
after := countRows(t, h)
if after.tokens >= before.tokens {
t.Error("cleanup removed nothing")
}
// The live refresh tokens must still work. This is the property that
// matters: cleanup must not disconnect anybody.
for name, token := range map[string]string{"live": live.RefreshToken, "old": old.RefreshToken} {
if _, err := h.store.RedeemRefreshToken(ctx, token); err != nil {
t.Errorf("the %s refresh token stopped working after cleanup: %v", name, err)
}
}
}
// The subtle one: a CONSUMED refresh token must survive until its expiry,
// because it is what reuse detection matches against. Delete it early and a
// replayed stolen token looks unknown rather than stolen, and the family is
// never revoked.
func TestCleanupKeepsConsumedRefreshTokensForReuseDetection(t *testing.T) {
h := newHarness(t)
ctx := context.Background()
clientID := h.register()
code := h.authorizeOK(clientID, verifier43)
first := decodeTokens(t, h.exchange(clientID, code, verifier43))
// Rotate: `first` is now consumed.
if _, err := h.store.RedeemRefreshToken(ctx, first.RefreshToken); err != nil {
t.Fatalf("rotate: %v", err)
}
// Cleanup well past the access token TTL, but inside the refresh TTL.
h.advance(AccessTokenTTL + CleanupGrace + time.Hour)
if _, err := h.store.Cleanup(ctx); err != nil {
t.Fatalf("Cleanup: %v", err)
}
// Replaying the consumed token must STILL be detected as reuse.
_, err := h.store.RedeemRefreshToken(ctx, first.RefreshToken)
if err != ErrRefreshReuse {
t.Errorf("err = %v, want ErrRefreshReuse — cleanup destroyed the evidence "+
"that makes theft detectable", err)
}
}
// Revoked rows are kept until expiry too: "revoked at 14:02 for
// refresh_token_reuse" is an answer somebody will eventually need.
func TestCleanupKeepsRevokedRowsUntilExpiry(t *testing.T) {
h := newHarness(t)
ctx := context.Background()
clientID := h.register()
code := h.authorizeOK(clientID, verifier43)
tokens := decodeTokens(t, h.exchange(clientID, code, verifier43))
if err := h.store.RevokeToken(ctx, tokens.AccessToken, "test"); err != nil {
t.Fatalf("revoke: %v", err)
}
// Just past the access TTL: the access row goes, the refresh row stays.
h.advance(AccessTokenTTL + CleanupGrace + time.Minute)
if _, err := h.store.Cleanup(ctx); err != nil {
t.Fatalf("Cleanup: %v", err)
}
var revokedRefresh int
if err := h.h.Pool.QueryRow(ctx,
`SELECT count(*) FROM oauth_tokens WHERE token_type='refresh' AND revoked_at IS NOT NULL`).
Scan(&revokedRefresh); err != nil {
t.Fatalf("count: %v", err)
}
if revokedRefresh != 1 {
t.Errorf("%d revoked refresh rows kept, want 1 — the audit trail was swept", revokedRefresh)
}
}
// Nothing is deleted before the grace period, so clock skew between instances
// cannot destroy a row another instance still considers live.
func TestCleanupHonoursTheGracePeriod(t *testing.T) {
h := newHarness(t)
ctx := context.Background()
clientID := h.register()
code := h.authorizeOK(clientID, verifier43)
decodeTokens(t, h.exchange(clientID, code, verifier43))
// Expired, but inside the grace.
h.advance(AccessTokenTTL + time.Minute)
result, err := h.store.Cleanup(ctx)
if err != nil {
t.Fatalf("Cleanup: %v", err)
}
if result.AccessTokens != 0 {
t.Errorf("removed %d access tokens inside the grace period, want 0", result.AccessTokens)
}
}
/* ── Safety ─────────────────────────────────────────────────────────────── */
func TestCleanupIsIdempotent(t *testing.T) {
h := newHarness(t)
ctx := context.Background()
clientID := h.register()
code := h.authorizeOK(clientID, verifier43)
decodeTokens(t, h.exchange(clientID, code, verifier43))
h.advance(AccessTokenTTL + CleanupGrace + time.Minute)
first, err := h.store.Cleanup(ctx)
if err != nil {
t.Fatalf("first: %v", err)
}
second, err := h.store.Cleanup(ctx)
if err != nil {
t.Fatalf("second: %v", err)
}
if second.AccessTokens != 0 || second.Grants != 0 || second.RefreshTokens != 0 {
t.Errorf("a second cleanup removed more rows: %+v (first was %+v)", second, first)
}
}
// Two workers running cleanup at once must not error and must not
// double-count. Run with -race.
func TestConcurrentCleanupIsSafe(t *testing.T) {
h := newHarness(t)
ctx := context.Background()
clientID := h.register()
for i := 0; i < 6; i++ {
code := h.authorizeOK(clientID, verifier43)
decodeTokens(t, h.exchange(clientID, code, verifier43))
}
h.advance(AccessTokenTTL + CleanupGrace + time.Minute)
const workers = 4
var wg sync.WaitGroup
var mu sync.Mutex
var total int64
errs := make([]error, 0, workers)
for i := 0; i < workers; i++ {
wg.Add(1)
go func() {
defer wg.Done()
r, err := h.store.Cleanup(ctx)
mu.Lock()
defer mu.Unlock()
if err != nil {
errs = append(errs, err)
return
}
total += r.AccessTokens
}()
}
wg.Wait()
if len(errs) > 0 {
t.Fatalf("concurrent cleanup errored: %v", errs)
}
// Six access tokens existed; between them the workers removed exactly six.
// More would mean a row was counted twice.
if total != 6 {
t.Errorf("workers removed %d access tokens between them, want 6", total)
}
}
func TestCleanupOnAnEmptyDatabaseIsHarmless(t *testing.T) {
h := newHarness(t)
result, err := h.store.Cleanup(context.Background())
if err != nil {
t.Fatalf("Cleanup on empty: %v", err)
}
if result.Grants != 0 || result.AccessTokens != 0 || result.RefreshTokens != 0 {
t.Errorf("cleanup on an empty database removed %+v", result)
}
}
/* ── Helpers ────────────────────────────────────────────────────────────── */
type rowCounts struct{ grants, tokens int }
func countRows(t *testing.T, h *harness) rowCounts {
t.Helper()
var c rowCounts
ctx := context.Background()
if err := h.h.Pool.QueryRow(ctx, `SELECT count(*) FROM oauth_grants`).Scan(&c.grants); err != nil {
t.Fatalf("count grants: %v", err)
}
if err := h.h.Pool.QueryRow(ctx, `SELECT count(*) FROM oauth_tokens`).Scan(&c.tokens); err != nil {
t.Fatalf("count tokens: %v", err)
}
return c
}

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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