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