mcp connection
This commit is contained in:
374
go-api/internal/oauth/abuse_test.go
Normal file
374
go-api/internal/oauth/abuse_test.go
Normal file
@@ -0,0 +1,374 @@
|
||||
package oauth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"net/url"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
)
|
||||
|
||||
// Abuse cases and input limits.
|
||||
//
|
||||
// The property every test here shares: a refusal must not teach the caller
|
||||
// anything. Not whether a token existed, not whether an account is suspended
|
||||
// rather than deleted, not what the database is called, and never the value of
|
||||
// a credential that was presented.
|
||||
|
||||
/* ── Nothing sensitive reaches a response ───────────────────────────────── */
|
||||
|
||||
// The broadest check in this file: drive every failure path with known secret
|
||||
// values, and assert none of them comes back.
|
||||
func TestNoSecretEverAppearsInAResponse(t *testing.T) {
|
||||
h := newHarness(t)
|
||||
clientID := h.register()
|
||||
|
||||
const (
|
||||
secretVerifier = "SENTINELverifier0123456789abcdefghijklmnop"
|
||||
secretCode = "SENTINELcodevalue"
|
||||
secretToken = "SENTINELtokenvalue"
|
||||
)
|
||||
|
||||
bodies := map[string]string{}
|
||||
|
||||
// A failed exchange, with a sentinel code and verifier.
|
||||
bodies["bad code"] = h.exchange(clientID, secretCode, secretVerifier).Body.String()
|
||||
|
||||
// A real code with the wrong verifier.
|
||||
real := h.authorizeOK(clientID, verifier43)
|
||||
bodies["bad verifier"] = h.exchange(clientID, real, secretVerifier).Body.String()
|
||||
|
||||
// A refresh with a sentinel token.
|
||||
bodies["bad refresh"] = h.token(url.Values{
|
||||
"grant_type": {"refresh_token"}, "refresh_token": {secretToken}, "client_id": {clientID},
|
||||
}).Body.String()
|
||||
|
||||
// Revocation of an unknown token.
|
||||
form := url.Values{"token": {secretToken}}
|
||||
req := httptest.NewRequest(http.MethodPost, "/oauth/revoke", strings.NewReader(form.Encode()))
|
||||
req.Header.Set("Content-Type", "application/x-www-form-urlencoded")
|
||||
rec := httptest.NewRecorder()
|
||||
h.server.RevokeHandler().ServeHTTP(rec, req)
|
||||
bodies["revoke unknown"] = rec.Body.String()
|
||||
|
||||
for where, body := range bodies {
|
||||
for what, secret := range map[string]string{
|
||||
"code_verifier": secretVerifier,
|
||||
"authorization code": secretCode,
|
||||
"token": secretToken,
|
||||
} {
|
||||
if strings.Contains(body, secret) {
|
||||
t.Errorf("%s: the response echoes the presented %s:\n%s", where, what, body)
|
||||
}
|
||||
}
|
||||
// Nor may it leak the shape of the system.
|
||||
for _, tell := range []string{"SQLSTATE", "pq:", "pgx", "oauth_tokens", "oauth_grants",
|
||||
"password", "Krow-force", "relation", "column"} {
|
||||
if strings.Contains(body, tell) {
|
||||
t.Errorf("%s: the response leaks an internal detail (%q):\n%s", where, tell, body)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/* ── Replay ─────────────────────────────────────────────────────────────── */
|
||||
|
||||
// Ten attempts to spend one code. Exactly one may succeed.
|
||||
func TestAuthorizationCodeReplayUnderLoad(t *testing.T) {
|
||||
h := newHarness(t)
|
||||
clientID := h.register()
|
||||
code := h.authorizeOK(clientID, verifier43)
|
||||
|
||||
succeeded := 0
|
||||
for i := 0; i < 10; i++ {
|
||||
if h.exchange(clientID, code, verifier43).Code == http.StatusOK {
|
||||
succeeded++
|
||||
}
|
||||
}
|
||||
if succeeded != 1 {
|
||||
t.Errorf("%d of 10 exchanges of the same code succeeded, want exactly 1", succeeded)
|
||||
}
|
||||
}
|
||||
|
||||
// The same, concurrently. A single-use credential redeemed by two racing
|
||||
// callers must be spent exactly once — this is the property that
|
||||
// UPDATE … RETURNING buys, and the one a SELECT-then-UPDATE would lose.
|
||||
// Run with -race.
|
||||
func TestConcurrentCodeRedemptionSpendsItOnce(t *testing.T) {
|
||||
h := newHarness(t)
|
||||
clientID := h.register()
|
||||
code := h.authorizeOK(clientID, verifier43)
|
||||
|
||||
const racers = 8
|
||||
var wg sync.WaitGroup
|
||||
var mu sync.Mutex
|
||||
succeeded := 0
|
||||
|
||||
for i := 0; i < racers; i++ {
|
||||
wg.Add(1)
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
if h.exchange(clientID, code, verifier43).Code == http.StatusOK {
|
||||
mu.Lock()
|
||||
succeeded++
|
||||
mu.Unlock()
|
||||
}
|
||||
}()
|
||||
}
|
||||
wg.Wait()
|
||||
|
||||
if succeeded != 1 {
|
||||
t.Errorf("%d of %d concurrent redemptions succeeded, want exactly 1", succeeded, racers)
|
||||
}
|
||||
}
|
||||
|
||||
// Concurrent refresh of the same token: one rotation, not several. Two
|
||||
// successes would mean two live families from one credential.
|
||||
// Run with -race.
|
||||
func TestConcurrentRefreshRotatesOnce(t *testing.T) {
|
||||
h := newHarness(t)
|
||||
clientID := h.register()
|
||||
code := h.authorizeOK(clientID, verifier43)
|
||||
tokens := decodeTokens(t, h.exchange(clientID, code, verifier43))
|
||||
|
||||
const racers = 8
|
||||
var wg sync.WaitGroup
|
||||
var mu sync.Mutex
|
||||
succeeded := 0
|
||||
|
||||
for i := 0; i < racers; i++ {
|
||||
wg.Add(1)
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
rec := h.token(url.Values{
|
||||
"grant_type": {"refresh_token"}, "refresh_token": {tokens.RefreshToken},
|
||||
"client_id": {clientID},
|
||||
})
|
||||
if rec.Code == http.StatusOK {
|
||||
mu.Lock()
|
||||
succeeded++
|
||||
mu.Unlock()
|
||||
}
|
||||
}()
|
||||
}
|
||||
wg.Wait()
|
||||
|
||||
if succeeded != 1 {
|
||||
t.Errorf("%d of %d concurrent refreshes succeeded, want exactly 1 — "+
|
||||
"a refresh token was double-spent", succeeded, racers)
|
||||
}
|
||||
}
|
||||
|
||||
/* ── Open redirect ──────────────────────────────────────────────────────── */
|
||||
|
||||
// Every shape of redirect tampering, each of which has been a real CVE
|
||||
// somewhere. None may be honoured, and none may be answered WITH a redirect —
|
||||
// redirecting an error to an unvalidated URI is the open redirect itself.
|
||||
func TestOpenRedirectAttempts(t *testing.T) {
|
||||
h := newHarness(t)
|
||||
clientID := h.register()
|
||||
|
||||
for name, redirect := range map[string]string{
|
||||
"different host": "https://attacker.example/cb",
|
||||
"prefix extension": testRedirect + ".attacker.example",
|
||||
"path traversal": testRedirect + "/../../evil",
|
||||
"userinfo trick": "https://claude.example.test@attacker.example/cb",
|
||||
"added query": testRedirect + "?next=https://attacker.example",
|
||||
"protocol swap": strings.Replace(testRedirect, "https", "http", 1),
|
||||
"case variation": strings.ToUpper(testRedirect),
|
||||
"trailing slash": testRedirect + "/",
|
||||
"double slash": "//attacker.example/cb",
|
||||
"encoded traversal": testRedirect + "/%2e%2e/evil",
|
||||
"null byte": testRedirect + "\x00.attacker.example",
|
||||
"newline injection": testRedirect + "\nLocation: https://attacker.example",
|
||||
"javascript": "javascript:alert(1)",
|
||||
"completely missing": "",
|
||||
} {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
rec := h.authorize(map[string]string{
|
||||
"client_id": clientID, "redirect_uri": redirect, "response_type": "code",
|
||||
"state": "xyz", "code_challenge": ChallengeFor(verifier43),
|
||||
"code_challenge_method": "S256", "resource": testResource,
|
||||
})
|
||||
|
||||
if rec.Code == http.StatusFound {
|
||||
location := rec.Header().Get("Location")
|
||||
t.Fatalf("answered with a redirect to %q — an unregistered target "+
|
||||
"must produce a direct error, never a redirect", location)
|
||||
}
|
||||
if rec.Code != http.StatusBadRequest {
|
||||
t.Errorf("status = %d, want 400", rec.Code)
|
||||
}
|
||||
if strings.Contains(rec.Body.String(), "attacker.example") {
|
||||
t.Error("the error echoes the attacker's host back")
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
/* ── Malformed and oversized input ──────────────────────────────────────── */
|
||||
|
||||
func TestRegistrationRejectsMalformedBodies(t *testing.T) {
|
||||
h := newHarness(t)
|
||||
|
||||
for name, body := range map[string]string{
|
||||
"not json": `not json at all`,
|
||||
"truncated": `{"client_name":`,
|
||||
"null": `null`,
|
||||
"array": `[]`,
|
||||
"deeply nested": `{"client_name":` + strings.Repeat(`[`, 2000) + strings.Repeat(`]`, 2000) + `}`,
|
||||
"empty": ``,
|
||||
"wrong types": `{"client_name":123,"redirect_uris":"not-an-array"}`,
|
||||
"huge name": `{"client_name":"` + strings.Repeat("A", 100_000) + `","redirect_uris":["https://a.test/cb"]}`,
|
||||
"too many uris": `{"client_name":"x","redirect_uris":[` + strings.TrimSuffix(strings.Repeat(`"https://a.test/cb",`, 50), ",") + `]}`,
|
||||
"oversized body": `{"client_name":"` + strings.Repeat("A", 20<<10) + `"}`,
|
||||
} {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
req := httptest.NewRequest(http.MethodPost, "/oauth/register", strings.NewReader(body))
|
||||
rec := httptest.NewRecorder()
|
||||
// A panic here fails the test by crashing it, which is the assertion.
|
||||
h.server.RegisterHandler().ServeHTTP(rec, req)
|
||||
|
||||
if rec.Code == http.StatusCreated {
|
||||
// Only the "huge name" case may legitimately succeed, truncated.
|
||||
if name != "huge name" {
|
||||
t.Errorf("status = %d; a malformed registration was accepted", rec.Code)
|
||||
}
|
||||
return
|
||||
}
|
||||
if rec.Code < 400 || rec.Code >= 500 {
|
||||
t.Errorf("status = %d, want a 4xx", rec.Code)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// A client name is stored, shown on a consent screen, and attacker-controlled.
|
||||
// It must be bounded, or registration becomes free storage.
|
||||
func TestClientNameIsBounded(t *testing.T) {
|
||||
h := newHarness(t)
|
||||
body := `{"client_name":"` + strings.Repeat("A", 5000) + `","redirect_uris":["https://a.test/cb"]}`
|
||||
|
||||
rec := httptest.NewRecorder()
|
||||
h.server.RegisterHandler().ServeHTTP(rec,
|
||||
httptest.NewRequest(http.MethodPost, "/oauth/register", strings.NewReader(body)))
|
||||
if rec.Code != http.StatusCreated {
|
||||
t.Fatalf("status = %d", rec.Code)
|
||||
}
|
||||
|
||||
var stored string
|
||||
if err := h.h.Pool.QueryRow(context.Background(),
|
||||
`SELECT client_name FROM oauth_clients ORDER BY created_date DESC LIMIT 1`).Scan(&stored); err != nil {
|
||||
t.Fatalf("read: %v", err)
|
||||
}
|
||||
if len(stored) > 200 {
|
||||
t.Errorf("stored client_name is %d characters; the column's CHECK allows 200", len(stored))
|
||||
}
|
||||
}
|
||||
|
||||
func TestTokenEndpointRejectsMalformedRequests(t *testing.T) {
|
||||
h := newHarness(t)
|
||||
|
||||
for name, tc := range map[string]struct {
|
||||
body string
|
||||
contentType string
|
||||
}{
|
||||
"no content type": {"grant_type=authorization_code", ""},
|
||||
"json body": {`{"grant_type":"authorization_code"}`, "application/json"},
|
||||
"empty": {"", "application/x-www-form-urlencoded"},
|
||||
"garbage": {"%%%%", "application/x-www-form-urlencoded"},
|
||||
"huge": {"grant_type=authorization_code&code=" + strings.Repeat("A", 200_000), "application/x-www-form-urlencoded"},
|
||||
"repeated params": {"grant_type=authorization_code&grant_type=password", "application/x-www-form-urlencoded"},
|
||||
"null grant": {"grant_type=", "application/x-www-form-urlencoded"},
|
||||
"unknown grant": {"grant_type=magic", "application/x-www-form-urlencoded"},
|
||||
"injection in code": {"grant_type=authorization_code&code=' OR 1=1 --&client_id=x&redirect_uri=y&code_verifier=" +
|
||||
strings.Repeat("a", 43), "application/x-www-form-urlencoded"},
|
||||
} {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
req := httptest.NewRequest(http.MethodPost, "/oauth/token", strings.NewReader(tc.body))
|
||||
if tc.contentType != "" {
|
||||
req.Header.Set("Content-Type", tc.contentType)
|
||||
}
|
||||
rec := httptest.NewRecorder()
|
||||
h.server.TokenHandler().ServeHTTP(rec, req)
|
||||
|
||||
if rec.Code == http.StatusOK {
|
||||
t.Errorf("a malformed token request succeeded: %s", rec.Body.String())
|
||||
}
|
||||
if rec.Code >= 500 {
|
||||
t.Errorf("status = %d; a malformed request must not be an internal error: %s",
|
||||
rec.Code, rec.Body.String())
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
/* ── Scope escalation ───────────────────────────────────────────────────── */
|
||||
|
||||
// krow.write must be unreachable from every angle: registration, authorization,
|
||||
// and the consent POST.
|
||||
func TestWriteScopeIsUnreachable(t *testing.T) {
|
||||
h := newHarness(t)
|
||||
|
||||
t.Run("at registration", func(t *testing.T) {
|
||||
body := `{"client_name":"x","redirect_uris":["` + testRedirect + `"],"scope":"` + ScopeWrite + `"}`
|
||||
rec := httptest.NewRecorder()
|
||||
h.server.RegisterHandler().ServeHTTP(rec,
|
||||
httptest.NewRequest(http.MethodPost, "/oauth/register", strings.NewReader(body)))
|
||||
if rec.Code == http.StatusCreated {
|
||||
t.Error("a client registered for krow.write")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("at authorization", func(t *testing.T) {
|
||||
clientID := h.register()
|
||||
rec := h.authorize(map[string]string{
|
||||
"client_id": clientID, "redirect_uri": testRedirect, "response_type": "code",
|
||||
"state": "xyz", "code_challenge": ChallengeFor(verifier43),
|
||||
"code_challenge_method": "S256", "resource": testResource,
|
||||
"scope": ScopeRead + " " + ScopeWrite,
|
||||
})
|
||||
loc, _ := url.Parse(rec.Header().Get("Location"))
|
||||
if loc.Query().Get("code") != "" {
|
||||
t.Error("an authorization requesting krow.write produced a code")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("no issued token carries it", func(t *testing.T) {
|
||||
clientID := h.register()
|
||||
code := h.authorizeOK(clientID, verifier43)
|
||||
tokens := decodeTokens(t, h.exchange(clientID, code, verifier43))
|
||||
if strings.Contains(tokens.Scope, ScopeWrite) {
|
||||
t.Errorf("an issued token carries %q", tokens.Scope)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
/* ── Cache and transport headers ────────────────────────────────────────── */
|
||||
|
||||
// A credential-bearing response must never be cached, and no endpoint may put
|
||||
// a token in a URL.
|
||||
func TestSensitiveResponsesAreNotCacheable(t *testing.T) {
|
||||
h := newHarness(t)
|
||||
clientID := h.register()
|
||||
code := h.authorizeOK(clientID, verifier43)
|
||||
|
||||
rec := h.exchange(clientID, code, verifier43)
|
||||
for header, want := range map[string]string{
|
||||
"Cache-Control": "no-store",
|
||||
"Pragma": "no-cache",
|
||||
} {
|
||||
if got := rec.Header().Get(header); !strings.Contains(got, want) {
|
||||
t.Errorf("%s = %q, want it to contain %q", header, got, want)
|
||||
}
|
||||
}
|
||||
|
||||
// The authorization redirect carries a code in its query — that is the
|
||||
// protocol — but it must never carry a token.
|
||||
approved := h.authorize(authorizeParamsFor(clientID, verifier43))
|
||||
if strings.Contains(approved.Header().Get("Location"), "access_token") {
|
||||
t.Error("an access token appeared in a redirect URL")
|
||||
}
|
||||
}
|
||||
172
go-api/internal/oauth/authenticator.go
Normal file
172
go-api/internal/oauth/authenticator.go
Normal file
@@ -0,0 +1,172 @@
|
||||
package oauth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/rand"
|
||||
"encoding/hex"
|
||||
"errors"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"strings"
|
||||
|
||||
"github.com/krow/krow-backend/go-api/internal/auth"
|
||||
"github.com/krow/krow-backend/go-api/internal/authctx"
|
||||
)
|
||||
|
||||
// Authenticator is the production implementation of
|
||||
// mcpserver.TokenAuthenticator.
|
||||
//
|
||||
// This is where Phase 2's seam is filled in, and the shape of it is the whole
|
||||
// argument for having defined the interface first: one method, taking a raw
|
||||
// token, returning the same authctx.Identity a cookie produces. Nothing
|
||||
// downstream — not tools.Context, not the policy table, not a single handler —
|
||||
// can tell which path built the identity, so authorization cannot drift between
|
||||
// them.
|
||||
//
|
||||
// THE IDENTITY IS BUILT FROM THE USER ROW, NOT FROM THE TOKEN.
|
||||
//
|
||||
// oauth_tokens carries org_id, and it would be cheaper to read it from there.
|
||||
// It is deliberately not: the token row records the tenant AT ISSUE TIME, and a
|
||||
// token can outlive the fact. A user moved to another organisation, or
|
||||
// suspended, would keep working against a stale claim until the token expired.
|
||||
// Re-reading the user costs one indexed lookup and makes suspension take effect
|
||||
// on the next call — which is exactly what httpserver/auth.go already does for
|
||||
// cookies, and the bearer path must not be weaker than the cookie path.
|
||||
type Authenticator struct {
|
||||
store *Store
|
||||
users UserLookup
|
||||
log *slog.Logger
|
||||
|
||||
// audience is this deployment's canonical MCP resource URI. A token whose
|
||||
// audience is anything else is refused — see the note in Authenticate.
|
||||
audience string
|
||||
}
|
||||
|
||||
// UserLookup is the subset of the existing user store this needs. auth.UserStore
|
||||
// satisfies it; nothing here builds a second user table or password store.
|
||||
type UserLookup interface {
|
||||
FindByID(ctx context.Context, id string) (auth.User, error)
|
||||
}
|
||||
|
||||
// NewAuthenticator builds the production token authenticator.
|
||||
func NewAuthenticator(store *Store, users UserLookup, audience string, log *slog.Logger) *Authenticator {
|
||||
if log == nil {
|
||||
log = slog.Default()
|
||||
}
|
||||
return &Authenticator{store: store, users: users, audience: audience, log: log}
|
||||
}
|
||||
|
||||
// ErrAudienceMismatch is internal. It never reaches a client — see the single
|
||||
// return below — but it is distinct so the log can say what happened.
|
||||
var ErrAudienceMismatch = errors.New("oauth: token audience does not match this resource")
|
||||
|
||||
// Authenticate resolves a bearer token into a KROW identity.
|
||||
//
|
||||
// EVERY failure returns the same error. Unknown, expired, revoked, wrong
|
||||
// audience, suspended user, deleted user — one answer, because a caller who can
|
||||
// tell them apart learns things they should not: that a token once existed,
|
||||
// that an account was suspended rather than deleted, that this server is not
|
||||
// the intended audience for a token they hold. Same discipline as
|
||||
// tools.Denied() and the session path's identical answer to "not found" and
|
||||
// "expired".
|
||||
//
|
||||
// The reason goes to the log, at warn, where the operator is.
|
||||
func (a *Authenticator) Authenticate(ctx context.Context, rawToken string) (authctx.Identity, error) {
|
||||
if strings.TrimSpace(rawToken) == "" {
|
||||
return authctx.Identity{}, ErrTokenUnusable
|
||||
}
|
||||
|
||||
// 1. The token must exist, be an access token, be unexpired and unrevoked.
|
||||
// All four are in the query's predicate.
|
||||
token, err := a.store.FindAccessToken(ctx, rawToken)
|
||||
if err != nil {
|
||||
a.log.Warn("mcp bearer refused", "reason", "token_unusable")
|
||||
return authctx.Identity{}, ErrTokenUnusable
|
||||
}
|
||||
|
||||
// 2. Audience. RFC 8707 and the MCP spec both require a server to verify
|
||||
// that a token was issued FOR IT. Without this check, a token minted by
|
||||
// this authorization server for some other resource would be spendable
|
||||
// here — the confused-deputy problem the spec calls out explicitly. The
|
||||
// comparison is against configuration, never against anything in the
|
||||
// request: a resource value supplied by the caller would let the caller
|
||||
// choose their own audience.
|
||||
if token.Audience != a.audience {
|
||||
a.log.Warn("mcp bearer refused",
|
||||
"reason", "audience_mismatch",
|
||||
"token_id", token.ID,
|
||||
"expected", a.audience,
|
||||
"presented", token.Audience)
|
||||
return authctx.Identity{}, ErrTokenUnusable
|
||||
}
|
||||
|
||||
// 3. Scope. krow.read is the only scope this phase issues, and the MCP
|
||||
// surface is read-only, so a token without it has no business here. The
|
||||
// check is present rather than implied so that adding krow.write later
|
||||
// is a change in one place.
|
||||
if !hasScope(token.Scopes, ScopeRead) {
|
||||
a.log.Warn("mcp bearer refused", "reason", "missing_scope", "token_id", token.ID)
|
||||
return authctx.Identity{}, ErrTokenUnusable
|
||||
}
|
||||
|
||||
// 4. The user, re-read live. See the type comment for why this is not taken
|
||||
// from the token row.
|
||||
user, err := a.users.FindByID(ctx, token.UserID)
|
||||
if err != nil {
|
||||
// The FK cascades, so a missing user should be unreachable. If it
|
||||
// happens the token is orphaned and worth killing.
|
||||
a.log.Warn("mcp bearer refused", "reason", "user_missing", "token_id", token.ID)
|
||||
_ = a.store.RevokeFamily(ctx, token.FamilyID, "user_missing")
|
||||
return authctx.Identity{}, ErrTokenUnusable
|
||||
}
|
||||
|
||||
// 5. Suspension revokes on contact, exactly as the cookie path does. Not
|
||||
// "the token stops working at expiry" — a suspended account must lose
|
||||
// access on its next request, and leaving the family alive would mean it
|
||||
// kept a working credential for up to thirty days.
|
||||
if !user.IsActive() {
|
||||
a.log.Warn("mcp bearer refused",
|
||||
"reason", "user_inactive", "user_id", user.ID, "status", user.Status)
|
||||
_ = a.store.RevokeFamily(ctx, token.FamilyID, "user_suspended")
|
||||
return authctx.Identity{}, ErrTokenUnusable
|
||||
}
|
||||
|
||||
// The same construction httpserver/auth.go performs for a cookie. SessionID
|
||||
// and ExpiresAt are deliberately left zero: there is no session row behind
|
||||
// this identity, and inventing one would make a token look like something
|
||||
// logout could end.
|
||||
return authctx.Identity{
|
||||
UserID: user.ID,
|
||||
OrgID: user.OrgID,
|
||||
Email: user.Email,
|
||||
FullName: user.FullName,
|
||||
Role: user.Role,
|
||||
AccountType: user.AccountType,
|
||||
Status: user.Status,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// hasScope reports whether a scope was granted.
|
||||
func hasScope(granted []string, want string) bool {
|
||||
for _, s := range granted {
|
||||
if s == want {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// newUUID returns a random UUID v4 string, for family ids.
|
||||
//
|
||||
// Hand-rolled rather than adding a dependency: the module is stdlib plus pgx,
|
||||
// and one 16-byte read with two bits set is not worth a third-party package.
|
||||
func newUUID() (string, error) {
|
||||
var b [16]byte
|
||||
if _, err := rand.Read(b[:]); err != nil {
|
||||
return "", fmt.Errorf("oauth: generate uuid: %w", err)
|
||||
}
|
||||
b[6] = (b[6] & 0x0f) | 0x40 // version 4
|
||||
b[8] = (b[8] & 0x3f) | 0x80 // variant 10
|
||||
h := hex.EncodeToString(b[:])
|
||||
return h[0:8] + "-" + h[8:12] + "-" + h[12:16] + "-" + h[16:20] + "-" + h[20:32], nil
|
||||
}
|
||||
857
go-api/internal/oauth/authserver.go
Normal file
857
go-api/internal/oauth/authserver.go
Normal file
@@ -0,0 +1,857 @@
|
||||
package oauth
|
||||
|
||||
import (
|
||||
"crypto/hmac"
|
||||
"crypto/rand"
|
||||
"crypto/sha256"
|
||||
"encoding/hex"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"log/slog"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"strings"
|
||||
|
||||
"github.com/krow/krow-backend/go-api/internal/authctx"
|
||||
)
|
||||
|
||||
// The authorization server's HTTP surface: register, authorize, token, revoke.
|
||||
//
|
||||
// HOW A PERSON IS AUTHENTICATED HERE
|
||||
//
|
||||
// They are not, by this package. The authorization endpoint requires a KROW
|
||||
// user to already be signed in, and it learns who that is from the SessionResolver
|
||||
// the server was built with — which the HTTP layer implements using the
|
||||
// existing cookie session. There is no second password store, no second login
|
||||
// form, and no credential of any kind in this package.
|
||||
//
|
||||
// That is also why the authorization endpoint is the only part of OAuth that
|
||||
// touches cookies: it runs in a browser, as a person, mid-redirect. Everything
|
||||
// after it — the token endpoint, the MCP endpoint — is a back-channel call from
|
||||
// the client and uses no cookie at all.
|
||||
|
||||
// SessionResolver reports who is signed in, for the authorization endpoint.
|
||||
//
|
||||
// Implemented by the HTTP layer over the existing session manager. An interface
|
||||
// rather than a direct dependency so this package does not reach into
|
||||
// httpserver, and so a test can drive the flow without a browser.
|
||||
type SessionResolver interface {
|
||||
// CurrentUser returns the signed-in identity, or false when there is none.
|
||||
CurrentUser(r *http.Request) (authctx.Identity, bool)
|
||||
}
|
||||
|
||||
// Server is the OAuth authorization server.
|
||||
type Server struct {
|
||||
cfg Config
|
||||
store *Store
|
||||
sessions SessionResolver
|
||||
log *slog.Logger
|
||||
|
||||
// loginPath is where an unauthenticated person is sent, with a return
|
||||
// target, so they can sign in and come back to the consent screen.
|
||||
loginPath string
|
||||
|
||||
// csrfKey signs consent-form tokens. Per-process and never persisted —
|
||||
// see csrfFor.
|
||||
csrfKey []byte
|
||||
}
|
||||
|
||||
// NewServer builds the authorization server.
|
||||
func NewServer(cfg Config, store *Store, sessions SessionResolver, loginPath string, log *slog.Logger) *Server {
|
||||
if log == nil {
|
||||
log = slog.Default()
|
||||
}
|
||||
if loginPath == "" {
|
||||
loginPath = "/login"
|
||||
}
|
||||
key := make([]byte, 32)
|
||||
if _, err := rand.Read(key); err != nil {
|
||||
// Unreachable short of the OS entropy source failing. Panicking is
|
||||
// correct: a server that cannot generate a CSRF key cannot render a
|
||||
// consent form safely, and starting without one would mean serving a
|
||||
// form nothing protects.
|
||||
panic("oauth: could not generate a consent CSRF key: " + err.Error())
|
||||
}
|
||||
return &Server{
|
||||
cfg: cfg.Normalise(),
|
||||
store: store,
|
||||
sessions: sessions,
|
||||
log: log,
|
||||
loginPath: loginPath,
|
||||
csrfKey: key,
|
||||
}
|
||||
}
|
||||
|
||||
/* ── Errors ─────────────────────────────────────────────────────────────── */
|
||||
|
||||
// oauthError is RFC 6749's error shape.
|
||||
type oauthError struct {
|
||||
Code string `json:"error"`
|
||||
Description string `json:"error_description,omitempty"`
|
||||
}
|
||||
|
||||
// Standard error codes. Kept to the set RFC 6749 and 7591 define, because a
|
||||
// client's error handling switches on these strings.
|
||||
const (
|
||||
errInvalidRequest = "invalid_request"
|
||||
errInvalidClient = "invalid_client"
|
||||
errInvalidGrant = "invalid_grant"
|
||||
errUnauthorizedClient = "unauthorized_client"
|
||||
errUnsupportedGrantType = "unsupported_grant_type"
|
||||
errInvalidScope = "invalid_scope"
|
||||
errInvalidRedirectURI = "invalid_redirect_uri"
|
||||
errInvalidTarget = "invalid_target" // RFC 8707, for a bad resource
|
||||
errServerError = "server_error"
|
||||
)
|
||||
|
||||
func writeOAuthError(w http.ResponseWriter, status int, code, description string) {
|
||||
w.Header().Set("Content-Type", "application/json; charset=utf-8")
|
||||
// A token or error response must never be cached: it is specific to one
|
||||
// request and may carry a credential.
|
||||
w.Header().Set("Cache-Control", "no-store")
|
||||
w.Header().Set("Pragma", "no-cache")
|
||||
writeJSONBody(w, status, oauthError{Code: code, Description: description})
|
||||
}
|
||||
|
||||
func writeJSONBody(w http.ResponseWriter, status int, payload any) {
|
||||
encoded, err := json.Marshal(payload)
|
||||
if err != nil {
|
||||
http.Error(w, "internal error", http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
w.WriteHeader(status)
|
||||
_, _ = w.Write(encoded)
|
||||
}
|
||||
|
||||
/* ── RFC 7591: Dynamic Client Registration ──────────────────────────────── */
|
||||
|
||||
type registrationRequest struct {
|
||||
ClientName string `json:"client_name"`
|
||||
RedirectURIs []string `json:"redirect_uris"`
|
||||
GrantTypes []string `json:"grant_types,omitempty"`
|
||||
ResponseTypes []string `json:"response_types,omitempty"`
|
||||
TokenEndpointAuthMethod string `json:"token_endpoint_auth_method,omitempty"`
|
||||
Scope string `json:"scope,omitempty"`
|
||||
}
|
||||
|
||||
type registrationResponse struct {
|
||||
ClientID string `json:"client_id"`
|
||||
ClientName string `json:"client_name,omitempty"`
|
||||
RedirectURIs []string `json:"redirect_uris"`
|
||||
GrantTypes []string `json:"grant_types"`
|
||||
ResponseTypes []string `json:"response_types"`
|
||||
TokenEndpointAuthMethod string `json:"token_endpoint_auth_method"`
|
||||
Scope string `json:"scope"`
|
||||
ClientIDIssuedAt int64 `json:"client_id_issued_at"`
|
||||
}
|
||||
|
||||
// maxRegistrationBytes bounds a registration body. A registration is a name and
|
||||
// a handful of URIs.
|
||||
const maxRegistrationBytes = 16 << 10
|
||||
|
||||
// RegisterHandler serves dynamic client registration.
|
||||
//
|
||||
// Open by necessity: a client that has never registered has no credential to
|
||||
// present, which is the entire point of RFC 7591 and what lets Claude connect
|
||||
// without anyone provisioning anything by hand.
|
||||
//
|
||||
// That openness is why redirect URI validation below is strict, and why
|
||||
// PHASE 5 MUST ADD RATE LIMITING HERE. This endpoint writes a row for any
|
||||
// caller that can reach it. It is structured for that — one handler, one
|
||||
// validation pass, nothing that would have to move — but today it has no limit,
|
||||
// and that is recorded as a known gap rather than quietly left unsaid.
|
||||
func (s *Server) RegisterHandler() http.Handler {
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodPost {
|
||||
w.Header().Set("Allow", http.MethodPost)
|
||||
writeOAuthError(w, http.StatusMethodNotAllowed, errInvalidRequest, "POST only")
|
||||
return
|
||||
}
|
||||
|
||||
var req registrationRequest
|
||||
if err := json.NewDecoder(http.MaxBytesReader(w, r.Body, maxRegistrationBytes)).Decode(&req); err != nil {
|
||||
writeOAuthError(w, http.StatusBadRequest, errInvalidRequest, "the request body was not valid JSON")
|
||||
return
|
||||
}
|
||||
|
||||
if len(req.RedirectURIs) == 0 {
|
||||
writeOAuthError(w, http.StatusBadRequest, errInvalidRedirectURI, "at least one redirect_uri is required")
|
||||
return
|
||||
}
|
||||
if len(req.RedirectURIs) > 10 {
|
||||
writeOAuthError(w, http.StatusBadRequest, errInvalidRedirectURI, "too many redirect_uris")
|
||||
return
|
||||
}
|
||||
for _, uri := range req.RedirectURIs {
|
||||
if err := validateRedirectURI(uri); err != nil {
|
||||
writeOAuthError(w, http.StatusBadRequest, errInvalidRedirectURI, err.Error())
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
// Only the scopes this server issues. A client asking for krow.write
|
||||
// is refused rather than quietly downgraded: silently granting less
|
||||
// than was asked for produces a client that believes it has a
|
||||
// capability and fails later, somewhere less obvious.
|
||||
scopes := []string{ScopeRead}
|
||||
if strings.TrimSpace(req.Scope) != "" {
|
||||
requested := strings.Fields(req.Scope)
|
||||
for _, sc := range requested {
|
||||
if sc != ScopeRead {
|
||||
writeOAuthError(w, http.StatusBadRequest, errInvalidScope,
|
||||
"the only scope available is "+ScopeRead)
|
||||
return
|
||||
}
|
||||
}
|
||||
scopes = requested
|
||||
}
|
||||
|
||||
clientID, err := newUUID()
|
||||
if err != nil {
|
||||
s.log.Error("oauth: client id generation failed", "error", err)
|
||||
writeOAuthError(w, http.StatusInternalServerError, errServerError, "")
|
||||
return
|
||||
}
|
||||
|
||||
name := strings.TrimSpace(req.ClientName)
|
||||
if len(name) > 200 {
|
||||
name = name[:200]
|
||||
}
|
||||
|
||||
client := Client{
|
||||
ClientID: clientID,
|
||||
ClientName: name,
|
||||
RedirectURIs: req.RedirectURIs,
|
||||
GrantTypes: []string{"authorization_code", "refresh_token"},
|
||||
Scopes: scopes,
|
||||
}
|
||||
if err := s.store.CreateClient(r.Context(), client); err != nil {
|
||||
s.log.Error("oauth: client registration failed", "error", err)
|
||||
writeOAuthError(w, http.StatusInternalServerError, errServerError, "")
|
||||
return
|
||||
}
|
||||
|
||||
s.log.Info("oauth client registered",
|
||||
"client_id", clientID, "client_name", name, "redirect_uris", len(req.RedirectURIs))
|
||||
|
||||
w.Header().Set("Content-Type", "application/json; charset=utf-8")
|
||||
w.Header().Set("Cache-Control", "no-store")
|
||||
writeJSONBody(w, http.StatusCreated, registrationResponse{
|
||||
ClientID: clientID,
|
||||
ClientName: name,
|
||||
RedirectURIs: req.RedirectURIs,
|
||||
GrantTypes: []string{"authorization_code", "refresh_token"},
|
||||
// No client_secret. A public client that was issued one would ship
|
||||
// it to every user's machine, and a secret everybody has is not a
|
||||
// secret — OAuth 2.1 handles public clients with PKCE instead.
|
||||
ResponseTypes: []string{"code"},
|
||||
TokenEndpointAuthMethod: "none",
|
||||
Scope: strings.Join(scopes, " "),
|
||||
ClientIDIssuedAt: s.store.now().Unix(),
|
||||
})
|
||||
})
|
||||
}
|
||||
|
||||
// validateRedirectURI refuses a redirect target that cannot be trusted.
|
||||
//
|
||||
// The rules, and why each one is here:
|
||||
//
|
||||
// - absolute, with a scheme and host — a relative URI has no meaning in a
|
||||
// redirect and a client sending one is confused about the flow.
|
||||
// - no fragment — RFC 6749 forbids it, and the authorization response appends
|
||||
// its own query parameters; a fragment would be silently dropped or would
|
||||
// mangle them.
|
||||
// - https, OR http on loopback only. Plain http anywhere else means the
|
||||
// authorization code travels in clear text. Loopback is the documented
|
||||
// exception for native clients (RFC 8252) and is safe because the traffic
|
||||
// never leaves the machine.
|
||||
//
|
||||
// Custom schemes (myapp://callback) are NOT accepted. They are legal per RFC
|
||||
// 8252 and are a real mechanism for native apps, but any application on the
|
||||
// machine can register the same scheme and steal the code. Claude's connectors
|
||||
// use https and loopback, so accepting custom schemes would widen the surface
|
||||
// for no caller that exists.
|
||||
func validateRedirectURI(raw string) error {
|
||||
parsed, err := url.Parse(raw)
|
||||
if err != nil {
|
||||
return errMsg("redirect_uri is not a valid URI")
|
||||
}
|
||||
if parsed.Scheme == "" || parsed.Host == "" {
|
||||
return errMsg("redirect_uri must be absolute, with a scheme and host")
|
||||
}
|
||||
if parsed.Fragment != "" || strings.Contains(raw, "#") {
|
||||
return errMsg("redirect_uri must not contain a fragment")
|
||||
}
|
||||
|
||||
switch strings.ToLower(parsed.Scheme) {
|
||||
case "https":
|
||||
return nil
|
||||
case "http":
|
||||
if isLoopbackHost(parsed.Hostname()) {
|
||||
return nil
|
||||
}
|
||||
return errMsg("http is only permitted for loopback redirect URIs")
|
||||
default:
|
||||
return errMsg("redirect_uri must use https, or http on loopback")
|
||||
}
|
||||
}
|
||||
|
||||
func isLoopbackHost(host string) bool {
|
||||
switch host {
|
||||
case "127.0.0.1", "::1", "localhost":
|
||||
return true
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
type errString string
|
||||
|
||||
func (e errString) Error() string { return string(e) }
|
||||
func errMsg(s string) error { return errString(s) }
|
||||
|
||||
/* ── Authorization endpoint ─────────────────────────────────────────────── */
|
||||
|
||||
// authorizeParams is a validated authorization request.
|
||||
type authorizeParams struct {
|
||||
ClientID string
|
||||
RedirectURI string
|
||||
ResponseType string
|
||||
Scopes []string
|
||||
State string
|
||||
CodeChallenge string
|
||||
CodeChallengeMethod string
|
||||
Resource string
|
||||
}
|
||||
|
||||
// AuthorizeHandler serves the authorization endpoint.
|
||||
//
|
||||
// THE ORDER OF VALIDATION IS A SECURITY PROPERTY, not a style choice.
|
||||
//
|
||||
// The client_id and redirect_uri are validated FIRST, against the registration,
|
||||
// before anything else is looked at. Only once the redirect target is known to
|
||||
// be one this client registered may an error be delivered BY REDIRECTING to it.
|
||||
// Getting this backwards — redirecting an error to an unvalidated URI — is an
|
||||
// open redirect, and it is the most common way this endpoint is got wrong.
|
||||
//
|
||||
// So: a bad client_id or a bad redirect_uri is answered as a direct HTTP error
|
||||
// that the browser displays. Everything after that is delivered as a redirect
|
||||
// with `error=` and the client's `state`, because by then the target is known
|
||||
// to be legitimate.
|
||||
func (s *Server) AuthorizeHandler() http.Handler {
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodGet && r.Method != http.MethodPost {
|
||||
w.Header().Set("Allow", "GET, POST")
|
||||
writeOAuthError(w, http.StatusMethodNotAllowed, errInvalidRequest, "GET or POST only")
|
||||
return
|
||||
}
|
||||
// A POST carries the decision and the flow's parameters in its body,
|
||||
// re-posted from the consent form's hidden fields. Merging them into
|
||||
// the query is what lets every validation below read from one place
|
||||
// regardless of method — and means the POST is validated exactly as
|
||||
// strictly as the GET that produced it, rather than trusting the form.
|
||||
q := r.URL.Query()
|
||||
if r.Method == http.MethodPost {
|
||||
if err := r.ParseForm(); err != nil {
|
||||
writeOAuthError(w, http.StatusBadRequest, errInvalidRequest,
|
||||
"the form could not be parsed")
|
||||
return
|
||||
}
|
||||
q = r.PostForm
|
||||
}
|
||||
|
||||
// ── Stage 1: the client and its redirect target. Errors here are
|
||||
// direct responses, never redirects.
|
||||
clientID := strings.TrimSpace(q.Get("client_id"))
|
||||
if clientID == "" {
|
||||
writeOAuthError(w, http.StatusBadRequest, errInvalidClient, "client_id is required")
|
||||
return
|
||||
}
|
||||
client, err := s.store.FindClient(r.Context(), clientID)
|
||||
if err != nil {
|
||||
s.log.Warn("oauth authorize refused", "reason", "unknown_client", "client_id", clientID)
|
||||
writeOAuthError(w, http.StatusBadRequest, errInvalidClient, "unknown client")
|
||||
return
|
||||
}
|
||||
|
||||
redirectURI := strings.TrimSpace(q.Get("redirect_uri"))
|
||||
if redirectURI == "" {
|
||||
writeOAuthError(w, http.StatusBadRequest, errInvalidRedirectURI, "redirect_uri is required")
|
||||
return
|
||||
}
|
||||
if !client.AllowsRedirect(redirectURI) {
|
||||
// Deliberately NOT redirected. This is the open-redirect guard.
|
||||
s.log.Warn("oauth authorize refused",
|
||||
"reason", "redirect_uri_mismatch", "client_id", clientID)
|
||||
writeOAuthError(w, http.StatusBadRequest, errInvalidRedirectURI,
|
||||
"redirect_uri does not match a registered URI for this client")
|
||||
return
|
||||
}
|
||||
|
||||
// ── Stage 2: everything else. The target is trusted now, so failures
|
||||
// are delivered to it.
|
||||
state := strings.TrimSpace(q.Get("state"))
|
||||
if state == "" {
|
||||
// Required, not optional. state is the client's CSRF defence for
|
||||
// the callback; a flow without one can be completed by an attacker
|
||||
// who injects their own authorization response.
|
||||
s.redirectError(w, r, redirectURI, "", errInvalidRequest, "state is required")
|
||||
return
|
||||
}
|
||||
|
||||
if rt := q.Get("response_type"); rt != "code" {
|
||||
s.redirectError(w, r, redirectURI, state, "unsupported_response_type",
|
||||
"only response_type=code is supported")
|
||||
return
|
||||
}
|
||||
|
||||
challenge := strings.TrimSpace(q.Get("code_challenge"))
|
||||
method := strings.TrimSpace(q.Get("code_challenge_method"))
|
||||
if challenge == "" {
|
||||
s.redirectError(w, r, redirectURI, state, errInvalidRequest,
|
||||
"code_challenge is required; this server requires PKCE")
|
||||
return
|
||||
}
|
||||
if method == "" {
|
||||
// RFC 7636 defaults an absent method to `plain`. This server does
|
||||
// not accept plain, so an absent method is an error rather than a
|
||||
// silent downgrade to the weaker mode.
|
||||
s.redirectError(w, r, redirectURI, state, errInvalidRequest,
|
||||
"code_challenge_method is required and must be S256")
|
||||
return
|
||||
}
|
||||
if err := ValidateChallenge(challenge, method); err != nil {
|
||||
s.redirectError(w, r, redirectURI, state, errInvalidRequest, err.Error())
|
||||
return
|
||||
}
|
||||
|
||||
// RFC 8707. The resource must be THIS server's canonical MCP URI. A
|
||||
// token is bound to it, so accepting an arbitrary value would let a
|
||||
// client mint a token aimed at something else.
|
||||
resource := strings.TrimSpace(q.Get("resource"))
|
||||
if resource == "" {
|
||||
s.redirectError(w, r, redirectURI, state, errInvalidTarget,
|
||||
"resource is required")
|
||||
return
|
||||
}
|
||||
if strings.TrimRight(resource, "/") != s.cfg.Resource {
|
||||
s.log.Warn("oauth authorize refused",
|
||||
"reason", "resource_mismatch", "client_id", clientID, "presented", resource)
|
||||
s.redirectError(w, r, redirectURI, state, errInvalidTarget,
|
||||
"resource is not a resource this server issues tokens for")
|
||||
return
|
||||
}
|
||||
|
||||
scopes := []string{ScopeRead}
|
||||
if raw := strings.TrimSpace(q.Get("scope")); raw != "" {
|
||||
scopes = strings.Fields(raw)
|
||||
for _, sc := range scopes {
|
||||
if sc != ScopeRead {
|
||||
s.redirectError(w, r, redirectURI, state, errInvalidScope,
|
||||
"the only scope available is "+ScopeRead)
|
||||
return
|
||||
}
|
||||
}
|
||||
}
|
||||
if !client.AllowsScopes(scopes) {
|
||||
s.redirectError(w, r, redirectURI, state, errInvalidScope,
|
||||
"this client is not registered for the requested scope")
|
||||
return
|
||||
}
|
||||
|
||||
params := authorizeParams{
|
||||
ClientID: clientID, RedirectURI: redirectURI, ResponseType: "code",
|
||||
Scopes: scopes, State: state, CodeChallenge: challenge,
|
||||
CodeChallengeMethod: method, Resource: resource,
|
||||
}
|
||||
|
||||
// ── Stage 3: who is this?
|
||||
identity, signedIn := s.sessions.CurrentUser(r)
|
||||
if !signedIn {
|
||||
// Not signed in. Send them to the existing login, with a return
|
||||
// target that brings them back to this exact authorization request.
|
||||
// No credential is handled here — the existing cookie login does
|
||||
// that, unchanged.
|
||||
s.redirectToLogin(w, r)
|
||||
return
|
||||
}
|
||||
|
||||
// ── Stage 4: consent.
|
||||
//
|
||||
// A GET renders the question. Only a POST carrying a session-bound
|
||||
// CSRF token answers it, so a cross-site navigation can show a person
|
||||
// the form but cannot approve on their behalf.
|
||||
csrf := s.csrfFor(identity)
|
||||
|
||||
if r.Method != http.MethodPost {
|
||||
s.renderConsent(w, r, params, identity, csrf)
|
||||
return
|
||||
}
|
||||
|
||||
if !s.csrfValid(identity, r.PostFormValue("csrf")) {
|
||||
// Not an OAuth protocol error — it is a request that did not come
|
||||
// from the form this server rendered. Answered directly rather
|
||||
// than redirected, because the client is not the party at fault
|
||||
// and telling it "access_denied" would be a lie.
|
||||
s.log.Warn("oauth consent refused", "reason", "csrf_mismatch",
|
||||
"client_id", params.ClientID, "user_id", identity.UserID)
|
||||
writeOAuthError(w, http.StatusForbidden, errInvalidRequest,
|
||||
"this consent form has expired; start the authorization again")
|
||||
return
|
||||
}
|
||||
|
||||
switch r.PostFormValue("decision") {
|
||||
case "approve":
|
||||
s.log.Info("oauth consent approved",
|
||||
"client_id", params.ClientID, "user_id", identity.UserID,
|
||||
"org_id", identity.OrgID, "scopes", params.Scopes)
|
||||
s.issueCode(w, r, params, identity)
|
||||
case "deny":
|
||||
// RFC 6749 section 4.1.2.1: a refusal is `access_denied`, returned
|
||||
// to the client at its registered redirect with the state intact.
|
||||
// NO CODE IS ISSUED — the deny path never reaches issueCode.
|
||||
s.log.Info("oauth consent denied",
|
||||
"client_id", params.ClientID, "user_id", identity.UserID)
|
||||
s.redirectError(w, r, params.RedirectURI, params.State,
|
||||
"access_denied", "the user declined this authorization")
|
||||
default:
|
||||
// A POST with neither decision. Re-render rather than guess: the
|
||||
// one thing that must not happen is inferring approval.
|
||||
s.renderConsent(w, r, params, identity, csrf)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
/* ── Consent CSRF ───────────────────────────────────────────────────────── */
|
||||
|
||||
// csrfFor derives a token binding the consent form to the signed-in user.
|
||||
//
|
||||
// An HMAC over the user id under a per-process key, rather than a random value
|
||||
// in server-side state. The property needed is only "this form was rendered by
|
||||
// this server for this user", and an HMAC gives that with nothing to store and
|
||||
// nothing to expire.
|
||||
//
|
||||
// The key is generated at startup and never leaves the process, so a token does
|
||||
// not survive a restart — which ends any consent form open at that moment. That
|
||||
// is acceptable: the window between rendering and deciding is seconds, and the
|
||||
// failure mode is a person clicking Approve and being asked to start again.
|
||||
func (s *Server) csrfFor(identity authctx.Identity) string {
|
||||
mac := hmac.New(sha256.New, s.csrfKey)
|
||||
mac.Write([]byte(identity.UserID))
|
||||
return hex.EncodeToString(mac.Sum(nil))
|
||||
}
|
||||
|
||||
// csrfValid checks a submitted token in constant time.
|
||||
func (s *Server) csrfValid(identity authctx.Identity, presented string) bool {
|
||||
if presented == "" {
|
||||
return false
|
||||
}
|
||||
return hmac.Equal([]byte(s.csrfFor(identity)), []byte(presented))
|
||||
}
|
||||
|
||||
// issueCode stores an authorization code and redirects it to the client.
|
||||
func (s *Server) issueCode(w http.ResponseWriter, r *http.Request, p authorizeParams, identity authctx.Identity) {
|
||||
code, err := s.store.CreateGrant(r.Context(), Grant{
|
||||
ClientID: p.ClientID,
|
||||
UserID: identity.UserID,
|
||||
OrgID: identity.OrgID,
|
||||
RedirectURI: p.RedirectURI,
|
||||
Scopes: p.Scopes,
|
||||
Resource: p.Resource,
|
||||
CodeChallenge: p.CodeChallenge,
|
||||
CodeChallengeMethod: p.CodeChallengeMethod,
|
||||
})
|
||||
if err != nil {
|
||||
s.log.Error("oauth: could not create grant", "error", err, "client_id", p.ClientID)
|
||||
s.redirectError(w, r, p.RedirectURI, p.State, errServerError, "")
|
||||
return
|
||||
}
|
||||
|
||||
// The code id is not logged, and neither is the code. What is logged is who
|
||||
// approved what, which is the audit question worth answering.
|
||||
s.log.Info("oauth code issued",
|
||||
"client_id", p.ClientID, "user_id", identity.UserID,
|
||||
"org_id", identity.OrgID, "scopes", p.Scopes, "resource", p.Resource)
|
||||
|
||||
target, err := url.Parse(p.RedirectURI)
|
||||
if err != nil {
|
||||
writeOAuthError(w, http.StatusBadRequest, errInvalidRedirectURI, "redirect_uri is not a valid URI")
|
||||
return
|
||||
}
|
||||
q := target.Query()
|
||||
q.Set("code", code)
|
||||
q.Set("state", p.State)
|
||||
target.RawQuery = q.Encode()
|
||||
|
||||
w.Header().Set("Cache-Control", "no-store")
|
||||
http.Redirect(w, r, target.String(), http.StatusFound)
|
||||
}
|
||||
|
||||
// redirectError delivers an error to a VALIDATED redirect target.
|
||||
//
|
||||
// Only ever called after the redirect_uri has been matched against the client's
|
||||
// registration. See the note on AuthorizeHandler.
|
||||
func (s *Server) redirectError(w http.ResponseWriter, r *http.Request, redirectURI, state, code, description string) {
|
||||
target, err := url.Parse(redirectURI)
|
||||
if err != nil {
|
||||
writeOAuthError(w, http.StatusBadRequest, errInvalidRedirectURI, "redirect_uri is not a valid URI")
|
||||
return
|
||||
}
|
||||
q := target.Query()
|
||||
q.Set("error", code)
|
||||
if description != "" {
|
||||
q.Set("error_description", description)
|
||||
}
|
||||
if state != "" {
|
||||
q.Set("state", state)
|
||||
}
|
||||
target.RawQuery = q.Encode()
|
||||
|
||||
w.Header().Set("Cache-Control", "no-store")
|
||||
http.Redirect(w, r, target.String(), http.StatusFound)
|
||||
}
|
||||
|
||||
// redirectToLogin sends an unauthenticated person to the existing login.
|
||||
//
|
||||
// The return target is this server's own path plus the original query, so the
|
||||
// authorization request survives the round trip. It is built from r.URL rather
|
||||
// than from anything the caller supplied, so it cannot be pointed elsewhere.
|
||||
func (s *Server) redirectToLogin(w http.ResponseWriter, r *http.Request) {
|
||||
returnTo := r.URL.Path
|
||||
if r.URL.RawQuery != "" {
|
||||
returnTo += "?" + r.URL.RawQuery
|
||||
}
|
||||
target := s.loginPath + "?returnTo=" + url.QueryEscape(returnTo)
|
||||
w.Header().Set("Cache-Control", "no-store")
|
||||
http.Redirect(w, r, target, http.StatusFound)
|
||||
}
|
||||
|
||||
/* ── Token endpoint ─────────────────────────────────────────────────────── */
|
||||
|
||||
type tokenResponse struct {
|
||||
AccessToken string `json:"access_token"`
|
||||
TokenType string `json:"token_type"`
|
||||
ExpiresIn int `json:"expires_in"`
|
||||
RefreshToken string `json:"refresh_token"`
|
||||
Scope string `json:"scope"`
|
||||
}
|
||||
|
||||
// TokenHandler serves the token endpoint: code exchange and refresh.
|
||||
func (s *Server) TokenHandler() http.Handler {
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodPost {
|
||||
w.Header().Set("Allow", http.MethodPost)
|
||||
writeOAuthError(w, http.StatusMethodNotAllowed, errInvalidRequest, "POST only")
|
||||
return
|
||||
}
|
||||
if err := r.ParseForm(); err != nil {
|
||||
writeOAuthError(w, http.StatusBadRequest, errInvalidRequest, "the request body could not be parsed")
|
||||
return
|
||||
}
|
||||
|
||||
switch r.PostFormValue("grant_type") {
|
||||
case "authorization_code":
|
||||
s.exchangeCode(w, r)
|
||||
case "refresh_token":
|
||||
s.refresh(w, r)
|
||||
case "":
|
||||
writeOAuthError(w, http.StatusBadRequest, errInvalidRequest, "grant_type is required")
|
||||
default:
|
||||
// password, client_credentials, implicit and anything else. Named
|
||||
// explicitly in the metadata as unsupported, and refused here.
|
||||
writeOAuthError(w, http.StatusBadRequest, errUnsupportedGrantType,
|
||||
"only authorization_code and refresh_token are supported")
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
// exchangeCode turns an authorization code into a token pair.
|
||||
//
|
||||
// Every binding recorded at authorization is re-verified. A code is not a
|
||||
// bearer credential on its own: it is a credential for one client, one redirect
|
||||
// target, one resource, and one PKCE verifier, and a mismatch on any of them
|
||||
// means the code is being spent by someone other than the client it was issued
|
||||
// to.
|
||||
func (s *Server) exchangeCode(w http.ResponseWriter, r *http.Request) {
|
||||
code := r.PostFormValue("code")
|
||||
clientID := r.PostFormValue("client_id")
|
||||
redirectURI := r.PostFormValue("redirect_uri")
|
||||
verifier := r.PostFormValue("code_verifier")
|
||||
resource := strings.TrimSpace(r.PostFormValue("resource"))
|
||||
|
||||
if code == "" || clientID == "" || redirectURI == "" {
|
||||
writeOAuthError(w, http.StatusBadRequest, errInvalidRequest,
|
||||
"code, client_id and redirect_uri are required")
|
||||
return
|
||||
}
|
||||
if verifier == "" {
|
||||
writeOAuthError(w, http.StatusBadRequest, errInvalidRequest,
|
||||
"code_verifier is required; this server requires PKCE")
|
||||
return
|
||||
}
|
||||
|
||||
// Redeeming CONSUMES the code, whatever happens next. That is deliberate:
|
||||
// if a later check fails, the code is still spent, so an attacker cannot
|
||||
// probe the remaining bindings by retrying the same code with different
|
||||
// values. One code, one attempt.
|
||||
grant, err := s.store.RedeemGrant(r.Context(), code)
|
||||
if err != nil {
|
||||
s.log.Warn("oauth token refused", "reason", "grant_unusable", "client_id", clientID)
|
||||
writeOAuthError(w, http.StatusBadRequest, errInvalidGrant,
|
||||
"the authorization code is invalid, expired or already used")
|
||||
return
|
||||
}
|
||||
|
||||
if grant.ClientID != clientID {
|
||||
s.log.Warn("oauth token refused", "reason", "client_mismatch", "client_id", clientID)
|
||||
writeOAuthError(w, http.StatusBadRequest, errInvalidGrant, "this code was not issued to this client")
|
||||
return
|
||||
}
|
||||
if grant.RedirectURI != redirectURI {
|
||||
s.log.Warn("oauth token refused", "reason", "redirect_uri_mismatch", "client_id", clientID)
|
||||
writeOAuthError(w, http.StatusBadRequest, errInvalidGrant, "redirect_uri does not match the authorization request")
|
||||
return
|
||||
}
|
||||
// The resource is optional at the token endpoint when the code already
|
||||
// carries one, but if it IS supplied it must agree.
|
||||
if resource != "" && strings.TrimRight(resource, "/") != grant.Resource {
|
||||
writeOAuthError(w, http.StatusBadRequest, errInvalidTarget, "resource does not match the authorization request")
|
||||
return
|
||||
}
|
||||
if err := VerifyChallenge(verifier, grant.CodeChallenge, grant.CodeChallengeMethod); err != nil {
|
||||
s.log.Warn("oauth token refused", "reason", "pkce_mismatch", "client_id", clientID)
|
||||
writeOAuthError(w, http.StatusBadRequest, errInvalidGrant, "code_verifier does not match")
|
||||
return
|
||||
}
|
||||
|
||||
pair, err := s.store.IssuePair(r.Context(), Token{
|
||||
ClientID: grant.ClientID,
|
||||
UserID: grant.UserID,
|
||||
OrgID: grant.OrgID,
|
||||
Scopes: grant.Scopes,
|
||||
Audience: grant.Resource,
|
||||
}, "")
|
||||
if err != nil {
|
||||
s.log.Error("oauth: could not issue tokens", "error", err)
|
||||
writeOAuthError(w, http.StatusInternalServerError, errServerError, "")
|
||||
return
|
||||
}
|
||||
|
||||
// The tokens themselves are NOT in this log line and never will be.
|
||||
s.log.Info("oauth tokens issued",
|
||||
"grant_type", "authorization_code", "client_id", grant.ClientID,
|
||||
"user_id", grant.UserID, "org_id", grant.OrgID, "family_id", pair.FamilyID)
|
||||
|
||||
writeTokenResponse(w, pair)
|
||||
}
|
||||
|
||||
// refresh rotates a refresh token.
|
||||
func (s *Server) refresh(w http.ResponseWriter, r *http.Request) {
|
||||
raw := r.PostFormValue("refresh_token")
|
||||
clientID := r.PostFormValue("client_id")
|
||||
|
||||
if raw == "" || clientID == "" {
|
||||
writeOAuthError(w, http.StatusBadRequest, errInvalidRequest,
|
||||
"refresh_token and client_id are required")
|
||||
return
|
||||
}
|
||||
|
||||
old, err := s.store.RedeemRefreshToken(r.Context(), raw)
|
||||
switch {
|
||||
case err == nil:
|
||||
// fall through
|
||||
case errors.Is(err, ErrRefreshReuse):
|
||||
// The family has already been revoked by the store. Logged at warn
|
||||
// because it is either a client bug or a stolen token, and both are
|
||||
// worth seeing. The CLIENT is told the same thing as for any other bad
|
||||
// token — distinguishing "reused" would confirm the token was once
|
||||
// real.
|
||||
s.log.Warn("oauth refresh refused", "reason", "reuse_detected", "client_id", clientID)
|
||||
writeOAuthError(w, http.StatusBadRequest, errInvalidGrant, "the refresh token is invalid")
|
||||
return
|
||||
default:
|
||||
s.log.Warn("oauth refresh refused", "reason", "token_unusable", "client_id", clientID)
|
||||
writeOAuthError(w, http.StatusBadRequest, errInvalidGrant, "the refresh token is invalid")
|
||||
return
|
||||
}
|
||||
|
||||
if old.ClientID != clientID {
|
||||
// Not this client's token. Revoke the family: a refresh token that has
|
||||
// reached the wrong client has leaked.
|
||||
_ = s.store.RevokeFamily(r.Context(), old.FamilyID, "client_mismatch_on_refresh")
|
||||
s.log.Warn("oauth refresh refused", "reason", "client_mismatch", "client_id", clientID)
|
||||
writeOAuthError(w, http.StatusBadRequest, errInvalidGrant, "the refresh token is invalid")
|
||||
return
|
||||
}
|
||||
|
||||
// Same family: the rotation continues the lineage, so reuse detection can
|
||||
// still revoke every descendant if an older token reappears.
|
||||
pair, err := s.store.IssuePair(r.Context(), Token{
|
||||
ClientID: old.ClientID,
|
||||
UserID: old.UserID,
|
||||
OrgID: old.OrgID,
|
||||
Scopes: old.Scopes,
|
||||
Audience: old.Audience,
|
||||
}, old.FamilyID)
|
||||
if err != nil {
|
||||
s.log.Error("oauth: could not rotate tokens", "error", err)
|
||||
writeOAuthError(w, http.StatusInternalServerError, errServerError, "")
|
||||
return
|
||||
}
|
||||
|
||||
s.log.Info("oauth tokens issued",
|
||||
"grant_type", "refresh_token", "client_id", old.ClientID,
|
||||
"user_id", old.UserID, "family_id", pair.FamilyID)
|
||||
|
||||
writeTokenResponse(w, pair)
|
||||
}
|
||||
|
||||
func writeTokenResponse(w http.ResponseWriter, pair TokenPair) {
|
||||
w.Header().Set("Content-Type", "application/json; charset=utf-8")
|
||||
// RFC 6749 section 5.1 requires both of these on a token response. The
|
||||
// body is a credential; nothing may cache it.
|
||||
w.Header().Set("Cache-Control", "no-store")
|
||||
w.Header().Set("Pragma", "no-cache")
|
||||
writeJSONBody(w, http.StatusOK, tokenResponse{
|
||||
AccessToken: pair.AccessToken,
|
||||
TokenType: "Bearer",
|
||||
ExpiresIn: pair.ExpiresIn,
|
||||
RefreshToken: pair.RefreshToken,
|
||||
Scope: strings.Join(pair.Scopes, " "),
|
||||
})
|
||||
}
|
||||
|
||||
/* ── Revocation (RFC 7009) ──────────────────────────────────────────────── */
|
||||
|
||||
// RevokeHandler serves token revocation.
|
||||
//
|
||||
// RFC 7009 requires 200 for an unknown token: answering 404 would turn this
|
||||
// into an oracle for whether a token exists. The store already behaves that
|
||||
// way; this handler just does not undo it.
|
||||
func (s *Server) RevokeHandler() http.Handler {
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodPost {
|
||||
w.Header().Set("Allow", http.MethodPost)
|
||||
writeOAuthError(w, http.StatusMethodNotAllowed, errInvalidRequest, "POST only")
|
||||
return
|
||||
}
|
||||
if err := r.ParseForm(); err != nil {
|
||||
writeOAuthError(w, http.StatusBadRequest, errInvalidRequest, "the request body could not be parsed")
|
||||
return
|
||||
}
|
||||
token := r.PostFormValue("token")
|
||||
if token == "" {
|
||||
writeOAuthError(w, http.StatusBadRequest, errInvalidRequest, "token is required")
|
||||
return
|
||||
}
|
||||
|
||||
if err := s.store.RevokeToken(r.Context(), token, "client_revocation"); err != nil {
|
||||
s.log.Error("oauth: revocation failed", "error", err)
|
||||
writeOAuthError(w, http.StatusInternalServerError, errServerError, "")
|
||||
return
|
||||
}
|
||||
s.log.Info("oauth token revoked", "client_id", r.PostFormValue("client_id"))
|
||||
|
||||
w.Header().Set("Cache-Control", "no-store")
|
||||
w.WriteHeader(http.StatusOK)
|
||||
})
|
||||
}
|
||||
130
go-api/internal/oauth/cleanup.go
Normal file
130
go-api/internal/oauth/cleanup.go
Normal file
@@ -0,0 +1,130 @@
|
||||
package oauth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"time"
|
||||
)
|
||||
|
||||
// Cleanup of spent and expired OAuth rows.
|
||||
//
|
||||
// WHAT IS DELETED, AND WHAT IS DELIBERATELY NOT
|
||||
//
|
||||
// Only rows that can no longer authenticate anything. Every predicate below
|
||||
// requires the row to be past its expiry — not merely consumed, not merely
|
||||
// revoked — because those two states are evidence, and evidence is worth
|
||||
// keeping until it stops being relevant.
|
||||
//
|
||||
// A consumed refresh token in particular must outlive its usefulness: it is
|
||||
// what REUSE DETECTION matches against. Delete it the moment it is spent and a
|
||||
// stolen token replayed a minute later looks like an unknown token rather than
|
||||
// a theft, and the family is never revoked. So a consumed refresh token is kept
|
||||
// until its original expiry, by which point replaying it proves nothing anyway.
|
||||
//
|
||||
// A revoked token is kept for the same reason plus one more: "this token was
|
||||
// revoked at 14:02 for refresh_token_reuse" is an answer to a question someone
|
||||
// will eventually ask.
|
||||
//
|
||||
// GRACE. Everything is deleted a grace period AFTER expiry rather than at it,
|
||||
// so a clock skewed between two instances cannot delete a row another instance
|
||||
// still considers live.
|
||||
//
|
||||
// SAFE TO RUN TWICE, AND SAFE TO RUN CONCURRENTLY. Every statement is a bounded
|
||||
// DELETE with a predicate that no longer matches once the row is gone. Two
|
||||
// workers running at once delete disjoint sets and neither errors.
|
||||
|
||||
// CleanupGrace is how long a dead row is kept past its expiry.
|
||||
//
|
||||
// An hour is far beyond any plausible clock skew between instances and short
|
||||
// enough that the tables do not accumulate. It also means a support question
|
||||
// asked within the hour can still see the row.
|
||||
const CleanupGrace = time.Hour
|
||||
|
||||
// CleanupBatch bounds one pass.
|
||||
//
|
||||
// Bounded because an unbounded DELETE holds locks for as long as it runs, and
|
||||
// on a table that every MCP request reads that is a latency spike nobody can
|
||||
// explain afterwards. Five thousand rows is milliseconds; if there is more, the
|
||||
// next pass takes it.
|
||||
const CleanupBatch = 5000
|
||||
|
||||
// CleanupResult reports what one pass removed.
|
||||
type CleanupResult struct {
|
||||
Grants int64
|
||||
AccessTokens int64
|
||||
RefreshTokens int64
|
||||
CompletedInOne bool // false when a batch filled, meaning more remains
|
||||
}
|
||||
|
||||
// Cleanup removes expired authorization codes and tokens.
|
||||
//
|
||||
// Returns counts rather than logging them, so the caller decides the level and
|
||||
// this function stays usable from a test.
|
||||
func (s *Store) Cleanup(ctx context.Context) (CleanupResult, error) {
|
||||
cutoff := s.now().Add(-CleanupGrace)
|
||||
var out CleanupResult
|
||||
|
||||
// Authorization codes. Sixty-second TTL, so almost every row here is
|
||||
// already dead; this is the highest-volume and cheapest of the three.
|
||||
//
|
||||
// ctid rather than id in the subquery because it is the physical row
|
||||
// address — the planner can go straight to it without a second index
|
||||
// lookup, which is what keeps a bounded delete genuinely cheap.
|
||||
tag, err := s.db.Exec(ctx,
|
||||
`DELETE FROM oauth_grants
|
||||
WHERE ctid IN (
|
||||
SELECT ctid FROM oauth_grants WHERE expires_at < $1 LIMIT $2
|
||||
)`, cutoff, CleanupBatch)
|
||||
if err != nil {
|
||||
return out, fmt.Errorf("oauth: cleanup grants: %w", err)
|
||||
}
|
||||
out.Grants = tag.RowsAffected()
|
||||
|
||||
// Access tokens. Fifteen-minute TTL. An expired one cannot authenticate —
|
||||
// FindAccessToken's predicate already excludes it — so deleting it removes
|
||||
// no capability.
|
||||
tag, err = s.db.Exec(ctx,
|
||||
`DELETE FROM oauth_tokens
|
||||
WHERE ctid IN (
|
||||
SELECT ctid FROM oauth_tokens
|
||||
WHERE token_type = 'access' AND expires_at < $1
|
||||
LIMIT $2
|
||||
)`, cutoff, CleanupBatch)
|
||||
if err != nil {
|
||||
return out, fmt.Errorf("oauth: cleanup access tokens: %w", err)
|
||||
}
|
||||
out.AccessTokens = tag.RowsAffected()
|
||||
|
||||
// Refresh tokens, and this is the one with a real constraint on it.
|
||||
//
|
||||
// EXPIRY ONLY — not `consumed_at IS NOT NULL`, and not `revoked_at IS NOT
|
||||
// NULL`. A consumed refresh token is what RedeemRefreshToken matches to
|
||||
// detect reuse; deleting it early turns a detectable theft into an
|
||||
// unremarkable "unknown token" and the family is never revoked. Thirty-day
|
||||
// TTL means these are the longest-lived rows in the schema, which is the
|
||||
// price of that detection and is worth paying.
|
||||
tag, err = s.db.Exec(ctx,
|
||||
`DELETE FROM oauth_tokens
|
||||
WHERE ctid IN (
|
||||
SELECT ctid FROM oauth_tokens
|
||||
WHERE token_type = 'refresh' AND expires_at < $1
|
||||
LIMIT $2
|
||||
)`, cutoff, CleanupBatch)
|
||||
if err != nil {
|
||||
return out, fmt.Errorf("oauth: cleanup refresh tokens: %w", err)
|
||||
}
|
||||
out.RefreshTokens = tag.RowsAffected()
|
||||
|
||||
out.CompletedInOne = out.Grants < CleanupBatch &&
|
||||
out.AccessTokens < CleanupBatch &&
|
||||
out.RefreshTokens < CleanupBatch
|
||||
return out, nil
|
||||
}
|
||||
|
||||
// RevokeExpiredFamilies is deliberately absent.
|
||||
//
|
||||
// It looks like it belongs here — "tidy up families whose tokens have all
|
||||
// lapsed" — and it would do nothing. Revocation is a state on a row, and a row
|
||||
// that has been deleted has no state to set. A family whose every token has
|
||||
// expired and been swept simply ceases to exist, which is the correct outcome
|
||||
// and requires no work.
|
||||
240
go-api/internal/oauth/cleanup_test.go
Normal file
240
go-api/internal/oauth/cleanup_test.go
Normal file
@@ -0,0 +1,240 @@
|
||||
package oauth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
/* ── What cleanup removes ───────────────────────────────────────────────── */
|
||||
|
||||
func TestCleanupRemovesOnlyDeadRows(t *testing.T) {
|
||||
h := newHarness(t)
|
||||
ctx := context.Background()
|
||||
clientID := h.register()
|
||||
|
||||
// A live pair, and a spent code.
|
||||
code := h.authorizeOK(clientID, verifier43)
|
||||
live := decodeTokens(t, h.exchange(clientID, code, verifier43))
|
||||
|
||||
// A second, which we let expire.
|
||||
oldCode := h.authorizeOK(clientID, verifier43)
|
||||
old := decodeTokens(t, h.exchange(clientID, oldCode, verifier43))
|
||||
|
||||
before := countRows(t, h)
|
||||
|
||||
// Past the access token TTL and the grace, but well inside the refresh
|
||||
// token's thirty days.
|
||||
h.advance(AccessTokenTTL + CleanupGrace + time.Minute)
|
||||
|
||||
result, err := h.store.Cleanup(ctx)
|
||||
if err != nil {
|
||||
t.Fatalf("Cleanup: %v", err)
|
||||
}
|
||||
|
||||
// Both access tokens and both codes are dead; both refresh tokens are not.
|
||||
if result.AccessTokens != 2 {
|
||||
t.Errorf("removed %d access tokens, want 2", result.AccessTokens)
|
||||
}
|
||||
if result.Grants != 2 {
|
||||
t.Errorf("removed %d grants, want 2", result.Grants)
|
||||
}
|
||||
if result.RefreshTokens != 0 {
|
||||
t.Errorf("removed %d refresh tokens, want 0 — they live thirty days", result.RefreshTokens)
|
||||
}
|
||||
if !result.CompletedInOne {
|
||||
t.Error("a small cleanup reported that more remained")
|
||||
}
|
||||
|
||||
after := countRows(t, h)
|
||||
if after.tokens >= before.tokens {
|
||||
t.Error("cleanup removed nothing")
|
||||
}
|
||||
|
||||
// The live refresh tokens must still work. This is the property that
|
||||
// matters: cleanup must not disconnect anybody.
|
||||
for name, token := range map[string]string{"live": live.RefreshToken, "old": old.RefreshToken} {
|
||||
if _, err := h.store.RedeemRefreshToken(ctx, token); err != nil {
|
||||
t.Errorf("the %s refresh token stopped working after cleanup: %v", name, err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// The subtle one: a CONSUMED refresh token must survive until its expiry,
|
||||
// because it is what reuse detection matches against. Delete it early and a
|
||||
// replayed stolen token looks unknown rather than stolen, and the family is
|
||||
// never revoked.
|
||||
func TestCleanupKeepsConsumedRefreshTokensForReuseDetection(t *testing.T) {
|
||||
h := newHarness(t)
|
||||
ctx := context.Background()
|
||||
clientID := h.register()
|
||||
|
||||
code := h.authorizeOK(clientID, verifier43)
|
||||
first := decodeTokens(t, h.exchange(clientID, code, verifier43))
|
||||
|
||||
// Rotate: `first` is now consumed.
|
||||
if _, err := h.store.RedeemRefreshToken(ctx, first.RefreshToken); err != nil {
|
||||
t.Fatalf("rotate: %v", err)
|
||||
}
|
||||
|
||||
// Cleanup well past the access token TTL, but inside the refresh TTL.
|
||||
h.advance(AccessTokenTTL + CleanupGrace + time.Hour)
|
||||
if _, err := h.store.Cleanup(ctx); err != nil {
|
||||
t.Fatalf("Cleanup: %v", err)
|
||||
}
|
||||
|
||||
// Replaying the consumed token must STILL be detected as reuse.
|
||||
_, err := h.store.RedeemRefreshToken(ctx, first.RefreshToken)
|
||||
if err != ErrRefreshReuse {
|
||||
t.Errorf("err = %v, want ErrRefreshReuse — cleanup destroyed the evidence "+
|
||||
"that makes theft detectable", err)
|
||||
}
|
||||
}
|
||||
|
||||
// Revoked rows are kept until expiry too: "revoked at 14:02 for
|
||||
// refresh_token_reuse" is an answer somebody will eventually need.
|
||||
func TestCleanupKeepsRevokedRowsUntilExpiry(t *testing.T) {
|
||||
h := newHarness(t)
|
||||
ctx := context.Background()
|
||||
clientID := h.register()
|
||||
|
||||
code := h.authorizeOK(clientID, verifier43)
|
||||
tokens := decodeTokens(t, h.exchange(clientID, code, verifier43))
|
||||
if err := h.store.RevokeToken(ctx, tokens.AccessToken, "test"); err != nil {
|
||||
t.Fatalf("revoke: %v", err)
|
||||
}
|
||||
|
||||
// Just past the access TTL: the access row goes, the refresh row stays.
|
||||
h.advance(AccessTokenTTL + CleanupGrace + time.Minute)
|
||||
if _, err := h.store.Cleanup(ctx); err != nil {
|
||||
t.Fatalf("Cleanup: %v", err)
|
||||
}
|
||||
|
||||
var revokedRefresh int
|
||||
if err := h.h.Pool.QueryRow(ctx,
|
||||
`SELECT count(*) FROM oauth_tokens WHERE token_type='refresh' AND revoked_at IS NOT NULL`).
|
||||
Scan(&revokedRefresh); err != nil {
|
||||
t.Fatalf("count: %v", err)
|
||||
}
|
||||
if revokedRefresh != 1 {
|
||||
t.Errorf("%d revoked refresh rows kept, want 1 — the audit trail was swept", revokedRefresh)
|
||||
}
|
||||
}
|
||||
|
||||
// Nothing is deleted before the grace period, so clock skew between instances
|
||||
// cannot destroy a row another instance still considers live.
|
||||
func TestCleanupHonoursTheGracePeriod(t *testing.T) {
|
||||
h := newHarness(t)
|
||||
ctx := context.Background()
|
||||
clientID := h.register()
|
||||
code := h.authorizeOK(clientID, verifier43)
|
||||
decodeTokens(t, h.exchange(clientID, code, verifier43))
|
||||
|
||||
// Expired, but inside the grace.
|
||||
h.advance(AccessTokenTTL + time.Minute)
|
||||
result, err := h.store.Cleanup(ctx)
|
||||
if err != nil {
|
||||
t.Fatalf("Cleanup: %v", err)
|
||||
}
|
||||
if result.AccessTokens != 0 {
|
||||
t.Errorf("removed %d access tokens inside the grace period, want 0", result.AccessTokens)
|
||||
}
|
||||
}
|
||||
|
||||
/* ── Safety ─────────────────────────────────────────────────────────────── */
|
||||
|
||||
func TestCleanupIsIdempotent(t *testing.T) {
|
||||
h := newHarness(t)
|
||||
ctx := context.Background()
|
||||
clientID := h.register()
|
||||
code := h.authorizeOK(clientID, verifier43)
|
||||
decodeTokens(t, h.exchange(clientID, code, verifier43))
|
||||
|
||||
h.advance(AccessTokenTTL + CleanupGrace + time.Minute)
|
||||
|
||||
first, err := h.store.Cleanup(ctx)
|
||||
if err != nil {
|
||||
t.Fatalf("first: %v", err)
|
||||
}
|
||||
second, err := h.store.Cleanup(ctx)
|
||||
if err != nil {
|
||||
t.Fatalf("second: %v", err)
|
||||
}
|
||||
if second.AccessTokens != 0 || second.Grants != 0 || second.RefreshTokens != 0 {
|
||||
t.Errorf("a second cleanup removed more rows: %+v (first was %+v)", second, first)
|
||||
}
|
||||
}
|
||||
|
||||
// Two workers running cleanup at once must not error and must not
|
||||
// double-count. Run with -race.
|
||||
func TestConcurrentCleanupIsSafe(t *testing.T) {
|
||||
h := newHarness(t)
|
||||
ctx := context.Background()
|
||||
clientID := h.register()
|
||||
|
||||
for i := 0; i < 6; i++ {
|
||||
code := h.authorizeOK(clientID, verifier43)
|
||||
decodeTokens(t, h.exchange(clientID, code, verifier43))
|
||||
}
|
||||
h.advance(AccessTokenTTL + CleanupGrace + time.Minute)
|
||||
|
||||
const workers = 4
|
||||
var wg sync.WaitGroup
|
||||
var mu sync.Mutex
|
||||
var total int64
|
||||
errs := make([]error, 0, workers)
|
||||
|
||||
for i := 0; i < workers; i++ {
|
||||
wg.Add(1)
|
||||
go func() {
|
||||
defer wg.Done()
|
||||
r, err := h.store.Cleanup(ctx)
|
||||
mu.Lock()
|
||||
defer mu.Unlock()
|
||||
if err != nil {
|
||||
errs = append(errs, err)
|
||||
return
|
||||
}
|
||||
total += r.AccessTokens
|
||||
}()
|
||||
}
|
||||
wg.Wait()
|
||||
|
||||
if len(errs) > 0 {
|
||||
t.Fatalf("concurrent cleanup errored: %v", errs)
|
||||
}
|
||||
// Six access tokens existed; between them the workers removed exactly six.
|
||||
// More would mean a row was counted twice.
|
||||
if total != 6 {
|
||||
t.Errorf("workers removed %d access tokens between them, want 6", total)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCleanupOnAnEmptyDatabaseIsHarmless(t *testing.T) {
|
||||
h := newHarness(t)
|
||||
result, err := h.store.Cleanup(context.Background())
|
||||
if err != nil {
|
||||
t.Fatalf("Cleanup on empty: %v", err)
|
||||
}
|
||||
if result.Grants != 0 || result.AccessTokens != 0 || result.RefreshTokens != 0 {
|
||||
t.Errorf("cleanup on an empty database removed %+v", result)
|
||||
}
|
||||
}
|
||||
|
||||
/* ── Helpers ────────────────────────────────────────────────────────────── */
|
||||
|
||||
type rowCounts struct{ grants, tokens int }
|
||||
|
||||
func countRows(t *testing.T, h *harness) rowCounts {
|
||||
t.Helper()
|
||||
var c rowCounts
|
||||
ctx := context.Background()
|
||||
if err := h.h.Pool.QueryRow(ctx, `SELECT count(*) FROM oauth_grants`).Scan(&c.grants); err != nil {
|
||||
t.Fatalf("count grants: %v", err)
|
||||
}
|
||||
if err := h.h.Pool.QueryRow(ctx, `SELECT count(*) FROM oauth_tokens`).Scan(&c.tokens); err != nil {
|
||||
t.Fatalf("count tokens: %v", err)
|
||||
}
|
||||
return c
|
||||
}
|
||||
346
go-api/internal/oauth/consent.go
Normal file
346
go-api/internal/oauth/consent.go
Normal file
@@ -0,0 +1,346 @@
|
||||
package oauth
|
||||
|
||||
import (
|
||||
"html/template"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"strings"
|
||||
|
||||
"github.com/krow/krow-backend/go-api/internal/authctx"
|
||||
)
|
||||
|
||||
// The consent step: the one place a person decides.
|
||||
//
|
||||
// Phase 3 approved a signed-in user's authorization immediately. That was
|
||||
// honest scaffolding and is not a flow anybody should ship: OAuth's entire
|
||||
// premise is that a RESOURCE OWNER grants access, and an authorization nobody
|
||||
// was asked about is a token minted on their behalf without their knowledge.
|
||||
// Any page on the internet could have linked a person to a crafted authorize
|
||||
// URL and had Claude connected to their workspace before they read anything.
|
||||
//
|
||||
// HOW THIS RESISTS THAT
|
||||
//
|
||||
// The consent form carries a CSRF token bound to the session, and approval is
|
||||
// a POST. A cross-site GET to /oauth/authorize can therefore render the form —
|
||||
// which is harmless, it is a question — but cannot answer it. Without the POST
|
||||
// and the token, an attacker who can make a browser navigate cannot make it
|
||||
// consent.
|
||||
//
|
||||
// WHAT IT SHOWS
|
||||
//
|
||||
// The client's self-declared name, the organisation being granted, the scope in
|
||||
// plain words, and the resource. The client name is UNTRUSTED — it is whatever
|
||||
// the registering client sent — so it is escaped by html/template and is never
|
||||
// the basis of a decision, only of a label. The organisation is read from the
|
||||
// signed-in identity, so a person can see which tenant they are about to hand
|
||||
// over even when they belong to more than one.
|
||||
|
||||
// consentTemplate is the approval page.
|
||||
//
|
||||
// Deliberately one self-contained page with inline styles: it renders before a
|
||||
// person is willing to trust anything, it must work with no stylesheet, no
|
||||
// script and no font available, and a consent screen that depends on assets is
|
||||
// a consent screen that can fail open into a blank page with two buttons.
|
||||
//
|
||||
// Every interpolation is escaped by html/template. The `.ClientName` in
|
||||
// particular is attacker-controlled — anyone may register a client called
|
||||
// `<script>…` — and the escaping is what makes displaying it safe.
|
||||
var consentTemplate = template.Must(template.New("consent").Parse(`<!DOCTYPE html>
|
||||
<html lang="en">
|
||||
<head>
|
||||
<meta charset="utf-8">
|
||||
<meta name="viewport" content="width=device-width, initial-scale=1">
|
||||
<title>Authorize access · Krow</title>
|
||||
<style>
|
||||
:root { color-scheme: light dark; }
|
||||
body { margin:0; min-height:100vh; display:flex; align-items:center;
|
||||
justify-content:center; background:#f4f5f7;
|
||||
font-family:-apple-system,BlinkMacSystemFont,"Segoe UI",Roboto,sans-serif;
|
||||
color:#14161a; padding:16px; box-sizing:border-box; }
|
||||
.card { background:#fff; border:1px solid #e3e5e8; border-radius:12px;
|
||||
max-width:440px; width:100%; padding:28px; box-sizing:border-box; }
|
||||
h1 { font-size:19px; margin:0 0 4px; }
|
||||
.sub { color:#5c6270; font-size:14px; margin:0 0 20px; }
|
||||
dl { margin:0 0 20px; border-top:1px solid #eceef0; }
|
||||
.row { display:flex; justify-content:space-between; gap:16px;
|
||||
padding:11px 0; border-bottom:1px solid #eceef0; font-size:14px; }
|
||||
dt { color:#5c6270; margin:0; flex:0 0 auto; }
|
||||
dd { margin:0; text-align:right; word-break:break-word; font-weight:500; }
|
||||
.grants { background:#f7f8f9; border-radius:8px; padding:14px 16px;
|
||||
font-size:14px; margin:0 0 20px; }
|
||||
.grants strong { display:block; margin-bottom:6px; font-size:13px;
|
||||
text-transform:uppercase; letter-spacing:.04em; color:#5c6270; }
|
||||
.grants ul { margin:0; padding-left:18px; }
|
||||
.grants li { margin:3px 0; }
|
||||
.actions { display:flex; gap:10px; }
|
||||
button { flex:1; padding:11px 16px; border-radius:8px; font-size:15px;
|
||||
font-weight:500; cursor:pointer; border:1px solid transparent; }
|
||||
.approve { background:#14161a; color:#fff; }
|
||||
.deny { background:#fff; color:#14161a; border-color:#d4d7dc; }
|
||||
.note { margin:16px 0 0; font-size:12.5px; color:#787e8a; line-height:1.5; }
|
||||
@media (prefers-color-scheme: dark) {
|
||||
body { background:#0e1013; color:#e9eaec; }
|
||||
.card { background:#16191d; border-color:#282c33; }
|
||||
dl,.row { border-color:#282c33; }
|
||||
.grants { background:#1c2026; }
|
||||
.approve { background:#e9eaec; color:#14161a; }
|
||||
.deny { background:#16191d; color:#e9eaec; border-color:#3a3f47; }
|
||||
dt,.sub,.note,.grants strong { color:#9aa1ad; }
|
||||
}
|
||||
</style>
|
||||
</head>
|
||||
<body>
|
||||
<main class="card">
|
||||
<h1>Authorize access to Krow</h1>
|
||||
<p class="sub"><strong>{{.ClientName}}</strong> is asking to connect to your Krow workspace.</p>
|
||||
|
||||
<dl>
|
||||
<div class="row"><dt>Application</dt><dd>{{.ClientName}}</dd></div>
|
||||
<div class="row"><dt>Signed in as</dt><dd>{{.UserEmail}}</dd></div>
|
||||
<div class="row"><dt>Organisation</dt><dd>{{.OrgName}}</dd></div>
|
||||
<div class="row"><dt>Connecting to</dt><dd>{{.Resource}}</dd></div>
|
||||
</dl>
|
||||
|
||||
<div class="grants">
|
||||
<strong>This will allow it to</strong>
|
||||
<ul>{{range .Grants}}<li>{{.}}</li>{{end}}</ul>
|
||||
</div>
|
||||
|
||||
<form method="POST" action="{{.FormAction}}">
|
||||
{{range $k, $v := .Hidden}}<input type="hidden" name="{{$k}}" value="{{$v}}">{{end}}
|
||||
<input type="hidden" name="csrf" value="{{.CSRF}}">
|
||||
<div class="actions">
|
||||
<button type="submit" name="decision" value="deny" class="deny">Deny</button>
|
||||
<button type="submit" name="decision" value="approve" class="approve">Approve</button>
|
||||
</div>
|
||||
</form>
|
||||
|
||||
<p class="note">Approving lets this application read Krow data that you can
|
||||
already see, as you, in this organisation. It cannot make changes. You can
|
||||
disconnect it at any time from your Krow settings.</p>
|
||||
</main>
|
||||
</body>
|
||||
</html>`))
|
||||
|
||||
// consentView is what the template renders.
|
||||
type consentView struct {
|
||||
ClientName string
|
||||
UserEmail string
|
||||
OrgName string
|
||||
Resource string
|
||||
Grants []string
|
||||
FormAction string
|
||||
Hidden map[string]string
|
||||
CSRF string
|
||||
}
|
||||
|
||||
// grantsFor renders scopes as sentences a person can act on.
|
||||
//
|
||||
// "krow.read" means nothing to the person being asked. A consent screen that
|
||||
// shows a scope identifier is a consent screen that has not obtained informed
|
||||
// consent — it has obtained a click.
|
||||
func grantsFor(scopes []string) []string {
|
||||
out := make([]string, 0, len(scopes))
|
||||
for _, scope := range scopes {
|
||||
switch scope {
|
||||
case ScopeRead:
|
||||
out = append(out,
|
||||
"Read workforce activity, staff, candidates and positions",
|
||||
"See only what your own Krow account can see",
|
||||
)
|
||||
case ScopeWrite:
|
||||
// Unreachable: krow.write is never issued and never registered.
|
||||
// Present so that if it ever is, it arrives with words attached
|
||||
// rather than as a bare identifier on a screen.
|
||||
out = append(out, "Make changes to your Krow data")
|
||||
default:
|
||||
out = append(out, scope)
|
||||
}
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// renderConsent shows the approval form.
|
||||
func (s *Server) renderConsent(w http.ResponseWriter, r *http.Request, p authorizeParams, identity authctx.Identity, csrf string) {
|
||||
client, err := s.store.FindClient(r.Context(), p.ClientID)
|
||||
if err != nil {
|
||||
writeOAuthError(w, http.StatusBadRequest, errInvalidClient, "unknown client")
|
||||
return
|
||||
}
|
||||
|
||||
name := strings.TrimSpace(client.ClientName)
|
||||
if name == "" {
|
||||
name = "An application"
|
||||
}
|
||||
|
||||
orgName := s.orgNameFor(r, identity.OrgID)
|
||||
|
||||
// Everything needed to complete the flow rides in hidden fields, so the
|
||||
// POST carries its own context and the server keeps no pending-request
|
||||
// state. State on the server would be state to expire and to clean up, for
|
||||
// a decision that is made in the next few seconds.
|
||||
hidden := map[string]string{
|
||||
"client_id": p.ClientID,
|
||||
"redirect_uri": p.RedirectURI,
|
||||
"response_type": p.ResponseType,
|
||||
"scope": strings.Join(p.Scopes, " "),
|
||||
"state": p.State,
|
||||
"code_challenge": p.CodeChallenge,
|
||||
"code_challenge_method": p.CodeChallengeMethod,
|
||||
"resource": p.Resource,
|
||||
}
|
||||
|
||||
w.Header().Set("Content-Type", "text/html; charset=utf-8")
|
||||
// A consent page names a client and an organisation and must never be
|
||||
// served from a cache to the next person on a shared machine.
|
||||
w.Header().Set("Cache-Control", "no-store, private")
|
||||
w.Header().Set("Pragma", "no-cache")
|
||||
// Defence in depth for a page that renders an attacker-supplied name:
|
||||
// no framing (so it cannot be clickjacked into an invisible overlay), no
|
||||
// referrer (so the query string does not leak to the client's site), and a
|
||||
// CSP that forbids script entirely — this page has none.
|
||||
w.Header().Set("X-Frame-Options", "DENY")
|
||||
w.Header().Set("Referrer-Policy", "no-referrer")
|
||||
w.Header().Set("Content-Security-Policy", consentCSP(client.RedirectURIs))
|
||||
w.Header().Set("X-Content-Type-Options", "nosniff")
|
||||
|
||||
w.WriteHeader(http.StatusOK)
|
||||
_ = consentTemplate.Execute(w, consentView{
|
||||
ClientName: name,
|
||||
UserEmail: identity.Email,
|
||||
OrgName: orgName,
|
||||
Resource: p.Resource,
|
||||
Grants: grantsFor(p.Scopes),
|
||||
FormAction: s.cfg.AuthorizePath,
|
||||
Hidden: hidden,
|
||||
CSRF: csrf,
|
||||
})
|
||||
}
|
||||
|
||||
// orgNameFor resolves an organisation's display name.
|
||||
//
|
||||
// Best effort: a missing name degrades to the id rather than failing the flow.
|
||||
// A consent screen that will not render because of a display lookup is worse
|
||||
// than one that shows a uuid.
|
||||
func (s *Server) orgNameFor(r *http.Request, orgID string) string {
|
||||
if orgID == "" {
|
||||
return "your organisation"
|
||||
}
|
||||
var name string
|
||||
if err := s.store.db.QueryRow(r.Context(),
|
||||
`SELECT name FROM organizations WHERE id = $1::uuid`, orgID).Scan(&name); err != nil {
|
||||
return orgID
|
||||
}
|
||||
if strings.TrimSpace(name) == "" {
|
||||
return orgID
|
||||
}
|
||||
return name
|
||||
}
|
||||
|
||||
/* ── The consent page's Content-Security-Policy ─────────────────────────── */
|
||||
|
||||
// consentCSP builds the policy for the consent page.
|
||||
//
|
||||
// WHY form-action CANNOT BE 'self' ALONE
|
||||
//
|
||||
// It was, and that was a real bug: the consent form is blocked in the browser
|
||||
// before it can submit. A consent form's successful submission ends, by
|
||||
// definition, at the OAuth client's registered redirect_uri — a third party's
|
||||
// callback, always cross-origin. Browsers enforce form-action across the whole
|
||||
// navigation chain including redirects (MDN carries an explicit warning that
|
||||
// this is inconsistent between engines; Chrome blocks, older Firefox did not),
|
||||
// so `form-action 'self'` makes the flow impossible to complete rather than
|
||||
// merely strict.
|
||||
//
|
||||
// The tests did not catch it because httptest executes no CSP. They asserted
|
||||
// the header's value, which was set exactly as intended; only a real browser
|
||||
// could show that what was intended was wrong.
|
||||
//
|
||||
// # WHAT IS ALLOWED INSTEAD
|
||||
//
|
||||
// 'self', plus the ORIGINS OF THIS CLIENT'S OWN REGISTERED REDIRECT URIs, and
|
||||
// nothing else. That is narrower than it may look:
|
||||
//
|
||||
// - The URIs were validated at registration — absolute, https (or http on
|
||||
// loopback), no fragment. That validation is untouched.
|
||||
// - The authorization endpoint still matches the presented redirect_uri
|
||||
// against the registration byte-for-byte. This policy does not widen what
|
||||
// a flow may redirect to; it only stops the browser blocking the redirect
|
||||
// the server was already going to permit.
|
||||
// - Each client gets its own policy, built from its own registration, so one
|
||||
// client's callback never appears in another's page.
|
||||
//
|
||||
// A URI that cannot be reduced to a safe origin is DROPPED rather than
|
||||
// broadened. The failure mode is a consent page whose form the browser blocks —
|
||||
// visible, and the safe direction — never a policy that permits more.
|
||||
func consentCSP(redirectURIs []string) string {
|
||||
directives := []string{
|
||||
"default-src 'none'",
|
||||
"style-src 'unsafe-inline'",
|
||||
"frame-ancestors 'none'",
|
||||
}
|
||||
|
||||
formAction := "form-action 'self'"
|
||||
for _, origin := range redirectOrigins(redirectURIs) {
|
||||
formAction += " " + origin
|
||||
}
|
||||
directives = append(directives, formAction)
|
||||
|
||||
return strings.Join(directives, "; ")
|
||||
}
|
||||
|
||||
// redirectOrigins reduces registered redirect URIs to CSP source expressions.
|
||||
//
|
||||
// A CSP source is an ORIGIN — scheme, host and port — never a path. Emitting
|
||||
// the full URI would be wrong twice: CSP would match it as a path prefix, and a
|
||||
// path is not what a form navigation is checked against.
|
||||
//
|
||||
// Every value is dropped unless it is unambiguously safe:
|
||||
//
|
||||
// unparseable → dropped (never widened to a bare scheme)
|
||||
// no scheme or no host → dropped
|
||||
// scheme other than
|
||||
// http/https → dropped; a custom scheme in a policy is a source
|
||||
// any app on the machine could claim
|
||||
// wildcard or separator → dropped; '*', ';' ',' or whitespace in a source
|
||||
// would either broaden the policy or split the
|
||||
// header. Registration already refuses these, so
|
||||
// this is the second lock on the same door.
|
||||
//
|
||||
// Duplicates are collapsed so two URIs on one host produce one source, and the
|
||||
// order registered is preserved so the header is stable and diffable.
|
||||
func redirectOrigins(redirectURIs []string) []string {
|
||||
seen := make(map[string]bool, len(redirectURIs))
|
||||
out := make([]string, 0, len(redirectURIs))
|
||||
|
||||
for _, raw := range redirectURIs {
|
||||
parsed, err := url.Parse(strings.TrimSpace(raw))
|
||||
if err != nil {
|
||||
continue
|
||||
}
|
||||
scheme := strings.ToLower(parsed.Scheme)
|
||||
if scheme != "http" && scheme != "https" {
|
||||
continue
|
||||
}
|
||||
// parsed.Host carries host and port together, which is exactly a CSP
|
||||
// source's host-part. Empty means the URI was relative or malformed.
|
||||
host := parsed.Host
|
||||
if host == "" {
|
||||
continue
|
||||
}
|
||||
|
||||
origin := scheme + "://" + host
|
||||
// Nothing that could broaden the policy or break the header out of its
|
||||
// directive. A registered URI cannot contain these — validateRedirectURI
|
||||
// rejects them — and this refuses to depend on that being true.
|
||||
if strings.ContainsAny(origin, "*; ,\t\r\n'\"") {
|
||||
continue
|
||||
}
|
||||
if seen[origin] {
|
||||
continue
|
||||
}
|
||||
seen[origin] = true
|
||||
out = append(out, origin)
|
||||
}
|
||||
return out
|
||||
}
|
||||
501
go-api/internal/oauth/consent_test.go
Normal file
501
go-api/internal/oauth/consent_test.go
Normal file
@@ -0,0 +1,501 @@
|
||||
package oauth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"net/url"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
/* ── Consent is required ────────────────────────────────────────────────── */
|
||||
|
||||
// A GET must ASK, not grant. This is the Phase 4 behaviour change, asserted
|
||||
// directly: before, a signed-in user's authorization was approved on sight.
|
||||
func TestAuthorizeRendersConsentRatherThanIssuingACode(t *testing.T) {
|
||||
h := newHarness(t)
|
||||
clientID := h.register()
|
||||
|
||||
rec := h.authorize(authorizeParamsFor(clientID, verifier43))
|
||||
|
||||
if rec.Code == http.StatusFound {
|
||||
t.Fatalf("a GET issued a code without asking: %s", rec.Header().Get("Location"))
|
||||
}
|
||||
if rec.Code != http.StatusOK {
|
||||
t.Fatalf("status = %d, want 200 with a consent page", rec.Code)
|
||||
}
|
||||
if ct := rec.Header().Get("Content-Type"); !strings.HasPrefix(ct, "text/html") {
|
||||
t.Errorf("Content-Type = %q, want text/html", ct)
|
||||
}
|
||||
|
||||
// No grant row may exist yet: rendering a question must not spend anything.
|
||||
var codes int
|
||||
if err := h.h.Pool.QueryRow(context.Background(),
|
||||
`SELECT count(*) FROM oauth_grants`).Scan(&codes); err != nil {
|
||||
t.Fatalf("count grants: %v", err)
|
||||
}
|
||||
if codes != 0 {
|
||||
t.Errorf("%d authorization codes exist after merely rendering consent", codes)
|
||||
}
|
||||
}
|
||||
|
||||
// The page must tell a person what they are agreeing to, in their terms.
|
||||
func TestConsentPageShowsWhatIsBeingGranted(t *testing.T) {
|
||||
h := newHarness(t)
|
||||
clientID := h.register()
|
||||
rec, _ := h.consent(authorizeParamsFor(clientID, verifier43))
|
||||
body := rec.Body.String()
|
||||
|
||||
for name, want := range map[string]string{
|
||||
"client name": "Test Client",
|
||||
"signed-in user": "oauth-user@example.test",
|
||||
"organisation": "OAuth Test",
|
||||
"resource": testResource,
|
||||
"approve control": "approve",
|
||||
"deny control": "deny",
|
||||
} {
|
||||
if !strings.Contains(body, want) {
|
||||
t.Errorf("the consent page does not show the %s (%q)", name, want)
|
||||
}
|
||||
}
|
||||
|
||||
// A person asked to approve "krow.read" has not been asked anything.
|
||||
if strings.Contains(body, ScopeRead) && !strings.Contains(body, "Read workforce activity") {
|
||||
t.Error("the page shows a raw scope identifier without explaining it")
|
||||
}
|
||||
// krow.write must never appear on a screen for a flow that cannot grant it.
|
||||
if strings.Contains(body, ScopeWrite) {
|
||||
t.Error("the consent page mentions krow.write")
|
||||
}
|
||||
}
|
||||
|
||||
// The client name is attacker-controlled: anyone may register a client called
|
||||
// <script>. It must be escaped, not rendered.
|
||||
func TestConsentPageEscapesTheClientName(t *testing.T) {
|
||||
h := newHarness(t)
|
||||
|
||||
const payload = `<script>alert('xss')</script>`
|
||||
body, _ := jsonMarshal(registrationRequest{
|
||||
ClientName: payload, RedirectURIs: []string{testRedirect},
|
||||
})
|
||||
rec := httptest.NewRecorder()
|
||||
h.server.RegisterHandler().ServeHTTP(rec,
|
||||
httptest.NewRequest(http.MethodPost, "/oauth/register", strings.NewReader(body)))
|
||||
var reg registrationResponse
|
||||
_ = jsonUnmarshal(rec.Body.Bytes(), ®)
|
||||
|
||||
page, _ := h.consent(authorizeParamsFor(reg.ClientID, verifier43))
|
||||
if strings.Contains(page.Body.String(), "<script>alert") {
|
||||
t.Fatal("a registered client name was rendered as live HTML")
|
||||
}
|
||||
if !strings.Contains(page.Body.String(), "<script>") {
|
||||
t.Error("the client name does not appear escaped; check it is shown at all")
|
||||
}
|
||||
}
|
||||
|
||||
/* ── Approve and deny ───────────────────────────────────────────────────── */
|
||||
|
||||
func TestConsentApproveIssuesACode(t *testing.T) {
|
||||
h := newHarness(t)
|
||||
clientID := h.register()
|
||||
params := authorizeParamsFor(clientID, verifier43)
|
||||
_, csrf := h.consent(params)
|
||||
|
||||
rec := h.decide(params, "approve", csrf)
|
||||
if rec.Code != http.StatusFound {
|
||||
t.Fatalf("status = %d, want 302", rec.Code)
|
||||
}
|
||||
loc, _ := url.Parse(rec.Header().Get("Location"))
|
||||
if loc.Query().Get("code") == "" {
|
||||
t.Error("approve issued no code")
|
||||
}
|
||||
if loc.Query().Get("state") != "xyz" {
|
||||
t.Errorf("state = %q, want xyz", loc.Query().Get("state"))
|
||||
}
|
||||
}
|
||||
|
||||
// Denial must reach the client as access_denied, at its registered redirect,
|
||||
// with state intact and NO code.
|
||||
func TestConsentDenyReturnsAccessDeniedAndNoCode(t *testing.T) {
|
||||
h := newHarness(t)
|
||||
clientID := h.register()
|
||||
params := authorizeParamsFor(clientID, verifier43)
|
||||
_, csrf := h.consent(params)
|
||||
|
||||
rec := h.decide(params, "deny", csrf)
|
||||
if rec.Code != http.StatusFound {
|
||||
t.Fatalf("status = %d, want 302", rec.Code)
|
||||
}
|
||||
|
||||
loc, err := url.Parse(rec.Header().Get("Location"))
|
||||
if err != nil {
|
||||
t.Fatalf("Location: %v", err)
|
||||
}
|
||||
if !strings.HasPrefix(loc.String(), testRedirect) {
|
||||
t.Fatalf("denial went to %q, not the registered redirect", loc)
|
||||
}
|
||||
if got := loc.Query().Get("error"); got != "access_denied" {
|
||||
t.Errorf("error = %q, want access_denied", got)
|
||||
}
|
||||
if got := loc.Query().Get("state"); got != "xyz" {
|
||||
t.Errorf("state = %q, want xyz — the client needs it to match the response", got)
|
||||
}
|
||||
if loc.Query().Get("code") != "" {
|
||||
t.Error("a denial returned an authorization code")
|
||||
}
|
||||
|
||||
// And nothing was written.
|
||||
var codes int
|
||||
_ = h.h.Pool.QueryRow(context.Background(), `SELECT count(*) FROM oauth_grants`).Scan(&codes)
|
||||
if codes != 0 {
|
||||
t.Errorf("%d authorization codes exist after a denial", codes)
|
||||
}
|
||||
}
|
||||
|
||||
// A POST with no decision must re-ask, never infer approval.
|
||||
func TestConsentWithNoDecisionDoesNotApprove(t *testing.T) {
|
||||
h := newHarness(t)
|
||||
clientID := h.register()
|
||||
params := authorizeParamsFor(clientID, verifier43)
|
||||
_, csrf := h.consent(params)
|
||||
|
||||
rec := h.decide(params, "", csrf)
|
||||
if rec.Code == http.StatusFound {
|
||||
t.Fatalf("a decision-less POST was treated as a decision: %s", rec.Header().Get("Location"))
|
||||
}
|
||||
var codes int
|
||||
_ = h.h.Pool.QueryRow(context.Background(), `SELECT count(*) FROM oauth_grants`).Scan(&codes)
|
||||
if codes != 0 {
|
||||
t.Error("a decision-less POST issued a code")
|
||||
}
|
||||
}
|
||||
|
||||
/* ── CSRF ───────────────────────────────────────────────────────────────── */
|
||||
|
||||
// Without the form's token, a cross-site POST must not be able to approve.
|
||||
// This is what stops a page on the internet connecting a client to somebody's
|
||||
// workspace while they are signed in.
|
||||
func TestConsentRequiresTheFormCSRFToken(t *testing.T) {
|
||||
h := newHarness(t)
|
||||
clientID := h.register()
|
||||
params := authorizeParamsFor(clientID, verifier43)
|
||||
_, valid := h.consent(params)
|
||||
|
||||
for name, token := range map[string]string{
|
||||
"absent": "",
|
||||
"garbage": "not-the-token",
|
||||
"flipped": strings.Repeat("0", len(valid)),
|
||||
} {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
rec := h.decide(params, "approve", token)
|
||||
if rec.Code == http.StatusFound {
|
||||
t.Fatalf("approval succeeded without a valid CSRF token: %s",
|
||||
rec.Header().Get("Location"))
|
||||
}
|
||||
if rec.Code != http.StatusForbidden {
|
||||
t.Errorf("status = %d, want 403", rec.Code)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
// The real token still works, or the test above would pass vacuously.
|
||||
if rec := h.decide(params, "approve", valid); rec.Code != http.StatusFound {
|
||||
t.Errorf("the valid CSRF token was rejected: %d", rec.Code)
|
||||
}
|
||||
}
|
||||
|
||||
/* ── Validation still applies on the POST ───────────────────────────────── */
|
||||
|
||||
// The POST must be validated as strictly as the GET. Trusting the form's
|
||||
// hidden fields would let a tampered POST change the redirect, the resource or
|
||||
// the PKCE challenge after the person read the page.
|
||||
func TestConsentPostRevalidatesEveryParameter(t *testing.T) {
|
||||
h := newHarness(t)
|
||||
clientID := h.register()
|
||||
params := authorizeParamsFor(clientID, verifier43)
|
||||
_, csrf := h.consent(params)
|
||||
|
||||
tamper := func(k, v string) map[string]string {
|
||||
out := map[string]string{}
|
||||
for key, val := range params {
|
||||
out[key] = val
|
||||
}
|
||||
out[k] = v
|
||||
return out
|
||||
}
|
||||
|
||||
t.Run("redirect swapped", func(t *testing.T) {
|
||||
rec := h.decide(tamper("redirect_uri", "https://attacker.example/steal"), "approve", csrf)
|
||||
if rec.Code == http.StatusFound {
|
||||
t.Fatalf("a tampered redirect_uri was honoured: %s", rec.Header().Get("Location"))
|
||||
}
|
||||
if rec.Code != http.StatusBadRequest {
|
||||
t.Errorf("status = %d, want 400", rec.Code)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("resource swapped", func(t *testing.T) {
|
||||
rec := h.decide(tamper("resource", "https://elsewhere.test/mcp"), "approve", csrf)
|
||||
loc, _ := url.Parse(rec.Header().Get("Location"))
|
||||
if loc.Query().Get("code") != "" {
|
||||
t.Error("a tampered resource still produced a code")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("pkce removed", func(t *testing.T) {
|
||||
rec := h.decide(tamper("code_challenge", ""), "approve", csrf)
|
||||
loc, _ := url.Parse(rec.Header().Get("Location"))
|
||||
if loc.Query().Get("code") != "" {
|
||||
t.Error("PKCE was dropped at the consent POST")
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("scope escalated", func(t *testing.T) {
|
||||
rec := h.decide(tamper("scope", ScopeWrite), "approve", csrf)
|
||||
loc, _ := url.Parse(rec.Header().Get("Location"))
|
||||
if loc.Query().Get("code") != "" {
|
||||
t.Error("krow.write was granted through the consent POST")
|
||||
}
|
||||
if got := loc.Query().Get("error"); got != errInvalidScope {
|
||||
t.Errorf("error = %q, want %q", got, errInvalidScope)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
// An anonymous visitor must be sent to the existing login, not shown a consent
|
||||
// screen for nobody.
|
||||
func TestConsentRequiresAuthentication(t *testing.T) {
|
||||
h := newHarness(t)
|
||||
clientID := h.register()
|
||||
h.session.signedIn = false
|
||||
|
||||
rec := h.authorize(authorizeParamsFor(clientID, verifier43))
|
||||
if rec.Code != http.StatusFound {
|
||||
t.Fatalf("status = %d, want a redirect to login", rec.Code)
|
||||
}
|
||||
if !strings.HasPrefix(rec.Header().Get("Location"), "/login?") {
|
||||
t.Errorf("Location = %q, want the existing login", rec.Header().Get("Location"))
|
||||
}
|
||||
}
|
||||
|
||||
// The consent page must never be cached: it names a client and an organisation,
|
||||
// and the next person on a shared machine must not see it.
|
||||
func TestConsentPageIsNotCacheable(t *testing.T) {
|
||||
h := newHarness(t)
|
||||
clientID := h.register()
|
||||
rec, _ := h.consent(authorizeParamsFor(clientID, verifier43))
|
||||
|
||||
if got := rec.Header().Get("Cache-Control"); !strings.Contains(got, "no-store") {
|
||||
t.Errorf("Cache-Control = %q, want no-store", got)
|
||||
}
|
||||
if got := rec.Header().Get("X-Frame-Options"); got != "DENY" {
|
||||
t.Errorf("X-Frame-Options = %q, want DENY — a consent screen must not be framed", got)
|
||||
}
|
||||
if !strings.Contains(rec.Header().Get("Content-Security-Policy"), "frame-ancestors 'none'") {
|
||||
t.Error("the CSP does not forbid framing")
|
||||
}
|
||||
}
|
||||
|
||||
// Thin wrappers so the test reads as prose rather than as error handling.
|
||||
func jsonMarshal(v any) (string, error) {
|
||||
b, err := json.Marshal(v)
|
||||
return string(b), err
|
||||
}
|
||||
|
||||
func jsonUnmarshal(b []byte, v any) error { return json.Unmarshal(b, v) }
|
||||
|
||||
/* ── The consent page's CSP ─────────────────────────────────────────────── */
|
||||
|
||||
// The regression test for the bug a real browser found and httptest could not.
|
||||
//
|
||||
// `form-action 'self'` blocked the consent form before it could submit, because
|
||||
// a consent form's successful submission ends at the client's registered
|
||||
// callback — always cross-origin. These tests assert the policy admits exactly
|
||||
// that callback and nothing else.
|
||||
|
||||
func TestConsentCSPAllowsTheRegisteredRedirectOrigin(t *testing.T) {
|
||||
h := newHarness(t)
|
||||
|
||||
body, _ := jsonMarshal(registrationRequest{
|
||||
ClientName: "Claude", RedirectURIs: []string{"https://claude.ai/api/mcp/auth_callback"},
|
||||
})
|
||||
rec := httptest.NewRecorder()
|
||||
h.server.RegisterHandler().ServeHTTP(rec,
|
||||
httptest.NewRequest(http.MethodPost, "/oauth/register", strings.NewReader(body)))
|
||||
var reg registrationResponse
|
||||
_ = jsonUnmarshal(rec.Body.Bytes(), ®)
|
||||
|
||||
page, _ := h.consent(authorizeParamsForClient(reg.ClientID, verifier43, "https://claude.ai/api/mcp/auth_callback"))
|
||||
csp := page.Header().Get("Content-Security-Policy")
|
||||
|
||||
// The registered ORIGIN — not the full URI. A CSP source is an origin; a
|
||||
// path would be matched as a prefix and is not what a navigation is
|
||||
// checked against.
|
||||
if !strings.Contains(csp, "https://claude.ai") {
|
||||
t.Errorf("CSP does not permit the registered redirect origin:\n %s", csp)
|
||||
}
|
||||
if strings.Contains(csp, "/api/mcp/auth_callback") {
|
||||
t.Errorf("CSP carries a path rather than an origin:\n %s", csp)
|
||||
}
|
||||
// Everything that must survive the change.
|
||||
for _, required := range []string{
|
||||
"default-src 'none'",
|
||||
"style-src 'unsafe-inline'",
|
||||
"frame-ancestors 'none'",
|
||||
"form-action 'self'",
|
||||
} {
|
||||
if !strings.Contains(csp, required) {
|
||||
t.Errorf("CSP lost %q:\n %s", required, csp)
|
||||
}
|
||||
}
|
||||
// And everything that must never appear.
|
||||
for _, forbidden := range []string{"form-action *", "'unsafe-eval'", "'unsafe-inline' 'unsafe", "*;", " *"} {
|
||||
if strings.Contains(csp, forbidden) {
|
||||
t.Errorf("CSP contains a broad source %q:\n %s", forbidden, csp)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// One client's callback must never appear in another client's policy.
|
||||
func TestConsentCSPDoesNotLeakBetweenClients(t *testing.T) {
|
||||
h := newHarness(t)
|
||||
|
||||
register := func(name, redirect string) string {
|
||||
t.Helper()
|
||||
body, _ := jsonMarshal(registrationRequest{ClientName: name, RedirectURIs: []string{redirect}})
|
||||
rec := httptest.NewRecorder()
|
||||
h.server.RegisterHandler().ServeHTTP(rec,
|
||||
httptest.NewRequest(http.MethodPost, "/oauth/register", strings.NewReader(body)))
|
||||
var reg registrationResponse
|
||||
_ = jsonUnmarshal(rec.Body.Bytes(), ®)
|
||||
return reg.ClientID
|
||||
}
|
||||
|
||||
first := register("First", "https://first.example.test/cb")
|
||||
second := register("Second", "https://second.example.test/cb")
|
||||
|
||||
firstPage, _ := h.consent(authorizeParamsForClient(first, verifier43, "https://first.example.test/cb"))
|
||||
firstCSP := firstPage.Header().Get("Content-Security-Policy")
|
||||
|
||||
if !strings.Contains(firstCSP, "https://first.example.test") {
|
||||
t.Errorf("the first client's own origin is missing:\n %s", firstCSP)
|
||||
}
|
||||
if strings.Contains(firstCSP, "second.example.test") {
|
||||
t.Errorf("another client's redirect origin leaked into this policy:\n %s", firstCSP)
|
||||
}
|
||||
|
||||
secondPage, _ := h.consent(authorizeParamsForClient(second, verifier43, "https://second.example.test/cb"))
|
||||
secondCSP := secondPage.Header().Get("Content-Security-Policy")
|
||||
if strings.Contains(secondCSP, "first.example.test") {
|
||||
t.Errorf("the first client's origin leaked into the second's policy:\n %s", secondCSP)
|
||||
}
|
||||
}
|
||||
|
||||
// An unrelated origin must never be permitted.
|
||||
func TestConsentCSPExcludesUnrelatedOrigins(t *testing.T) {
|
||||
h := newHarness(t)
|
||||
clientID := h.register() // registered for testRedirect only
|
||||
page, _ := h.consent(authorizeParamsFor(clientID, verifier43))
|
||||
csp := page.Header().Get("Content-Security-Policy")
|
||||
|
||||
for _, unrelated := range []string{"https://evil.test", "https://attacker.example", "https://google.com"} {
|
||||
if strings.Contains(csp, unrelated) {
|
||||
t.Errorf("CSP permits an unrelated origin %q:\n %s", unrelated, csp)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/* ── redirectOrigins, directly ──────────────────────────────────────────── */
|
||||
|
||||
// The helper carries the whole safety argument, so it is tested on its own
|
||||
// rather than only through a rendered page.
|
||||
func TestRedirectOrigins(t *testing.T) {
|
||||
for name, tc := range map[string]struct {
|
||||
in []string
|
||||
want []string
|
||||
}{
|
||||
"https with path": {
|
||||
[]string{"https://claude.ai/api/mcp/auth_callback"},
|
||||
[]string{"https://claude.ai"},
|
||||
},
|
||||
"port preserved": {
|
||||
[]string{"https://app.example.test:8443/cb"},
|
||||
[]string{"https://app.example.test:8443"},
|
||||
},
|
||||
"loopback http is kept, per registration rules": {
|
||||
[]string{"http://127.0.0.1:33418/callback"},
|
||||
[]string{"http://127.0.0.1:33418"},
|
||||
},
|
||||
"localhost loopback": {
|
||||
[]string{"http://localhost:3000/cb"},
|
||||
[]string{"http://localhost:3000"},
|
||||
},
|
||||
"multiple registered URIs": {
|
||||
[]string{"https://claude.ai/cb", "http://127.0.0.1:33418/callback"},
|
||||
[]string{"https://claude.ai", "http://127.0.0.1:33418"},
|
||||
},
|
||||
"duplicates collapse to one source": {
|
||||
[]string{"https://claude.ai/one", "https://claude.ai/two", "https://claude.ai/three"},
|
||||
[]string{"https://claude.ai"},
|
||||
},
|
||||
"order is the order registered": {
|
||||
[]string{"https://b.test/cb", "https://a.test/cb"},
|
||||
[]string{"https://b.test", "https://a.test"},
|
||||
},
|
||||
// Everything below must be DROPPED, never broadened.
|
||||
"relative uri": {[]string{"/callback"}, nil},
|
||||
"no host": {[]string{"https://"}, nil},
|
||||
"custom scheme": {[]string{"myapp://callback"}, nil},
|
||||
"javascript scheme": {[]string{"javascript:alert(1)"}, nil},
|
||||
"data scheme": {[]string{"data:text/html,x"}, nil},
|
||||
"wildcard host": {[]string{"https://*.evil.test/cb"}, nil},
|
||||
"semicolon injection": {[]string{"https://evil.test;form-action *"}, nil},
|
||||
"space injection": {[]string{"https://evil.test /cb"}, nil},
|
||||
"empty": {[]string{""}, nil},
|
||||
"whitespace only": {[]string{" "}, nil},
|
||||
} {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
got := redirectOrigins(tc.in)
|
||||
if len(got) != len(tc.want) {
|
||||
t.Fatalf("redirectOrigins(%q) = %q, want %q", tc.in, got, tc.want)
|
||||
}
|
||||
for i := range tc.want {
|
||||
if got[i] != tc.want[i] {
|
||||
t.Errorf("origin[%d] = %q, want %q", i, got[i], tc.want[i])
|
||||
}
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// A dropped URI must never widen the policy — the page still renders, and the
|
||||
// form-action list is simply shorter.
|
||||
func TestABadRedirectURICannotWidenTheCSP(t *testing.T) {
|
||||
csp := consentCSP([]string{"https://evil.test;form-action *", "myapp://cb", "https://*.evil.test"})
|
||||
|
||||
if strings.Contains(csp, "*") {
|
||||
t.Errorf("a malformed redirect URI introduced a wildcard:\n %s", csp)
|
||||
}
|
||||
if strings.Count(csp, ";") != 3 {
|
||||
t.Errorf("the header has %d separators, want 3 — a URI broke out of its directive:\n %s",
|
||||
strings.Count(csp, ";"), csp)
|
||||
}
|
||||
if !strings.Contains(csp, "form-action 'self'") {
|
||||
t.Errorf("form-action lost 'self':\n %s", csp)
|
||||
}
|
||||
// With every URI dropped, the policy is exactly the strict one — which
|
||||
// blocks the flow visibly rather than permitting more.
|
||||
if strings.Contains(csp, "evil.test") {
|
||||
t.Errorf("a dropped URI still reached the policy:\n %s", csp)
|
||||
}
|
||||
}
|
||||
|
||||
// authorizeParamsForClient is authorizeParamsFor with an explicit redirect, so
|
||||
// a test can drive a client registered for something other than testRedirect.
|
||||
func authorizeParamsForClient(clientID, verifier, redirect string) map[string]string {
|
||||
p := authorizeParamsFor(clientID, verifier)
|
||||
p["redirect_uri"] = redirect
|
||||
return p
|
||||
}
|
||||
337
go-api/internal/oauth/mcp_integration_test.go
Normal file
337
go-api/internal/oauth/mcp_integration_test.go
Normal file
@@ -0,0 +1,337 @@
|
||||
package oauth_test
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"io"
|
||||
"log/slog"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"net/url"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"github.com/krow/krow-backend/go-api/internal/auth"
|
||||
"github.com/krow/krow-backend/go-api/internal/authctx"
|
||||
"github.com/krow/krow-backend/go-api/internal/mcpserver"
|
||||
"github.com/krow/krow-backend/go-api/internal/oauth"
|
||||
"github.com/krow/krow-backend/go-api/internal/runtime"
|
||||
"github.com/krow/krow-backend/go-api/internal/testutil"
|
||||
)
|
||||
|
||||
// The seam, joined.
|
||||
//
|
||||
// This is the test that matters most in Phase 3, and it is in an EXTERNAL test
|
||||
// package (oauth_test) on purpose: it may use only the exported surface, which
|
||||
// is exactly what the HTTP layer will use when it wires these two packages
|
||||
// together in a later phase. If this compiles and passes, the wiring is a
|
||||
// constructor call and nothing else.
|
||||
//
|
||||
// What it proves end to end, with a real database and no fakes anywhere:
|
||||
//
|
||||
// OAuth authorization code flow
|
||||
// → access token
|
||||
// → mcpserver.TokenAuthenticator (the PRODUCTION implementation)
|
||||
// → authctx.Identity built from the live user row
|
||||
// → tools.Registry.Dispatch
|
||||
// → the existing policy table and org pre-filter
|
||||
// → real rows from Postgres
|
||||
|
||||
const (
|
||||
itIssuer = "https://api.example.test"
|
||||
itResource = "https://api.example.test/mcp"
|
||||
itRedirect = "https://claude.example.test/callback"
|
||||
itVerifier = "abcdefghijklmnopqrstuvwxyz0123456789ABCDEFG"
|
||||
)
|
||||
|
||||
// itSession is the SessionResolver a signed-in browser satisfies with a cookie.
|
||||
// The real implementation lives in the HTTP layer; this stands in for it so the
|
||||
// flow can be driven without a browser.
|
||||
type itSession struct{ id authctx.Identity }
|
||||
|
||||
func (s itSession) CurrentUser(*http.Request) (authctx.Identity, bool) { return s.id, true }
|
||||
|
||||
func sessionFor(userID, orgID string) itSession {
|
||||
return itSession{id: authctx.Identity{
|
||||
UserID: userID, OrgID: orgID, Role: "admin",
|
||||
Email: "a@example.test", Status: "active", AccountType: "employer",
|
||||
}}
|
||||
}
|
||||
|
||||
// itRegister performs dynamic client registration over the real handler.
|
||||
func itRegister(t *testing.T, as *oauth.Server) string {
|
||||
t.Helper()
|
||||
body := `{"client_name":"Integration Client","redirect_uris":["` + itRedirect + `"]}`
|
||||
req := httptest.NewRequest(http.MethodPost, "/oauth/register", strings.NewReader(body))
|
||||
rec := httptest.NewRecorder()
|
||||
as.RegisterHandler().ServeHTTP(rec, req)
|
||||
if rec.Code != http.StatusCreated {
|
||||
t.Fatalf("register: %d %s", rec.Code, rec.Body.String())
|
||||
}
|
||||
var out struct {
|
||||
ClientID string `json:"client_id"`
|
||||
}
|
||||
if err := json.Unmarshal(rec.Body.Bytes(), &out); err != nil {
|
||||
t.Fatalf("register response: %v", err)
|
||||
}
|
||||
return out.ClientID
|
||||
}
|
||||
|
||||
// itAuthorizeParams is a well-formed authorization request.
|
||||
func itAuthorizeParams(clientID string) url.Values {
|
||||
return url.Values{
|
||||
"client_id": {clientID}, "redirect_uri": {itRedirect}, "response_type": {"code"},
|
||||
"state": {"st8"}, "code_challenge": {oauth.ChallengeFor(itVerifier)},
|
||||
"code_challenge_method": {"S256"}, "resource": {itResource}, "scope": {oauth.ScopeRead},
|
||||
}
|
||||
}
|
||||
|
||||
// itCSRF pulls the consent form's token out of the rendered page.
|
||||
func itCSRF(t *testing.T, body string) string {
|
||||
t.Helper()
|
||||
const marker = `name="csrf" value="`
|
||||
i := strings.Index(body, marker)
|
||||
if i < 0 {
|
||||
t.Fatalf("no csrf field in the consent page")
|
||||
}
|
||||
rest := body[i+len(marker):]
|
||||
return rest[:strings.Index(rest, `"`)]
|
||||
}
|
||||
|
||||
// itDecide posts an approve/deny decision.
|
||||
func itDecide(t *testing.T, as *oauth.Server, params url.Values, decision, csrf string) *httptest.ResponseRecorder {
|
||||
t.Helper()
|
||||
form := url.Values{}
|
||||
for k, v := range params {
|
||||
form[k] = v
|
||||
}
|
||||
form.Set("decision", decision)
|
||||
form.Set("csrf", csrf)
|
||||
req := httptest.NewRequest(http.MethodPost, "/oauth/authorize", strings.NewReader(form.Encode()))
|
||||
req.Header.Set("Content-Type", "application/x-www-form-urlencoded")
|
||||
rec := httptest.NewRecorder()
|
||||
as.AuthorizeHandler().ServeHTTP(rec, req)
|
||||
return rec
|
||||
}
|
||||
|
||||
// itAuthorize drives the authorization endpoint through CONSENT and returns
|
||||
// the code.
|
||||
func itAuthorize(t *testing.T, as *oauth.Server, clientID string) string {
|
||||
t.Helper()
|
||||
q := itAuthorizeParams(clientID)
|
||||
|
||||
// The consent page first — a GET no longer issues a code.
|
||||
page := httptest.NewRecorder()
|
||||
as.AuthorizeHandler().ServeHTTP(page,
|
||||
httptest.NewRequest(http.MethodGet, "/oauth/authorize?"+q.Encode(), nil))
|
||||
if page.Code != http.StatusOK {
|
||||
t.Fatalf("consent page: %d %s", page.Code, page.Body.String())
|
||||
}
|
||||
|
||||
rec := itDecide(t, as, q, "approve", itCSRF(t, page.Body.String()))
|
||||
if rec.Code != http.StatusFound {
|
||||
t.Fatalf("authorize: %d %s", rec.Code, rec.Body.String())
|
||||
}
|
||||
loc, err := url.Parse(rec.Header().Get("Location"))
|
||||
if err != nil {
|
||||
t.Fatalf("Location: %v", err)
|
||||
}
|
||||
code := loc.Query().Get("code")
|
||||
if code == "" {
|
||||
t.Fatalf("no code: %s", loc)
|
||||
}
|
||||
return code
|
||||
}
|
||||
|
||||
// itExchange redeems the code for an access token.
|
||||
func itExchange(t *testing.T, as *oauth.Server, clientID, code string) string {
|
||||
t.Helper()
|
||||
form := url.Values{
|
||||
"grant_type": {"authorization_code"}, "code": {code}, "client_id": {clientID},
|
||||
"redirect_uri": {itRedirect}, "code_verifier": {itVerifier},
|
||||
}
|
||||
req := httptest.NewRequest(http.MethodPost, "/oauth/token", strings.NewReader(form.Encode()))
|
||||
req.Header.Set("Content-Type", "application/x-www-form-urlencoded")
|
||||
rec := httptest.NewRecorder()
|
||||
as.TokenHandler().ServeHTTP(rec, req)
|
||||
if rec.Code != http.StatusOK {
|
||||
t.Fatalf("token: %d %s", rec.Code, rec.Body.String())
|
||||
}
|
||||
var out struct {
|
||||
AccessToken string `json:"access_token"`
|
||||
}
|
||||
if err := json.Unmarshal(rec.Body.Bytes(), &out); err != nil {
|
||||
t.Fatalf("token response: %v", err)
|
||||
}
|
||||
return out.AccessToken
|
||||
}
|
||||
|
||||
func TestOAuthTokenReachesMCPToolsAndRealAuthorization(t *testing.T) {
|
||||
db := testutil.New(t)
|
||||
ctx := context.Background()
|
||||
log := slog.New(slog.NewTextHandler(io.Discard, nil))
|
||||
|
||||
// Two tenants with different volumes, so a leak is visible as a number.
|
||||
orgA := mustOrg(t, db, "it-org-a")
|
||||
orgB := mustOrg(t, db, "it-org-b")
|
||||
userA := mustUser(t, db, orgA, "a@example.test", "admin")
|
||||
seedActivity(t, db, orgA, 7, "a@example.test")
|
||||
seedActivity(t, db, orgB, 55, "b@example.test")
|
||||
|
||||
store := oauth.NewStore(db.Pool)
|
||||
|
||||
// ── Register, authorize, exchange: the real flow, over the real handlers.
|
||||
as := oauth.NewServer(
|
||||
oauth.Config{Issuer: itIssuer, Resource: itResource},
|
||||
store,
|
||||
sessionFor(userA, orgA),
|
||||
"/login", log,
|
||||
)
|
||||
|
||||
clientID := itRegister(t, as)
|
||||
code := itAuthorize(t, as, clientID)
|
||||
accessToken := itExchange(t, as, clientID, code)
|
||||
|
||||
// ── The production authenticator, plugged into the Phase 2 seam.
|
||||
authenticator := oauth.NewAuthenticator(store, auth.NewPGUserStore(db.Pool), itResource, log)
|
||||
mcp := mcpserver.New(runtime.DefaultTools(db.Pool, nil), authenticator, log)
|
||||
|
||||
// ── A real MCP tool call, carrying a real OAuth token.
|
||||
rec := itCall(t, mcp, "Bearer "+accessToken,
|
||||
`{"jsonrpc":"2.0","id":1,"method":"tools/call","params":
|
||||
{"name":"activity_breakdown","arguments":{}}}`)
|
||||
|
||||
if rec.Code != http.StatusOK {
|
||||
t.Fatalf("MCP call with an OAuth token: %d %s", rec.Code, rec.Body.String())
|
||||
}
|
||||
total := itTotalEvents(t, rec)
|
||||
if total != 7 {
|
||||
t.Errorf("totalEvents = %d, want 7 (org A only). Org B has 55; a wrong "+
|
||||
"number here means the OAuth identity did not scope the query", total)
|
||||
}
|
||||
|
||||
// ── Without the token, the same call must be refused.
|
||||
if rec := itCall(t, mcp, "",
|
||||
`{"jsonrpc":"2.0","id":1,"method":"tools/list"}`); rec.Code != http.StatusUnauthorized {
|
||||
t.Errorf("unauthenticated MCP call = %d, want 401", rec.Code)
|
||||
}
|
||||
|
||||
// ── Revoking disconnects: the same token must stop working immediately,
|
||||
// not at expiry.
|
||||
if err := store.RevokeToken(ctx, accessToken, "test_disconnect"); err != nil {
|
||||
t.Fatalf("revoke: %v", err)
|
||||
}
|
||||
if rec := itCall(t, mcp, "Bearer "+accessToken,
|
||||
`{"jsonrpc":"2.0","id":1,"method":"tools/list"}`); rec.Code != http.StatusUnauthorized {
|
||||
t.Errorf("a revoked token still reached MCP: %d", rec.Code)
|
||||
}
|
||||
}
|
||||
|
||||
// A token for another resource must not open the MCP surface, even though this
|
||||
// same server minted it. The confused-deputy case, end to end.
|
||||
func TestTokenForAnotherResourceCannotReachMCP(t *testing.T) {
|
||||
db := testutil.New(t)
|
||||
log := slog.New(slog.NewTextHandler(io.Discard, nil))
|
||||
|
||||
org := mustOrg(t, db, "it-aud-org")
|
||||
user := mustUser(t, db, org, "aud@example.test", "admin")
|
||||
store := oauth.NewStore(db.Pool)
|
||||
|
||||
as := oauth.NewServer(
|
||||
oauth.Config{Issuer: itIssuer, Resource: itResource},
|
||||
store, sessionFor(user, org), "/login", log)
|
||||
clientID := itRegister(t, as)
|
||||
|
||||
pair, err := store.IssuePair(context.Background(), oauth.Token{
|
||||
ClientID: clientID, UserID: user, OrgID: org,
|
||||
Scopes: []string{oauth.ScopeRead}, Audience: "https://a-different-service.test/mcp",
|
||||
}, "")
|
||||
if err != nil {
|
||||
t.Fatalf("issue: %v", err)
|
||||
}
|
||||
|
||||
mcp := mcpserver.New(
|
||||
runtime.DefaultTools(db.Pool, nil),
|
||||
oauth.NewAuthenticator(store, auth.NewPGUserStore(db.Pool), itResource, log),
|
||||
log)
|
||||
|
||||
if rec := itCall(t, mcp, "Bearer "+pair.AccessToken,
|
||||
`{"jsonrpc":"2.0","id":1,"method":"tools/list"}`); rec.Code != http.StatusUnauthorized {
|
||||
t.Errorf("a token for another resource reached MCP: %d", rec.Code)
|
||||
}
|
||||
}
|
||||
|
||||
/* ── Helpers ────────────────────────────────────────────────────────────── */
|
||||
|
||||
func itCall(t *testing.T, s *mcpserver.Server, authHeader, body string) *httptest.ResponseRecorder {
|
||||
t.Helper()
|
||||
req := httptest.NewRequest(http.MethodPost, "/mcp", strings.NewReader(body))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
if authHeader != "" {
|
||||
req.Header.Set("Authorization", authHeader)
|
||||
}
|
||||
rec := httptest.NewRecorder()
|
||||
s.Handler().ServeHTTP(rec, req)
|
||||
return rec
|
||||
}
|
||||
|
||||
func itTotalEvents(t *testing.T, rec *httptest.ResponseRecorder) int {
|
||||
t.Helper()
|
||||
var envelope struct {
|
||||
Result struct {
|
||||
Content []struct {
|
||||
Text string `json:"text"`
|
||||
} `json:"content"`
|
||||
IsError bool `json:"isError"`
|
||||
} `json:"result"`
|
||||
}
|
||||
if err := json.Unmarshal(rec.Body.Bytes(), &envelope); err != nil {
|
||||
t.Fatalf("response: %v", err)
|
||||
}
|
||||
if envelope.Result.IsError || len(envelope.Result.Content) == 0 {
|
||||
t.Fatalf("tool call failed: %s", rec.Body.String())
|
||||
}
|
||||
var payload struct {
|
||||
Data struct {
|
||||
TotalEvents int `json:"totalEvents"`
|
||||
} `json:"data"`
|
||||
}
|
||||
if err := json.Unmarshal([]byte(envelope.Result.Content[0].Text), &payload); err != nil {
|
||||
t.Fatalf("tool payload: %v", err)
|
||||
}
|
||||
return payload.Data.TotalEvents
|
||||
}
|
||||
|
||||
func mustOrg(t *testing.T, db *testutil.Harness, slug string) string {
|
||||
t.Helper()
|
||||
var id string
|
||||
if err := db.Pool.QueryRow(context.Background(),
|
||||
`INSERT INTO organizations (name, slug) VALUES ($1, $2) RETURNING id::text`,
|
||||
slug, slug).Scan(&id); err != nil {
|
||||
t.Fatalf("org %s: %v", slug, err)
|
||||
}
|
||||
return id
|
||||
}
|
||||
|
||||
func mustUser(t *testing.T, db *testutil.Harness, orgID, email, role string) string {
|
||||
t.Helper()
|
||||
var id string
|
||||
if err := db.Pool.QueryRow(context.Background(),
|
||||
`INSERT INTO users (org_id, email, full_name, role, account_type, status)
|
||||
VALUES ($1::uuid, $2, 'Test', $3, 'employer', 'active') RETURNING id::text`,
|
||||
orgID, email, role).Scan(&id); err != nil {
|
||||
t.Fatalf("user %s: %v", email, err)
|
||||
}
|
||||
return id
|
||||
}
|
||||
|
||||
func seedActivity(t *testing.T, db *testutil.Harness, orgID string, n int, email string) {
|
||||
t.Helper()
|
||||
for i := 0; i < n; i++ {
|
||||
if _, err := db.Pool.Exec(context.Background(),
|
||||
`INSERT INTO user_activity (org_id, event_type, user_email, user_name)
|
||||
VALUES ($1::uuid, 'login', $2, 'Someone')`, orgID, email); err != nil {
|
||||
t.Fatalf("seed: %v", err)
|
||||
}
|
||||
}
|
||||
}
|
||||
186
go-api/internal/oauth/metadata.go
Normal file
186
go-api/internal/oauth/metadata.go
Normal file
@@ -0,0 +1,186 @@
|
||||
package oauth
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// Discovery: the two documents an MCP client reads before it can authenticate.
|
||||
//
|
||||
// The MCP authorization flow starts with the client calling the MCP endpoint
|
||||
// with no token, getting a 401, and following its way to an authorization
|
||||
// server. Two RFCs define the path:
|
||||
//
|
||||
// RFC 9728 Protected Resource Metadata — served BY THE RESOURCE (the MCP
|
||||
// server). Answers "which authorization server issues tokens for
|
||||
// you". The 401's WWW-Authenticate header points here.
|
||||
// RFC 8414 Authorization Server Metadata — served by the AS. Answers "where
|
||||
// are your authorize, token and registration endpoints, and what do
|
||||
// you support".
|
||||
//
|
||||
// Both are unauthenticated by necessity: a client that cannot authenticate yet
|
||||
// has to be able to read them. Neither contains a secret — they are a map of
|
||||
// public endpoints, which is exactly what discovery means.
|
||||
//
|
||||
// NO URL IS GUESSED OR HARDCODED. Every value comes from configuration, so a
|
||||
// deployment on a different host is a config change and not a code change, and
|
||||
// so this file contains no production domain.
|
||||
|
||||
// Scopes this server issues.
|
||||
//
|
||||
// ScopeWrite is DECLARED and never granted. Naming it here means the constant
|
||||
// exists for a future phase to use deliberately, rather than being invented at
|
||||
// the point somebody is trying to make a write work. It appears in no
|
||||
// scopes_supported list and no issued token.
|
||||
const (
|
||||
ScopeRead = "krow.read"
|
||||
ScopeWrite = "krow.write" // reserved; not issued, not advertised
|
||||
)
|
||||
|
||||
// Config is the deployment's OAuth identity.
|
||||
//
|
||||
// Issuer and Resource are separate values that will often look similar, and
|
||||
// conflating them is a real mistake: the ISSUER identifies the authorization
|
||||
// server, the RESOURCE identifies the thing a token is good for. A token's
|
||||
// audience is checked against Resource, and its origin against Issuer.
|
||||
type Config struct {
|
||||
// Issuer is the authorization server's identity, e.g.
|
||||
// https://api.example.com. No trailing slash.
|
||||
Issuer string
|
||||
|
||||
// Resource is the canonical MCP endpoint URI, e.g.
|
||||
// https://api.example.com/mcp. This is what a client puts in its
|
||||
// `resource` parameter and what an issued token's audience is set to.
|
||||
Resource string
|
||||
|
||||
// The paths, relative to Issuer. Defaults are applied by Normalise.
|
||||
AuthorizePath string
|
||||
TokenPath string
|
||||
RegistrationPath string
|
||||
RevocationPath string
|
||||
}
|
||||
|
||||
// Normalise fills defaults and trims trailing slashes.
|
||||
//
|
||||
// The canonical form of a resource URI has no trailing slash — RFC 8707 says
|
||||
// implementations SHOULD use that form — and a mismatch here is a token that
|
||||
// validates everywhere except the one place it was minted for.
|
||||
func (c Config) Normalise() Config {
|
||||
c.Issuer = strings.TrimRight(strings.TrimSpace(c.Issuer), "/")
|
||||
c.Resource = strings.TrimRight(strings.TrimSpace(c.Resource), "/")
|
||||
if c.AuthorizePath == "" {
|
||||
c.AuthorizePath = "/oauth/authorize"
|
||||
}
|
||||
if c.TokenPath == "" {
|
||||
c.TokenPath = "/oauth/token"
|
||||
}
|
||||
if c.RegistrationPath == "" {
|
||||
c.RegistrationPath = "/oauth/register"
|
||||
}
|
||||
if c.RevocationPath == "" {
|
||||
c.RevocationPath = "/oauth/revoke"
|
||||
}
|
||||
return c
|
||||
}
|
||||
|
||||
// Valid reports whether this configuration can serve discovery at all.
|
||||
func (c Config) Valid() bool {
|
||||
return c.Issuer != "" && c.Resource != ""
|
||||
}
|
||||
|
||||
func (c Config) authorizeURL() string { return c.Issuer + c.AuthorizePath }
|
||||
func (c Config) tokenURL() string { return c.Issuer + c.TokenPath }
|
||||
func (c Config) registrationURL() string { return c.Issuer + c.RegistrationPath }
|
||||
func (c Config) revocationURL() string { return c.Issuer + c.RevocationPath }
|
||||
|
||||
/* ── RFC 9728: Protected Resource Metadata ──────────────────────────────── */
|
||||
|
||||
type protectedResourceMetadata struct {
|
||||
Resource string `json:"resource"`
|
||||
AuthorizationServers []string `json:"authorization_servers"`
|
||||
ScopesSupported []string `json:"scopes_supported"`
|
||||
BearerMethodsSupported []string `json:"bearer_methods_supported"`
|
||||
}
|
||||
|
||||
// ProtectedResourceHandler serves /.well-known/oauth-protected-resource.
|
||||
func (c Config) ProtectedResourceHandler() http.Handler {
|
||||
cfg := c.Normalise()
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodGet {
|
||||
w.Header().Set("Allow", http.MethodGet)
|
||||
http.Error(w, "method not allowed", http.StatusMethodNotAllowed)
|
||||
return
|
||||
}
|
||||
writeMetadata(w, protectedResourceMetadata{
|
||||
Resource: cfg.Resource,
|
||||
AuthorizationServers: []string{cfg.Issuer},
|
||||
ScopesSupported: []string{ScopeRead},
|
||||
// header only. RFC 6750 also defines a form-encoded body parameter
|
||||
// and a query parameter; the MCP spec forbids the query form and
|
||||
// this server accepts neither.
|
||||
BearerMethodsSupported: []string{"header"},
|
||||
})
|
||||
})
|
||||
}
|
||||
|
||||
/* ── RFC 8414: Authorization Server Metadata ────────────────────────────── */
|
||||
|
||||
type authorizationServerMetadata struct {
|
||||
Issuer string `json:"issuer"`
|
||||
AuthorizationEndpoint string `json:"authorization_endpoint"`
|
||||
TokenEndpoint string `json:"token_endpoint"`
|
||||
RegistrationEndpoint string `json:"registration_endpoint"`
|
||||
RevocationEndpoint string `json:"revocation_endpoint"`
|
||||
ScopesSupported []string `json:"scopes_supported"`
|
||||
ResponseTypesSupported []string `json:"response_types_supported"`
|
||||
GrantTypesSupported []string `json:"grant_types_supported"`
|
||||
CodeChallengeMethodsSupported []string `json:"code_challenge_methods_supported"`
|
||||
TokenEndpointAuthMethodsSupported []string `json:"token_endpoint_auth_methods_supported"`
|
||||
ResourceIndicatorsSupported bool `json:"resource_indicators_supported"`
|
||||
}
|
||||
|
||||
// AuthorizationServerHandler serves /.well-known/oauth-authorization-server.
|
||||
//
|
||||
// Every list below is a promise, so each one names only what is implemented:
|
||||
//
|
||||
// - response_types: `code`. No `token`, because implicit is gone from OAuth
|
||||
// 2.1 and advertising it would invite a flow this server refuses.
|
||||
// - grant_types: authorization_code and refresh_token. No password, no
|
||||
// client_credentials — neither has a caller here, and both would be a way
|
||||
// to get a token without a person approving anything.
|
||||
// - code_challenge_methods: S256 only. Listing `plain` would tell a client it
|
||||
// may use the method this server rejects.
|
||||
// - token_endpoint_auth_methods: `none`, which is the correct declaration
|
||||
// for public clients. They authenticate with PKCE, not a secret.
|
||||
func (c Config) AuthorizationServerHandler() http.Handler {
|
||||
cfg := c.Normalise()
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
if r.Method != http.MethodGet {
|
||||
w.Header().Set("Allow", http.MethodGet)
|
||||
http.Error(w, "method not allowed", http.StatusMethodNotAllowed)
|
||||
return
|
||||
}
|
||||
writeMetadata(w, authorizationServerMetadata{
|
||||
Issuer: cfg.Issuer,
|
||||
AuthorizationEndpoint: cfg.authorizeURL(),
|
||||
TokenEndpoint: cfg.tokenURL(),
|
||||
RegistrationEndpoint: cfg.registrationURL(),
|
||||
RevocationEndpoint: cfg.revocationURL(),
|
||||
ScopesSupported: []string{ScopeRead},
|
||||
ResponseTypesSupported: []string{"code"},
|
||||
GrantTypesSupported: []string{"authorization_code", "refresh_token"},
|
||||
CodeChallengeMethodsSupported: []string{MethodS256},
|
||||
TokenEndpointAuthMethodsSupported: []string{"none"},
|
||||
ResourceIndicatorsSupported: true,
|
||||
})
|
||||
})
|
||||
}
|
||||
|
||||
func writeMetadata(w http.ResponseWriter, payload any) {
|
||||
w.Header().Set("Content-Type", "application/json; charset=utf-8")
|
||||
// Discovery documents change only with a deployment, and a client that
|
||||
// re-reads them on every connection costs nothing to serve. Five minutes
|
||||
// keeps a stale document from outliving a config change by long.
|
||||
w.Header().Set("Cache-Control", "public, max-age=300")
|
||||
writeJSONBody(w, http.StatusOK, payload)
|
||||
}
|
||||
952
go-api/internal/oauth/oauth_test.go
Normal file
952
go-api/internal/oauth/oauth_test.go
Normal file
@@ -0,0 +1,952 @@
|
||||
package oauth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"io"
|
||||
"log/slog"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"net/url"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/krow/krow-backend/go-api/internal/auth"
|
||||
"github.com/krow/krow-backend/go-api/internal/authctx"
|
||||
"github.com/krow/krow-backend/go-api/internal/testutil"
|
||||
)
|
||||
|
||||
/* ── Fixtures ───────────────────────────────────────────────────────────── */
|
||||
|
||||
const (
|
||||
testIssuer = "https://api.example.test"
|
||||
testResource = "https://api.example.test/mcp"
|
||||
testRedirect = "https://claude.example.test/callback"
|
||||
)
|
||||
|
||||
func testConfig() Config {
|
||||
return Config{Issuer: testIssuer, Resource: testResource}.Normalise()
|
||||
}
|
||||
|
||||
func discard() *slog.Logger { return slog.New(slog.NewTextHandler(io.Discard, nil)) }
|
||||
|
||||
// fakeSession is the SessionResolver a browser would satisfy with a cookie.
|
||||
type fakeSession struct {
|
||||
identity authctx.Identity
|
||||
signedIn bool
|
||||
}
|
||||
|
||||
func (f *fakeSession) CurrentUser(*http.Request) (authctx.Identity, bool) {
|
||||
return f.identity, f.signedIn
|
||||
}
|
||||
|
||||
// harness wires a real database to a real authorization server.
|
||||
type harness struct {
|
||||
t *testing.T
|
||||
h *testutil.Harness
|
||||
store *Store
|
||||
server *Server
|
||||
session *fakeSession
|
||||
userID string
|
||||
orgID string
|
||||
clock time.Time
|
||||
}
|
||||
|
||||
func newHarness(t *testing.T) *harness {
|
||||
t.Helper()
|
||||
db := testutil.New(t)
|
||||
ctx := context.Background()
|
||||
|
||||
var orgID string
|
||||
if err := db.Pool.QueryRow(ctx,
|
||||
`INSERT INTO organizations (name, slug) VALUES ('OAuth Test', 'oauth-test')
|
||||
RETURNING id::text`).Scan(&orgID); err != nil {
|
||||
t.Fatalf("create org: %v", err)
|
||||
}
|
||||
var userID string
|
||||
if err := db.Pool.QueryRow(ctx,
|
||||
`INSERT INTO users (org_id, email, full_name, role, account_type, status)
|
||||
VALUES ($1::uuid, 'oauth-user@example.test', 'OAuth User', 'admin', 'employer', 'active')
|
||||
RETURNING id::text`, orgID).Scan(&userID); err != nil {
|
||||
t.Fatalf("create user: %v", err)
|
||||
}
|
||||
|
||||
clock := time.Now()
|
||||
store := NewStore(db.Pool).WithClock(func() time.Time { return clock })
|
||||
session := &fakeSession{
|
||||
identity: authctx.Identity{UserID: userID, OrgID: orgID, Role: "admin",
|
||||
Email: "oauth-user@example.test", Status: "active"},
|
||||
signedIn: true,
|
||||
}
|
||||
|
||||
hs := &harness{
|
||||
t: t, h: db, store: store, session: session,
|
||||
userID: userID, orgID: orgID, clock: clock,
|
||||
}
|
||||
hs.server = NewServer(testConfig(), store, session, "/login", discard())
|
||||
return hs
|
||||
}
|
||||
|
||||
// advance moves the store's clock, so expiry is tested without sleeping.
|
||||
func (h *harness) advance(d time.Duration) {
|
||||
h.clock = h.clock.Add(d)
|
||||
h.store.WithClock(func() time.Time { return h.clock })
|
||||
}
|
||||
|
||||
// register performs dynamic client registration and returns the client_id.
|
||||
func (h *harness) register(redirectURIs ...string) string {
|
||||
h.t.Helper()
|
||||
if len(redirectURIs) == 0 {
|
||||
redirectURIs = []string{testRedirect}
|
||||
}
|
||||
body, _ := json.Marshal(registrationRequest{
|
||||
ClientName: "Test Client", RedirectURIs: redirectURIs,
|
||||
})
|
||||
rec := httptest.NewRecorder()
|
||||
h.server.RegisterHandler().ServeHTTP(rec,
|
||||
httptest.NewRequest(http.MethodPost, "/oauth/register", strings.NewReader(string(body))))
|
||||
if rec.Code != http.StatusCreated {
|
||||
h.t.Fatalf("registration failed: %d %s", rec.Code, rec.Body.String())
|
||||
}
|
||||
var out registrationResponse
|
||||
if err := json.Unmarshal(rec.Body.Bytes(), &out); err != nil {
|
||||
h.t.Fatalf("registration response: %v", err)
|
||||
}
|
||||
return out.ClientID
|
||||
}
|
||||
|
||||
// authorize drives the authorization endpoint and returns the response recorder.
|
||||
func (h *harness) authorize(params map[string]string) *httptest.ResponseRecorder {
|
||||
h.t.Helper()
|
||||
q := url.Values{}
|
||||
for k, v := range params {
|
||||
if v != "" {
|
||||
q.Set(k, v)
|
||||
}
|
||||
}
|
||||
rec := httptest.NewRecorder()
|
||||
h.server.AuthorizeHandler().ServeHTTP(rec,
|
||||
httptest.NewRequest(http.MethodGet, "/oauth/authorize?"+q.Encode(), nil))
|
||||
return rec
|
||||
}
|
||||
|
||||
// authorizeParamsFor is a well-formed authorization request.
|
||||
func authorizeParamsFor(clientID, verifier string) map[string]string {
|
||||
return map[string]string{
|
||||
"client_id": clientID, "redirect_uri": testRedirect, "response_type": "code",
|
||||
"state": "xyz", "code_challenge": ChallengeFor(verifier),
|
||||
"code_challenge_method": "S256", "resource": testResource, "scope": ScopeRead,
|
||||
}
|
||||
}
|
||||
|
||||
// csrfFromConsentPage pulls the token out of the rendered form.
|
||||
//
|
||||
// PHASE 4: a GET now RENDERS a consent page rather than issuing a code. This
|
||||
// helper and decide() below are how a test plays the part of the person
|
||||
// clicking a button. Not one assertion in this file changed — the flow gained
|
||||
// a step, and the tests walk through it.
|
||||
func csrfFromConsentPage(t *testing.T, body string) string {
|
||||
t.Helper()
|
||||
const marker = `name="csrf" value="`
|
||||
i := strings.Index(body, marker)
|
||||
if i < 0 {
|
||||
t.Fatalf("no csrf field in the consent page:\n%s", body)
|
||||
}
|
||||
rest := body[i+len(marker):]
|
||||
j := strings.Index(rest, `"`)
|
||||
if j < 0 {
|
||||
t.Fatal("malformed csrf field")
|
||||
}
|
||||
return rest[:j]
|
||||
}
|
||||
|
||||
// decide posts an approve/deny decision to the authorization endpoint.
|
||||
func (h *harness) decide(params map[string]string, decision, csrf string) *httptest.ResponseRecorder {
|
||||
h.t.Helper()
|
||||
form := url.Values{}
|
||||
for k, v := range params {
|
||||
if v != "" {
|
||||
form.Set(k, v)
|
||||
}
|
||||
}
|
||||
form.Set("decision", decision)
|
||||
form.Set("csrf", csrf)
|
||||
|
||||
req := httptest.NewRequest(http.MethodPost, "/oauth/authorize", strings.NewReader(form.Encode()))
|
||||
req.Header.Set("Content-Type", "application/x-www-form-urlencoded")
|
||||
rec := httptest.NewRecorder()
|
||||
h.server.AuthorizeHandler().ServeHTTP(rec, req)
|
||||
return rec
|
||||
}
|
||||
|
||||
// consent renders the form and returns the page plus its CSRF token.
|
||||
func (h *harness) consent(params map[string]string) (*httptest.ResponseRecorder, string) {
|
||||
h.t.Helper()
|
||||
rec := h.authorize(params)
|
||||
if rec.Code != http.StatusOK {
|
||||
h.t.Fatalf("consent page: %d %s", rec.Code, rec.Body.String())
|
||||
}
|
||||
return rec, csrfFromConsentPage(h.t, rec.Body.String())
|
||||
}
|
||||
|
||||
// authorizeOK runs a well-formed authorization, APPROVES it, and returns the
|
||||
// code.
|
||||
func (h *harness) authorizeOK(clientID, verifier string) string {
|
||||
h.t.Helper()
|
||||
params := authorizeParamsFor(clientID, verifier)
|
||||
_, csrf := h.consent(params)
|
||||
rec := h.decide(params, "approve", csrf)
|
||||
if rec.Code != http.StatusFound {
|
||||
h.t.Fatalf("authorize: %d %s", rec.Code, rec.Body.String())
|
||||
}
|
||||
loc, err := url.Parse(rec.Header().Get("Location"))
|
||||
if err != nil {
|
||||
h.t.Fatalf("bad Location: %v", err)
|
||||
}
|
||||
if e := loc.Query().Get("error"); e != "" {
|
||||
h.t.Fatalf("authorize returned error=%s (%s)", e, loc.Query().Get("error_description"))
|
||||
}
|
||||
code := loc.Query().Get("code")
|
||||
if code == "" {
|
||||
h.t.Fatalf("no code in %s", loc)
|
||||
}
|
||||
if got := loc.Query().Get("state"); got != "xyz" {
|
||||
h.t.Errorf("state = %q, want xyz — the client's CSRF defence must be echoed", got)
|
||||
}
|
||||
return code
|
||||
}
|
||||
|
||||
// token posts to the token endpoint.
|
||||
func (h *harness) token(form url.Values) *httptest.ResponseRecorder {
|
||||
h.t.Helper()
|
||||
req := httptest.NewRequest(http.MethodPost, "/oauth/token", strings.NewReader(form.Encode()))
|
||||
req.Header.Set("Content-Type", "application/x-www-form-urlencoded")
|
||||
rec := httptest.NewRecorder()
|
||||
h.server.TokenHandler().ServeHTTP(rec, req)
|
||||
return rec
|
||||
}
|
||||
|
||||
func (h *harness) exchange(clientID, code, verifier string) *httptest.ResponseRecorder {
|
||||
return h.token(url.Values{
|
||||
"grant_type": {"authorization_code"}, "code": {code}, "client_id": {clientID},
|
||||
"redirect_uri": {testRedirect}, "code_verifier": {verifier},
|
||||
})
|
||||
}
|
||||
|
||||
func decodeTokens(t *testing.T, rec *httptest.ResponseRecorder) tokenResponse {
|
||||
t.Helper()
|
||||
if rec.Code != http.StatusOK {
|
||||
t.Fatalf("token endpoint: %d %s", rec.Code, rec.Body.String())
|
||||
}
|
||||
var out tokenResponse
|
||||
if err := json.Unmarshal(rec.Body.Bytes(), &out); err != nil {
|
||||
t.Fatalf("token response: %v", err)
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func oauthErrorCode(t *testing.T, rec *httptest.ResponseRecorder) string {
|
||||
t.Helper()
|
||||
var out oauthError
|
||||
_ = json.Unmarshal(rec.Body.Bytes(), &out)
|
||||
return out.Code
|
||||
}
|
||||
|
||||
// verifier43 is a legal code_verifier.
|
||||
const verifier43 = "abcdefghijklmnopqrstuvwxyz0123456789ABCDEFG"
|
||||
|
||||
/* ── The happy path ─────────────────────────────────────────────────────── */
|
||||
|
||||
func TestFullAuthorizationCodeFlow(t *testing.T) {
|
||||
h := newHarness(t)
|
||||
clientID := h.register()
|
||||
code := h.authorizeOK(clientID, verifier43)
|
||||
tokens := decodeTokens(t, h.exchange(clientID, code, verifier43))
|
||||
|
||||
if tokens.TokenType != "Bearer" {
|
||||
t.Errorf("token_type = %q, want Bearer", tokens.TokenType)
|
||||
}
|
||||
if tokens.AccessToken == "" || tokens.RefreshToken == "" {
|
||||
t.Fatal("a token response must carry both tokens")
|
||||
}
|
||||
if tokens.AccessToken == tokens.RefreshToken {
|
||||
t.Error("access and refresh tokens are identical")
|
||||
}
|
||||
// 15 minutes, as committed in the plan.
|
||||
if tokens.ExpiresIn != int(AccessTokenTTL.Seconds()) {
|
||||
t.Errorf("expires_in = %d, want %d", tokens.ExpiresIn, int(AccessTokenTTL.Seconds()))
|
||||
}
|
||||
if tokens.Scope != ScopeRead {
|
||||
t.Errorf("scope = %q, want %q", tokens.Scope, ScopeRead)
|
||||
}
|
||||
}
|
||||
|
||||
// The single most important storage property: a dump of these tables must not
|
||||
// be replayable. Asserted by searching every text column for the raw values.
|
||||
func TestPlaintextTokensAreNeverStored(t *testing.T) {
|
||||
h := newHarness(t)
|
||||
clientID := h.register()
|
||||
code := h.authorizeOK(clientID, verifier43)
|
||||
tokens := decodeTokens(t, h.exchange(clientID, code, verifier43))
|
||||
|
||||
for name, secret := range map[string]string{
|
||||
"authorization code": code,
|
||||
"access token": tokens.AccessToken,
|
||||
"refresh token": tokens.RefreshToken,
|
||||
} {
|
||||
for _, table := range []string{"oauth_grants", "oauth_tokens"} {
|
||||
var found int
|
||||
// Cast the whole row to text and search it. Cruder than naming
|
||||
// columns and much harder to fool: a future column that stored a
|
||||
// raw token would be caught without anyone updating this test.
|
||||
if err := h.h.Pool.QueryRow(context.Background(),
|
||||
`SELECT count(*) FROM `+table+` t WHERE t::text LIKE '%' || $1 || '%'`,
|
||||
secret).Scan(&found); err != nil {
|
||||
t.Fatalf("scan %s: %v", table, err)
|
||||
}
|
||||
if found != 0 {
|
||||
t.Errorf("the %s appears in PLAINTEXT in %s (%d rows)", name, table, found)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/* ── Authorization endpoint rejection ───────────────────────────────────── */
|
||||
|
||||
func TestAuthorizeRejections(t *testing.T) {
|
||||
h := newHarness(t)
|
||||
clientID := h.register()
|
||||
base := map[string]string{
|
||||
"client_id": clientID, "redirect_uri": testRedirect, "response_type": "code",
|
||||
"state": "xyz", "code_challenge": ChallengeFor(verifier43),
|
||||
"code_challenge_method": "S256", "resource": testResource,
|
||||
}
|
||||
with := func(changes map[string]string) map[string]string {
|
||||
out := map[string]string{}
|
||||
for k, v := range base {
|
||||
out[k] = v
|
||||
}
|
||||
for k, v := range changes {
|
||||
out[k] = v
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// These are answered DIRECTLY, never by redirecting — redirecting an error
|
||||
// to an unvalidated URI is an open redirect.
|
||||
t.Run("direct errors", func(t *testing.T) {
|
||||
for name, changes := range map[string]map[string]string{
|
||||
"unknown client": {"client_id": "00000000-0000-4000-8000-000000000000"},
|
||||
"missing client": {"client_id": ""},
|
||||
"missing redirect": {"redirect_uri": ""},
|
||||
"unregistered redirect": {"redirect_uri": "https://attacker.example/steal"},
|
||||
"redirect near-miss": {"redirect_uri": testRedirect + "/../evil"},
|
||||
} {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
rec := h.authorize(with(changes))
|
||||
if rec.Code == http.StatusFound {
|
||||
t.Fatalf("answered with a REDIRECT to %q — this must be a direct error",
|
||||
rec.Header().Get("Location"))
|
||||
}
|
||||
if rec.Code != http.StatusBadRequest {
|
||||
t.Errorf("status = %d, want 400", rec.Code)
|
||||
}
|
||||
})
|
||||
}
|
||||
})
|
||||
|
||||
// The redirect target is validated by now, so errors go to it.
|
||||
t.Run("redirected errors", func(t *testing.T) {
|
||||
for name, tc := range map[string]struct {
|
||||
changes map[string]string
|
||||
want string
|
||||
}{
|
||||
"missing state": {map[string]string{"state": ""}, errInvalidRequest},
|
||||
"missing pkce": {map[string]string{"code_challenge": ""}, errInvalidRequest},
|
||||
"missing method": {map[string]string{"code_challenge_method": ""}, errInvalidRequest},
|
||||
"plain pkce": {map[string]string{"code_challenge_method": "plain"}, errInvalidRequest},
|
||||
"bad challenge": {map[string]string{"code_challenge": "too-short"}, errInvalidRequest},
|
||||
"implicit flow": {map[string]string{"response_type": "token"}, "unsupported_response_type"},
|
||||
"missing resource": {map[string]string{"resource": ""}, errInvalidTarget},
|
||||
"wrong resource": {map[string]string{"resource": "https://elsewhere.test/mcp"}, errInvalidTarget},
|
||||
"write scope denied": {map[string]string{"scope": ScopeWrite}, errInvalidScope},
|
||||
} {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
rec := h.authorize(with(tc.changes))
|
||||
if rec.Code != http.StatusFound {
|
||||
t.Fatalf("status = %d, want a 302 carrying the error", rec.Code)
|
||||
}
|
||||
loc, _ := url.Parse(rec.Header().Get("Location"))
|
||||
if !strings.HasPrefix(loc.String(), testRedirect) {
|
||||
t.Fatalf("error went to %q, not the registered redirect", loc)
|
||||
}
|
||||
if got := loc.Query().Get("error"); got != tc.want {
|
||||
t.Errorf("error = %q, want %q", got, tc.want)
|
||||
}
|
||||
if loc.Query().Get("code") != "" {
|
||||
t.Error("a failed authorization returned a code")
|
||||
}
|
||||
})
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
// An unauthenticated person is sent to the existing login, not refused and not
|
||||
// asked for a password by this package.
|
||||
func TestAuthorizeRedirectsAnonymousToLogin(t *testing.T) {
|
||||
h := newHarness(t)
|
||||
clientID := h.register()
|
||||
h.session.signedIn = false
|
||||
|
||||
rec := h.authorize(map[string]string{
|
||||
"client_id": clientID, "redirect_uri": testRedirect, "response_type": "code",
|
||||
"state": "xyz", "code_challenge": ChallengeFor(verifier43),
|
||||
"code_challenge_method": "S256", "resource": testResource,
|
||||
})
|
||||
if rec.Code != http.StatusFound {
|
||||
t.Fatalf("status = %d, want 302 to login", rec.Code)
|
||||
}
|
||||
loc := rec.Header().Get("Location")
|
||||
if !strings.HasPrefix(loc, "/login?returnTo=") {
|
||||
t.Fatalf("Location = %q, want a redirect to /login carrying returnTo", loc)
|
||||
}
|
||||
// The authorization request must survive the round trip, or the person
|
||||
// signs in and lands nowhere.
|
||||
if !strings.Contains(loc, url.QueryEscape("client_id="+clientID)) {
|
||||
t.Error("returnTo does not preserve the authorization request")
|
||||
}
|
||||
}
|
||||
|
||||
/* ── Registration ───────────────────────────────────────────────────────── */
|
||||
|
||||
func TestRegistrationRejectsUnsafeRedirectURIs(t *testing.T) {
|
||||
h := newHarness(t)
|
||||
for name, uri := range map[string]string{
|
||||
"plain http": "http://attacker.example/cb",
|
||||
"relative": "/callback",
|
||||
"no host": "https://",
|
||||
"with fragment": "https://ok.example/cb#frag",
|
||||
"custom scheme": "myapp://callback",
|
||||
"javascript": "javascript:alert(1)",
|
||||
"data uri": "data:text/html,hi",
|
||||
"missing scheme": "ok.example/cb",
|
||||
} {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
body, _ := json.Marshal(registrationRequest{
|
||||
ClientName: "x", RedirectURIs: []string{uri},
|
||||
})
|
||||
rec := httptest.NewRecorder()
|
||||
h.server.RegisterHandler().ServeHTTP(rec,
|
||||
httptest.NewRequest(http.MethodPost, "/oauth/register", strings.NewReader(string(body))))
|
||||
if rec.Code != http.StatusBadRequest {
|
||||
t.Errorf("status = %d, want 400 for redirect_uri %q", rec.Code, uri)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// http on loopback is the documented exception for native clients (RFC 8252):
|
||||
// the traffic never leaves the machine.
|
||||
func TestRegistrationAllowsLoopbackHTTP(t *testing.T) {
|
||||
for _, uri := range []string{
|
||||
"http://127.0.0.1:8765/callback",
|
||||
"http://localhost:3000/cb",
|
||||
"https://claude.example.test/cb",
|
||||
} {
|
||||
if err := validateRedirectURI(uri); err != nil {
|
||||
t.Errorf("%q was rejected: %v", uri, err)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// A public client must not be issued a secret.
|
||||
func TestRegistrationIssuesNoClientSecret(t *testing.T) {
|
||||
h := newHarness(t)
|
||||
body, _ := json.Marshal(registrationRequest{ClientName: "x", RedirectURIs: []string{testRedirect}})
|
||||
rec := httptest.NewRecorder()
|
||||
h.server.RegisterHandler().ServeHTTP(rec,
|
||||
httptest.NewRequest(http.MethodPost, "/oauth/register", strings.NewReader(string(body))))
|
||||
|
||||
if strings.Contains(strings.ToLower(rec.Body.String()), "client_secret") {
|
||||
t.Errorf("a public client was issued a secret: %s", rec.Body.String())
|
||||
}
|
||||
var out registrationResponse
|
||||
_ = json.Unmarshal(rec.Body.Bytes(), &out)
|
||||
if out.TokenEndpointAuthMethod != "none" {
|
||||
t.Errorf("token_endpoint_auth_method = %q, want none", out.TokenEndpointAuthMethod)
|
||||
}
|
||||
}
|
||||
|
||||
/* ── Token endpoint ─────────────────────────────────────────────────────── */
|
||||
|
||||
func TestTokenEndpointRejections(t *testing.T) {
|
||||
h := newHarness(t)
|
||||
clientID := h.register()
|
||||
|
||||
t.Run("wrong verifier", func(t *testing.T) {
|
||||
code := h.authorizeOK(clientID, verifier43)
|
||||
rec := h.exchange(clientID, code, strings.Repeat("z", 43))
|
||||
if rec.Code != http.StatusBadRequest || oauthErrorCode(t, rec) != errInvalidGrant {
|
||||
t.Errorf("status=%d error=%q, want 400 invalid_grant", rec.Code, oauthErrorCode(t, rec))
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("missing verifier", func(t *testing.T) {
|
||||
code := h.authorizeOK(clientID, verifier43)
|
||||
rec := h.token(url.Values{
|
||||
"grant_type": {"authorization_code"}, "code": {code},
|
||||
"client_id": {clientID}, "redirect_uri": {testRedirect},
|
||||
})
|
||||
if rec.Code != http.StatusBadRequest {
|
||||
t.Errorf("status = %d, want 400 — PKCE is mandatory", rec.Code)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("wrong client", func(t *testing.T) {
|
||||
other := h.register("https://other.example.test/cb")
|
||||
code := h.authorizeOK(clientID, verifier43)
|
||||
rec := h.exchange(other, code, verifier43)
|
||||
if rec.Code != http.StatusBadRequest {
|
||||
t.Errorf("status = %d, want 400 — a code is bound to its client", rec.Code)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("wrong redirect_uri", func(t *testing.T) {
|
||||
code := h.authorizeOK(clientID, verifier43)
|
||||
rec := h.token(url.Values{
|
||||
"grant_type": {"authorization_code"}, "code": {code}, "client_id": {clientID},
|
||||
"redirect_uri": {"https://attacker.example/steal"}, "code_verifier": {verifier43},
|
||||
})
|
||||
if rec.Code != http.StatusBadRequest {
|
||||
t.Errorf("status = %d, want 400", rec.Code)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("unknown code", func(t *testing.T) {
|
||||
rec := h.exchange(clientID, "not-a-real-code", verifier43)
|
||||
if rec.Code != http.StatusBadRequest {
|
||||
t.Errorf("status = %d, want 400", rec.Code)
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
// A code is single-use. The second attempt must fail even with everything else
|
||||
// correct — this is replay protection.
|
||||
func TestAuthorizationCodeIsSingleUse(t *testing.T) {
|
||||
h := newHarness(t)
|
||||
clientID := h.register()
|
||||
code := h.authorizeOK(clientID, verifier43)
|
||||
|
||||
if rec := h.exchange(clientID, code, verifier43); rec.Code != http.StatusOK {
|
||||
t.Fatalf("first exchange failed: %d %s", rec.Code, rec.Body.String())
|
||||
}
|
||||
rec := h.exchange(clientID, code, verifier43)
|
||||
if rec.Code != http.StatusBadRequest {
|
||||
t.Errorf("status = %d, want 400 — a code must not be redeemable twice", rec.Code)
|
||||
}
|
||||
}
|
||||
|
||||
// A failed exchange still spends the code, so an attacker cannot probe the
|
||||
// remaining bindings by retrying with different values.
|
||||
func TestAFailedExchangeStillConsumesTheCode(t *testing.T) {
|
||||
h := newHarness(t)
|
||||
clientID := h.register()
|
||||
code := h.authorizeOK(clientID, verifier43)
|
||||
|
||||
if rec := h.exchange(clientID, code, strings.Repeat("z", 43)); rec.Code != http.StatusBadRequest {
|
||||
t.Fatalf("expected the wrong verifier to fail, got %d", rec.Code)
|
||||
}
|
||||
if rec := h.exchange(clientID, code, verifier43); rec.Code != http.StatusBadRequest {
|
||||
t.Error("the code was still usable after a failed exchange")
|
||||
}
|
||||
}
|
||||
|
||||
func TestExpiredAuthorizationCodeIsRejected(t *testing.T) {
|
||||
h := newHarness(t)
|
||||
clientID := h.register()
|
||||
code := h.authorizeOK(clientID, verifier43)
|
||||
|
||||
h.advance(GrantTTL + time.Second)
|
||||
|
||||
if rec := h.exchange(clientID, code, verifier43); rec.Code != http.StatusBadRequest {
|
||||
t.Errorf("status = %d, want 400 for an expired code", rec.Code)
|
||||
}
|
||||
}
|
||||
|
||||
func TestUnsupportedGrantTypesAreRejected(t *testing.T) {
|
||||
h := newHarness(t)
|
||||
for _, grant := range []string{"password", "client_credentials", "implicit", "device_code", "nonsense"} {
|
||||
t.Run(grant, func(t *testing.T) {
|
||||
rec := h.token(url.Values{
|
||||
"grant_type": {grant}, "username": {"a"}, "password": {"b"},
|
||||
})
|
||||
if rec.Code != http.StatusBadRequest {
|
||||
t.Fatalf("status = %d, want 400", rec.Code)
|
||||
}
|
||||
if got := oauthErrorCode(t, rec); got != errUnsupportedGrantType {
|
||||
t.Errorf("error = %q, want %q", got, errUnsupportedGrantType)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
/* ── Refresh rotation and reuse detection ───────────────────────────────── */
|
||||
|
||||
func TestRefreshRotatesAndInvalidatesTheOldToken(t *testing.T) {
|
||||
h := newHarness(t)
|
||||
clientID := h.register()
|
||||
code := h.authorizeOK(clientID, verifier43)
|
||||
first := decodeTokens(t, h.exchange(clientID, code, verifier43))
|
||||
|
||||
second := decodeTokens(t, h.token(url.Values{
|
||||
"grant_type": {"refresh_token"}, "refresh_token": {first.RefreshToken}, "client_id": {clientID},
|
||||
}))
|
||||
|
||||
if second.RefreshToken == first.RefreshToken {
|
||||
t.Error("the refresh token was not rotated")
|
||||
}
|
||||
if second.AccessToken == first.AccessToken {
|
||||
t.Error("refresh returned the same access token")
|
||||
}
|
||||
|
||||
// The rotated-away token must be dead. Presenting it again is also the
|
||||
// reuse signal — see the next test.
|
||||
rec := h.token(url.Values{
|
||||
"grant_type": {"refresh_token"}, "refresh_token": {first.RefreshToken}, "client_id": {clientID},
|
||||
})
|
||||
if rec.Code != http.StatusBadRequest {
|
||||
t.Errorf("the old refresh token still worked: %d", rec.Code)
|
||||
}
|
||||
}
|
||||
|
||||
// Replaying a consumed refresh token means either a client bug or a stolen
|
||||
// token, and there is no way to tell. OAuth 2.1's answer is to assume theft and
|
||||
// revoke the whole family — so the attacker AND the legitimate holder both lose
|
||||
// access, and the legitimate one reauthorizes.
|
||||
func TestRefreshReuseRevokesTheWholeFamily(t *testing.T) {
|
||||
h := newHarness(t)
|
||||
clientID := h.register()
|
||||
code := h.authorizeOK(clientID, verifier43)
|
||||
first := decodeTokens(t, h.exchange(clientID, code, verifier43))
|
||||
|
||||
second := decodeTokens(t, h.token(url.Values{
|
||||
"grant_type": {"refresh_token"}, "refresh_token": {first.RefreshToken}, "client_id": {clientID},
|
||||
}))
|
||||
|
||||
// The attacker replays the stolen (already rotated) token.
|
||||
if rec := h.token(url.Values{
|
||||
"grant_type": {"refresh_token"}, "refresh_token": {first.RefreshToken}, "client_id": {clientID},
|
||||
}); rec.Code != http.StatusBadRequest {
|
||||
t.Fatalf("reuse was accepted: %d", rec.Code)
|
||||
}
|
||||
|
||||
// Now the LEGITIMATE current token must also be dead.
|
||||
if rec := h.token(url.Values{
|
||||
"grant_type": {"refresh_token"}, "refresh_token": {second.RefreshToken}, "client_id": {clientID},
|
||||
}); rec.Code != http.StatusBadRequest {
|
||||
t.Error("the family was not revoked after reuse — the thief keeps access")
|
||||
}
|
||||
|
||||
// And so must the access token it minted.
|
||||
if _, err := h.store.FindAccessToken(context.Background(), second.AccessToken); !errors.Is(err, ErrTokenUnusable) {
|
||||
t.Error("an access token in the revoked family still validates")
|
||||
}
|
||||
}
|
||||
|
||||
func TestRefreshWithTheWrongClientIsRejected(t *testing.T) {
|
||||
h := newHarness(t)
|
||||
clientID := h.register()
|
||||
other := h.register("https://other.example.test/cb")
|
||||
code := h.authorizeOK(clientID, verifier43)
|
||||
tokens := decodeTokens(t, h.exchange(clientID, code, verifier43))
|
||||
|
||||
if rec := h.token(url.Values{
|
||||
"grant_type": {"refresh_token"}, "refresh_token": {tokens.RefreshToken}, "client_id": {other},
|
||||
}); rec.Code != http.StatusBadRequest {
|
||||
t.Errorf("another client refreshed this token: %d", rec.Code)
|
||||
}
|
||||
}
|
||||
|
||||
func TestExpiredRefreshTokenIsRejected(t *testing.T) {
|
||||
h := newHarness(t)
|
||||
clientID := h.register()
|
||||
code := h.authorizeOK(clientID, verifier43)
|
||||
tokens := decodeTokens(t, h.exchange(clientID, code, verifier43))
|
||||
|
||||
h.advance(RefreshTokenTTL + time.Hour)
|
||||
|
||||
if rec := h.token(url.Values{
|
||||
"grant_type": {"refresh_token"}, "refresh_token": {tokens.RefreshToken}, "client_id": {clientID},
|
||||
}); rec.Code != http.StatusBadRequest {
|
||||
t.Errorf("an expired refresh token was accepted: %d", rec.Code)
|
||||
}
|
||||
}
|
||||
|
||||
/* ── Revocation ─────────────────────────────────────────────────────────── */
|
||||
|
||||
// Revoking must disconnect, which means killing the refresh token too.
|
||||
// Revoking only the access token would leave the client able to mint another
|
||||
// within seconds — so the button marked "disconnect" would not disconnect.
|
||||
func TestRevocationKillsTheWholeFamily(t *testing.T) {
|
||||
h := newHarness(t)
|
||||
clientID := h.register()
|
||||
code := h.authorizeOK(clientID, verifier43)
|
||||
tokens := decodeTokens(t, h.exchange(clientID, code, verifier43))
|
||||
|
||||
form := url.Values{"token": {tokens.AccessToken}, "client_id": {clientID}}
|
||||
req := httptest.NewRequest(http.MethodPost, "/oauth/revoke", strings.NewReader(form.Encode()))
|
||||
req.Header.Set("Content-Type", "application/x-www-form-urlencoded")
|
||||
rec := httptest.NewRecorder()
|
||||
h.server.RevokeHandler().ServeHTTP(rec, req)
|
||||
if rec.Code != http.StatusOK {
|
||||
t.Fatalf("revocation: %d %s", rec.Code, rec.Body.String())
|
||||
}
|
||||
|
||||
if _, err := h.store.FindAccessToken(context.Background(), tokens.AccessToken); !errors.Is(err, ErrTokenUnusable) {
|
||||
t.Error("the access token still validates after revocation")
|
||||
}
|
||||
if r := h.token(url.Values{
|
||||
"grant_type": {"refresh_token"}, "refresh_token": {tokens.RefreshToken}, "client_id": {clientID},
|
||||
}); r.Code != http.StatusBadRequest {
|
||||
t.Error("the refresh token survived revocation — this is not a disconnect")
|
||||
}
|
||||
}
|
||||
|
||||
// RFC 7009: revoking an unknown token is a success, or the endpoint becomes a
|
||||
// way to test whether a token exists.
|
||||
func TestRevokingAnUnknownTokenSucceeds(t *testing.T) {
|
||||
h := newHarness(t)
|
||||
form := url.Values{"token": {"not-a-real-token"}}
|
||||
req := httptest.NewRequest(http.MethodPost, "/oauth/revoke", strings.NewReader(form.Encode()))
|
||||
req.Header.Set("Content-Type", "application/x-www-form-urlencoded")
|
||||
rec := httptest.NewRecorder()
|
||||
h.server.RevokeHandler().ServeHTTP(rec, req)
|
||||
if rec.Code != http.StatusOK {
|
||||
t.Errorf("status = %d, want 200 per RFC 7009", rec.Code)
|
||||
}
|
||||
}
|
||||
|
||||
/* ── The authenticator: audience, scope, suspension ─────────────────────── */
|
||||
|
||||
func newAuthenticator(h *harness) *Authenticator {
|
||||
return NewAuthenticator(h.store, auth.NewPGUserStore(h.h.Pool), testResource, discard())
|
||||
}
|
||||
|
||||
func TestAuthenticatorProducesTheExistingIdentity(t *testing.T) {
|
||||
h := newHarness(t)
|
||||
clientID := h.register()
|
||||
code := h.authorizeOK(clientID, verifier43)
|
||||
tokens := decodeTokens(t, h.exchange(clientID, code, verifier43))
|
||||
|
||||
identity, err := newAuthenticator(h).Authenticate(context.Background(), tokens.AccessToken)
|
||||
if err != nil {
|
||||
t.Fatalf("a freshly issued token was refused: %v", err)
|
||||
}
|
||||
if identity.UserID != h.userID {
|
||||
t.Errorf("UserID = %q, want %q", identity.UserID, h.userID)
|
||||
}
|
||||
// The tenant must come from the USER ROW, which is what makes a moved or
|
||||
// suspended user take effect immediately.
|
||||
if identity.OrgID != h.orgID {
|
||||
t.Errorf("OrgID = %q, want %q", identity.OrgID, h.orgID)
|
||||
}
|
||||
if identity.Role != "admin" {
|
||||
t.Errorf("Role = %q, want admin", identity.Role)
|
||||
}
|
||||
// No session behind a bearer identity; inventing one would make a token
|
||||
// look like something logout could end.
|
||||
if identity.SessionID != "" {
|
||||
t.Errorf("SessionID = %q, want empty for a bearer identity", identity.SessionID)
|
||||
}
|
||||
}
|
||||
|
||||
// Audience confusion: a token minted by THIS server, for a DIFFERENT resource,
|
||||
// must not be spendable here. This is the confused-deputy case the MCP spec
|
||||
// calls out explicitly.
|
||||
func TestAudienceConfusionIsRejected(t *testing.T) {
|
||||
h := newHarness(t)
|
||||
pair, err := h.store.IssuePair(context.Background(), Token{
|
||||
ClientID: h.register(), UserID: h.userID, OrgID: h.orgID,
|
||||
Scopes: []string{ScopeRead}, Audience: "https://some-other-service.test/mcp",
|
||||
}, "")
|
||||
if err != nil {
|
||||
t.Fatalf("issue: %v", err)
|
||||
}
|
||||
|
||||
if _, err := newAuthenticator(h).Authenticate(context.Background(), pair.AccessToken); err == nil {
|
||||
t.Fatal("a token for another resource was accepted here")
|
||||
}
|
||||
}
|
||||
|
||||
func TestMissingScopeIsRejected(t *testing.T) {
|
||||
h := newHarness(t)
|
||||
pair, err := h.store.IssuePair(context.Background(), Token{
|
||||
ClientID: h.register(), UserID: h.userID, OrgID: h.orgID,
|
||||
Scopes: []string{"some.other.scope"}, Audience: testResource,
|
||||
}, "")
|
||||
if err != nil {
|
||||
t.Fatalf("issue: %v", err)
|
||||
}
|
||||
if _, err := newAuthenticator(h).Authenticate(context.Background(), pair.AccessToken); err == nil {
|
||||
t.Fatal("a token without krow.read was accepted")
|
||||
}
|
||||
}
|
||||
|
||||
// Suspension must take effect on the NEXT CALL, not at token expiry. Fifteen
|
||||
// minutes of access for a suspended account is fifteen minutes too many.
|
||||
func TestSuspendedUserLosesAccessImmediately(t *testing.T) {
|
||||
h := newHarness(t)
|
||||
clientID := h.register()
|
||||
code := h.authorizeOK(clientID, verifier43)
|
||||
tokens := decodeTokens(t, h.exchange(clientID, code, verifier43))
|
||||
authr := newAuthenticator(h)
|
||||
|
||||
if _, err := authr.Authenticate(context.Background(), tokens.AccessToken); err != nil {
|
||||
t.Fatalf("token should work while the user is active: %v", err)
|
||||
}
|
||||
|
||||
if _, err := h.h.Pool.Exec(context.Background(),
|
||||
`UPDATE users SET status = 'suspended' WHERE id = $1::uuid`, h.userID); err != nil {
|
||||
t.Fatalf("suspend: %v", err)
|
||||
}
|
||||
|
||||
if _, err := authr.Authenticate(context.Background(), tokens.AccessToken); err == nil {
|
||||
t.Fatal("a suspended user's token still authenticated")
|
||||
}
|
||||
// And the family must be revoked, not merely refused once.
|
||||
if r := h.token(url.Values{
|
||||
"grant_type": {"refresh_token"}, "refresh_token": {tokens.RefreshToken}, "client_id": {clientID},
|
||||
}); r.Code != http.StatusBadRequest {
|
||||
t.Error("a suspended user could still refresh")
|
||||
}
|
||||
}
|
||||
|
||||
func TestExpiredAccessTokenIsRejected(t *testing.T) {
|
||||
h := newHarness(t)
|
||||
clientID := h.register()
|
||||
code := h.authorizeOK(clientID, verifier43)
|
||||
tokens := decodeTokens(t, h.exchange(clientID, code, verifier43))
|
||||
|
||||
h.advance(AccessTokenTTL + time.Minute)
|
||||
|
||||
if _, err := newAuthenticator(h).Authenticate(context.Background(), tokens.AccessToken); err == nil {
|
||||
t.Fatal("an expired access token authenticated")
|
||||
}
|
||||
}
|
||||
|
||||
func TestRefreshTokenCannotBeUsedAsAnAccessToken(t *testing.T) {
|
||||
h := newHarness(t)
|
||||
clientID := h.register()
|
||||
code := h.authorizeOK(clientID, verifier43)
|
||||
tokens := decodeTokens(t, h.exchange(clientID, code, verifier43))
|
||||
|
||||
if _, err := newAuthenticator(h).Authenticate(context.Background(), tokens.RefreshToken); err == nil {
|
||||
t.Fatal("a refresh token authenticated an MCP request")
|
||||
}
|
||||
}
|
||||
|
||||
func TestGarbageTokensAreRejected(t *testing.T) {
|
||||
h := newHarness(t)
|
||||
authr := newAuthenticator(h)
|
||||
for name, token := range map[string]string{
|
||||
"empty": "",
|
||||
"whitespace": " ",
|
||||
"random": "not-a-token",
|
||||
"sql-ish": "' OR 1=1 --",
|
||||
"very long": strings.Repeat("a", 5000),
|
||||
} {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
if _, err := authr.Authenticate(context.Background(), token); err == nil {
|
||||
t.Errorf("%q authenticated", name)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
/* ── Discovery metadata ─────────────────────────────────────────────────── */
|
||||
|
||||
func TestProtectedResourceMetadata(t *testing.T) {
|
||||
rec := httptest.NewRecorder()
|
||||
testConfig().ProtectedResourceHandler().ServeHTTP(rec,
|
||||
httptest.NewRequest(http.MethodGet, "/.well-known/oauth-protected-resource", nil))
|
||||
|
||||
if rec.Code != http.StatusOK {
|
||||
t.Fatalf("status = %d, want 200", rec.Code)
|
||||
}
|
||||
var out protectedResourceMetadata
|
||||
if err := json.Unmarshal(rec.Body.Bytes(), &out); err != nil {
|
||||
t.Fatalf("metadata did not decode: %v", err)
|
||||
}
|
||||
if out.Resource != testResource {
|
||||
t.Errorf("resource = %q, want %q", out.Resource, testResource)
|
||||
}
|
||||
if len(out.AuthorizationServers) != 1 || out.AuthorizationServers[0] != testIssuer {
|
||||
t.Errorf("authorization_servers = %v, want [%q]", out.AuthorizationServers, testIssuer)
|
||||
}
|
||||
// The MCP spec forbids a token in the query string; advertising anything
|
||||
// but "header" would tell a client otherwise.
|
||||
if len(out.BearerMethodsSupported) != 1 || out.BearerMethodsSupported[0] != "header" {
|
||||
t.Errorf("bearer_methods_supported = %v, want [header]", out.BearerMethodsSupported)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAuthorizationServerMetadata(t *testing.T) {
|
||||
rec := httptest.NewRecorder()
|
||||
testConfig().AuthorizationServerHandler().ServeHTTP(rec,
|
||||
httptest.NewRequest(http.MethodGet, "/.well-known/oauth-authorization-server", nil))
|
||||
|
||||
var out authorizationServerMetadata
|
||||
if err := json.Unmarshal(rec.Body.Bytes(), &out); err != nil {
|
||||
t.Fatalf("metadata did not decode: %v", err)
|
||||
}
|
||||
|
||||
if out.Issuer != testIssuer {
|
||||
t.Errorf("issuer = %q, want %q", out.Issuer, testIssuer)
|
||||
}
|
||||
// Every list is a promise. Each must name only what is implemented.
|
||||
if strings.Join(out.ResponseTypesSupported, ",") != "code" {
|
||||
t.Errorf("response_types_supported = %v; implicit must not be advertised", out.ResponseTypesSupported)
|
||||
}
|
||||
if strings.Join(out.CodeChallengeMethodsSupported, ",") != MethodS256 {
|
||||
t.Errorf("code_challenge_methods_supported = %v, want [S256]", out.CodeChallengeMethodsSupported)
|
||||
}
|
||||
for _, forbidden := range []string{"password", "client_credentials", "implicit"} {
|
||||
for _, advertised := range out.GrantTypesSupported {
|
||||
if advertised == forbidden {
|
||||
t.Errorf("grant_types_supported advertises %q, which is refused", forbidden)
|
||||
}
|
||||
}
|
||||
}
|
||||
for _, advertised := range out.ScopesSupported {
|
||||
if advertised == ScopeWrite {
|
||||
t.Error("scopes_supported advertises krow.write, which is not issued")
|
||||
}
|
||||
}
|
||||
if !out.ResourceIndicatorsSupported {
|
||||
t.Error("resource_indicators_supported must be true — RFC 8707 is required by MCP")
|
||||
}
|
||||
// Every endpoint comes from configuration, never a hardcoded domain.
|
||||
for name, got := range map[string]string{
|
||||
"authorization_endpoint": out.AuthorizationEndpoint,
|
||||
"token_endpoint": out.TokenEndpoint,
|
||||
"registration_endpoint": out.RegistrationEndpoint,
|
||||
} {
|
||||
if !strings.HasPrefix(got, testIssuer) {
|
||||
t.Errorf("%s = %q, want it under the configured issuer", name, got)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// A token response must never be cached: the body is a credential.
|
||||
func TestTokenResponsesAreNotCacheable(t *testing.T) {
|
||||
h := newHarness(t)
|
||||
clientID := h.register()
|
||||
code := h.authorizeOK(clientID, verifier43)
|
||||
rec := h.exchange(clientID, code, verifier43)
|
||||
|
||||
if got := rec.Header().Get("Cache-Control"); !strings.Contains(got, "no-store") {
|
||||
t.Errorf("Cache-Control = %q, want no-store", got)
|
||||
}
|
||||
}
|
||||
130
go-api/internal/oauth/pkce.go
Normal file
130
go-api/internal/oauth/pkce.go
Normal file
@@ -0,0 +1,130 @@
|
||||
// Package oauth is KROW's OAuth 2.1 authorization server and the token store
|
||||
// behind it.
|
||||
//
|
||||
// It exists for one caller: the MCP surface, which needs a way to authenticate
|
||||
// a client that cannot hold a cookie. Everything here is in service of turning
|
||||
// a browser-based approval into an opaque bearer token that
|
||||
// mcpserver.TokenAuthenticator can resolve back into the SAME
|
||||
// authctx.Identity the cookie path produces.
|
||||
//
|
||||
// # WHAT THIS PACKAGE DOES NOT DO
|
||||
//
|
||||
// It does not authorize anything. It establishes WHO is calling; what they may
|
||||
// then read is decided by the existing policy table, in the existing tool
|
||||
// layer, exactly as it is for a cookie session. There is no OAuth scope that
|
||||
// grants access to a row. `krow.read` says "this client may use the read
|
||||
// tools"; whether this user may see a particular row is a question
|
||||
// tools/scope.go answers and this package never touches.
|
||||
//
|
||||
// It also does not store a password, check one, or keep a second user table.
|
||||
// The authorization endpoint authenticates the person using the session they
|
||||
// already have — see authserver.go.
|
||||
package oauth
|
||||
|
||||
import (
|
||||
"crypto/sha256"
|
||||
"crypto/subtle"
|
||||
"encoding/base64"
|
||||
"errors"
|
||||
"regexp"
|
||||
)
|
||||
|
||||
// PKCE — Proof Key for Code Exchange, RFC 7636.
|
||||
//
|
||||
// The problem it solves: an authorization code travels back through a browser
|
||||
// redirect, which is the least trustworthy hop in the flow. Anything that can
|
||||
// observe that redirect — a malicious app registered for the same custom URL
|
||||
// scheme, a proxy, a shoulder — can steal the code. For a confidential client
|
||||
// that does not matter, because redeeming the code also requires a client
|
||||
// secret. A public client has no secret, so the code alone would be enough.
|
||||
//
|
||||
// PKCE gives the client a per-request secret instead. It invents a random
|
||||
// `code_verifier`, sends only SHA-256 of it with the authorization request, and
|
||||
// presents the verifier itself at the token endpoint. A stolen code is useless
|
||||
// without the verifier, which never travelled through the browser.
|
||||
//
|
||||
// S256 ONLY. RFC 7636 also defines `plain`, where the challenge IS the
|
||||
// verifier. That protects against nothing — anyone who stole the code from the
|
||||
// redirect also stole the challenge, and the challenge is the verifier — and
|
||||
// OAuth 2.1 forbids it for public clients. It is refused here, and refused
|
||||
// again by a CHECK constraint in migration 000013, so no code path can relax it.
|
||||
|
||||
// MethodS256 is the only code_challenge_method this server accepts.
|
||||
const MethodS256 = "S256"
|
||||
|
||||
var (
|
||||
// ErrUnsupportedChallengeMethod covers `plain` and anything else.
|
||||
ErrUnsupportedChallengeMethod = errors.New("oauth: code_challenge_method must be S256")
|
||||
|
||||
// ErrMalformedChallenge covers a challenge that is not base64url of a
|
||||
// SHA-256 digest.
|
||||
ErrMalformedChallenge = errors.New("oauth: malformed code_challenge")
|
||||
|
||||
// ErrMalformedVerifier covers a verifier outside RFC 7636's length or
|
||||
// character set.
|
||||
ErrMalformedVerifier = errors.New("oauth: malformed code_verifier")
|
||||
|
||||
// ErrVerifierMismatch is the one that matters: a verifier that does not
|
||||
// hash to the stored challenge.
|
||||
ErrVerifierMismatch = errors.New("oauth: code_verifier does not match code_challenge")
|
||||
)
|
||||
|
||||
// challengePattern is base64url of a 32-byte digest: 43 characters, unpadded.
|
||||
// The same pattern migration 000013 enforces in oauth_grants_challenge_shape.
|
||||
var challengePattern = regexp.MustCompile(`^[A-Za-z0-9_-]{43}$`)
|
||||
|
||||
// verifierPattern is RFC 7636 section 4.1's `code_verifier` grammar:
|
||||
// unreserved characters only, 43 to 128 of them.
|
||||
var verifierPattern = regexp.MustCompile(`^[A-Za-z0-9._~-]{43,128}$`)
|
||||
|
||||
// ValidateChallenge checks a code_challenge and its method at the authorization
|
||||
// endpoint, before any row is written.
|
||||
//
|
||||
// Rejecting a malformed challenge here rather than at the token endpoint means
|
||||
// the failure lands where the client can act on it — on its own authorization
|
||||
// request — instead of after a person has been walked through a consent screen
|
||||
// for a flow that was never going to complete.
|
||||
func ValidateChallenge(challenge, method string) error {
|
||||
if method != MethodS256 {
|
||||
return ErrUnsupportedChallengeMethod
|
||||
}
|
||||
if !challengePattern.MatchString(challenge) {
|
||||
return ErrMalformedChallenge
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// VerifyChallenge reports whether a verifier matches a stored challenge.
|
||||
//
|
||||
// The comparison is constant-time. A byte-by-byte comparison that returned
|
||||
// early would leak, through timing, how much of a guessed verifier was correct
|
||||
// — which turns an infeasible search into a feasible one, one character at a
|
||||
// time. The values being compared are both base64url text of the same fixed
|
||||
// length, so subtle.ConstantTimeCompare is exactly the right tool.
|
||||
func VerifyChallenge(verifier, challenge, method string) error {
|
||||
if method != MethodS256 {
|
||||
return ErrUnsupportedChallengeMethod
|
||||
}
|
||||
if !verifierPattern.MatchString(verifier) {
|
||||
return ErrMalformedVerifier
|
||||
}
|
||||
if !challengePattern.MatchString(challenge) {
|
||||
return ErrMalformedChallenge
|
||||
}
|
||||
|
||||
computed := ChallengeFor(verifier)
|
||||
if subtle.ConstantTimeCompare([]byte(computed), []byte(challenge)) != 1 {
|
||||
return ErrVerifierMismatch
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// ChallengeFor derives the S256 challenge for a verifier.
|
||||
//
|
||||
// base64url WITHOUT padding, per RFC 7636 appendix A. Padding would add a '='
|
||||
// that has to be escaped in a query string, and a client that padded would
|
||||
// produce a challenge this server did not recognise.
|
||||
func ChallengeFor(verifier string) string {
|
||||
sum := sha256.Sum256([]byte(verifier))
|
||||
return base64.RawURLEncoding.EncodeToString(sum[:])
|
||||
}
|
||||
107
go-api/internal/oauth/pkce_test.go
Normal file
107
go-api/internal/oauth/pkce_test.go
Normal file
@@ -0,0 +1,107 @@
|
||||
package oauth
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
// The RFC 7636 appendix B worked example. Using the spec's own vector rather
|
||||
// than a value this implementation produced means the test would catch an
|
||||
// encoding mistake that is self-consistent — base64 standard instead of
|
||||
// base64url, say, or padded instead of raw — which a round-trip test could not.
|
||||
const (
|
||||
specVerifier = "dBjftJeZ4CVP-mB92K27uhbUJU1p1r_wW1gFWFOEjXk"
|
||||
specChallenge = "E9Melhoa2OwvFrEMTJguCHaoeK1t8URWbuGJSstw-cM"
|
||||
)
|
||||
|
||||
func TestChallengeForMatchesTheRFCVector(t *testing.T) {
|
||||
if got := ChallengeFor(specVerifier); got != specChallenge {
|
||||
t.Errorf("ChallengeFor(RFC 7636 verifier) = %q, want %q", got, specChallenge)
|
||||
}
|
||||
}
|
||||
|
||||
func TestVerifyChallengeAcceptsTheCorrectVerifier(t *testing.T) {
|
||||
if err := VerifyChallenge(specVerifier, specChallenge, MethodS256); err != nil {
|
||||
t.Errorf("the RFC's own verifier was rejected: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestVerifyChallengeRejectsAWrongVerifier(t *testing.T) {
|
||||
// Same length and character set, one character different. A comparison
|
||||
// that was accidentally checking length or prefix would let this through.
|
||||
wrong := "dBjftJeZ4CVP-mB92K27uhbUJU1p1r_wW1gFWFOEjXX"
|
||||
err := VerifyChallenge(wrong, specChallenge, MethodS256)
|
||||
if !errors.Is(err, ErrVerifierMismatch) {
|
||||
t.Errorf("err = %v, want ErrVerifierMismatch", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestVerifyChallengeRejectsAMissingVerifier(t *testing.T) {
|
||||
if err := VerifyChallenge("", specChallenge, MethodS256); !errors.Is(err, ErrMalformedVerifier) {
|
||||
t.Errorf("err = %v, want ErrMalformedVerifier", err)
|
||||
}
|
||||
}
|
||||
|
||||
// `plain` must be refused wherever it appears. It is legal in RFC 7636 and
|
||||
// forbidden by OAuth 2.1 for public clients, because the challenge IS the
|
||||
// verifier and anyone who stole one stole both.
|
||||
func TestPlainMethodIsRejected(t *testing.T) {
|
||||
for _, method := range []string{"plain", "PLAIN", "", "s256", "S512"} {
|
||||
t.Run("method="+method, func(t *testing.T) {
|
||||
if err := ValidateChallenge(specChallenge, method); !errors.Is(err, ErrUnsupportedChallengeMethod) {
|
||||
t.Errorf("ValidateChallenge: err = %v, want ErrUnsupportedChallengeMethod", err)
|
||||
}
|
||||
if err := VerifyChallenge(specVerifier, specChallenge, method); !errors.Is(err, ErrUnsupportedChallengeMethod) {
|
||||
t.Errorf("VerifyChallenge: err = %v, want ErrUnsupportedChallengeMethod", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestMalformedChallengeIsRejected(t *testing.T) {
|
||||
for name, challenge := range map[string]string{
|
||||
"empty": "",
|
||||
"too short": "abc",
|
||||
"too long": strings.Repeat("a", 44),
|
||||
"padded base64": "E9Melhoa2OwvFrEMTJguCHaoeK1t8URWbuGJSstw-cM=",
|
||||
"standard base64": "E9Melhoa2OwvFrEMTJguCHaoeK1t8URWbuGJSstw+cM",
|
||||
"illegal char": "E9Melhoa2OwvFrEMTJguCHaoeK1t8URWbuGJSstw!cM",
|
||||
"whitespace": "E9Melhoa2OwvFrEMTJguCHaoeK1t8URWbuGJSstw cM",
|
||||
"newline injected": "E9Melhoa2OwvFrEMTJguCHaoeK1t8URWbuGJSstw\ncM",
|
||||
} {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
if err := ValidateChallenge(challenge, MethodS256); !errors.Is(err, ErrMalformedChallenge) {
|
||||
t.Errorf("err = %v, want ErrMalformedChallenge", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// RFC 7636 section 4.1 constrains the verifier to 43–128 unreserved characters.
|
||||
// A verifier outside that range is malformed regardless of what it hashes to.
|
||||
func TestMalformedVerifierIsRejected(t *testing.T) {
|
||||
for name, verifier := range map[string]string{
|
||||
"too short (42)": strings.Repeat("a", 42),
|
||||
"too long (129)": strings.Repeat("a", 129),
|
||||
"illegal char": strings.Repeat("a", 42) + "!",
|
||||
"whitespace": strings.Repeat("a", 42) + " ",
|
||||
} {
|
||||
t.Run(name, func(t *testing.T) {
|
||||
if err := VerifyChallenge(verifier, specChallenge, MethodS256); !errors.Is(err, ErrMalformedVerifier) {
|
||||
t.Errorf("err = %v, want ErrMalformedVerifier", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// A verifier at each end of the legal range must be accepted, or clients
|
||||
// generating the maximum length would fail against this server.
|
||||
func TestVerifierBoundariesAreAccepted(t *testing.T) {
|
||||
for _, length := range []int{43, 128} {
|
||||
verifier := strings.Repeat("a", length)
|
||||
if err := VerifyChallenge(verifier, ChallengeFor(verifier), MethodS256); err != nil {
|
||||
t.Errorf("a %d-character verifier was rejected: %v", length, err)
|
||||
}
|
||||
}
|
||||
}
|
||||
444
go-api/internal/oauth/store.go
Normal file
444
go-api/internal/oauth/store.go
Normal file
@@ -0,0 +1,444 @@
|
||||
package oauth
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
"github.com/krow/krow-backend/go-api/internal/auth"
|
||||
"github.com/krow/krow-backend/go-api/internal/repo"
|
||||
)
|
||||
|
||||
// The persistence layer for clients, authorization codes and tokens.
|
||||
//
|
||||
// Two rules hold throughout this file and are worth stating once:
|
||||
//
|
||||
// 1. NO RAW CREDENTIAL IS EVER WRITTEN. Every code and token is hashed with
|
||||
// auth.HashToken before it reaches a statement. The CHECK constraints in
|
||||
// migrations 000013 and 000014 refuse anything that is not 64 hex
|
||||
// characters, so this is enforced twice — in Go where the mistake would be
|
||||
// made, and in the schema where it would land.
|
||||
//
|
||||
// 2. SINGLE-USE IS ENFORCED BY THE UPDATE, NOT BY A READ. Redeeming a code or
|
||||
// a refresh token is one statement that marks the row consumed and returns
|
||||
// it in the same breath. A SELECT followed by an UPDATE has a window
|
||||
// between them where two concurrent requests both see an unspent row and
|
||||
// both proceed, which is precisely the replay the single-use rule exists to
|
||||
// prevent.
|
||||
|
||||
var (
|
||||
// ErrNotFound covers a client, code or token that does not exist.
|
||||
ErrNotFound = errors.New("oauth: not found")
|
||||
|
||||
// ErrGrantUnusable covers a code that is expired, already consumed, or
|
||||
// simply absent. ONE error for all three: distinguishing them tells a
|
||||
// caller whether a code they hold was ever real, which is an oracle.
|
||||
ErrGrantUnusable = errors.New("oauth: authorization code is not usable")
|
||||
|
||||
// ErrTokenUnusable covers a token that is unknown, expired, revoked or
|
||||
// consumed. One error, same reasoning.
|
||||
ErrTokenUnusable = errors.New("oauth: token is not usable")
|
||||
|
||||
// ErrRefreshReuse is raised when a CONSUMED refresh token is presented
|
||||
// again. It is distinct from ErrTokenUnusable internally because it
|
||||
// triggers family revocation — but the caller must still answer the client
|
||||
// with an indistinguishable error.
|
||||
ErrRefreshReuse = errors.New("oauth: refresh token reuse detected")
|
||||
)
|
||||
|
||||
// Store is the database-backed persistence for this package.
|
||||
type Store struct {
|
||||
db repo.Querier
|
||||
now func() time.Time
|
||||
}
|
||||
|
||||
// NewStore builds a store over the existing pool.
|
||||
func NewStore(db repo.Querier) *Store {
|
||||
return &Store{db: db, now: time.Now}
|
||||
}
|
||||
|
||||
// WithClock replaces the clock, so expiry can be tested without sleeping.
|
||||
func (s *Store) WithClock(now func() time.Time) *Store {
|
||||
s.now = now
|
||||
return s
|
||||
}
|
||||
|
||||
/* ── Clients ────────────────────────────────────────────────────────────── */
|
||||
|
||||
// Client is a registered OAuth client.
|
||||
type Client struct {
|
||||
ClientID string
|
||||
ClientName string
|
||||
RedirectURIs []string
|
||||
GrantTypes []string
|
||||
Scopes []string
|
||||
DisabledAt *time.Time
|
||||
}
|
||||
|
||||
// AllowsRedirect reports whether a redirect_uri is registered to this client.
|
||||
//
|
||||
// EXACT string equality. Not a prefix match, not a normalised comparison, not
|
||||
// "same host and port". Every relaxation of this check is an open redirect: a
|
||||
// prefix match lets `https://good.example/cb.attacker.com` through, and
|
||||
// normalising lets encoding tricks through. RFC 6749 section 3.1.2.3 says
|
||||
// exact, and exact is what this is.
|
||||
func (c Client) AllowsRedirect(uri string) bool {
|
||||
for _, registered := range c.RedirectURIs {
|
||||
if registered == uri {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// AllowsScopes reports whether every requested scope is registered.
|
||||
func (c Client) AllowsScopes(requested []string) bool {
|
||||
for _, want := range requested {
|
||||
found := false
|
||||
for _, have := range c.Scopes {
|
||||
if have == want {
|
||||
found = true
|
||||
break
|
||||
}
|
||||
}
|
||||
if !found {
|
||||
return false
|
||||
}
|
||||
}
|
||||
return true
|
||||
}
|
||||
|
||||
// CreateClient registers a new public client.
|
||||
func (s *Store) CreateClient(ctx context.Context, c Client) error {
|
||||
_, err := s.db.Exec(ctx,
|
||||
`INSERT INTO oauth_clients (client_id, client_name, redirect_uris, grant_types, scopes)
|
||||
VALUES ($1, $2, $3, $4, $5)`,
|
||||
c.ClientID, c.ClientName, c.RedirectURIs, c.GrantTypes, c.Scopes)
|
||||
if err != nil {
|
||||
return fmt.Errorf("oauth: create client: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// FindClient resolves a client_id. A disabled client is reported as not found:
|
||||
// whether it once existed is not the caller's business.
|
||||
func (s *Store) FindClient(ctx context.Context, clientID string) (Client, error) {
|
||||
var c Client
|
||||
err := s.db.QueryRow(ctx,
|
||||
`SELECT client_id, client_name, redirect_uris, grant_types, scopes, disabled_at
|
||||
FROM oauth_clients
|
||||
WHERE client_id = $1 AND disabled_at IS NULL`,
|
||||
clientID).Scan(&c.ClientID, &c.ClientName, &c.RedirectURIs, &c.GrantTypes, &c.Scopes, &c.DisabledAt)
|
||||
if err != nil {
|
||||
return Client{}, ErrNotFound
|
||||
}
|
||||
return c, nil
|
||||
}
|
||||
|
||||
/* ── Authorization codes ────────────────────────────────────────────────── */
|
||||
|
||||
// Grant is an issued authorization code, as stored.
|
||||
type Grant struct {
|
||||
ID string
|
||||
ClientID string
|
||||
UserID string
|
||||
OrgID string
|
||||
RedirectURI string
|
||||
Scopes []string
|
||||
Resource string
|
||||
CodeChallenge string
|
||||
CodeChallengeMethod string
|
||||
ExpiresAt time.Time
|
||||
}
|
||||
|
||||
// GrantTTL is how long an authorization code stays redeemable.
|
||||
//
|
||||
// Sixty seconds. RFC 6749 recommends a maximum of ten minutes and "a maximum
|
||||
// of 60 seconds is RECOMMENDED" for the code's lifetime in OAuth 2.1 guidance.
|
||||
// The code is in transit through a browser redirect and is exchanged
|
||||
// immediately by a client that is already waiting for it; a longer window buys
|
||||
// nothing and widens the replay opportunity.
|
||||
const GrantTTL = 60 * time.Second
|
||||
|
||||
// CreateGrant stores an authorization code, returning the RAW code exactly
|
||||
// once.
|
||||
//
|
||||
// The raw value is returned and never persisted. The caller puts it in a
|
||||
// redirect and forgets it.
|
||||
func (s *Store) CreateGrant(ctx context.Context, g Grant) (rawCode string, err error) {
|
||||
rawCode, err = auth.GenerateToken()
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("oauth: generate code: %w", err)
|
||||
}
|
||||
|
||||
_, err = s.db.Exec(ctx,
|
||||
`INSERT INTO oauth_grants
|
||||
(code_hash, client_id, user_id, org_id, redirect_uri, scopes, resource,
|
||||
code_challenge, code_challenge_method, expires_at)
|
||||
VALUES ($1, $2, $3::uuid, $4::uuid, $5, $6, $7, $8, $9, $10)`,
|
||||
auth.HashToken(rawCode), g.ClientID, g.UserID, g.OrgID, g.RedirectURI,
|
||||
g.Scopes, g.Resource, g.CodeChallenge, g.CodeChallengeMethod,
|
||||
s.now().Add(GrantTTL))
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("oauth: create grant: %w", err)
|
||||
}
|
||||
return rawCode, nil
|
||||
}
|
||||
|
||||
// RedeemGrant consumes an authorization code and returns what it was bound to.
|
||||
//
|
||||
// ONE STATEMENT. The UPDATE marks the row consumed and RETURNS it, so the read
|
||||
// and the write cannot be interleaved by a concurrent request. The predicate
|
||||
// carries the whole single-use rule: `consumed_at IS NULL` means a spent code
|
||||
// matches nothing, and `expires_at > now()` means an old one does too. A second
|
||||
// redemption of the same code updates zero rows and therefore fails, which is
|
||||
// what replay protection looks like when the database enforces it.
|
||||
func (s *Store) RedeemGrant(ctx context.Context, rawCode string) (Grant, error) {
|
||||
var g Grant
|
||||
err := s.db.QueryRow(ctx,
|
||||
`UPDATE oauth_grants
|
||||
SET consumed_at = now()
|
||||
WHERE code_hash = $1
|
||||
AND consumed_at IS NULL
|
||||
AND expires_at > $2
|
||||
RETURNING id::text, client_id, user_id::text, org_id::text, redirect_uri,
|
||||
scopes, resource, code_challenge, code_challenge_method, expires_at`,
|
||||
auth.HashToken(rawCode), s.now()).
|
||||
Scan(&g.ID, &g.ClientID, &g.UserID, &g.OrgID, &g.RedirectURI,
|
||||
&g.Scopes, &g.Resource, &g.CodeChallenge, &g.CodeChallengeMethod, &g.ExpiresAt)
|
||||
if err != nil {
|
||||
// No row: unknown, expired or already spent. Indistinguishable on
|
||||
// purpose — see ErrGrantUnusable.
|
||||
return Grant{}, ErrGrantUnusable
|
||||
}
|
||||
return g, nil
|
||||
}
|
||||
|
||||
/* ── Tokens ─────────────────────────────────────────────────────────────── */
|
||||
|
||||
// Token is an issued access or refresh token, as stored.
|
||||
type Token struct {
|
||||
ID string
|
||||
Type string
|
||||
FamilyID string
|
||||
ClientID string
|
||||
UserID string
|
||||
OrgID string
|
||||
Scopes []string
|
||||
Audience string
|
||||
ExpiresAt time.Time
|
||||
}
|
||||
|
||||
// Token lifetimes.
|
||||
//
|
||||
// Fifteen minutes for an access token is the number the MCP plan committed to,
|
||||
// and the reasoning is that an access token travels on every single request: it
|
||||
// is the most exposed credential in the system and the one with the least need
|
||||
// to be long-lived, because a refresh token exists precisely so the client can
|
||||
// get another without troubling the user.
|
||||
//
|
||||
// Thirty days for a refresh token matches the session's own "remember me"
|
||||
// ceiling, so a connected client and a remembered browser lapse on the same
|
||||
// schedule rather than on two different ones nobody can remember.
|
||||
const (
|
||||
AccessTokenTTL = 15 * time.Minute
|
||||
RefreshTokenTTL = 30 * 24 * time.Hour
|
||||
)
|
||||
|
||||
// TokenPair is what a successful token request produces.
|
||||
//
|
||||
// The raw values are here and nowhere else: they are returned to the client in
|
||||
// the token response and are never stored, logged or re-derivable.
|
||||
type TokenPair struct {
|
||||
AccessToken string
|
||||
RefreshToken string
|
||||
ExpiresIn int
|
||||
Scopes []string
|
||||
FamilyID string
|
||||
}
|
||||
|
||||
// IssuePair mints an access and refresh token in one family.
|
||||
//
|
||||
// familyID empty starts a new lineage; a supplied one continues an existing
|
||||
// lineage through a rotation, which is what lets reuse detection revoke every
|
||||
// descendant of a stolen token.
|
||||
func (s *Store) IssuePair(ctx context.Context, t Token, familyID string) (TokenPair, error) {
|
||||
if familyID == "" {
|
||||
generated, err := newUUID()
|
||||
if err != nil {
|
||||
return TokenPair{}, err
|
||||
}
|
||||
familyID = generated
|
||||
}
|
||||
|
||||
access, err := auth.GenerateToken()
|
||||
if err != nil {
|
||||
return TokenPair{}, fmt.Errorf("oauth: generate access token: %w", err)
|
||||
}
|
||||
refresh, err := auth.GenerateToken()
|
||||
if err != nil {
|
||||
return TokenPair{}, fmt.Errorf("oauth: generate refresh token: %w", err)
|
||||
}
|
||||
|
||||
now := s.now()
|
||||
for _, row := range []struct {
|
||||
raw string
|
||||
kind string
|
||||
expires time.Time
|
||||
}{
|
||||
{access, "access", now.Add(AccessTokenTTL)},
|
||||
{refresh, "refresh", now.Add(RefreshTokenTTL)},
|
||||
} {
|
||||
if _, err := s.db.Exec(ctx,
|
||||
`INSERT INTO oauth_tokens
|
||||
(token_hash, token_type, family_id, client_id, user_id, org_id,
|
||||
scopes, audience, expires_at)
|
||||
VALUES ($1, $2, $3::uuid, $4, $5::uuid, $6::uuid, $7, $8, $9)`,
|
||||
auth.HashToken(row.raw), row.kind, familyID, t.ClientID, t.UserID,
|
||||
t.OrgID, t.Scopes, t.Audience, row.expires); err != nil {
|
||||
return TokenPair{}, fmt.Errorf("oauth: store %s token: %w", row.kind, err)
|
||||
}
|
||||
}
|
||||
|
||||
return TokenPair{
|
||||
AccessToken: access,
|
||||
RefreshToken: refresh,
|
||||
ExpiresIn: int(AccessTokenTTL.Seconds()),
|
||||
Scopes: t.Scopes,
|
||||
FamilyID: familyID,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// FindAccessToken resolves a raw access token for validation.
|
||||
//
|
||||
// Read-only: validation happens on every MCP request and must not write. The
|
||||
// predicate does the whole job — unknown, expired, revoked and wrong-type all
|
||||
// return no row and therefore the same error.
|
||||
func (s *Store) FindAccessToken(ctx context.Context, raw string) (Token, error) {
|
||||
var t Token
|
||||
err := s.db.QueryRow(ctx,
|
||||
`SELECT id::text, token_type, family_id::text, client_id, user_id::text,
|
||||
org_id::text, scopes, audience, expires_at
|
||||
FROM oauth_tokens
|
||||
WHERE token_hash = $1
|
||||
AND token_type = 'access'
|
||||
AND revoked_at IS NULL
|
||||
AND expires_at > $2`,
|
||||
auth.HashToken(raw), s.now()).
|
||||
Scan(&t.ID, &t.Type, &t.FamilyID, &t.ClientID, &t.UserID, &t.OrgID,
|
||||
&t.Scopes, &t.Audience, &t.ExpiresAt)
|
||||
if err != nil {
|
||||
return Token{}, ErrTokenUnusable
|
||||
}
|
||||
return t, nil
|
||||
}
|
||||
|
||||
// RedeemRefreshToken consumes a refresh token, or detects its reuse.
|
||||
//
|
||||
// The two-step here is deliberate and is the heart of reuse detection:
|
||||
//
|
||||
// 1. Try to consume an unspent, unexpired, unrevoked refresh token. One
|
||||
// statement, same single-use reasoning as RedeemGrant.
|
||||
// 2. If that matched nothing, look again WITHOUT the `consumed_at IS NULL`
|
||||
// predicate. A row that exists but was already consumed is not an ordinary
|
||||
// failure — it means someone presented a token that had already been
|
||||
// rotated away, and there is no way to tell the legitimate client retrying
|
||||
// from an attacker replaying a stolen token.
|
||||
//
|
||||
// OAuth 2.1's answer to that ambiguity is to assume the worse case and revoke
|
||||
// the whole family. The attacker loses access; the legitimate client is pushed
|
||||
// through a fresh authorization it can complete. Doing nothing would leave a
|
||||
// thief with a working credential.
|
||||
func (s *Store) RedeemRefreshToken(ctx context.Context, raw string) (Token, error) {
|
||||
hash := auth.HashToken(raw)
|
||||
|
||||
var t Token
|
||||
err := s.db.QueryRow(ctx,
|
||||
`UPDATE oauth_tokens
|
||||
SET consumed_at = now(), last_used_at = now()
|
||||
WHERE token_hash = $1
|
||||
AND token_type = 'refresh'
|
||||
AND consumed_at IS NULL
|
||||
AND revoked_at IS NULL
|
||||
AND expires_at > $2
|
||||
RETURNING id::text, token_type, family_id::text, client_id, user_id::text,
|
||||
org_id::text, scopes, audience, expires_at`,
|
||||
hash, s.now()).
|
||||
Scan(&t.ID, &t.Type, &t.FamilyID, &t.ClientID, &t.UserID, &t.OrgID,
|
||||
&t.Scopes, &t.Audience, &t.ExpiresAt)
|
||||
if err == nil {
|
||||
return t, nil
|
||||
}
|
||||
|
||||
// Step 2: was this a token that HAD been valid and is now spent?
|
||||
var familyID string
|
||||
if probeErr := s.db.QueryRow(ctx,
|
||||
`SELECT family_id::text FROM oauth_tokens
|
||||
WHERE token_hash = $1 AND token_type = 'refresh' AND consumed_at IS NOT NULL`,
|
||||
hash).Scan(&familyID); probeErr == nil {
|
||||
// Reuse. Revoke the lineage and report it, so the caller can log it at
|
||||
// a level that gets noticed — while still answering the client with an
|
||||
// indistinguishable error.
|
||||
_ = s.RevokeFamily(ctx, familyID, "refresh_token_reuse")
|
||||
return Token{}, ErrRefreshReuse
|
||||
}
|
||||
|
||||
return Token{}, ErrTokenUnusable
|
||||
}
|
||||
|
||||
/* ── Revocation ─────────────────────────────────────────────────────────── */
|
||||
|
||||
// RevokeFamily revokes every token in a rotation lineage.
|
||||
//
|
||||
// Idempotent, and it does not care whether the rows were already revoked: the
|
||||
// predicate narrows to unrevoked rows so a second call is a no-op rather than
|
||||
// an error, which matters because this is called from an error path.
|
||||
func (s *Store) RevokeFamily(ctx context.Context, familyID, reason string) error {
|
||||
_, err := s.db.Exec(ctx,
|
||||
`UPDATE oauth_tokens
|
||||
SET revoked_at = now(), revoked_reason = $2
|
||||
WHERE family_id = $1::uuid AND revoked_at IS NULL`,
|
||||
familyID, reason)
|
||||
if err != nil {
|
||||
return fmt.Errorf("oauth: revoke family: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// RevokeToken revokes one token by its raw value, and its family with it.
|
||||
//
|
||||
// Revoking the family rather than the single row is what makes "disconnect"
|
||||
// mean what a person expects. Revoking one access token would leave the
|
||||
// refresh token alive to mint another within seconds, so the button that says
|
||||
// "disconnect Claude" would not disconnect Claude.
|
||||
func (s *Store) RevokeToken(ctx context.Context, raw, reason string) error {
|
||||
var familyID string
|
||||
if err := s.db.QueryRow(ctx,
|
||||
`SELECT family_id::text FROM oauth_tokens WHERE token_hash = $1`,
|
||||
auth.HashToken(raw)).Scan(&familyID); err != nil {
|
||||
// RFC 7009: revoking an unknown token is a success. Saying otherwise
|
||||
// turns the revocation endpoint into a way to test whether a token
|
||||
// exists.
|
||||
return nil
|
||||
}
|
||||
return s.RevokeFamily(ctx, familyID, reason)
|
||||
}
|
||||
|
||||
// RevokeAllForUser revokes every token a user holds.
|
||||
//
|
||||
// Called when an account is suspended or a person disconnects every app. Token
|
||||
// validation already re-reads the user and refuses a suspended one, so this is
|
||||
// belt to that braces: it stops the tokens existing rather than relying on
|
||||
// every future validation to notice.
|
||||
func (s *Store) RevokeAllForUser(ctx context.Context, userID, reason string) error {
|
||||
_, err := s.db.Exec(ctx,
|
||||
`UPDATE oauth_tokens
|
||||
SET revoked_at = now(), revoked_reason = $2
|
||||
WHERE user_id = $1::uuid AND revoked_at IS NULL`,
|
||||
userID, reason)
|
||||
if err != nil {
|
||||
return fmt.Errorf("oauth: revoke user tokens: %w", err)
|
||||
}
|
||||
return nil
|
||||
}
|
||||
Reference in New Issue
Block a user