refactor(10.2-01): nest framework packages under modules
- Move remaining beach packages and embedded admin assets\n- Rewrite framework, example, build, and gate paths
This commit is contained in:
272
modules/wristband/authorize.go
Normal file
272
modules/wristband/authorize.go
Normal file
@@ -0,0 +1,272 @@
|
||||
// RFC 6749 authorization endpoint for MCP OAuth, ported from PHP
|
||||
// OAuthAuthorizeController::authorize byte-for-byte including its
|
||||
// validation order (08-CONTEXT.md D-02/D-04/D-05; canonical PHP source:
|
||||
// OAuthAuthorizeController.php).
|
||||
//
|
||||
// D-02: query-only parsing (no body is ever read). The client and the exact
|
||||
// registered redirect URI are validated before any redirect response is
|
||||
// constructed (T-08-OPEN-REDIRECT): an unknown client or unregistered
|
||||
// redirect is a local text/plain 400 with no Location header. Every later
|
||||
// failure redirects to the now-trusted redirect_uri with an ordered
|
||||
// error/error_description/iss[/state] query built through an RFC 3986
|
||||
// encoder, never url.Values.Encode (08-RESEARCH.md Pattern 3/Pitfall 5).
|
||||
package wristband
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"net/http"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// authorizeAllowedScopes mirrors PHP OAuthAuthorizeController::ALLOWED_SCOPES.
|
||||
var authorizeAllowedScopes = map[string]bool{
|
||||
"read": true, "write": true, "ai": true, "offline_access": true,
|
||||
}
|
||||
|
||||
// Authorize handles GET /oauth/mcp/authorize (D-09: raw route, no
|
||||
// middleware). It validates the client and exact redirect before any
|
||||
// redirect response, enforces S256 PKCE syntax, the client's scope ceiling,
|
||||
// and the RFC 8707 resource check, then persists an opaque pending request
|
||||
// and redirects to the app's /connect handoff.
|
||||
func (s *Server) Authorize(w http.ResponseWriter, r *http.Request) {
|
||||
if s.backend == nil {
|
||||
w.WriteHeader(http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
|
||||
ctx := r.Context()
|
||||
q := r.URL.Query()
|
||||
|
||||
clientID := q.Get("client_id")
|
||||
var client *ClientRecord
|
||||
if clientID != "" {
|
||||
err := s.backend.WithinTx(ctx, func(tx Tx) error {
|
||||
c, err := tx.ByClientID(ctx, clientID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
client = c
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
w.WriteHeader(http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
}
|
||||
if client == nil || client.RevokedAt != nil {
|
||||
writeAuthorizeLocalError(w, "Unknown client.")
|
||||
return
|
||||
}
|
||||
|
||||
redirectURI := q.Get("redirect_uri")
|
||||
if redirectURI == "" || !stringSliceContains(client.RedirectURIs, redirectURI) {
|
||||
writeAuthorizeLocalError(w, "Unregistered redirect URI.")
|
||||
return
|
||||
}
|
||||
|
||||
var state *string
|
||||
if v := q.Get("state"); v != "" {
|
||||
state = &v
|
||||
}
|
||||
|
||||
if q.Get("response_type") != "code" {
|
||||
s.authorizeErrorRedirect(w, redirectURI, "unsupported_response_type", "response_type must be code", state)
|
||||
return
|
||||
}
|
||||
|
||||
if q.Get("code_challenge_method") != "S256" {
|
||||
s.authorizeErrorRedirect(w, redirectURI, "invalid_request", "code_challenge_method must be S256", state)
|
||||
return
|
||||
}
|
||||
|
||||
challenge := q.Get("code_challenge")
|
||||
if l := len(challenge); l < 43 || l > 128 {
|
||||
s.authorizeErrorRedirect(w, redirectURI, "invalid_request", "code_challenge is required", state)
|
||||
return
|
||||
}
|
||||
|
||||
scopes, scopeErr := parseAuthorizeScopes(q.Get("scope"))
|
||||
if scopeErr != "" {
|
||||
s.authorizeErrorRedirect(w, redirectURI, "invalid_scope", scopeErr, state)
|
||||
return
|
||||
}
|
||||
|
||||
// Phase 07.13 (D-10/D-11/D-12 in the PHP source comment): a nil ceiling
|
||||
// is a no-op; offline_access always survives truncation because it is
|
||||
// peeled into a boolean below, not a data scope; truncating here (not
|
||||
// rejecting) means an over-broad request silently narrows instead of
|
||||
// failing, except when nothing data-bearing survives.
|
||||
if client.ScopeCeiling != nil {
|
||||
kept := make([]string, 0, len(scopes))
|
||||
for _, sc := range scopes {
|
||||
if sc == "offline_access" || stringSliceContains(client.ScopeCeiling, sc) {
|
||||
kept = append(kept, sc)
|
||||
}
|
||||
}
|
||||
dataScopes := make([]string, 0, len(kept))
|
||||
for _, sc := range kept {
|
||||
if sc != "offline_access" {
|
||||
dataScopes = append(dataScopes, sc)
|
||||
}
|
||||
}
|
||||
if len(dataScopes) == 0 {
|
||||
s.authorizeErrorRedirect(w, redirectURI, "invalid_scope", "requested scope is outside this client's ceiling", state)
|
||||
return
|
||||
}
|
||||
scopes = kept
|
||||
}
|
||||
|
||||
var resource *string
|
||||
if v := q.Get("resource"); v != "" {
|
||||
resource = &v
|
||||
}
|
||||
if resource != nil && *resource != s.opts.Resource {
|
||||
s.authorizeErrorRedirect(w, redirectURI, "invalid_target", "resource does not match this server", state)
|
||||
return
|
||||
}
|
||||
|
||||
// createPendingRequest (PHP OAuthCodeManager): peel offline_access into
|
||||
// a boolean flag; the persisted scopes list never contains it.
|
||||
offline := false
|
||||
dataScopes := make([]string, 0, len(scopes))
|
||||
for _, sc := range scopes {
|
||||
if sc == "offline_access" {
|
||||
offline = true
|
||||
continue
|
||||
}
|
||||
dataScopes = append(dataScopes, sc)
|
||||
}
|
||||
|
||||
requestID, err := s.randomBytes(32)
|
||||
if err != nil {
|
||||
w.WriteHeader(http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
|
||||
rec := &AuthCodeRecord{
|
||||
RequestID: &requestID,
|
||||
ClientID: client.ClientID,
|
||||
RedirectURI: redirectURI,
|
||||
Scopes: dataScopes,
|
||||
CodeChallenge: challenge,
|
||||
CodeChallengeMethod: "S256",
|
||||
Resource: resource,
|
||||
State: state,
|
||||
ExpiresAt: s.now().Add(s.opts.PendingRequestTTL),
|
||||
OfflineAccess: offline,
|
||||
}
|
||||
err = s.backend.WithinTx(ctx, func(tx Tx) error {
|
||||
return tx.CreatePending(ctx, rec)
|
||||
})
|
||||
if err != nil {
|
||||
w.WriteHeader(http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
|
||||
spa := s.opts.Issuer + "/connect?request=" + rfc3986Escape(requestID)
|
||||
// Laravel appends ", private" to every explicit Cache-Control this
|
||||
// endpoint sets (session-cookie default merge); recorded PHP traffic is
|
||||
// "no-store, private", never a bare "no-store" (D-04, live-recorded byte
|
||||
// contract, 08-09-PLAN.md Task 2).
|
||||
w.Header().Set("Cache-Control", "no-store, private")
|
||||
writeRedirectHTML(w, http.StatusFound, spa)
|
||||
}
|
||||
|
||||
// writeAuthorizeLocalError writes the PHP localError() response: a bare
|
||||
// text/plain 400 with no Location and no house envelope
|
||||
// (T-08-OPEN-REDIRECT).
|
||||
func writeAuthorizeLocalError(w http.ResponseWriter, message string) {
|
||||
w.Header().Set("Cache-Control", "no-store, private")
|
||||
w.Header().Set("Content-Type", "text/plain; charset=UTF-8")
|
||||
w.WriteHeader(http.StatusBadRequest)
|
||||
_, _ = w.Write([]byte(message))
|
||||
}
|
||||
|
||||
// authorizeErrorRedirect ports OAuthAuthorizeController::errorRedirect: an
|
||||
// ordered error/error_description/iss[/state] query appended to the
|
||||
// already-trusted redirectURI via RFC 3986 encoding (D-04/Pattern 3).
|
||||
func (s *Server) authorizeErrorRedirect(w http.ResponseWriter, redirectURI, errCode, description string, state *string) {
|
||||
pairs := [][2]string{
|
||||
{"error", errCode},
|
||||
{"error_description", description},
|
||||
{"iss", s.opts.Issuer},
|
||||
}
|
||||
if state != nil {
|
||||
pairs = append(pairs, [2]string{"state", *state})
|
||||
}
|
||||
w.Header().Set("Cache-Control", "no-store, private")
|
||||
writeRedirectHTML(w, http.StatusFound, appendOrderedQuery(redirectURI, pairs))
|
||||
}
|
||||
|
||||
// parseAuthorizeScopes ports OAuthAuthorizeController::parseScopes. An
|
||||
// absent or empty scope query value defaults to ["read"]; otherwise every
|
||||
// whitespace-separated token must be one of authorizeAllowedScopes.
|
||||
func parseAuthorizeScopes(raw string) ([]string, string) {
|
||||
parts := strings.Fields(strings.TrimSpace(raw))
|
||||
if len(parts) == 0 {
|
||||
return []string{"read"}, ""
|
||||
}
|
||||
for _, p := range parts {
|
||||
if !authorizeAllowedScopes[p] {
|
||||
return nil, "scope contains an unsupported value"
|
||||
}
|
||||
}
|
||||
return parts, ""
|
||||
}
|
||||
|
||||
// stringSliceContains reports whether ss contains v.
|
||||
func stringSliceContains(ss []string, v string) bool {
|
||||
for _, s := range ss {
|
||||
if s == v {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// rfc3986Escape percent-encodes s per RFC 3986 (PHP rawurlencode/
|
||||
// PHP_QUERY_RFC3986): unreserved characters A-Z a-z 0-9 - _ . ~ pass
|
||||
// through unescaped; everything else, including the space byte, becomes an
|
||||
// uppercase %XX triplet. This intentionally differs from
|
||||
// net/url.QueryEscape (form-style '+' for space, different reserved-char
|
||||
// handling), which would corrupt the exact PHP redirect-byte contract
|
||||
// (08-RESEARCH.md Pattern 3/Pitfall 5).
|
||||
func rfc3986Escape(s string) string {
|
||||
var b strings.Builder
|
||||
b.Grow(len(s))
|
||||
for i := 0; i < len(s); i++ {
|
||||
c := s[i]
|
||||
if isRFC3986Unreserved(c) {
|
||||
b.WriteByte(c)
|
||||
} else {
|
||||
fmt.Fprintf(&b, "%%%02X", c)
|
||||
}
|
||||
}
|
||||
return b.String()
|
||||
}
|
||||
|
||||
func isRFC3986Unreserved(c byte) bool {
|
||||
return (c >= 'A' && c <= 'Z') || (c >= 'a' && c <= 'z') || (c >= '0' && c <= '9') ||
|
||||
c == '-' || c == '_' || c == '.' || c == '~'
|
||||
}
|
||||
|
||||
// buildOrderedQuery RFC3986-encodes and joins pairs in the given order
|
||||
// (never net/url.Values.Encode, which sorts keys and space-encodes as '+').
|
||||
func buildOrderedQuery(pairs [][2]string) string {
|
||||
parts := make([]string, len(pairs))
|
||||
for i, p := range pairs {
|
||||
parts[i] = rfc3986Escape(p[0]) + "=" + rfc3986Escape(p[1])
|
||||
}
|
||||
return strings.Join(parts, "&")
|
||||
}
|
||||
|
||||
// appendOrderedQuery ports OAuthAuthorizeController::appendQuery: appends an
|
||||
// ordered RFC 3986 query to uri, using '&' when uri already carries a query
|
||||
// string and '?' otherwise.
|
||||
func appendOrderedQuery(uri string, pairs [][2]string) string {
|
||||
sep := "?"
|
||||
if strings.Contains(uri, "?") {
|
||||
sep = "&"
|
||||
}
|
||||
return uri + sep + buildOrderedQuery(pairs)
|
||||
}
|
||||
730
modules/wristband/authorize_test.go
Normal file
730
modules/wristband/authorize_test.go
Normal file
@@ -0,0 +1,730 @@
|
||||
package wristband
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"net/url"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
// insertAuthorizeTestClient inserts an already-usable ClientRecord directly
|
||||
// into backend (bypassing CreateWithCap's cap/sweep policy, which this
|
||||
// plan's tests do not exercise) and returns it.
|
||||
func insertAuthorizeTestClient(backend *memoryBackend, clientID string, redirectURIs []string, ceiling []string) *ClientRecord {
|
||||
backend.mu.Lock()
|
||||
defer backend.mu.Unlock()
|
||||
backend.nextID++
|
||||
rec := &ClientRecord{
|
||||
ID: backend.nextID,
|
||||
ClientID: clientID,
|
||||
ClientName: "Test Client",
|
||||
RedirectURIs: redirectURIs,
|
||||
GrantTypes: []string{"authorization_code", "refresh_token"},
|
||||
TokenEndpointAuthMethod: "client_secret_post",
|
||||
ScopeCeiling: ceiling,
|
||||
CreatedAt: time.Now(),
|
||||
}
|
||||
backend.clients = append(backend.clients, rec)
|
||||
return rec
|
||||
}
|
||||
|
||||
// s256Pair returns a random PKCE verifier and its S256 challenge, matching
|
||||
// the byte transform every Phase 8 PHP/Go PKCE fixture shares.
|
||||
func s256Pair(t *testing.T) (verifier, challenge string) {
|
||||
t.Helper()
|
||||
v, err := randomBase64URL(32)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
return v, s256Challenge(v)
|
||||
}
|
||||
|
||||
// authorizeQuery builds a GET /oauth/mcp/authorize request from an ordered
|
||||
// param map (net/url.Values handles encoding fine for *requests*: only
|
||||
// response Location construction is bound by the RFC3986 ordered-pair
|
||||
// requirement).
|
||||
func authorizeRequest(params map[string]string) *http.Request {
|
||||
q := url.Values{}
|
||||
for k, v := range params {
|
||||
q.Set(k, v)
|
||||
}
|
||||
return httptest.NewRequest(http.MethodGet, "/oauth/mcp/authorize?"+q.Encode(), nil)
|
||||
}
|
||||
|
||||
// queryOf parses the query component of a redirect Location header into a
|
||||
// flat map (every Phase 8 authorize fixture uses at most one value per key).
|
||||
func queryOf(t *testing.T, location string) map[string]string {
|
||||
t.Helper()
|
||||
u, err := url.Parse(location)
|
||||
if err != nil {
|
||||
t.Fatalf("parse Location %q: %v", location, err)
|
||||
}
|
||||
out := map[string]string{}
|
||||
for k, v := range u.Query() {
|
||||
if len(v) > 0 {
|
||||
out[k] = v[0]
|
||||
}
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
// TestPhase8RedAuthorize is the Phase 8 Wave 3 RED anchor (08-03-PLAN.md
|
||||
// Task 1, D-02/D-04/D-05). It drives one valid S256 authorize request
|
||||
// through the real (in-memory-backed) Server.Authorize and asserts the
|
||||
// exact success contract: 302 to <issuer>/connect?request=<opaque> with no
|
||||
// state/code leaked onto our own redirect and Cache-Control: no-store. It
|
||||
// fails with the PHASE8_RED:authorize sentinel while Authorize is the 501
|
||||
// stub; scripts/check-phase8-red.sh verifies this failure is fail-closed.
|
||||
func TestPhase8RedAuthorize(t *testing.T) {
|
||||
backend := newMemoryBackend()
|
||||
srv := newTestServer(backend)
|
||||
insertAuthorizeTestClient(backend, "cli-red", []string{"https://chatgpt.com/connector/oauth/cb"}, nil)
|
||||
_, challenge := s256Pair(t)
|
||||
|
||||
req := authorizeRequest(map[string]string{
|
||||
"client_id": "cli-red",
|
||||
"redirect_uri": "https://chatgpt.com/connector/oauth/cb",
|
||||
"response_type": "code",
|
||||
"code_challenge": challenge,
|
||||
"code_challenge_method": "S256",
|
||||
"scope": "read write",
|
||||
"state": "must-not-appear-on-our-url",
|
||||
"resource": "https://mcp.plytarium.com/mcp",
|
||||
})
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Authorize(rec, req)
|
||||
|
||||
if rec.Code != http.StatusFound {
|
||||
t.Fatalf("PHASE8_RED:authorize: status = %d, want %d (body=%s)", rec.Code, http.StatusFound, rec.Body.String())
|
||||
}
|
||||
loc := rec.Header().Get("Location")
|
||||
if !strings.HasPrefix(loc, "https://plytarium.com/connect?") {
|
||||
t.Fatalf("PHASE8_RED:authorize: Location = %q, want prefix \"https://plytarium.com/connect?\"", loc)
|
||||
}
|
||||
q := queryOf(t, loc)
|
||||
if q["request"] == "" {
|
||||
t.Fatalf("PHASE8_RED:authorize: Location %q missing non-empty request param", loc)
|
||||
}
|
||||
if _, has := q["state"]; has {
|
||||
t.Fatalf("PHASE8_RED:authorize: Location %q leaks state onto our own redirect", loc)
|
||||
}
|
||||
if _, has := q["code"]; has {
|
||||
t.Fatalf("PHASE8_RED:authorize: Location %q leaks a code onto our own redirect", loc)
|
||||
}
|
||||
if cc := rec.Header().Get("Cache-Control"); cc != "no-store, private" {
|
||||
t.Fatalf("PHASE8_RED:authorize: Cache-Control = %q, want \"no-store, private\"", cc)
|
||||
}
|
||||
}
|
||||
|
||||
const authorizeTestRedirect = "https://chatgpt.com/connector/oauth/cb"
|
||||
|
||||
func newAuthorizeTestServer() (*Server, *memoryBackend) {
|
||||
backend := newMemoryBackend()
|
||||
return newTestServer(backend), backend
|
||||
}
|
||||
|
||||
func pendingCount(backend *memoryBackend) int {
|
||||
backend.mu.Lock()
|
||||
defer backend.mu.Unlock()
|
||||
return len(backend.codes)
|
||||
}
|
||||
|
||||
func TestAuthorizeUnknownClientReturnsLocal400NoLocation(t *testing.T) {
|
||||
srv, _ := newAuthorizeTestServer()
|
||||
req := authorizeRequest(map[string]string{
|
||||
"client_id": "does-not-exist",
|
||||
"redirect_uri": authorizeTestRedirect,
|
||||
})
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Authorize(rec, req)
|
||||
|
||||
if rec.Code != http.StatusBadRequest {
|
||||
t.Fatalf("status = %d, want %d (body=%s)", rec.Code, http.StatusBadRequest, rec.Body.String())
|
||||
}
|
||||
if loc := rec.Header().Get("Location"); loc != "" {
|
||||
t.Fatalf("Location = %q, want none", loc)
|
||||
}
|
||||
if ct := rec.Header().Get("Content-Type"); ct != "text/plain; charset=UTF-8" {
|
||||
t.Fatalf("Content-Type = %q, want text/plain; charset=UTF-8", ct)
|
||||
}
|
||||
if body := rec.Body.String(); body != "Unknown client." {
|
||||
t.Fatalf("body = %q, want %q", body, "Unknown client.")
|
||||
}
|
||||
if cc := rec.Header().Get("Cache-Control"); cc != "no-store, private" {
|
||||
t.Fatalf("Cache-Control = %q, want no-store, private", cc)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAuthorizeRevokedClientIsUnknown(t *testing.T) {
|
||||
srv, backend := newAuthorizeTestServer()
|
||||
rec := insertAuthorizeTestClient(backend, "cli-revoked", []string{authorizeTestRedirect}, nil)
|
||||
now := time.Now()
|
||||
rec.RevokedAt = &now
|
||||
|
||||
req := authorizeRequest(map[string]string{
|
||||
"client_id": "cli-revoked",
|
||||
"redirect_uri": authorizeTestRedirect,
|
||||
})
|
||||
w := httptest.NewRecorder()
|
||||
srv.Authorize(w, req)
|
||||
if w.Code != http.StatusBadRequest {
|
||||
t.Fatalf("status = %d, want %d", w.Code, http.StatusBadRequest)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAuthorizeUnregisteredRedirectURIReturnsLocal400NoLocation(t *testing.T) {
|
||||
srv, backend := newAuthorizeTestServer()
|
||||
insertAuthorizeTestClient(backend, "cli-redirect", []string{authorizeTestRedirect}, nil)
|
||||
|
||||
req := authorizeRequest(map[string]string{
|
||||
"client_id": "cli-redirect",
|
||||
"redirect_uri": "https://evil.test/cb",
|
||||
})
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Authorize(rec, req)
|
||||
|
||||
if rec.Code != http.StatusBadRequest {
|
||||
t.Fatalf("status = %d, want %d", rec.Code, http.StatusBadRequest)
|
||||
}
|
||||
if loc := rec.Header().Get("Location"); loc != "" {
|
||||
t.Fatalf("Location = %q, want none", loc)
|
||||
}
|
||||
if strings.Contains(rec.Body.String(), "https://evil.test/cb") {
|
||||
t.Fatal("body echoes the untrusted redirect URI")
|
||||
}
|
||||
}
|
||||
|
||||
func TestAuthorizeTrailingSlashMismatchIsUnregistered(t *testing.T) {
|
||||
srv, backend := newAuthorizeTestServer()
|
||||
insertAuthorizeTestClient(backend, "cli-slash", []string{authorizeTestRedirect}, nil)
|
||||
|
||||
req := authorizeRequest(map[string]string{
|
||||
"client_id": "cli-slash",
|
||||
"redirect_uri": authorizeTestRedirect + "/",
|
||||
})
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Authorize(rec, req)
|
||||
|
||||
if rec.Code != http.StatusBadRequest {
|
||||
t.Fatalf("status = %d, want %d", rec.Code, http.StatusBadRequest)
|
||||
}
|
||||
if loc := rec.Header().Get("Location"); loc != "" {
|
||||
t.Fatalf("Location = %q, want none", loc)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAuthorizeUnsupportedResponseTypeRedirectsWithIssAndState(t *testing.T) {
|
||||
srv, backend := newAuthorizeTestServer()
|
||||
insertAuthorizeTestClient(backend, "cli-resptype", []string{authorizeTestRedirect}, nil)
|
||||
|
||||
req := authorizeRequest(map[string]string{
|
||||
"client_id": "cli-resptype",
|
||||
"redirect_uri": authorizeTestRedirect,
|
||||
"response_type": "token",
|
||||
"code_challenge": strings.Repeat("b", 43),
|
||||
"code_challenge_method": "S256",
|
||||
"state": "iss-state",
|
||||
})
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Authorize(rec, req)
|
||||
|
||||
if rec.Code != http.StatusFound {
|
||||
t.Fatalf("status = %d, want %d", rec.Code, http.StatusFound)
|
||||
}
|
||||
loc := rec.Header().Get("Location")
|
||||
if !strings.HasPrefix(loc, authorizeTestRedirect) {
|
||||
t.Fatalf("Location = %q, want prefix %q", loc, authorizeTestRedirect)
|
||||
}
|
||||
q := queryOf(t, loc)
|
||||
if q["error"] != "unsupported_response_type" {
|
||||
t.Fatalf("error = %q, want unsupported_response_type", q["error"])
|
||||
}
|
||||
if q["iss"] != "https://plytarium.com" {
|
||||
t.Fatalf("iss = %q, want https://plytarium.com", q["iss"])
|
||||
}
|
||||
if q["state"] != "iss-state" {
|
||||
t.Fatalf("state = %q, want iss-state", q["state"])
|
||||
}
|
||||
}
|
||||
|
||||
func TestPKCEChallengeMethodMustBeS256(t *testing.T) {
|
||||
srv, backend := newAuthorizeTestServer()
|
||||
insertAuthorizeTestClient(backend, "cli-plain", []string{authorizeTestRedirect}, nil)
|
||||
state := "state-plain"
|
||||
|
||||
req := authorizeRequest(map[string]string{
|
||||
"client_id": "cli-plain",
|
||||
"redirect_uri": authorizeTestRedirect,
|
||||
"response_type": "code",
|
||||
"code_challenge": strings.Repeat("a", 43),
|
||||
"code_challenge_method": "plain",
|
||||
"state": state,
|
||||
})
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Authorize(rec, req)
|
||||
|
||||
if rec.Code != http.StatusFound {
|
||||
t.Fatalf("status = %d, want %d", rec.Code, http.StatusFound)
|
||||
}
|
||||
q := queryOf(t, rec.Header().Get("Location"))
|
||||
if q["error"] != "invalid_request" {
|
||||
t.Fatalf("error = %q, want invalid_request", q["error"])
|
||||
}
|
||||
if q["state"] != state {
|
||||
t.Fatalf("state = %q, want %q", q["state"], state)
|
||||
}
|
||||
if q["iss"] != "https://plytarium.com" {
|
||||
t.Fatalf("iss = %q, want https://plytarium.com", q["iss"])
|
||||
}
|
||||
if cc := rec.Header().Get("Cache-Control"); cc != "no-store, private" {
|
||||
t.Fatalf("Cache-Control = %q, want no-store, private", cc)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPKCEChallengeLengthBounds(t *testing.T) {
|
||||
cases := []struct {
|
||||
name string
|
||||
challenge string
|
||||
}{
|
||||
{"too-short", strings.Repeat("a", 42)},
|
||||
{"too-long", strings.Repeat("a", 129)},
|
||||
{"empty", ""},
|
||||
}
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
srv, backend := newAuthorizeTestServer()
|
||||
insertAuthorizeTestClient(backend, "cli-len-"+tc.name, []string{authorizeTestRedirect}, nil)
|
||||
req := authorizeRequest(map[string]string{
|
||||
"client_id": "cli-len-" + tc.name,
|
||||
"redirect_uri": authorizeTestRedirect,
|
||||
"response_type": "code",
|
||||
"code_challenge": tc.challenge,
|
||||
"code_challenge_method": "S256",
|
||||
})
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Authorize(rec, req)
|
||||
if rec.Code != http.StatusFound {
|
||||
t.Fatalf("status = %d, want %d", rec.Code, http.StatusFound)
|
||||
}
|
||||
q := queryOf(t, rec.Header().Get("Location"))
|
||||
if q["error"] != "invalid_request" {
|
||||
t.Fatalf("error = %q, want invalid_request", q["error"])
|
||||
}
|
||||
if q["error_description"] != "code_challenge is required" {
|
||||
t.Fatalf("error_description = %q, want %q", q["error_description"], "code_challenge is required")
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
func TestAuthorizeValidRequestRedirectsToConnectWithOpaqueHandleOnly(t *testing.T) {
|
||||
srv, backend := newAuthorizeTestServer()
|
||||
insertAuthorizeTestClient(backend, "cli-valid", []string{authorizeTestRedirect}, nil)
|
||||
_, challenge := s256Pair(t)
|
||||
state := "must-not-appear-on-our-url"
|
||||
|
||||
req := authorizeRequest(map[string]string{
|
||||
"client_id": "cli-valid",
|
||||
"redirect_uri": authorizeTestRedirect,
|
||||
"response_type": "code",
|
||||
"code_challenge": challenge,
|
||||
"code_challenge_method": "S256",
|
||||
"scope": "read write",
|
||||
"state": state,
|
||||
"resource": "https://mcp.plytarium.com/mcp",
|
||||
})
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Authorize(rec, req)
|
||||
|
||||
if rec.Code != http.StatusFound {
|
||||
t.Fatalf("status = %d, want %d (body=%s)", rec.Code, http.StatusFound, rec.Body.String())
|
||||
}
|
||||
loc := rec.Header().Get("Location")
|
||||
if !strings.HasPrefix(loc, "https://plytarium.com/connect?") {
|
||||
t.Fatalf("Location = %q, want prefix https://plytarium.com/connect?", loc)
|
||||
}
|
||||
q := queryOf(t, loc)
|
||||
if q["request"] == "" {
|
||||
t.Fatal("Location missing non-empty request param")
|
||||
}
|
||||
if _, has := q["code"]; has {
|
||||
t.Fatal("Location leaks a code")
|
||||
}
|
||||
if _, has := q["state"]; has {
|
||||
t.Fatal("Location leaks state")
|
||||
}
|
||||
if _, has := q["client_secret"]; has {
|
||||
t.Fatal("Location leaks client_secret")
|
||||
}
|
||||
if cc := rec.Header().Get("Cache-Control"); cc != "no-store, private" {
|
||||
t.Fatalf("Cache-Control = %q, want no-store, private", cc)
|
||||
}
|
||||
|
||||
backend.mu.Lock()
|
||||
defer backend.mu.Unlock()
|
||||
if len(backend.codes) != 1 {
|
||||
t.Fatalf("pending rows = %d, want 1", len(backend.codes))
|
||||
}
|
||||
pending := backend.codes[0]
|
||||
if pending.RequestID == nil || *pending.RequestID != q["request"] {
|
||||
t.Fatalf("pending.RequestID = %v, want %q", pending.RequestID, q["request"])
|
||||
}
|
||||
if pending.CodeHash != nil {
|
||||
t.Fatal("pending row already has a code hash")
|
||||
}
|
||||
if pending.UserID != nil {
|
||||
t.Fatal("pending row already has a user id")
|
||||
}
|
||||
if pending.State == nil || *pending.State != state {
|
||||
t.Fatalf("pending.State = %v, want %q", pending.State, state)
|
||||
}
|
||||
if pending.ExpiresAt.Before(time.Now().Add(500*time.Second)) || pending.ExpiresAt.After(time.Now().Add(700*time.Second)) {
|
||||
t.Fatalf("pending.ExpiresAt = %v, want ~600s from now", pending.ExpiresAt)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAuthorizeIssIsPresentOnEveryErrorRedirect(t *testing.T) {
|
||||
srv, backend := newAuthorizeTestServer()
|
||||
insertAuthorizeTestClient(backend, "cli-iss", []string{authorizeTestRedirect}, nil)
|
||||
|
||||
req := authorizeRequest(map[string]string{
|
||||
"client_id": "cli-iss",
|
||||
"redirect_uri": authorizeTestRedirect,
|
||||
"response_type": "token",
|
||||
"code_challenge": strings.Repeat("b", 43),
|
||||
"code_challenge_method": "S256",
|
||||
"state": "iss-state",
|
||||
})
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Authorize(rec, req)
|
||||
q := queryOf(t, rec.Header().Get("Location"))
|
||||
if q["iss"] != "https://plytarium.com" {
|
||||
t.Fatalf("iss = %q, want https://plytarium.com", q["iss"])
|
||||
}
|
||||
if q["state"] != "iss-state" {
|
||||
t.Fatalf("state = %q, want iss-state", q["state"])
|
||||
}
|
||||
if q["error"] != "unsupported_response_type" {
|
||||
t.Fatalf("error = %q, want unsupported_response_type", q["error"])
|
||||
}
|
||||
}
|
||||
|
||||
func TestAuthorizeScopeDefaultsToRead(t *testing.T) {
|
||||
srv, backend := newAuthorizeTestServer()
|
||||
insertAuthorizeTestClient(backend, "cli-default-scope", []string{authorizeTestRedirect}, nil)
|
||||
_, challenge := s256Pair(t)
|
||||
|
||||
req := authorizeRequest(map[string]string{
|
||||
"client_id": "cli-default-scope",
|
||||
"redirect_uri": authorizeTestRedirect,
|
||||
"response_type": "code",
|
||||
"code_challenge": challenge,
|
||||
"code_challenge_method": "S256",
|
||||
})
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Authorize(rec, req)
|
||||
if rec.Code != http.StatusFound {
|
||||
t.Fatalf("status = %d, want %d", rec.Code, http.StatusFound)
|
||||
}
|
||||
|
||||
backend.mu.Lock()
|
||||
defer backend.mu.Unlock()
|
||||
if len(backend.codes) != 1 {
|
||||
t.Fatalf("pending rows = %d, want 1", len(backend.codes))
|
||||
}
|
||||
if got := backend.codes[0].Scopes; len(got) != 1 || got[0] != "read" {
|
||||
t.Fatalf("scopes = %v, want [read]", got)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAuthorizeScopeInvalidValueRedirectsInvalidScope(t *testing.T) {
|
||||
srv, backend := newAuthorizeTestServer()
|
||||
insertAuthorizeTestClient(backend, "cli-bad-scope", []string{authorizeTestRedirect}, nil)
|
||||
_, challenge := s256Pair(t)
|
||||
before := pendingCount(backend)
|
||||
|
||||
req := authorizeRequest(map[string]string{
|
||||
"client_id": "cli-bad-scope",
|
||||
"redirect_uri": authorizeTestRedirect,
|
||||
"response_type": "code",
|
||||
"code_challenge": challenge,
|
||||
"code_challenge_method": "S256",
|
||||
"scope": "read bogus",
|
||||
})
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Authorize(rec, req)
|
||||
if rec.Code != http.StatusFound {
|
||||
t.Fatalf("status = %d, want %d", rec.Code, http.StatusFound)
|
||||
}
|
||||
q := queryOf(t, rec.Header().Get("Location"))
|
||||
if q["error"] != "invalid_scope" {
|
||||
t.Fatalf("error = %q, want invalid_scope", q["error"])
|
||||
}
|
||||
if q["error_description"] != "scope contains an unsupported value" {
|
||||
t.Fatalf("error_description = %q", q["error_description"])
|
||||
}
|
||||
if got := pendingCount(backend); got != before {
|
||||
t.Fatalf("pending rows = %d, want unchanged %d", got, before)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAuthorizeNullCeilingPreservesAiScope(t *testing.T) {
|
||||
srv, backend := newAuthorizeTestServer()
|
||||
insertAuthorizeTestClient(backend, "cli-null-ceiling", []string{authorizeTestRedirect}, nil)
|
||||
_, challenge := s256Pair(t)
|
||||
|
||||
req := authorizeRequest(map[string]string{
|
||||
"client_id": "cli-null-ceiling",
|
||||
"redirect_uri": authorizeTestRedirect,
|
||||
"response_type": "code",
|
||||
"code_challenge": challenge,
|
||||
"code_challenge_method": "S256",
|
||||
"scope": "read write ai",
|
||||
})
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Authorize(rec, req)
|
||||
if rec.Code != http.StatusFound {
|
||||
t.Fatalf("status = %d, body=%s", rec.Code, rec.Body.String())
|
||||
}
|
||||
|
||||
backend.mu.Lock()
|
||||
defer backend.mu.Unlock()
|
||||
got := backend.codes[len(backend.codes)-1].Scopes
|
||||
want := []string{"read", "write", "ai"}
|
||||
if len(got) != len(want) {
|
||||
t.Fatalf("scopes = %v, want %v", got, want)
|
||||
}
|
||||
for i := range want {
|
||||
if got[i] != want[i] {
|
||||
t.Fatalf("scopes = %v, want %v", got, want)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestAuthorizeCeilingTruncatesAiWithoutError(t *testing.T) {
|
||||
srv, backend := newAuthorizeTestServer()
|
||||
insertAuthorizeTestClient(backend, "cli-ceiling-trunc", []string{authorizeTestRedirect}, []string{"read", "write"})
|
||||
_, challenge := s256Pair(t)
|
||||
|
||||
req := authorizeRequest(map[string]string{
|
||||
"client_id": "cli-ceiling-trunc",
|
||||
"redirect_uri": authorizeTestRedirect,
|
||||
"response_type": "code",
|
||||
"code_challenge": challenge,
|
||||
"code_challenge_method": "S256",
|
||||
"scope": "read write ai",
|
||||
})
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Authorize(rec, req)
|
||||
if rec.Code != http.StatusFound {
|
||||
t.Fatalf("status = %d, body=%s", rec.Code, rec.Body.String())
|
||||
}
|
||||
|
||||
backend.mu.Lock()
|
||||
defer backend.mu.Unlock()
|
||||
got := backend.codes[len(backend.codes)-1].Scopes
|
||||
if stringSliceContains(got, "ai") {
|
||||
t.Fatalf("scopes = %v, still contains ai", got)
|
||||
}
|
||||
want := []string{"read", "write"}
|
||||
if len(got) != len(want) || got[0] != want[0] || got[1] != want[1] {
|
||||
t.Fatalf("scopes = %v, want %v", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAuthorizeCeilingPreservesOfflineAccessFlag(t *testing.T) {
|
||||
srv, backend := newAuthorizeTestServer()
|
||||
insertAuthorizeTestClient(backend, "cli-ceiling-offline", []string{authorizeTestRedirect}, []string{"read", "write"})
|
||||
_, challenge := s256Pair(t)
|
||||
|
||||
req := authorizeRequest(map[string]string{
|
||||
"client_id": "cli-ceiling-offline",
|
||||
"redirect_uri": authorizeTestRedirect,
|
||||
"response_type": "code",
|
||||
"code_challenge": challenge,
|
||||
"code_challenge_method": "S256",
|
||||
"scope": "read write ai offline_access",
|
||||
})
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Authorize(rec, req)
|
||||
if rec.Code != http.StatusFound {
|
||||
t.Fatalf("status = %d, body=%s", rec.Code, rec.Body.String())
|
||||
}
|
||||
|
||||
backend.mu.Lock()
|
||||
defer backend.mu.Unlock()
|
||||
pending := backend.codes[len(backend.codes)-1]
|
||||
want := []string{"read", "write"}
|
||||
if len(pending.Scopes) != len(want) || pending.Scopes[0] != want[0] || pending.Scopes[1] != want[1] {
|
||||
t.Fatalf("scopes = %v, want %v", pending.Scopes, want)
|
||||
}
|
||||
if !pending.OfflineAccess {
|
||||
t.Fatal("OfflineAccess = false, want true")
|
||||
}
|
||||
}
|
||||
|
||||
func TestAuthorizeCeilingRejectsAiOnlyAsInvalidScopeOnTrustedRedirect(t *testing.T) {
|
||||
srv, backend := newAuthorizeTestServer()
|
||||
insertAuthorizeTestClient(backend, "cli-ceiling-ai-only", []string{authorizeTestRedirect}, []string{"read", "write"})
|
||||
_, challenge := s256Pair(t)
|
||||
state := "state-ai-only"
|
||||
before := pendingCount(backend)
|
||||
|
||||
req := authorizeRequest(map[string]string{
|
||||
"client_id": "cli-ceiling-ai-only",
|
||||
"redirect_uri": authorizeTestRedirect,
|
||||
"response_type": "code",
|
||||
"code_challenge": challenge,
|
||||
"code_challenge_method": "S256",
|
||||
"scope": "ai",
|
||||
"state": state,
|
||||
})
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Authorize(rec, req)
|
||||
if rec.Code != http.StatusFound {
|
||||
t.Fatalf("status = %d, body=%s", rec.Code, rec.Body.String())
|
||||
}
|
||||
loc := rec.Header().Get("Location")
|
||||
if !strings.HasPrefix(loc, authorizeTestRedirect) {
|
||||
t.Fatalf("Location = %q, want prefix %q", loc, authorizeTestRedirect)
|
||||
}
|
||||
q := queryOf(t, loc)
|
||||
if q["error"] != "invalid_scope" {
|
||||
t.Fatalf("error = %q, want invalid_scope", q["error"])
|
||||
}
|
||||
if q["state"] != state {
|
||||
t.Fatalf("state = %q, want %q", q["state"], state)
|
||||
}
|
||||
if q["iss"] != "https://plytarium.com" {
|
||||
t.Fatalf("iss = %q, want https://plytarium.com", q["iss"])
|
||||
}
|
||||
if got := pendingCount(backend); got != before {
|
||||
t.Fatalf("pending rows = %d, want unchanged %d", got, before)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAuthorizeCeilingRejectsAiPlusOfflineAccessAsInvalidScope(t *testing.T) {
|
||||
srv, backend := newAuthorizeTestServer()
|
||||
insertAuthorizeTestClient(backend, "cli-ceiling-ai-offline", []string{authorizeTestRedirect}, []string{"read", "write"})
|
||||
_, challenge := s256Pair(t)
|
||||
state := "state-ai-offline"
|
||||
before := pendingCount(backend)
|
||||
|
||||
req := authorizeRequest(map[string]string{
|
||||
"client_id": "cli-ceiling-ai-offline",
|
||||
"redirect_uri": authorizeTestRedirect,
|
||||
"response_type": "code",
|
||||
"code_challenge": challenge,
|
||||
"code_challenge_method": "S256",
|
||||
"scope": "ai offline_access",
|
||||
"state": state,
|
||||
})
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Authorize(rec, req)
|
||||
if rec.Code != http.StatusFound {
|
||||
t.Fatalf("status = %d, body=%s", rec.Code, rec.Body.String())
|
||||
}
|
||||
q := queryOf(t, rec.Header().Get("Location"))
|
||||
if q["error"] != "invalid_scope" {
|
||||
t.Fatalf("error = %q, want invalid_scope", q["error"])
|
||||
}
|
||||
if q["state"] != state {
|
||||
t.Fatalf("state = %q, want %q", q["state"], state)
|
||||
}
|
||||
if got := pendingCount(backend); got != before {
|
||||
t.Fatalf("pending rows = %d, want unchanged %d", got, before)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAuthorizeResourceMismatchRedirectsInvalidTarget(t *testing.T) {
|
||||
srv, backend := newAuthorizeTestServer()
|
||||
insertAuthorizeTestClient(backend, "cli-resource", []string{authorizeTestRedirect}, nil)
|
||||
_, challenge := s256Pair(t)
|
||||
before := pendingCount(backend)
|
||||
|
||||
req := authorizeRequest(map[string]string{
|
||||
"client_id": "cli-resource",
|
||||
"redirect_uri": authorizeTestRedirect,
|
||||
"response_type": "code",
|
||||
"code_challenge": challenge,
|
||||
"code_challenge_method": "S256",
|
||||
"resource": "https://wrong.example.test/mcp",
|
||||
})
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Authorize(rec, req)
|
||||
if rec.Code != http.StatusFound {
|
||||
t.Fatalf("status = %d, body=%s", rec.Code, rec.Body.String())
|
||||
}
|
||||
q := queryOf(t, rec.Header().Get("Location"))
|
||||
if q["error"] != "invalid_target" {
|
||||
t.Fatalf("error = %q, want invalid_target", q["error"])
|
||||
}
|
||||
if q["error_description"] != "resource does not match this server" {
|
||||
t.Fatalf("error_description = %q", q["error_description"])
|
||||
}
|
||||
if got := pendingCount(backend); got != before {
|
||||
t.Fatalf("pending rows = %d, want unchanged %d", got, before)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAuthorizeResourceOmittedIsAccepted(t *testing.T) {
|
||||
srv, backend := newAuthorizeTestServer()
|
||||
insertAuthorizeTestClient(backend, "cli-resource-omitted", []string{authorizeTestRedirect}, nil)
|
||||
_, challenge := s256Pair(t)
|
||||
|
||||
req := authorizeRequest(map[string]string{
|
||||
"client_id": "cli-resource-omitted",
|
||||
"redirect_uri": authorizeTestRedirect,
|
||||
"response_type": "code",
|
||||
"code_challenge": challenge,
|
||||
"code_challenge_method": "S256",
|
||||
})
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Authorize(rec, req)
|
||||
if rec.Code != http.StatusFound {
|
||||
t.Fatalf("status = %d, body=%s", rec.Code, rec.Body.String())
|
||||
}
|
||||
}
|
||||
|
||||
// TestOrderedRedirectQueryEncodingIsRFC3986 proves the redirect query
|
||||
// encoder differs from net/url.Values.Encode exactly where PHP
|
||||
// http_build_query(..., PHP_QUERY_RFC3986) does: spaces become %20 (not
|
||||
// '+'), reserved characters are percent-encoded, and key order is
|
||||
// preserved (not sorted) — Pitfall 5.
|
||||
func TestOrderedRedirectQueryEncodingIsRFC3986(t *testing.T) {
|
||||
got := buildOrderedQuery([][2]string{
|
||||
{"error", "invalid_scope"},
|
||||
{"error_description", "requested scope is outside this client's ceiling"},
|
||||
{"iss", "https://plytarium.com"},
|
||||
{"state", "a b&c=d"},
|
||||
})
|
||||
want := "error=invalid_scope&error_description=requested%20scope%20is%20outside%20this%20client%27s%20ceiling&iss=https%3A%2F%2Fplytarium.com&state=a%20b%26c%3Dd"
|
||||
if got != want {
|
||||
t.Fatalf("got: %s\nwant: %s", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestOrderedRedirectAppendsToExistingQuery(t *testing.T) {
|
||||
got := appendOrderedQuery("https://client.example.test/cb?already=here", [][2]string{
|
||||
{"error", "invalid_request"},
|
||||
{"iss", "https://plytarium.com"},
|
||||
})
|
||||
want := "https://client.example.test/cb?already=here&error=invalid_request&iss=https%3A%2F%2Fplytarium.com"
|
||||
if got != want {
|
||||
t.Fatalf("got: %s\nwant: %s", got, want)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAuthorizeBackendUnavailableIsOpaque500(t *testing.T) {
|
||||
opts := DefaultOptions()
|
||||
opts.Issuer = "https://plytarium.com"
|
||||
srv := NewServer(opts)
|
||||
req := authorizeRequest(map[string]string{"client_id": "anything"})
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Authorize(rec, req)
|
||||
if rec.Code != http.StatusInternalServerError {
|
||||
t.Fatalf("status = %d, want %d", rec.Code, http.StatusInternalServerError)
|
||||
}
|
||||
}
|
||||
40
modules/wristband/client_issue.go
Normal file
40
modules/wristband/client_issue.go
Normal file
@@ -0,0 +1,40 @@
|
||||
// Shared client-issuing primitives used by both RFC 7591 registration
|
||||
// (register.go) and the app-owned fonoteka:oauth-client operator command
|
||||
// (08-CONTEXT.md D-19; T-08-SECRET-TIMING: "Shared hash/validation path and
|
||||
// one-time secret"). Exporting these rather than letting the command
|
||||
// re-derive its own random-id/hash/redirect-URI-validation logic keeps
|
||||
// exactly one code path responsible for how an OAuth client secret is
|
||||
// generated and hashed.
|
||||
package wristband
|
||||
|
||||
// IssueClientCredentials mints a random opaque client_id (16 raw bytes,
|
||||
// base64url) and, for every token_endpoint_auth_method other than "none", a
|
||||
// client_secret (32 raw bytes) plus its sha256 hex hash -- the identical
|
||||
// fixed transform RFC 7591 registration uses (D-04). secret is "" and
|
||||
// secretHash is nil for a public ("none") client. The raw secret is
|
||||
// returned exactly once; only secretHash is meant to be persisted.
|
||||
func IssueClientCredentials(authMethod string) (clientID, secret string, secretHash *string, err error) {
|
||||
clientID, err = randomBase64URL(16)
|
||||
if err != nil {
|
||||
return "", "", nil, err
|
||||
}
|
||||
if authMethod == "none" {
|
||||
return clientID, "", nil, nil
|
||||
}
|
||||
secret, err = randomBase64URL(32)
|
||||
if err != nil {
|
||||
return "", "", nil, err
|
||||
}
|
||||
h := sha256Hex(secret)
|
||||
return clientID, secret, &h, nil
|
||||
}
|
||||
|
||||
// RejectRedirectURI is the exported form of the redirect-URI validation
|
||||
// RFC 7591 registration already enforces (PHP OAuthClient::rejectRedirectUri):
|
||||
// at most 512 characters, a valid URL with scheme+host, https:// or loopback
|
||||
// http://127.0.0.1 / http://localhost. It returns "" when uri is accepted,
|
||||
// or a human-readable rejection reason otherwise. The operator command
|
||||
// shares this exact rule with DCR rather than re-deriving it (D-19).
|
||||
func RejectRedirectURI(uri string) string {
|
||||
return rejectRedirectURI(uri)
|
||||
}
|
||||
167
modules/wristband/consent.go
Normal file
167
modules/wristband/consent.go
Normal file
@@ -0,0 +1,167 @@
|
||||
// JWT-group consent operations for MCP OAuth, ported from PHP
|
||||
// OAuthConsentController::show/store/deny and OAuthCodeManager::issueCode
|
||||
// (08-CONTEXT.md D-08; canonical PHP source: OAuthConsentController.php,
|
||||
// OAuthCodeManager.php).
|
||||
//
|
||||
// Unlike Metadata/Authorize/Token, consent has no HTTP handler here: the
|
||||
// browser-facing JWT surface, request validation, MINTABLE_SCOPES
|
||||
// intersection and ActiveCollectionResolver pin are app-owned (D-08). This
|
||||
// file exposes the protocol-level operations the app's consent controller
|
||||
// calls instead of touching wristband.AuthCodeStore/ClientStore rows
|
||||
// directly: resolving an owner-bound pending request, issuing its code, and
|
||||
// denying it. Every one of "missing", "foreign", "already used", "expired",
|
||||
// and "already issued" collapses to the identical ErrPendingNotFound so a
|
||||
// caller cannot distinguish them (T-08-CROSS-USER/T-08-REQUEST-LEAK).
|
||||
package wristband
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"net/url"
|
||||
"time"
|
||||
)
|
||||
|
||||
// ErrPendingNotFound is returned by PendingRequest, IssueCode and
|
||||
// DenyPending when requestID does not resolve to a live pending row owned
|
||||
// by userID: missing, foreign, already issued (non-nil CodeHash), used,
|
||||
// expired, or bound to a revoked client all collapse to this single error
|
||||
// (PHP OAuthConsentController::pendingFor).
|
||||
var ErrPendingNotFound = errors.New("wristband: pending request not found")
|
||||
|
||||
// ErrNoGrantableScopes is returned by IssueCode when grantedScopes is empty
|
||||
// (PHP: "No grantable scopes", 422).
|
||||
var ErrNoGrantableScopes = errors.New("wristband: no grantable scopes")
|
||||
|
||||
// PendingRequestView is consent's GET show payload ingredients. It
|
||||
// deliberately omits collection_name (server-resolved via the app's own
|
||||
// ActiveCollectionResolver, D-08) and any collection id.
|
||||
type PendingRequestView struct {
|
||||
ClientName string
|
||||
RedirectHost string
|
||||
ScopesRequested []string // the pending row's own (already ceiling-truncated) scopes, unfiltered by mintability
|
||||
ExpiresAt time.Time
|
||||
}
|
||||
|
||||
// PendingRequest resolves requestID for GET .../oauth/request/{request_id}
|
||||
// (D-08). userID is the acting JWT principal.
|
||||
func (s *Server) PendingRequest(ctx context.Context, requestID string, userID uint) (PendingRequestView, error) {
|
||||
pending, client, err := s.lookupOwnedPending(ctx, requestID, userID)
|
||||
if err != nil {
|
||||
return PendingRequestView{}, err
|
||||
}
|
||||
host := ""
|
||||
if u, perr := url.Parse(pending.RedirectURI); perr == nil {
|
||||
host = u.Hostname()
|
||||
}
|
||||
return PendingRequestView{
|
||||
ClientName: client.ClientName,
|
||||
RedirectHost: host,
|
||||
ScopesRequested: pending.Scopes,
|
||||
ExpiresAt: pending.ExpiresAt,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// IssueCode grants consent (D-08). It trusts the caller's already-computed
|
||||
// grantedScopes (submitted ∩ pending's own scopes ∩ the app's mintable set)
|
||||
// and collectionIDs (server-resolved, never client-supplied): it does not
|
||||
// recompute either intersection, matching the PHP boundary where
|
||||
// OAuthConsentController::store owns the scope/collection policy and
|
||||
// OAuthCodeManager::issueCode is a dumb persist-and-redirect step. It
|
||||
// persists the granted scopes/collection ids onto the pending row alongside
|
||||
// a fresh one-time code and a fresh CodeTTL expiry, consumes the request
|
||||
// handle, stamps the client's ConsentedAt once, and returns the ordered
|
||||
// redirect_to URL with `code`, `iss`, and the pending's own `state`.
|
||||
func (s *Server) IssueCode(ctx context.Context, requestID string, userID uint, grantedScopes []string, collectionIDs []uint) (string, error) {
|
||||
if len(grantedScopes) == 0 {
|
||||
return "", ErrNoGrantableScopes
|
||||
}
|
||||
pending, client, err := s.lookupOwnedPending(ctx, requestID, userID)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
code, err := s.randomBytes(32)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
expiresAt := s.now().Add(s.opts.CodeTTL)
|
||||
err = s.backend.WithinTx(ctx, func(tx Tx) error {
|
||||
if err := tx.MarkIssued(ctx, pending.ID, sha256Hex(code), userID, grantedScopes, collectionIDs, expiresAt); err != nil {
|
||||
return err
|
||||
}
|
||||
return tx.MarkConsented(ctx, client.ClientID)
|
||||
})
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
pairs := [][2]string{{"code", code}, {"iss", s.opts.Issuer}}
|
||||
if pending.State != nil && *pending.State != "" {
|
||||
pairs = append(pairs, [2]string{"state", *pending.State})
|
||||
}
|
||||
return appendOrderedQuery(pending.RedirectURI, pairs), nil
|
||||
}
|
||||
|
||||
// DenyPending consumes requestID without issuing a code (D-08) and returns
|
||||
// the ordered redirect_to URL with error=access_denied, iss, and the
|
||||
// pending's own state (PHP OAuthConsentController::deny).
|
||||
func (s *Server) DenyPending(ctx context.Context, requestID string, userID uint) (string, error) {
|
||||
pending, _, err := s.lookupOwnedPending(ctx, requestID, userID)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
err = s.backend.WithinTx(ctx, func(tx Tx) error {
|
||||
return tx.MarkUsed(ctx, pending.ID)
|
||||
})
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
pairs := [][2]string{{"error", "access_denied"}, {"iss", s.opts.Issuer}}
|
||||
if pending.State != nil && *pending.State != "" {
|
||||
pairs = append(pairs, [2]string{"state", *pending.State})
|
||||
}
|
||||
return appendOrderedQuery(pending.RedirectURI, pairs), nil
|
||||
}
|
||||
|
||||
// lookupOwnedPending ports PHP OAuthConsentController::pendingFor exactly:
|
||||
// missing, used, expired, already-issued (non-nil CodeHash), foreign-owner,
|
||||
// or bound to an unusable/revoked client all collapse to ErrPendingNotFound
|
||||
// (T-08-CROSS-USER/T-08-REQUEST-LEAK). A nil pending.UserID (not yet
|
||||
// consented by anyone) is owned by every caller, matching PHP's
|
||||
// `$pending->user_id !== null && ... !== $user->id` guard.
|
||||
func (s *Server) lookupOwnedPending(ctx context.Context, requestID string, userID uint) (*AuthCodeRecord, *ClientRecord, error) {
|
||||
if requestID == "" {
|
||||
return nil, nil, ErrPendingNotFound
|
||||
}
|
||||
if s.backend == nil {
|
||||
return nil, nil, errors.New("wristband: backend is not configured")
|
||||
}
|
||||
var pending *AuthCodeRecord
|
||||
var client *ClientRecord
|
||||
err := s.backend.WithinTx(ctx, func(tx Tx) error {
|
||||
rec, err := tx.ByRequestID(ctx, requestID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
pending = rec
|
||||
if rec == nil {
|
||||
return nil
|
||||
}
|
||||
c, err := tx.ByClientID(ctx, rec.ClientID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
client = c
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
return nil, nil, err
|
||||
}
|
||||
if pending == nil ||
|
||||
pending.UsedAt != nil ||
|
||||
!pending.ExpiresAt.After(s.now()) ||
|
||||
pending.CodeHash != nil ||
|
||||
(pending.UserID != nil && *pending.UserID != userID) ||
|
||||
client == nil || client.RevokedAt != nil {
|
||||
return nil, nil, ErrPendingNotFound
|
||||
}
|
||||
return pending, client, nil
|
||||
}
|
||||
139
modules/wristband/consent_test.go
Normal file
139
modules/wristband/consent_test.go
Normal file
@@ -0,0 +1,139 @@
|
||||
package wristband
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
// newConsentFixture builds a Server over the in-memory backend with one
|
||||
// usable client and one pending request owned by no one yet (userID 0
|
||||
// means "unconsented", matching Authorize's freshly created row).
|
||||
func newConsentFixture(t *testing.T) (*Server, string) {
|
||||
t.Helper()
|
||||
s := NewServer(DefaultOptions())
|
||||
s.opts.Issuer = "https://plytarium-consent-test.example"
|
||||
b := newMemoryBackend()
|
||||
s.SetBackend(b)
|
||||
ctx := context.Background()
|
||||
|
||||
err := b.WithinTx(ctx, func(tx Tx) error {
|
||||
client := &ClientRecord{
|
||||
ClientID: "cli-consent-test",
|
||||
ClientName: "Test Client",
|
||||
RedirectURIs: []string{"https://client.example.test/cb"},
|
||||
TokenEndpointAuthMethod: "none",
|
||||
}
|
||||
if err := tx.CreateWithCap(ctx, client, 200); err != nil {
|
||||
return err
|
||||
}
|
||||
state := "consent-test-state"
|
||||
pending := &AuthCodeRecord{
|
||||
ClientID: client.ClientID,
|
||||
RedirectURI: "https://client.example.test/cb",
|
||||
Scopes: []string{"read", "write"},
|
||||
CodeChallenge: strings.Repeat("a", 43),
|
||||
CodeChallengeMethod: "S256",
|
||||
State: &state,
|
||||
ExpiresAt: s.now().Add(s.opts.PendingRequestTTL),
|
||||
}
|
||||
requestID, err := s.randomBytes(32)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
pending.RequestID = &requestID
|
||||
return tx.CreatePending(ctx, pending)
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
|
||||
var requestID string
|
||||
for _, c := range b.codes {
|
||||
if c.RequestID != nil {
|
||||
requestID = *c.RequestID
|
||||
}
|
||||
}
|
||||
return s, requestID
|
||||
}
|
||||
|
||||
func TestPendingRequestReturnsClientAndScopes(t *testing.T) {
|
||||
s, requestID := newConsentFixture(t)
|
||||
view, err := s.PendingRequest(context.Background(), requestID, 1)
|
||||
if err != nil {
|
||||
t.Fatalf("PendingRequest: %v", err)
|
||||
}
|
||||
if view.ClientName != "Test Client" {
|
||||
t.Fatalf("ClientName = %q, want %q", view.ClientName, "Test Client")
|
||||
}
|
||||
if view.RedirectHost != "client.example.test" {
|
||||
t.Fatalf("RedirectHost = %q, want client.example.test", view.RedirectHost)
|
||||
}
|
||||
if len(view.ScopesRequested) != 2 || view.ScopesRequested[0] != "read" || view.ScopesRequested[1] != "write" {
|
||||
t.Fatalf("ScopesRequested = %v, want [read write]", view.ScopesRequested)
|
||||
}
|
||||
}
|
||||
|
||||
func TestPendingRequestMissingHandleIsNotFound(t *testing.T) {
|
||||
s, _ := newConsentFixture(t)
|
||||
_, err := s.PendingRequest(context.Background(), "does-not-exist", 1)
|
||||
if !errors.Is(err, ErrPendingNotFound) {
|
||||
t.Fatalf("err = %v, want ErrPendingNotFound", err)
|
||||
}
|
||||
}
|
||||
|
||||
// TestPendingRequestForeignOwnerIsNotFound proves T-08-CROSS-USER: once a
|
||||
// pending row is bound to one user (by a prior consent attempt), a
|
||||
// different user's lookup is indistinguishable from missing.
|
||||
func TestPendingRequestForeignOwnerIsNotFound(t *testing.T) {
|
||||
s, requestID := newConsentFixture(t)
|
||||
ctx := context.Background()
|
||||
if _, err := s.IssueCode(ctx, requestID, 1, []string{"read"}, []uint{9}); err != nil {
|
||||
t.Fatalf("IssueCode: %v", err)
|
||||
}
|
||||
// The code is now issued (CodeHash set); a second lookup by anyone,
|
||||
// including the original owner, must miss (single-use).
|
||||
if _, err := s.PendingRequest(ctx, requestID, 1); !errors.Is(err, ErrPendingNotFound) {
|
||||
t.Fatalf("err = %v, want ErrPendingNotFound after issuance", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestIssueCodeGrantsOnlySubmittedScopesAndReturnsOrderedRedirect(t *testing.T) {
|
||||
s, requestID := newConsentFixture(t)
|
||||
ctx := context.Background()
|
||||
redirectTo, err := s.IssueCode(ctx, requestID, 42, []string{"read"}, []uint{7})
|
||||
if err != nil {
|
||||
t.Fatalf("IssueCode: %v", err)
|
||||
}
|
||||
if !strings.HasPrefix(redirectTo, "https://client.example.test/cb?code=") {
|
||||
t.Fatalf("redirectTo = %q, want code= prefix", redirectTo)
|
||||
}
|
||||
if !strings.Contains(redirectTo, "&iss=") || !strings.Contains(redirectTo, "&state=consent-test-state") {
|
||||
t.Fatalf("redirectTo = %q, want ordered iss/state", redirectTo)
|
||||
}
|
||||
}
|
||||
|
||||
func TestIssueCodeEmptyGrantedScopesIsRejected(t *testing.T) {
|
||||
s, requestID := newConsentFixture(t)
|
||||
if _, err := s.IssueCode(context.Background(), requestID, 1, nil, nil); !errors.Is(err, ErrNoGrantableScopes) {
|
||||
t.Fatalf("err = %v, want ErrNoGrantableScopes", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestDenyPendingConsumesAndReturnsAccessDeniedRedirect(t *testing.T) {
|
||||
s, requestID := newConsentFixture(t)
|
||||
ctx := context.Background()
|
||||
redirectTo, err := s.DenyPending(ctx, requestID, 1)
|
||||
if err != nil {
|
||||
t.Fatalf("DenyPending: %v", err)
|
||||
}
|
||||
if !strings.Contains(redirectTo, "error=access_denied") {
|
||||
t.Fatalf("redirectTo = %q, want error=access_denied", redirectTo)
|
||||
}
|
||||
// A second action on the same handle (allow or deny) is not found
|
||||
// (single-use), never a distinguishable 403.
|
||||
if _, err := s.DenyPending(ctx, requestID, 1); !errors.Is(err, ErrPendingNotFound) {
|
||||
t.Fatalf("second deny err = %v, want ErrPendingNotFound", err)
|
||||
}
|
||||
}
|
||||
41
modules/wristband/crypto.go
Normal file
41
modules/wristband/crypto.go
Normal file
@@ -0,0 +1,41 @@
|
||||
package wristband
|
||||
|
||||
import (
|
||||
"crypto/rand"
|
||||
"crypto/sha256"
|
||||
"crypto/subtle"
|
||||
"encoding/base64"
|
||||
"encoding/hex"
|
||||
)
|
||||
|
||||
// randomBase64URL returns n cryptographically random bytes, base64url
|
||||
// (RawURLEncoding, no padding) encoded, matching PHP's
|
||||
// rtrim(strtr(base64_encode(random_bytes(n)), '+/', '-_'), '=') byte for
|
||||
// byte (D-01/D-04).
|
||||
func randomBase64URL(n int) (string, error) {
|
||||
buf := make([]byte, n)
|
||||
if _, err := rand.Read(buf); err != nil {
|
||||
return "", err
|
||||
}
|
||||
return base64.RawURLEncoding.EncodeToString(buf), nil
|
||||
}
|
||||
|
||||
// sha256Hex is the fixed transform every opaque secret (client secret, code,
|
||||
// refresh token) is compared and persisted through: never the raw variable-
|
||||
// length secret (D-04).
|
||||
func sha256Hex(raw string) string {
|
||||
sum := sha256.Sum256([]byte(raw))
|
||||
return hex.EncodeToString(sum[:])
|
||||
}
|
||||
|
||||
// constantEqual compares two fixed-transform strings (sha256 hex digests or
|
||||
// PKCE S256 challenges) in constant time (D-04; RFC 7636 verifier compare).
|
||||
func constantEqual(a, b string) bool {
|
||||
return subtle.ConstantTimeCompare([]byte(a), []byte(b)) == 1
|
||||
}
|
||||
|
||||
// s256Challenge is the RFC 7636 S256 transform: base64url(sha256(verifier)).
|
||||
func s256Challenge(verifier string) string {
|
||||
sum := sha256.Sum256([]byte(verifier))
|
||||
return base64.RawURLEncoding.EncodeToString(sum[:])
|
||||
}
|
||||
151
modules/wristband/phase08_coverage_test.go
Normal file
151
modules/wristband/phase08_coverage_test.go
Normal file
@@ -0,0 +1,151 @@
|
||||
package wristband
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
// This file closes the last named-Go-evidence gaps 08-10-PLAN.md Task 1's
|
||||
// 103-method PHP audit found in wristband: cases the existing 08-01..08-09
|
||||
// test files exercised only partially, or not at all. See
|
||||
// .planning/phases/08-oauth2-1-authorization-server/08-PHP-TEST-MAP.md for
|
||||
// the full map; each test below is named after (and comments cite) the PHP
|
||||
// method it closes.
|
||||
|
||||
// TestPhase08Coverage is the 08-10-PLAN.md Task 1 focused-verify aggregator
|
||||
// (`go test ./wristband -run '^Test(Phase08Coverage|PHPTestMap)'`): each
|
||||
// individually-named test in this file also runs directly and is the
|
||||
// evidence 08-PHP-TEST-MAP.md cites by name, but this wrapper gives the
|
||||
// task's own prescribed fast command something to match.
|
||||
func TestPhase08Coverage(t *testing.T) {
|
||||
t.Run("TestRegisterLoopbackHTTPIsAccepted", TestRegisterLoopbackHTTPIsAccepted)
|
||||
t.Run("TestRegisterJavascriptURIIsRejected", TestRegisterJavascriptURIIsRejected)
|
||||
t.Run("TestRegisterErrorBodyHasNoSecretOrStackTrace", TestRegisterErrorBodyHasNoSecretOrStackTrace)
|
||||
t.Run("TestTokenPlainPKCEIsRejectedEvenWhenVerifierEqualsChallenge", TestTokenPlainPKCEIsRejectedEvenWhenVerifierEqualsChallenge)
|
||||
}
|
||||
|
||||
// TestRegisterLoopbackHTTPIsAccepted ports
|
||||
// OAuthRegisterTest::test_loopback_http_is_accepted: an http:// redirect
|
||||
// URI on 127.0.0.1 or localhost is not rejected the way any other
|
||||
// http:// URI is.
|
||||
func TestRegisterLoopbackHTTPIsAccepted(t *testing.T) {
|
||||
for _, host := range []string{"127.0.0.1", "localhost"} {
|
||||
t.Run(host, func(t *testing.T) {
|
||||
srv := newTestServer(newMemoryBackend())
|
||||
body := `{"redirect_uris":["http://` + host + `:8080/callback"]}`
|
||||
req := httptest.NewRequest(http.MethodPost, "/oauth/mcp/register", strings.NewReader(body))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Register(rec, req)
|
||||
if rec.Code != http.StatusCreated {
|
||||
t.Fatalf("status = %d, want %d (body=%s)", rec.Code, http.StatusCreated, rec.Body.String())
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestRegisterJavascriptURIIsRejected ports the "javascript uris" half of
|
||||
// OAuthRegisterTest::test_http_non_loopback_and_javascript_uris_are_rejected
|
||||
// (the http-non-loopback half is TestRegisterRedirectURIBounds/http-non-loopback).
|
||||
func TestRegisterJavascriptURIIsRejected(t *testing.T) {
|
||||
srv := newTestServer(newMemoryBackend())
|
||||
body := `{"redirect_uris":["javascript:alert(1)"]}`
|
||||
req := httptest.NewRequest(http.MethodPost, "/oauth/mcp/register", strings.NewReader(body))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Register(rec, req)
|
||||
// javascript: has no host component, so it fails the same "not a valid
|
||||
// URL" branch a hostless URI does (rejectRedirectURI), not the
|
||||
// scheme-specific message -- still rejected, never accepted.
|
||||
assertRegisterError(t, rec, http.StatusBadRequest, "invalid_redirect_uri",
|
||||
"Redirect URI is not a valid URL: javascript:alert(1)")
|
||||
}
|
||||
|
||||
// TestRegisterErrorBodyHasNoSecretOrStackTrace ports
|
||||
// OAuthRegisterTest::test_error_body_has_no_secret_or_stack: every DCR
|
||||
// error path returns exactly {"error","error_description"} -- no stray
|
||||
// field, no client_secret, no Go internal type/path/stack leakage -- and a
|
||||
// still-later valid registration is unaffected by the earlier failures.
|
||||
func TestRegisterErrorBodyHasNoSecretOrStackTrace(t *testing.T) {
|
||||
backend := newMemoryBackend()
|
||||
srv := newTestServer(backend)
|
||||
|
||||
badBodies := []string{
|
||||
`{not json`,
|
||||
`{"redirect_uris":[]}`,
|
||||
`{"redirect_uris":["not-a-url"]}`,
|
||||
`{"redirect_uris":["https://client.example.test/cb"],"token_endpoint_auth_method":"bogus"}`,
|
||||
}
|
||||
for _, body := range badBodies {
|
||||
req := httptest.NewRequest(http.MethodPost, "/oauth/mcp/register", strings.NewReader(body))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Register(rec, req)
|
||||
|
||||
if rec.Code != http.StatusBadRequest {
|
||||
t.Fatalf("body %q: status = %d, want %d", body, rec.Code, http.StatusBadRequest)
|
||||
}
|
||||
var got map[string]any
|
||||
if err := json.Unmarshal(rec.Body.Bytes(), &got); err != nil {
|
||||
t.Fatalf("body %q: decode: %v", body, err)
|
||||
}
|
||||
if len(got) != 2 {
|
||||
t.Fatalf("body %q: error body has %d fields, want exactly 2 (error, error_description): %v", body, len(got), got)
|
||||
}
|
||||
if _, ok := got["error"]; !ok {
|
||||
t.Fatalf("body %q: missing error field: %v", body, got)
|
||||
}
|
||||
if _, ok := got["error_description"]; !ok {
|
||||
t.Fatalf("body %q: missing error_description field: %v", body, got)
|
||||
}
|
||||
raw := rec.Body.String()
|
||||
for _, forbidden := range []string{"client_secret", "runtime error", "goroutine", ".go:", "panic"} {
|
||||
if strings.Contains(raw, forbidden) {
|
||||
t.Fatalf("body %q: error response leaked %q: %s", body, forbidden, raw)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// A still-later valid registration is unaffected by the prior failures
|
||||
// (no partial state, no leaked cap consumption).
|
||||
okReq := httptest.NewRequest(http.MethodPost, "/oauth/mcp/register", strings.NewReader(
|
||||
`{"redirect_uris":["https://client.example.test/cb"]}`))
|
||||
okReq.Header.Set("Content-Type", "application/json")
|
||||
okRec := httptest.NewRecorder()
|
||||
srv.Register(okRec, okReq)
|
||||
if okRec.Code != http.StatusCreated {
|
||||
t.Fatalf("valid registration after failures: status = %d, want %d (body=%s)", okRec.Code, http.StatusCreated, okRec.Body.String())
|
||||
}
|
||||
}
|
||||
|
||||
// TestTokenPlainPKCEIsRejectedEvenWhenVerifierEqualsChallenge ports
|
||||
// security/OAuthRefreshRotationTest::test_plain_pkce_is_rejected_even_when_verifier_equals_challenge.
|
||||
// Authorize itself never persists a "plain" code_challenge_method
|
||||
// (TestPKCEChallengeMethodMustBeS256), so this seeds a code row directly to
|
||||
// prove the defense also holds at the token endpoint: verifyPkce rejects
|
||||
// any non-S256 method outright, never falling back to a naive
|
||||
// verifier == stored-challenge string compare.
|
||||
func TestTokenPlainPKCEIsRejectedEvenWhenVerifierEqualsChallenge(t *testing.T) {
|
||||
backend := newMemoryBackend()
|
||||
srv := newTestServer(backend)
|
||||
insertTokenTestClient(backend, "cli-plain-exchange", "none", nil, nil)
|
||||
const verifier = "same-value-used-as-both-verifier-and-challenge-01234"
|
||||
rawCode, _ := insertTokenTestCode(t, backend, "cli-plain-exchange", verifier, func(rec *AuthCodeRecord) {
|
||||
rec.CodeChallengeMethod = "plain"
|
||||
})
|
||||
|
||||
form := validExchangeForm(rawCode, verifier, "cli-plain-exchange")
|
||||
req := tokenRequest(form, "")
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Token(rec, req)
|
||||
assertTokenError(t, rec, http.StatusBadRequest, "invalid_grant")
|
||||
|
||||
backend.mu.Lock()
|
||||
defer backend.mu.Unlock()
|
||||
if len(backend.tokens) != 0 {
|
||||
t.Fatal("a plain-method exchange must not mint an access token even when verifier equals the stored challenge")
|
||||
}
|
||||
}
|
||||
60
modules/wristband/redirect_html.go
Normal file
60
modules/wristband/redirect_html.go
Normal file
@@ -0,0 +1,60 @@
|
||||
package wristband
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// htmlEscapePHP ports PHP's htmlspecialchars($s, ENT_QUOTES, 'UTF-8') byte
|
||||
// for byte: Go's stdlib html.EscapeString differs on the quote entities
|
||||
// ('/" vs PHP's '/"), which would diverge from the
|
||||
// recorded redirect body whenever a redirect_uri or state value contains a
|
||||
// quote character.
|
||||
func htmlEscapePHP(s string) string {
|
||||
var b strings.Builder
|
||||
b.Grow(len(s))
|
||||
for _, r := range s {
|
||||
switch r {
|
||||
case '&':
|
||||
b.WriteString("&")
|
||||
case '"':
|
||||
b.WriteString(""")
|
||||
case '\'':
|
||||
b.WriteString("'")
|
||||
case '<':
|
||||
b.WriteString("<")
|
||||
case '>':
|
||||
b.WriteString(">")
|
||||
default:
|
||||
b.WriteRune(r)
|
||||
}
|
||||
}
|
||||
return b.String()
|
||||
}
|
||||
|
||||
// writeRedirectHTML ports Laravel/Symfony's RedirectResponse default HTML
|
||||
// body byte-for-byte (T-08-OPEN-REDIRECT/D-04). Go's net/http never emits a
|
||||
// body for a 3xx Location redirect; every wristband redirect needs this
|
||||
// exact body plus Content-Type because real recorded PHP traffic includes
|
||||
// it, and an unchanged browser-based client (the Nuxt /connect handoff)
|
||||
// observes it. Cache-Control is the caller's responsibility -- callers set
|
||||
// it before invoking this helper because its value differs between the
|
||||
// authorize success path and error redirects versus other endpoints.
|
||||
func writeRedirectHTML(w http.ResponseWriter, status int, target string) {
|
||||
escaped := htmlEscapePHP(target)
|
||||
var b strings.Builder
|
||||
b.WriteString("<!DOCTYPE html>\n<html>\n <head>\n <meta charset=\"UTF-8\" />\n <meta http-equiv=\"refresh\" content=\"0;url='")
|
||||
b.WriteString(escaped)
|
||||
b.WriteString("'\" />\n\n <title>Redirecting to ")
|
||||
b.WriteString(escaped)
|
||||
b.WriteString("</title>\n </head>\n <body>\n Redirecting to <a href=\"")
|
||||
b.WriteString(escaped)
|
||||
b.WriteString("\">")
|
||||
b.WriteString(escaped)
|
||||
b.WriteString("</a>.\n </body>\n</html>")
|
||||
|
||||
w.Header().Set("Content-Type", "text/html; charset=utf-8")
|
||||
w.Header().Set("Location", target)
|
||||
w.WriteHeader(status)
|
||||
_, _ = w.Write([]byte(b.String()))
|
||||
}
|
||||
332
modules/wristband/register.go
Normal file
332
modules/wristband/register.go
Normal file
@@ -0,0 +1,332 @@
|
||||
// RFC 7591 Dynamic Client Registration, ported from PHP
|
||||
// OAuthRegisterController::register byte-for-byte including its validation
|
||||
// order (08-CONTEXT.md D-02/D-05/D-06/D-21; canonical PHP source:
|
||||
// OAuthRegisterController.php).
|
||||
package wristband
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"strings"
|
||||
)
|
||||
|
||||
var supportedAuthMethods = map[string]bool{
|
||||
"none": true,
|
||||
"client_secret_post": true,
|
||||
"client_secret_basic": true,
|
||||
}
|
||||
|
||||
var supportedGrantTypes = []string{"authorization_code", "refresh_token"}
|
||||
|
||||
// registerRequestBody is the RFC 7591 JSON request document. Only the
|
||||
// fields PHP reads are decoded; unknown fields are ignored (PHP's
|
||||
// $request->input() does the same).
|
||||
type registerRequestBody struct {
|
||||
RedirectURIs []string `json:"redirect_uris"`
|
||||
TokenEndpointAuthMethod string `json:"token_endpoint_auth_method"`
|
||||
GrantTypes []string `json:"grant_types"`
|
||||
ResponseTypes []string `json:"response_types"`
|
||||
ClientName string `json:"client_name"`
|
||||
}
|
||||
|
||||
type registerResponseBody struct {
|
||||
ClientID string `json:"client_id"`
|
||||
ClientIDIssuedAt int64 `json:"client_id_issued_at"`
|
||||
ClientName string `json:"client_name"`
|
||||
RedirectURIs []string `json:"redirect_uris"`
|
||||
GrantTypes []string `json:"grant_types"`
|
||||
ResponseTypes []string `json:"response_types"`
|
||||
TokenEndpointAuthMethod string `json:"token_endpoint_auth_method"`
|
||||
ClientSecret string `json:"client_secret,omitempty"`
|
||||
ClientSecretExpiresAt *int64 `json:"client_secret_expires_at,omitempty"`
|
||||
}
|
||||
|
||||
type rfcErrorBody struct {
|
||||
Error string `json:"error"`
|
||||
ErrorDescription string `json:"error_description"`
|
||||
}
|
||||
|
||||
// Register handles POST /oauth/mcp/register (RFC 7591). D-02: JSON-only,
|
||||
// PHP isJson() semantics ("contains /json"). D-21: the body is bounded to
|
||||
// Options.RegisterMaxBodyBytes before decoding, and an oversized body gets
|
||||
// the endpoint's ordinary invalid_client_metadata response, not a generic
|
||||
// error. Sweep, the cap check and the create all run inside one transaction
|
||||
// (T-08-DCR-FLOOD).
|
||||
func (s *Server) Register(w http.ResponseWriter, r *http.Request) {
|
||||
if !isJSONContentType(r.Header.Get("Content-Type")) {
|
||||
writeRegisterError(w, http.StatusBadRequest, "invalid_client_metadata", "Request must be application/json")
|
||||
return
|
||||
}
|
||||
if s.backend == nil {
|
||||
w.WriteHeader(http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
|
||||
r.Body = http.MaxBytesReader(w, r.Body, s.opts.RegisterMaxBodyBytes)
|
||||
var body registerRequestBody
|
||||
dec := json.NewDecoder(r.Body)
|
||||
if err := dec.Decode(&body); err != nil {
|
||||
// D-21: an oversized (http.MaxBytesReader) or otherwise malformed
|
||||
// JSON body both collapse to the endpoint-native error, never a
|
||||
// generic 413/house body.
|
||||
writeRegisterError(w, http.StatusBadRequest, "invalid_client_metadata", "Request must be application/json")
|
||||
return
|
||||
}
|
||||
|
||||
uris, errMsg := parseRedirectURIs(body.RedirectURIs)
|
||||
if errMsg != "" {
|
||||
writeRegisterError(w, http.StatusBadRequest, "invalid_redirect_uri", errMsg)
|
||||
return
|
||||
}
|
||||
|
||||
authMethod := body.TokenEndpointAuthMethod
|
||||
if authMethod == "" {
|
||||
authMethod = "none"
|
||||
}
|
||||
if !supportedAuthMethods[authMethod] {
|
||||
writeRegisterError(w, http.StatusBadRequest, "invalid_client_metadata", "token_endpoint_auth_method is not supported")
|
||||
return
|
||||
}
|
||||
|
||||
grants, errMsg := parseGrantTypes(body.GrantTypes)
|
||||
if errMsg != "" {
|
||||
writeRegisterError(w, http.StatusBadRequest, "invalid_client_metadata", errMsg)
|
||||
return
|
||||
}
|
||||
|
||||
responseTypes, errMsg := parseResponseTypes(body.ResponseTypes)
|
||||
if errMsg != "" {
|
||||
writeRegisterError(w, http.StatusBadRequest, "invalid_client_metadata", errMsg)
|
||||
return
|
||||
}
|
||||
|
||||
name := truncateRunes(strings.TrimSpace(stripControlChars(body.ClientName)), 255)
|
||||
if name == "" {
|
||||
name = "MCP client"
|
||||
}
|
||||
|
||||
clientID, err := s.randomBytes(16)
|
||||
if err != nil {
|
||||
w.WriteHeader(http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
var secret string
|
||||
var secretHash *string
|
||||
if authMethod != "none" {
|
||||
secret, err = s.randomBytes(32)
|
||||
if err != nil {
|
||||
w.WriteHeader(http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
h := sha256Hex(secret)
|
||||
secretHash = &h
|
||||
}
|
||||
|
||||
ip := clientIP(r)
|
||||
rec := &ClientRecord{
|
||||
ClientID: clientID,
|
||||
ClientSecretHash: secretHash,
|
||||
ClientName: name,
|
||||
RedirectURIs: uris,
|
||||
GrantTypes: grants,
|
||||
TokenEndpointAuthMethod: authMethod,
|
||||
RegistrationIP: &ip,
|
||||
}
|
||||
|
||||
ctx := r.Context()
|
||||
sweepBefore := s.now().Add(-s.opts.DCRUnconsentedSweepAge)
|
||||
err = s.backend.WithinTx(ctx, func(tx Tx) error {
|
||||
if err := tx.SweepUnconsented(ctx, sweepBefore); err != nil {
|
||||
return err
|
||||
}
|
||||
return tx.CreateWithCap(ctx, rec, s.opts.DCRClientCap)
|
||||
})
|
||||
if errors.Is(err, ErrClientCapReached) {
|
||||
writeRegisterError(w, http.StatusBadRequest, "invalid_client_metadata", "Registration temporarily unavailable")
|
||||
return
|
||||
}
|
||||
if err != nil {
|
||||
w.WriteHeader(http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
|
||||
resp := registerResponseBody{
|
||||
ClientID: clientID,
|
||||
ClientIDIssuedAt: rec.CreatedAt.Unix(),
|
||||
ClientName: name,
|
||||
RedirectURIs: uris,
|
||||
GrantTypes: grants,
|
||||
ResponseTypes: responseTypes,
|
||||
TokenEndpointAuthMethod: authMethod,
|
||||
}
|
||||
if secret != "" {
|
||||
resp.ClientSecret = secret
|
||||
zero := int64(0)
|
||||
resp.ClientSecretExpiresAt = &zero
|
||||
}
|
||||
writeExactJSON(w, http.StatusCreated, resp, map[string]string{"Cache-Control": "no-store, private"})
|
||||
}
|
||||
|
||||
func writeRegisterError(w http.ResponseWriter, status int, code, description string) {
|
||||
writeExactJSON(w, status, rfcErrorBody{Error: code, ErrorDescription: description}, map[string]string{"Cache-Control": "no-store, private"})
|
||||
}
|
||||
|
||||
// isJSONContentType mirrors Laravel's Request::isJson(): the Content-Type
|
||||
// header contains "/json" or "+json" anywhere (D-02).
|
||||
func isJSONContentType(ct string) bool {
|
||||
ct = strings.ToLower(ct)
|
||||
return strings.Contains(ct, "/json") || strings.Contains(ct, "+json")
|
||||
}
|
||||
|
||||
// parseRedirectURIs ports OAuthRegisterController::parseRedirectUris: a
|
||||
// required, non-empty array of unique URI strings, at most 5, each passing
|
||||
// rejectRedirectURI.
|
||||
func parseRedirectURIs(raw []string) ([]string, string) {
|
||||
if len(raw) == 0 {
|
||||
return nil, "redirect_uris is required"
|
||||
}
|
||||
seen := make(map[string]bool, len(raw))
|
||||
uris := make([]string, 0, len(raw))
|
||||
for _, entry := range raw {
|
||||
if entry == "" {
|
||||
return nil, "redirect_uris must be an array of URI strings"
|
||||
}
|
||||
if seen[entry] {
|
||||
continue
|
||||
}
|
||||
seen[entry] = true
|
||||
uris = append(uris, entry)
|
||||
}
|
||||
if len(uris) > 5 {
|
||||
return nil, "A client may have at most 5 redirect URIs"
|
||||
}
|
||||
for _, u := range uris {
|
||||
if reason := rejectRedirectURI(u); reason != "" {
|
||||
return nil, reason
|
||||
}
|
||||
}
|
||||
return uris, ""
|
||||
}
|
||||
|
||||
// rejectRedirectURI ports OAuthClient::rejectRedirectUri: at most 512
|
||||
// characters, a valid URL with scheme+host, https:// or loopback
|
||||
// http://127.0.0.1 / http://localhost. Returns "" when the URI is accepted.
|
||||
func rejectRedirectURI(uri string) string {
|
||||
if len(uri) > 512 {
|
||||
return "Each redirect URI must be at most 512 characters."
|
||||
}
|
||||
parsed, err := url.Parse(uri)
|
||||
if err != nil || parsed.Scheme == "" || parsed.Hostname() == "" {
|
||||
return "Redirect URI is not a valid URL: " + uri
|
||||
}
|
||||
scheme := strings.ToLower(parsed.Scheme)
|
||||
host := strings.ToLower(parsed.Hostname())
|
||||
if scheme == "https" {
|
||||
return ""
|
||||
}
|
||||
if scheme == "http" && (host == "127.0.0.1" || host == "localhost") {
|
||||
return ""
|
||||
}
|
||||
return "Redirect URI must be https:// or loopback http://127.0.0.1 / http://localhost: " + uri
|
||||
}
|
||||
|
||||
// parseGrantTypes ports OAuthRegisterController::parseGrantTypes: nil means
|
||||
// the default (both supported grants); otherwise a non-empty array of
|
||||
// supported values that must include authorization_code.
|
||||
func parseGrantTypes(raw []string) ([]string, string) {
|
||||
if raw == nil {
|
||||
return append([]string(nil), supportedGrantTypes...), ""
|
||||
}
|
||||
if len(raw) == 0 {
|
||||
return nil, "grant_types is invalid"
|
||||
}
|
||||
seen := make(map[string]bool, len(raw))
|
||||
grants := make([]string, 0, len(raw))
|
||||
hasAuthCode := false
|
||||
for _, entry := range raw {
|
||||
if !isSupportedGrantType(entry) {
|
||||
return nil, "grant_types contains an unsupported value"
|
||||
}
|
||||
if seen[entry] {
|
||||
continue
|
||||
}
|
||||
seen[entry] = true
|
||||
grants = append(grants, entry)
|
||||
if entry == "authorization_code" {
|
||||
hasAuthCode = true
|
||||
}
|
||||
}
|
||||
if !hasAuthCode {
|
||||
return nil, "grant_types must include authorization_code"
|
||||
}
|
||||
return grants, ""
|
||||
}
|
||||
|
||||
func isSupportedGrantType(v string) bool {
|
||||
for _, g := range supportedGrantTypes {
|
||||
if g == v {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// parseResponseTypes ports OAuthRegisterController::parseResponseTypes:
|
||||
// nil means the default ["code"]; otherwise every entry must be "code".
|
||||
func parseResponseTypes(raw []string) ([]string, string) {
|
||||
if raw == nil {
|
||||
return []string{"code"}, ""
|
||||
}
|
||||
if len(raw) == 0 {
|
||||
return nil, "response_types is invalid"
|
||||
}
|
||||
for _, entry := range raw {
|
||||
if entry != "code" {
|
||||
return nil, "response_types must be code"
|
||||
}
|
||||
}
|
||||
return []string{"code"}, ""
|
||||
}
|
||||
|
||||
// stripControlChars removes C0 control characters and DEL, mirroring the
|
||||
// PHP preg_replace('/[\x00-\x1F\x7F]/u', ”, ...) regex ConnectedApp/Consent
|
||||
// controllers apply to client_name at display time (08-CONTEXT.md
|
||||
// discretion note). Registration applies the same filter at capture time so
|
||||
// no untrusted control byte from an issuer's declared name is ever stored.
|
||||
func stripControlChars(s string) string {
|
||||
return strings.Map(func(r rune) rune {
|
||||
if r <= 0x1F || r == 0x7F {
|
||||
return -1
|
||||
}
|
||||
return r
|
||||
}, s)
|
||||
}
|
||||
|
||||
// truncateRunes returns s truncated to at most n runes (PHP mb_substr parity).
|
||||
func truncateRunes(s string, n int) string {
|
||||
r := []rune(s)
|
||||
if len(r) <= n {
|
||||
return s
|
||||
}
|
||||
return string(r[:n])
|
||||
}
|
||||
|
||||
// clientIP mirrors PHP's substr((string) $request->ip(), 0, 45): the
|
||||
// direct remote address (no trusted-proxy chain here; DCR is unauthenticated
|
||||
// and app-specific proxy trust stays out of the app-agnostic framework).
|
||||
func clientIP(r *http.Request) string {
|
||||
host := r.RemoteAddr
|
||||
if idx := strings.LastIndex(host, ":"); idx != -1 && !strings.Contains(host, "]") {
|
||||
host = host[:idx]
|
||||
} else if strings.HasPrefix(host, "[") {
|
||||
if end := strings.Index(host, "]"); end != -1 {
|
||||
host = host[1:end]
|
||||
}
|
||||
}
|
||||
if len(host) > 45 {
|
||||
host = host[:45]
|
||||
}
|
||||
return host
|
||||
}
|
||||
620
modules/wristband/registration_test.go
Normal file
620
modules/wristband/registration_test.go
Normal file
@@ -0,0 +1,620 @@
|
||||
package wristband
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
// memoryBackend is wristband's own in-memory Backend/Tx implementation
|
||||
// (08-CONTEXT.md D-07: "wristband ships an in-memory store for its own
|
||||
// tests"). One mutex guards the whole transaction closure so concurrency
|
||||
// tests model the seam, not individual map operations (08-PATTERNS.md).
|
||||
type memoryBackend struct {
|
||||
mu sync.Mutex
|
||||
clients []*ClientRecord
|
||||
codes []*AuthCodeRecord
|
||||
refresh []*RefreshTokenRecord
|
||||
tokens []*IssuedToken
|
||||
// revoked tracks IssuedToken.ID -> revoked, since IssuedToken itself
|
||||
// carries no status field (it is the one-time mint result, not a
|
||||
// queryable row). 08-06-PLAN.md's rotation/replay/revoke tests need to
|
||||
// observe access-token revocation the same way real Postgres tests
|
||||
// observe models.ApiToken.RevokedAt.
|
||||
revoked map[uint]bool
|
||||
nextID uint
|
||||
}
|
||||
|
||||
func newMemoryBackend() *memoryBackend {
|
||||
return &memoryBackend{}
|
||||
}
|
||||
|
||||
func (b *memoryBackend) WithinTx(ctx context.Context, fn func(Tx) error) error {
|
||||
b.mu.Lock()
|
||||
defer b.mu.Unlock()
|
||||
return fn(&memoryTx{b: b})
|
||||
}
|
||||
|
||||
type memoryTx struct{ b *memoryBackend }
|
||||
|
||||
var _ Tx = (*memoryTx)(nil)
|
||||
|
||||
func (t *memoryTx) ByClientID(ctx context.Context, clientID string) (*ClientRecord, error) {
|
||||
for _, c := range t.b.clients {
|
||||
if c.ClientID == clientID {
|
||||
cp := *c
|
||||
return &cp, nil
|
||||
}
|
||||
}
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
func (t *memoryTx) CreateWithCap(ctx context.Context, rec *ClientRecord, cap int) error {
|
||||
count := 0
|
||||
for _, c := range t.b.clients {
|
||||
if c.RevokedAt == nil {
|
||||
count++
|
||||
}
|
||||
}
|
||||
if count >= cap {
|
||||
return ErrClientCapReached
|
||||
}
|
||||
t.b.nextID++
|
||||
rec.ID = t.b.nextID
|
||||
rec.CreatedAt = time.Now()
|
||||
cp := *rec
|
||||
t.b.clients = append(t.b.clients, &cp)
|
||||
return nil
|
||||
}
|
||||
|
||||
func (t *memoryTx) SweepUnconsented(ctx context.Context, olderThan time.Time) error {
|
||||
kept := t.b.clients[:0:0]
|
||||
for _, c := range t.b.clients {
|
||||
if c.ConsentedAt == nil && c.RegistrationIP != nil && c.CreatedAt.Before(olderThan) {
|
||||
continue
|
||||
}
|
||||
kept = append(kept, c)
|
||||
}
|
||||
t.b.clients = kept
|
||||
return nil
|
||||
}
|
||||
|
||||
func (t *memoryTx) CreatePending(ctx context.Context, rec *AuthCodeRecord) error {
|
||||
t.b.nextID++
|
||||
rec.ID = t.b.nextID
|
||||
cp := *rec
|
||||
t.b.codes = append(t.b.codes, &cp)
|
||||
return nil
|
||||
}
|
||||
|
||||
func (t *memoryTx) ByRequestID(ctx context.Context, requestID string) (*AuthCodeRecord, error) {
|
||||
for _, c := range t.b.codes {
|
||||
if c.RequestID != nil && *c.RequestID == requestID {
|
||||
cp := *c
|
||||
return &cp, nil
|
||||
}
|
||||
}
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
func (t *memoryTx) ByCodeHashForUpdate(ctx context.Context, codeHash string) (*AuthCodeRecord, error) {
|
||||
for _, c := range t.b.codes {
|
||||
if c.CodeHash != nil && *c.CodeHash == codeHash {
|
||||
cp := *c
|
||||
return &cp, nil
|
||||
}
|
||||
}
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
func (t *memoryTx) MarkIssued(ctx context.Context, id uint, codeHash string, userID uint, scopes []string, collectionIDs []uint, expiresAt time.Time) error {
|
||||
for _, c := range t.b.codes {
|
||||
if c.ID == id {
|
||||
c.RequestID = nil
|
||||
c.CodeHash = &codeHash
|
||||
c.UserID = &userID
|
||||
c.Scopes = scopes
|
||||
c.CollectionIDs = collectionIDs
|
||||
c.ExpiresAt = expiresAt
|
||||
return nil
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (t *memoryTx) MarkConsented(ctx context.Context, clientID string) error {
|
||||
for _, c := range t.b.clients {
|
||||
if c.ClientID == clientID {
|
||||
if c.ConsentedAt == nil {
|
||||
now := time.Now()
|
||||
c.ConsentedAt = &now
|
||||
}
|
||||
return nil
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (t *memoryTx) MarkUsed(ctx context.Context, id uint) error {
|
||||
now := time.Now()
|
||||
for _, c := range t.b.codes {
|
||||
if c.ID == id {
|
||||
c.UsedAt = &now
|
||||
return nil
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (t *memoryTx) Create(ctx context.Context, rec *RefreshTokenRecord) error {
|
||||
t.b.nextID++
|
||||
rec.ID = t.b.nextID
|
||||
cp := *rec
|
||||
t.b.refresh = append(t.b.refresh, &cp)
|
||||
return nil
|
||||
}
|
||||
|
||||
func (t *memoryTx) ByTokenHashForUpdate(ctx context.Context, tokenHash string) (*RefreshTokenRecord, error) {
|
||||
for _, r := range t.b.refresh {
|
||||
if r.TokenHash == tokenHash {
|
||||
cp := *r
|
||||
return &cp, nil
|
||||
}
|
||||
}
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
func (t *memoryTx) ByAPITokenIDForUpdate(ctx context.Context, apiTokenID uint) (*RefreshTokenRecord, error) {
|
||||
for _, r := range t.b.refresh {
|
||||
if r.APITokenID != nil && *r.APITokenID == apiTokenID {
|
||||
cp := *r
|
||||
return &cp, nil
|
||||
}
|
||||
}
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
func (t *memoryTx) MarkRotated(ctx context.Context, id uint, successorID uint) error {
|
||||
for _, r := range t.b.refresh {
|
||||
if r.ID == id {
|
||||
sid := successorID
|
||||
r.RotatedToID = &sid
|
||||
return nil
|
||||
}
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// RevokeLineage walks forward through RotatedToID starting at startID,
|
||||
// stamping RevokedAt on every visited refresh row and marking each row's
|
||||
// linked access token revoked too (08-06-PLAN.md: mirrors the GORM
|
||||
// adapter's RevokeLineage so T-08-REFRESH-REPLAY's "kill the whole lineage"
|
||||
// contract is provable against the in-memory backend, not just real
|
||||
// Postgres).
|
||||
func (t *memoryTx) RevokeLineage(ctx context.Context, startID uint) error {
|
||||
now := time.Now()
|
||||
id := startID
|
||||
for {
|
||||
var found *RefreshTokenRecord
|
||||
for _, r := range t.b.refresh {
|
||||
if r.ID == id {
|
||||
found = r
|
||||
break
|
||||
}
|
||||
}
|
||||
if found == nil {
|
||||
return nil
|
||||
}
|
||||
if found.RevokedAt == nil {
|
||||
found.RevokedAt = &now
|
||||
}
|
||||
if found.APITokenID != nil {
|
||||
if t.b.revoked == nil {
|
||||
t.b.revoked = map[uint]bool{}
|
||||
}
|
||||
t.b.revoked[*found.APITokenID] = true
|
||||
}
|
||||
if found.RotatedToID == nil {
|
||||
return nil
|
||||
}
|
||||
id = *found.RotatedToID
|
||||
}
|
||||
}
|
||||
|
||||
func (t *memoryTx) DeleteExpiredCodes(ctx context.Context, now time.Time) error {
|
||||
kept := t.b.codes[:0:0]
|
||||
for _, c := range t.b.codes {
|
||||
if c.ExpiresAt.Before(now) {
|
||||
continue
|
||||
}
|
||||
kept = append(kept, c)
|
||||
}
|
||||
t.b.codes = kept
|
||||
return nil
|
||||
}
|
||||
|
||||
func (t *memoryTx) DeleteExpiredRefreshTokens(ctx context.Context, now time.Time) error {
|
||||
kept := t.b.refresh[:0:0]
|
||||
for _, r := range t.b.refresh {
|
||||
if r.ExpiresAt.Before(now) {
|
||||
continue
|
||||
}
|
||||
kept = append(kept, r)
|
||||
}
|
||||
t.b.refresh = kept
|
||||
return nil
|
||||
}
|
||||
|
||||
// Mint's secret embeds the freshly-allocated id so two mints for the same
|
||||
// client name (e.g. across a rotation) never collide on an identical
|
||||
// "mem_<name>" string -- 08-06-PLAN.md's rotation tests distinguish the old
|
||||
// and new access tokens by their returned secret.
|
||||
func (t *memoryTx) Mint(ctx context.Context, userID uint, name string, scopes []string, expiresAt time.Time, collectionIDs []uint, clientID string) (IssuedToken, error) {
|
||||
t.b.nextID++
|
||||
tok := IssuedToken{ID: t.b.nextID, Secret: fmt.Sprintf("mem_%d_%s", t.b.nextID, name)}
|
||||
t.b.tokens = append(t.b.tokens, &tok)
|
||||
return tok, nil
|
||||
}
|
||||
|
||||
func (t *memoryTx) Revoke(ctx context.Context, tokenID uint) error {
|
||||
if t.b.revoked == nil {
|
||||
t.b.revoked = map[uint]bool{}
|
||||
}
|
||||
t.b.revoked[tokenID] = true
|
||||
return nil
|
||||
}
|
||||
|
||||
func newTestServer(backend Backend) *Server {
|
||||
opts := DefaultOptions()
|
||||
opts.Issuer = "https://plytarium.com"
|
||||
srv := NewServer(opts)
|
||||
srv.SetBackend(backend)
|
||||
return srv
|
||||
}
|
||||
|
||||
func registerRequest(t *testing.T, body string, contentType string) *httptest.ResponseRecorder {
|
||||
t.Helper()
|
||||
req := httptest.NewRequest(http.MethodPost, "/oauth/mcp/register", strings.NewReader(body))
|
||||
req.Header.Set("Content-Type", contentType)
|
||||
rec := httptest.NewRecorder()
|
||||
return rec
|
||||
}
|
||||
|
||||
// TestPhase8RedRegistration is the Phase 8 Wave 2 RED anchor (08-02-PLAN.md
|
||||
// Task 2, D-02/D-05/D-06/D-21). It asserts the exact RFC 7591 public-client
|
||||
// success contract and fails with the PHASE8_RED:registration sentinel
|
||||
// while Server.Register is a 501 stub; scripts/check-phase8-red.sh verifies
|
||||
// this failure is fail-closed.
|
||||
func TestPhase8RedRegistration(t *testing.T) {
|
||||
srv := newTestServer(newMemoryBackend())
|
||||
|
||||
body := `{"redirect_uris":["https://client.example.test/callback"]}`
|
||||
req := httptest.NewRequest(http.MethodPost, "/oauth/mcp/register", strings.NewReader(body))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Register(rec, req)
|
||||
|
||||
if rec.Code != http.StatusCreated {
|
||||
t.Fatalf("PHASE8_RED:registration: status = %d, want %d (body=%s)", rec.Code, http.StatusCreated, rec.Body.String())
|
||||
}
|
||||
var got map[string]any
|
||||
if err := json.Unmarshal(rec.Body.Bytes(), &got); err != nil {
|
||||
t.Fatalf("PHASE8_RED:registration: decode response: %v", err)
|
||||
}
|
||||
if got["client_name"] != "MCP client" {
|
||||
t.Fatalf("PHASE8_RED:registration: client_name = %v, want default \"MCP client\"", got["client_name"])
|
||||
}
|
||||
if _, hasSecret := got["client_secret"]; hasSecret {
|
||||
t.Fatal("PHASE8_RED:registration: public client response has a client_secret, want none")
|
||||
}
|
||||
if cc := rec.Header().Get("Cache-Control"); cc != "no-store, private" {
|
||||
t.Fatalf("PHASE8_RED:registration: Cache-Control = %q, want \"no-store, private\"", cc)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRegisterPublicClientResponse(t *testing.T) {
|
||||
srv := newTestServer(newMemoryBackend())
|
||||
body := `{"redirect_uris":["https://client.example.test/callback"],"client_name":"My Client"}`
|
||||
req := httptest.NewRequest(http.MethodPost, "/oauth/mcp/register", strings.NewReader(body))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Register(rec, req)
|
||||
|
||||
if rec.Code != http.StatusCreated {
|
||||
t.Fatalf("status = %d, body=%s", rec.Code, rec.Body.String())
|
||||
}
|
||||
var got map[string]any
|
||||
if err := json.Unmarshal(rec.Body.Bytes(), &got); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if got["client_name"] != "My Client" {
|
||||
t.Fatalf("client_name = %v", got["client_name"])
|
||||
}
|
||||
if got["token_endpoint_auth_method"] != "none" {
|
||||
t.Fatalf("token_endpoint_auth_method = %v, want none", got["token_endpoint_auth_method"])
|
||||
}
|
||||
if _, hasSecret := got["client_secret"]; hasSecret {
|
||||
t.Fatal("public client response carries client_secret")
|
||||
}
|
||||
if _, hasExpires := got["client_secret_expires_at"]; hasExpires {
|
||||
t.Fatal("public client response carries client_secret_expires_at")
|
||||
}
|
||||
if strings.HasSuffix(rec.Body.String(), "\n") {
|
||||
t.Fatal("body has a trailing newline")
|
||||
}
|
||||
if _, hasData := got["data"]; hasData {
|
||||
t.Fatal("body has a house \"data\" envelope")
|
||||
}
|
||||
}
|
||||
|
||||
func TestRegisterConfidentialClientReturnsSecretOnce(t *testing.T) {
|
||||
backend := newMemoryBackend()
|
||||
srv := newTestServer(backend)
|
||||
body := `{"redirect_uris":["https://client.example.test/callback"],"token_endpoint_auth_method":"client_secret_post"}`
|
||||
req := httptest.NewRequest(http.MethodPost, "/oauth/mcp/register", strings.NewReader(body))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Register(rec, req)
|
||||
|
||||
if rec.Code != http.StatusCreated {
|
||||
t.Fatalf("status = %d, body=%s", rec.Code, rec.Body.String())
|
||||
}
|
||||
var got map[string]any
|
||||
if err := json.Unmarshal(rec.Body.Bytes(), &got); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
secret, _ := got["client_secret"].(string)
|
||||
if secret == "" {
|
||||
t.Fatal("confidential client response has no client_secret")
|
||||
}
|
||||
if got["client_secret_expires_at"] != float64(0) {
|
||||
t.Fatalf("client_secret_expires_at = %v, want 0", got["client_secret_expires_at"])
|
||||
}
|
||||
|
||||
backend.mu.Lock()
|
||||
defer backend.mu.Unlock()
|
||||
if len(backend.clients) != 1 {
|
||||
t.Fatalf("stored clients = %d, want 1", len(backend.clients))
|
||||
}
|
||||
stored := backend.clients[0]
|
||||
if stored.ClientSecretHash == nil || *stored.ClientSecretHash == secret {
|
||||
t.Fatal("stored record must hold a hash, never the raw secret")
|
||||
}
|
||||
if *stored.ClientSecretHash != sha256Hex(secret) {
|
||||
t.Fatal("stored hash does not match sha256Hex(secret)")
|
||||
}
|
||||
if strings.Contains(rec.Body.String(), *stored.ClientSecretHash) {
|
||||
t.Fatal("response body must not also contain the persisted hash")
|
||||
}
|
||||
}
|
||||
|
||||
func TestRegisterWrongContentType(t *testing.T) {
|
||||
srv := newTestServer(newMemoryBackend())
|
||||
body := `redirect_uris=https://client.example.test/callback`
|
||||
req := httptest.NewRequest(http.MethodPost, "/oauth/mcp/register", strings.NewReader(body))
|
||||
req.Header.Set("Content-Type", "application/x-www-form-urlencoded")
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Register(rec, req)
|
||||
|
||||
assertRegisterError(t, rec, http.StatusBadRequest, "invalid_client_metadata", "Request must be application/json")
|
||||
}
|
||||
|
||||
func TestRegisterMalformedJSON(t *testing.T) {
|
||||
srv := newTestServer(newMemoryBackend())
|
||||
req := httptest.NewRequest(http.MethodPost, "/oauth/mcp/register", strings.NewReader(`{not json`))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Register(rec, req)
|
||||
|
||||
assertRegisterError(t, rec, http.StatusBadRequest, "invalid_client_metadata", "Request must be application/json")
|
||||
}
|
||||
|
||||
func TestRegisterOversizedBody(t *testing.T) {
|
||||
srv := newTestServer(newMemoryBackend())
|
||||
srv.opts.RegisterMaxBodyBytes = 64 * 1024
|
||||
|
||||
// D-21: exactly one byte past the 64 KiB bound.
|
||||
padding := strings.Repeat("a", 65537)
|
||||
body := `{"redirect_uris":["https://client.example.test/callback"],"client_name":"` + padding + `"}`
|
||||
req := httptest.NewRequest(http.MethodPost, "/oauth/mcp/register", strings.NewReader(body))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Register(rec, req)
|
||||
|
||||
assertRegisterError(t, rec, http.StatusBadRequest, "invalid_client_metadata", "Request must be application/json")
|
||||
}
|
||||
|
||||
func TestRegisterRedirectURIBounds(t *testing.T) {
|
||||
cases := []struct {
|
||||
name string
|
||||
uris string
|
||||
want string
|
||||
}{
|
||||
{"missing", `[]`, "redirect_uris is required"},
|
||||
{"empty-entry", `["https://a.test/cb",""]`, "redirect_uris must be an array of URI strings"},
|
||||
{"too-many", `["https://a.test/1","https://a.test/2","https://a.test/3","https://a.test/4","https://a.test/5","https://a.test/6"]`, "A client may have at most 5 redirect URIs"},
|
||||
{"http-non-loopback", `["http://evil.example.test/cb"]`, "Redirect URI must be https:// or loopback http://127.0.0.1 / http://localhost: http://evil.example.test/cb"},
|
||||
{"not-a-url", `["not-a-url"]`, "Redirect URI is not a valid URL: not-a-url"},
|
||||
}
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
srv := newTestServer(newMemoryBackend())
|
||||
body := `{"redirect_uris":` + tc.uris + `}`
|
||||
req := httptest.NewRequest(http.MethodPost, "/oauth/mcp/register", strings.NewReader(body))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Register(rec, req)
|
||||
assertRegisterError(t, rec, http.StatusBadRequest, "invalid_redirect_uri", tc.want)
|
||||
})
|
||||
}
|
||||
|
||||
t.Run("too-long", func(t *testing.T) {
|
||||
srv := newTestServer(newMemoryBackend())
|
||||
longURI := "https://a.test/" + strings.Repeat("x", 512)
|
||||
bodyBytes, _ := json.Marshal(map[string][]string{"redirect_uris": {longURI}})
|
||||
req := httptest.NewRequest(http.MethodPost, "/oauth/mcp/register", bytes.NewReader(bodyBytes))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Register(rec, req)
|
||||
assertRegisterError(t, rec, http.StatusBadRequest, "invalid_redirect_uri", "Each redirect URI must be at most 512 characters.")
|
||||
})
|
||||
}
|
||||
|
||||
func TestRegisterUnsupportedAuthMethod(t *testing.T) {
|
||||
srv := newTestServer(newMemoryBackend())
|
||||
body := `{"redirect_uris":["https://client.example.test/callback"],"token_endpoint_auth_method":"bogus"}`
|
||||
req := httptest.NewRequest(http.MethodPost, "/oauth/mcp/register", strings.NewReader(body))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Register(rec, req)
|
||||
assertRegisterError(t, rec, http.StatusBadRequest, "invalid_client_metadata", "token_endpoint_auth_method is not supported")
|
||||
}
|
||||
|
||||
func TestRegisterUnsupportedGrantType(t *testing.T) {
|
||||
srv := newTestServer(newMemoryBackend())
|
||||
body := `{"redirect_uris":["https://client.example.test/callback"],"grant_types":["client_credentials"]}`
|
||||
req := httptest.NewRequest(http.MethodPost, "/oauth/mcp/register", strings.NewReader(body))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Register(rec, req)
|
||||
assertRegisterError(t, rec, http.StatusBadRequest, "invalid_client_metadata", "grant_types contains an unsupported value")
|
||||
}
|
||||
|
||||
func TestRegisterGrantTypesMustIncludeAuthorizationCode(t *testing.T) {
|
||||
srv := newTestServer(newMemoryBackend())
|
||||
body := `{"redirect_uris":["https://client.example.test/callback"],"grant_types":["refresh_token"]}`
|
||||
req := httptest.NewRequest(http.MethodPost, "/oauth/mcp/register", strings.NewReader(body))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Register(rec, req)
|
||||
assertRegisterError(t, rec, http.StatusBadRequest, "invalid_client_metadata", "grant_types must include authorization_code")
|
||||
}
|
||||
|
||||
func TestRegisterUnsupportedResponseType(t *testing.T) {
|
||||
srv := newTestServer(newMemoryBackend())
|
||||
body := `{"redirect_uris":["https://client.example.test/callback"],"response_types":["token"]}`
|
||||
req := httptest.NewRequest(http.MethodPost, "/oauth/mcp/register", strings.NewReader(body))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Register(rec, req)
|
||||
assertRegisterError(t, rec, http.StatusBadRequest, "invalid_client_metadata", "response_types must be code")
|
||||
}
|
||||
|
||||
func TestRegisterControlCharacterNameCleaning(t *testing.T) {
|
||||
srv := newTestServer(newMemoryBackend())
|
||||
dirty := "Evil\x00Name\x1FWith\x7FControls"
|
||||
bodyBytes, _ := json.Marshal(map[string]any{
|
||||
"redirect_uris": []string{"https://client.example.test/callback"},
|
||||
"client_name": dirty,
|
||||
})
|
||||
req := httptest.NewRequest(http.MethodPost, "/oauth/mcp/register", bytes.NewReader(bodyBytes))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Register(rec, req)
|
||||
|
||||
if rec.Code != http.StatusCreated {
|
||||
t.Fatalf("status = %d, body=%s", rec.Code, rec.Body.String())
|
||||
}
|
||||
var got map[string]any
|
||||
if err := json.Unmarshal(rec.Body.Bytes(), &got); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
name, _ := got["client_name"].(string)
|
||||
if strings.ContainsAny(name, "\x00\x1f\x7f") {
|
||||
t.Fatalf("client_name %q still contains control characters", name)
|
||||
}
|
||||
if name != "EvilNameWithControls" {
|
||||
t.Fatalf("client_name = %q, want stripped \"EvilNameWithControls\"", name)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRegisterCapReached(t *testing.T) {
|
||||
backend := newMemoryBackend()
|
||||
srv := newTestServer(backend)
|
||||
srv.opts.DCRClientCap = 1
|
||||
|
||||
body := `{"redirect_uris":["https://client.example.test/callback"]}`
|
||||
|
||||
req1 := httptest.NewRequest(http.MethodPost, "/oauth/mcp/register", strings.NewReader(body))
|
||||
req1.Header.Set("Content-Type", "application/json")
|
||||
rec1 := httptest.NewRecorder()
|
||||
srv.Register(rec1, req1)
|
||||
if rec1.Code != http.StatusCreated {
|
||||
t.Fatalf("first registration status = %d, body=%s", rec1.Code, rec1.Body.String())
|
||||
}
|
||||
|
||||
req2 := httptest.NewRequest(http.MethodPost, "/oauth/mcp/register", strings.NewReader(body))
|
||||
req2.Header.Set("Content-Type", "application/json")
|
||||
rec2 := httptest.NewRecorder()
|
||||
srv.Register(rec2, req2)
|
||||
assertRegisterError(t, rec2, http.StatusBadRequest, "invalid_client_metadata", "Registration temporarily unavailable")
|
||||
}
|
||||
|
||||
func TestRegisterSweepsStaleUnconsentedButKeepsArtisanClients(t *testing.T) {
|
||||
backend := newMemoryBackend()
|
||||
srv := newTestServer(backend)
|
||||
|
||||
staleIP := "203.0.113.9"
|
||||
backend.clients = append(backend.clients,
|
||||
&ClientRecord{ClientID: "stale-dynamic", RegistrationIP: &staleIP, CreatedAt: time.Now().Add(-48 * time.Hour)},
|
||||
&ClientRecord{ClientID: "artisan-client", RegistrationIP: nil, CreatedAt: time.Now().Add(-48 * time.Hour)},
|
||||
)
|
||||
|
||||
body := `{"redirect_uris":["https://client.example.test/callback"]}`
|
||||
req := httptest.NewRequest(http.MethodPost, "/oauth/mcp/register", strings.NewReader(body))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Register(rec, req)
|
||||
if rec.Code != http.StatusCreated {
|
||||
t.Fatalf("status = %d, body=%s", rec.Code, rec.Body.String())
|
||||
}
|
||||
|
||||
backend.mu.Lock()
|
||||
defer backend.mu.Unlock()
|
||||
var ids []string
|
||||
for _, c := range backend.clients {
|
||||
ids = append(ids, c.ClientID)
|
||||
}
|
||||
if containsString(ids, "stale-dynamic") {
|
||||
t.Fatalf("stale unconsented dynamic client survived sweep: %v", ids)
|
||||
}
|
||||
if !containsString(ids, "artisan-client") {
|
||||
t.Fatalf("artisan client (nil RegistrationIP) was swept: %v", ids)
|
||||
}
|
||||
}
|
||||
|
||||
func containsString(ss []string, v string) bool {
|
||||
for _, s := range ss {
|
||||
if s == v {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
func assertRegisterError(t *testing.T, rec *httptest.ResponseRecorder, status int, code, description string) {
|
||||
t.Helper()
|
||||
if rec.Code != status {
|
||||
t.Fatalf("status = %d, want %d (body=%s)", rec.Code, status, rec.Body.String())
|
||||
}
|
||||
var got map[string]any
|
||||
if err := json.Unmarshal(rec.Body.Bytes(), &got); err != nil {
|
||||
t.Fatalf("decode error body: %v (body=%s)", err, rec.Body.String())
|
||||
}
|
||||
if got["error"] != code {
|
||||
t.Fatalf("error = %v, want %q", got["error"], code)
|
||||
}
|
||||
if got["error_description"] != description {
|
||||
t.Fatalf("error_description = %v, want %q", got["error_description"], description)
|
||||
}
|
||||
if cc := rec.Header().Get("Cache-Control"); cc != "no-store, private" {
|
||||
t.Fatalf("Cache-Control = %q, want \"no-store, private\"", cc)
|
||||
}
|
||||
}
|
||||
199
modules/wristband/server.go
Normal file
199
modules/wristband/server.go
Normal file
@@ -0,0 +1,199 @@
|
||||
// Package wristband implements the app-agnostic RFC 8414 / OAuth
|
||||
// authorization-server surface ported from Płytarium's hand-rolled PHP OAuth
|
||||
// server (08-CONTEXT.md D-05). It never imports an application package, a
|
||||
// GORM type, or any fonoteka model: every deployment-specific value (issuer,
|
||||
// scopes, endpoint paths, TTLs) arrives through Options, and every app-owned
|
||||
// concern (users, collections, persistence) stays out of this package.
|
||||
//
|
||||
// D-06: PHP's RFC-minimal response shapes are wristband's defaults. There
|
||||
// are no response hooks; callers cannot alter the wire bytes beyond the
|
||||
// values exposed on Options.
|
||||
package wristband
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"time"
|
||||
)
|
||||
|
||||
// Options configures a Server's advertised endpoints and metadata values.
|
||||
// Every field has a PHP-parity default via DefaultOptions except Issuer,
|
||||
// which the caller must set from app.url with its trailing slash trimmed
|
||||
// exactly once (D-03). wristband never hardcodes an app's issuer.
|
||||
type Options struct {
|
||||
// Issuer is app.url with exactly one trailing slash trimmed by the
|
||||
// caller. Every metadata endpoint URL is built by appending a fixed
|
||||
// RFC path suffix to Issuer.
|
||||
Issuer string
|
||||
|
||||
// ServiceDocumentationPath is appended to Issuer for the metadata
|
||||
// service_documentation field. PHP default: "/help".
|
||||
ServiceDocumentationPath string
|
||||
|
||||
// ScopesSupported is the RFC 8414 scopes_supported list. PHP default:
|
||||
// ["read","write","ai","offline_access"].
|
||||
ScopesSupported []string
|
||||
|
||||
// TokenEndpointAuthMethodsSupported is the RFC 8414
|
||||
// token_endpoint_auth_methods_supported list. PHP default:
|
||||
// ["none","client_secret_post","client_secret_basic"].
|
||||
TokenEndpointAuthMethodsSupported []string
|
||||
|
||||
// AuthorizationResponseIssParameterSupported is the RFC 9207 metadata
|
||||
// capability flag. PHP default: true.
|
||||
AuthorizationResponseIssParameterSupported bool
|
||||
|
||||
// DCRClientCap is the maximum number of unrevoked OAuth clients RFC
|
||||
// 7591 registration allows (PHP OAuthRegisterController::MAX_CLIENTS).
|
||||
// PHP default: 200 (D-03).
|
||||
DCRClientCap int
|
||||
|
||||
// DCRUnconsentedSweepAge is how old an unconsented, dynamically
|
||||
// registered client (non-nil RegistrationIP) must be before
|
||||
// registration sweeps it. PHP default: 24h (D-03).
|
||||
DCRUnconsentedSweepAge time.Duration
|
||||
|
||||
// RegisterMaxBodyBytes bounds the RFC 7591 registration request body
|
||||
// before JSON decoding (D-21, T-08-DCR-FLOOD). PHP default: 65536 (64 KiB).
|
||||
RegisterMaxBodyBytes int64
|
||||
|
||||
// Resource is the expected RFC 8707 resource indicator value authorize
|
||||
// checks an optional resource query parameter against (PHP
|
||||
// config('fonoteka.mcp.resource'), D-03). PHP default:
|
||||
// "https://mcp.plytarium.com/mcp".
|
||||
Resource string
|
||||
|
||||
// PendingRequestTTL is how long a pre-consent pending authorization row
|
||||
// created by authorize stays valid (PHP
|
||||
// OAuthCodeManager::PENDING_TTL_SECONDS, D-03). PHP default: 600s.
|
||||
PendingRequestTTL time.Duration
|
||||
|
||||
// CodeTTL is how long an issued authorization code stays valid after
|
||||
// consent (PHP OAuthCodeManager::CODE_TTL_SECONDS, D-03). PHP default:
|
||||
// 600s. Consent issuance always sets a fresh expiry from this TTL
|
||||
// rather than reusing the pending row's original expiry.
|
||||
CodeTTL time.Duration
|
||||
|
||||
// AccessTokenTTL is how long an inv_ access token minted by a successful
|
||||
// code exchange or refresh rotation stays valid (PHP
|
||||
// OAuthCodeManager::ACCESS_TTL_SECONDS, D-03). PHP default: 3600s (1h).
|
||||
AccessTokenTTL time.Duration
|
||||
|
||||
// RefreshTokenTTL is how long a refresh-token lineage row stays valid
|
||||
// from issuance (PHP OAuthCodeManager::REFRESH_TTL_DAYS, D-03). PHP
|
||||
// default: 30 days.
|
||||
RefreshTokenTTL time.Duration
|
||||
}
|
||||
|
||||
// DefaultOptions returns PHP-parity defaults for every metadata option
|
||||
// other than Issuer, which the caller must set from app.url.
|
||||
func DefaultOptions() Options {
|
||||
return Options{
|
||||
ServiceDocumentationPath: "/help",
|
||||
ScopesSupported: []string{"read", "write", "ai", "offline_access"},
|
||||
TokenEndpointAuthMethodsSupported: []string{"none", "client_secret_post", "client_secret_basic"},
|
||||
AuthorizationResponseIssParameterSupported: true,
|
||||
DCRClientCap: 200,
|
||||
DCRUnconsentedSweepAge: 24 * time.Hour,
|
||||
RegisterMaxBodyBytes: 65536,
|
||||
Resource: "https://mcp.plytarium.com/mcp",
|
||||
PendingRequestTTL: 600 * time.Second,
|
||||
CodeTTL: 600 * time.Second,
|
||||
AccessTokenTTL: 3600 * time.Second,
|
||||
RefreshTokenTTL: 30 * 24 * time.Hour,
|
||||
}
|
||||
}
|
||||
|
||||
// Server is the app-agnostic wristband authorization-server surface. It is
|
||||
// constructed with Options and never imports an application package.
|
||||
type Server struct {
|
||||
opts Options
|
||||
backend Backend
|
||||
|
||||
// now and randomBytes are deterministic clock/entropy seams so tests
|
||||
// can control timestamps and generated secrets without depending on
|
||||
// wall-clock time or true randomness (08-02-PLAN.md Task 2).
|
||||
now func() time.Time
|
||||
randomBytes func(n int) (string, error)
|
||||
}
|
||||
|
||||
// NewServer constructs a Server from Options. The backend is nil until
|
||||
// SetBackend is called (D-09: the metadata route needs no backend at all,
|
||||
// so plugin boot can construct a Server before a *gorm.DB is available).
|
||||
func NewServer(opts Options) *Server {
|
||||
return &Server{
|
||||
opts: opts,
|
||||
now: time.Now,
|
||||
randomBytes: randomBase64URL,
|
||||
}
|
||||
}
|
||||
|
||||
// SetBackend attaches the app's transaction-scoped store bundle. Handlers
|
||||
// that need persistence (Register) return an opaque 500 until this is
|
||||
// called.
|
||||
func (s *Server) SetBackend(b Backend) {
|
||||
s.backend = b
|
||||
}
|
||||
|
||||
// metadataDocument is the exact unwrapped RFC 8414 body. Field order matches
|
||||
// the PHP array literal in OAuthMetadataController::show() byte for byte;
|
||||
// encoding/json preserves struct declaration order, so this struct is the
|
||||
// single source of truth for the wire order.
|
||||
type metadataDocument struct {
|
||||
Issuer string `json:"issuer"`
|
||||
AuthorizationEndpoint string `json:"authorization_endpoint"`
|
||||
TokenEndpoint string `json:"token_endpoint"`
|
||||
RegistrationEndpoint string `json:"registration_endpoint"`
|
||||
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"`
|
||||
ScopesSupported []string `json:"scopes_supported"`
|
||||
ServiceDocumentation string `json:"service_documentation"`
|
||||
AuthorizationResponseIssParameterSupported bool `json:"authorization_response_iss_parameter_supported"`
|
||||
}
|
||||
|
||||
// Metadata handles GET /.well-known/oauth-authorization-server, writing the
|
||||
// exact unwrapped RFC 8414 document (D-06). response_types_supported,
|
||||
// grant_types_supported and code_challenge_methods_supported are fixed
|
||||
// protocol constants, not Options: this phase's authorization server only
|
||||
// ever supports the authorization_code/refresh_token grants with S256 PKCE
|
||||
// (D-01), so there is nothing app-specific to configure there.
|
||||
func (s *Server) Metadata(w http.ResponseWriter, r *http.Request) {
|
||||
doc := metadataDocument{
|
||||
Issuer: s.opts.Issuer,
|
||||
AuthorizationEndpoint: s.opts.Issuer + "/oauth/mcp/authorize",
|
||||
TokenEndpoint: s.opts.Issuer + "/oauth/mcp/token",
|
||||
RegistrationEndpoint: s.opts.Issuer + "/oauth/mcp/register",
|
||||
ResponseTypesSupported: []string{"code"},
|
||||
GrantTypesSupported: []string{"authorization_code", "refresh_token"},
|
||||
CodeChallengeMethodsSupported: []string{"S256"},
|
||||
TokenEndpointAuthMethodsSupported: s.opts.TokenEndpointAuthMethodsSupported,
|
||||
ScopesSupported: s.opts.ScopesSupported,
|
||||
ServiceDocumentation: s.opts.Issuer + s.opts.ServiceDocumentationPath,
|
||||
AuthorizationResponseIssParameterSupported: s.opts.AuthorizationResponseIssParameterSupported,
|
||||
}
|
||||
writeExactJSON(w, http.StatusOK, doc, map[string]string{"Cache-Control": "no-cache, private"})
|
||||
}
|
||||
|
||||
// writeExactJSON writes v as an unwrapped, no-trailing-newline JSON document
|
||||
// (matching the wire/response.go WriteJSON technique) but never falls back to
|
||||
// the house opaque-500 envelope: raw RFC responses must never acquire a
|
||||
// house-shaped body (D-06/D-09).
|
||||
func writeExactJSON(w http.ResponseWriter, status int, v any, extraHeaders map[string]string) {
|
||||
var buf bytes.Buffer
|
||||
enc := json.NewEncoder(&buf)
|
||||
enc.SetEscapeHTML(false)
|
||||
if err := enc.Encode(v); err != nil {
|
||||
w.WriteHeader(http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
h := w.Header()
|
||||
h.Set("Content-Type", "application/json")
|
||||
for k, v := range extraHeaders {
|
||||
h.Set(k, v)
|
||||
}
|
||||
w.WriteHeader(status)
|
||||
_, _ = w.Write(bytes.TrimSuffix(buf.Bytes(), []byte("\n")))
|
||||
}
|
||||
123
modules/wristband/server_test.go
Normal file
123
modules/wristband/server_test.go
Normal file
@@ -0,0 +1,123 @@
|
||||
package wristband
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"testing"
|
||||
)
|
||||
|
||||
// TestPhase8RedMetadata is the Phase 8 Wave 1 RED anchor (08-CONTEXT.md
|
||||
// D-06). It asserts the exact unwrapped RFC 8414 metadata document recorded
|
||||
// from the live PHP fixture (parity/fixtures/routes/GET__.well-known_oauth-authorization-server_oauth.yaml)
|
||||
// and fails with the PHASE8_RED:metadata sentinel while Server.Metadata is a
|
||||
// stub. scripts/check-phase8-red.sh verifies this failure is fail-closed.
|
||||
func TestPhase8RedMetadata(t *testing.T) {
|
||||
opts := DefaultOptions()
|
||||
opts.Issuer = "https://plytarium.com"
|
||||
srv := NewServer(opts)
|
||||
|
||||
req := httptest.NewRequest(http.MethodGet, "/.well-known/oauth-authorization-server", nil)
|
||||
req.Header.Set("Accept", "application/json")
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Metadata(rec, req)
|
||||
|
||||
if rec.Code != http.StatusOK {
|
||||
t.Fatalf("PHASE8_RED:metadata: status = %d, want %d", rec.Code, http.StatusOK)
|
||||
}
|
||||
|
||||
const want = `{"issuer":"https://plytarium.com","authorization_endpoint":"https://plytarium.com/oauth/mcp/authorize","token_endpoint":"https://plytarium.com/oauth/mcp/token","registration_endpoint":"https://plytarium.com/oauth/mcp/register","response_types_supported":["code"],"grant_types_supported":["authorization_code","refresh_token"],"code_challenge_methods_supported":["S256"],"token_endpoint_auth_methods_supported":["none","client_secret_post","client_secret_basic"],"scopes_supported":["read","write","ai","offline_access"],"service_documentation":"https://plytarium.com/help","authorization_response_iss_parameter_supported":true}`
|
||||
|
||||
if got := rec.Body.String(); got != want {
|
||||
t.Fatalf("PHASE8_RED:metadata: body mismatch\n got: %s\nwant: %s", got, want)
|
||||
}
|
||||
if ct := rec.Header().Get("Content-Type"); ct != "application/json" {
|
||||
t.Fatalf("PHASE8_RED:metadata: Content-Type = %q, want application/json", ct)
|
||||
}
|
||||
if cc := rec.Header().Get("Cache-Control"); cc != "no-cache, private" {
|
||||
t.Fatalf("PHASE8_RED:metadata: Cache-Control = %q, want \"no-cache, private\"", cc)
|
||||
}
|
||||
}
|
||||
|
||||
// TestMetadataExactBytes is the GREEN direct-handler regression for
|
||||
// TestPhase8RedMetadata: same PHP-fixture bytes, now expected to pass.
|
||||
func TestMetadataExactBytes(t *testing.T) {
|
||||
opts := DefaultOptions()
|
||||
opts.Issuer = "https://plytarium.com"
|
||||
srv := NewServer(opts)
|
||||
|
||||
req := httptest.NewRequest(http.MethodGet, "/.well-known/oauth-authorization-server", nil)
|
||||
req.Header.Set("Accept", "application/json")
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Metadata(rec, req)
|
||||
|
||||
if rec.Code != http.StatusOK {
|
||||
t.Fatalf("status = %d, want %d", rec.Code, http.StatusOK)
|
||||
}
|
||||
|
||||
const want = `{"issuer":"https://plytarium.com","authorization_endpoint":"https://plytarium.com/oauth/mcp/authorize","token_endpoint":"https://plytarium.com/oauth/mcp/token","registration_endpoint":"https://plytarium.com/oauth/mcp/register","response_types_supported":["code"],"grant_types_supported":["authorization_code","refresh_token"],"code_challenge_methods_supported":["S256"],"token_endpoint_auth_methods_supported":["none","client_secret_post","client_secret_basic"],"scopes_supported":["read","write","ai","offline_access"],"service_documentation":"https://plytarium.com/help","authorization_response_iss_parameter_supported":true}`
|
||||
|
||||
if got := rec.Body.String(); got != want {
|
||||
t.Fatalf("body mismatch\n got: %s\nwant: %s", got, want)
|
||||
}
|
||||
if strings.HasSuffix(rec.Body.String(), "\n") {
|
||||
t.Fatal("body has a trailing newline, want none")
|
||||
}
|
||||
if _, hasData := decodeAsMap(t, rec.Body.Bytes())["data"]; hasData {
|
||||
t.Fatal("body has a house \"data\" envelope, want unwrapped RFC 8414 document")
|
||||
}
|
||||
if ct := rec.Header().Get("Content-Type"); ct != "application/json" {
|
||||
t.Fatalf("Content-Type = %q, want application/json", ct)
|
||||
}
|
||||
if cc := rec.Header().Get("Cache-Control"); cc != "no-cache, private" {
|
||||
t.Fatalf("Cache-Control = %q, want %q", cc, "no-cache, private")
|
||||
}
|
||||
}
|
||||
|
||||
// TestMetadataUsesConfiguredOptions proves service_documentation,
|
||||
// scopes_supported and the auth-methods list come from Options, not a
|
||||
// hardcoded literal, while the three protocol-constant fields never change
|
||||
// (D-06: only those four fields are configurable).
|
||||
func TestMetadataUsesConfiguredOptions(t *testing.T) {
|
||||
opts := Options{
|
||||
Issuer: "https://example.test",
|
||||
ServiceDocumentationPath: "/docs/oauth",
|
||||
ScopesSupported: []string{"read"},
|
||||
TokenEndpointAuthMethodsSupported: []string{"client_secret_post"},
|
||||
AuthorizationResponseIssParameterSupported: false,
|
||||
}
|
||||
srv := NewServer(opts)
|
||||
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Metadata(rec, httptest.NewRequest(http.MethodGet, "/.well-known/oauth-authorization-server", nil))
|
||||
|
||||
var got map[string]any
|
||||
if err := json.Unmarshal(rec.Body.Bytes(), &got); err != nil {
|
||||
t.Fatalf("decode response: %v", err)
|
||||
}
|
||||
if got["service_documentation"] != "https://example.test/docs/oauth" {
|
||||
t.Fatalf("service_documentation = %v, want configured path", got["service_documentation"])
|
||||
}
|
||||
if scopes, _ := got["scopes_supported"].([]any); len(scopes) != 1 || scopes[0] != "read" {
|
||||
t.Fatalf("scopes_supported = %v, want [read]", got["scopes_supported"])
|
||||
}
|
||||
if methods, _ := got["token_endpoint_auth_methods_supported"].([]any); len(methods) != 1 || methods[0] != "client_secret_post" {
|
||||
t.Fatalf("token_endpoint_auth_methods_supported = %v, want [client_secret_post]", got["token_endpoint_auth_methods_supported"])
|
||||
}
|
||||
if got["authorization_response_iss_parameter_supported"] != false {
|
||||
t.Fatalf("authorization_response_iss_parameter_supported = %v, want false", got["authorization_response_iss_parameter_supported"])
|
||||
}
|
||||
if rts, _ := got["response_types_supported"].([]any); len(rts) != 1 || rts[0] != "code" {
|
||||
t.Fatalf("response_types_supported = %v, want the fixed [code] constant", got["response_types_supported"])
|
||||
}
|
||||
}
|
||||
|
||||
func decodeAsMap(t *testing.T, body []byte) map[string]any {
|
||||
t.Helper()
|
||||
var m map[string]any
|
||||
if err := json.Unmarshal(body, &m); err != nil {
|
||||
t.Fatalf("decode response: %v", err)
|
||||
}
|
||||
return m
|
||||
}
|
||||
177
modules/wristband/stores.go
Normal file
177
modules/wristband/stores.go
Normal file
@@ -0,0 +1,177 @@
|
||||
package wristband
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"time"
|
||||
)
|
||||
|
||||
// ErrClientCapReached is returned by ClientStore.CreateWithCap when the
|
||||
// unrevoked client count is already at or above the configured cap
|
||||
// (D-03/D-21, T-08-DCR-FLOOD). Register translates it into the exact PHP
|
||||
// invalid_client_metadata "Registration temporarily unavailable" body.
|
||||
var ErrClientCapReached = errors.New("wristband: dynamic client registration cap reached")
|
||||
|
||||
// ClientRecord is the app-agnostic persisted shape of an OAuth client row.
|
||||
// fonoteka's GORM adapter (classes/auth/oauth_store.go) converts to/from
|
||||
// its models.OAuthClient; wristband never imports GORM or fonoteka.
|
||||
type ClientRecord struct {
|
||||
ID uint
|
||||
ClientID string
|
||||
ClientSecretHash *string // nil for public (auth method "none") clients
|
||||
ClientName string
|
||||
RedirectURIs []string
|
||||
GrantTypes []string
|
||||
TokenEndpointAuthMethod string
|
||||
RegistrationIP *string // nil for artisan-issued clients (D-19); never swept
|
||||
ConsentedAt *time.Time
|
||||
RevokedAt *time.Time
|
||||
ScopeCeiling []string // nil/empty means no ceiling
|
||||
CreatedAt time.Time
|
||||
}
|
||||
|
||||
// AuthCodeRecord is the app-agnostic persisted shape of a pending or issued
|
||||
// authorization row. Its full lifecycle (issue, exchange, replay) belongs to
|
||||
// a later Phase 8 plan (08-03/08-04); this plan only needs the shape and the
|
||||
// row-lock read so the Tx bundle is complete for T-08-CODE-REPLAY.
|
||||
type AuthCodeRecord struct {
|
||||
ID uint
|
||||
RequestID *string // non-nil while pending; nulled once a code is issued
|
||||
CodeHash *string // nil while pending; set once a code is issued
|
||||
ClientID string
|
||||
UserID *uint // nil until consent
|
||||
RedirectURI string
|
||||
Scopes []string
|
||||
CollectionIDs []uint
|
||||
CodeChallenge string
|
||||
CodeChallengeMethod string
|
||||
Resource *string
|
||||
State *string
|
||||
ExpiresAt time.Time
|
||||
UsedAt *time.Time
|
||||
OfflineAccess bool
|
||||
}
|
||||
|
||||
// RefreshTokenRecord is the app-agnostic persisted shape of a refresh-token
|
||||
// lineage row. Rotation/replay semantics belong to a later plan (08-04);
|
||||
// this plan only needs the shape and the row-lock read for T-08-REFRESH-REPLAY.
|
||||
type RefreshTokenRecord struct {
|
||||
ID uint
|
||||
TokenHash string
|
||||
APITokenID *uint
|
||||
ClientID string
|
||||
UserID uint
|
||||
Scopes []string
|
||||
CollectionIDs []uint
|
||||
ExpiresAt time.Time
|
||||
RevokedAt *time.Time
|
||||
RotatedToID *uint
|
||||
OfflineAccess bool
|
||||
}
|
||||
|
||||
// IssuedToken is what an AccessTokenIssuer mints: the one-time raw secret
|
||||
// plus the persisted row id (for later Revoke calls).
|
||||
type IssuedToken struct {
|
||||
ID uint
|
||||
Secret string
|
||||
}
|
||||
|
||||
// ClientStore persists OAuthClient rows. This plan (08-02) exercises
|
||||
// ByClientID, CreateWithCap and SweepUnconsented through Register; Revoke
|
||||
// belongs to a later connected-apps plan.
|
||||
type ClientStore interface {
|
||||
ByClientID(ctx context.Context, clientID string) (*ClientRecord, error)
|
||||
// CreateWithCap creates rec only when the unrevoked client count is
|
||||
// below cap, atomically with the count check (T-08-DCR-FLOOD). It
|
||||
// returns ErrClientCapReached, leaving no row created, when the cap is
|
||||
// already reached. On success it fills rec.ID and rec.CreatedAt.
|
||||
CreateWithCap(ctx context.Context, rec *ClientRecord, cap int) error
|
||||
// SweepUnconsented deletes dynamically-registered (non-nil
|
||||
// RegistrationIP), still-unconsented clients created before olderThan.
|
||||
// Artisan-issued clients (nil RegistrationIP) are never swept (D-19).
|
||||
SweepUnconsented(ctx context.Context, olderThan time.Time) error
|
||||
// MarkConsented stamps ConsentedAt once for clientID unless it is
|
||||
// already set (idempotent; PHP OAuthConsentController::store's "if
|
||||
// ($client->consented_at === null)" guard, 08-05-PLAN.md D-08). A
|
||||
// consented client is never later swept by SweepUnconsented.
|
||||
MarkConsented(ctx context.Context, clientID string) error
|
||||
}
|
||||
|
||||
// AuthCodeStore persists pending/issued authorization rows. ByCodeHashForUpdate
|
||||
// is the row-lock read a later plan's code-exchange/replay-kill logic needs
|
||||
// (T-08-CODE-REPLAY); it ships now so the Tx bundle does not change shape
|
||||
// later.
|
||||
type AuthCodeStore interface {
|
||||
CreatePending(ctx context.Context, rec *AuthCodeRecord) error
|
||||
ByRequestID(ctx context.Context, requestID string) (*AuthCodeRecord, error)
|
||||
ByCodeHashForUpdate(ctx context.Context, codeHash string) (*AuthCodeRecord, error)
|
||||
// MarkIssued turns a pending row into an issued authorization code
|
||||
// (PHP OAuthCodeManager::issueCode, 08-05-PLAN.md D-08): it nulls
|
||||
// RequestID, sets CodeHash/UserID, overwrites Scopes/CollectionIDs
|
||||
// with the consent-granted values (never the originally requested
|
||||
// ones), and extends ExpiresAt to the fresh code TTL.
|
||||
MarkIssued(ctx context.Context, id uint, codeHash string, userID uint, scopes []string, collectionIDs []uint, expiresAt time.Time) error
|
||||
MarkUsed(ctx context.Context, id uint) error
|
||||
// DeleteExpiredCodes removes pending and issued authorization-code rows
|
||||
// whose ExpiresAt is before now (D-17: wristband adds an expiry sweep
|
||||
// PHP lacks). used_at/request_id status is irrelevant to the decision:
|
||||
// only expiry drives deletion, so unexpired issued-but-unused rows and
|
||||
// unexpired used rows both survive untouched.
|
||||
DeleteExpiredCodes(ctx context.Context, now time.Time) error
|
||||
}
|
||||
|
||||
// RefreshTokenStore persists refresh-token lineage rows. ByTokenHashForUpdate
|
||||
// is the row-lock read a later plan's rotation/replay-kill logic needs
|
||||
// (T-08-REFRESH-REPLAY).
|
||||
type RefreshTokenStore interface {
|
||||
Create(ctx context.Context, rec *RefreshTokenRecord) error
|
||||
ByTokenHashForUpdate(ctx context.Context, tokenHash string) (*RefreshTokenRecord, error)
|
||||
// ByAPITokenIDForUpdate row-locks the refresh row currently linked to
|
||||
// apiTokenID, if any (08-06-PLAN.md D-08). It is the seam a
|
||||
// connected-app revoke uses to find the lineage to kill without
|
||||
// wristband inventing its own SQL join; a nil result (no linked
|
||||
// refresh row) is not an error.
|
||||
ByAPITokenIDForUpdate(ctx context.Context, apiTokenID uint) (*RefreshTokenRecord, error)
|
||||
// MarkRotated links a spent-by-rotation predecessor to its successor
|
||||
// (PHP OAuthCodeManager::rotateRefresh's `$record->rotated_to_id =
|
||||
// $successor->id`). The predecessor's own RevokedAt stays nil: a
|
||||
// rotated-but-not-yet-replayed row remains retrievable as replay
|
||||
// evidence (D-17); "already rotated" is signaled by RotatedToID, not
|
||||
// RevokedAt.
|
||||
MarkRotated(ctx context.Context, id uint, successorID uint) error
|
||||
// RevokeLineage stamps RevokedAt on startID and every row it was
|
||||
// rotated to (walking forward through RotatedToID), and revokes each
|
||||
// visited row's linked access token too (T-08-REFRESH-REPLAY). Rows
|
||||
// are not deleted: unexpired revoked rows stay as replay evidence
|
||||
// (D-17).
|
||||
RevokeLineage(ctx context.Context, startID uint) error
|
||||
// DeleteExpiredRefreshTokens removes refresh rows whose ExpiresAt is
|
||||
// before now (D-17). Revoked/rotated-but-unexpired rows are untouched
|
||||
// so replay detection and connected-apps history stay correct.
|
||||
DeleteExpiredRefreshTokens(ctx context.Context, now time.Time) error
|
||||
}
|
||||
|
||||
// AccessTokenIssuer mints/revokes the app's ordinary personal access token
|
||||
// (fonoteka: an inv_ token via ApiTokenManager) and stamps the owning OAuth
|
||||
// client id.
|
||||
type AccessTokenIssuer interface {
|
||||
Mint(ctx context.Context, userID uint, name string, scopes []string, expiresAt time.Time, collectionIDs []uint, clientID string) (IssuedToken, error)
|
||||
Revoke(ctx context.Context, tokenID uint) error
|
||||
}
|
||||
|
||||
// Tx bundles every store/issuer onto one transaction-scoped handle so
|
||||
// sweep+cap+create (this plan) and later code-exchange/refresh-rotation
|
||||
// cannot straddle two transactions (08-RESEARCH.md Pattern 1).
|
||||
type Tx interface {
|
||||
ClientStore
|
||||
AuthCodeStore
|
||||
RefreshTokenStore
|
||||
AccessTokenIssuer
|
||||
}
|
||||
|
||||
// Backend opens one transaction-scoped Tx per call. The GORM adapter lives
|
||||
// in fonoteka.go's classes/auth package (D-07): wristband never imports
|
||||
// gorm.io/gorm or a fonoteka model.
|
||||
type Backend interface {
|
||||
WithinTx(ctx context.Context, fn func(Tx) error) error
|
||||
}
|
||||
417
modules/wristband/token.go
Normal file
417
modules/wristband/token.go
Normal file
@@ -0,0 +1,417 @@
|
||||
// RFC 6749 token endpoint for MCP OAuth, ported from PHP
|
||||
// OAuthTokenController::token / OAuthCodeManager::exchangeCode/rotateRefresh
|
||||
// byte-for-byte including their validation order (08-CONTEXT.md
|
||||
// D-02/D-04/D-05/D-07; canonical PHP source: OAuthTokenController.php,
|
||||
// OAuthCodeManager.php).
|
||||
//
|
||||
// 08-06-PLAN.md completes grant_type=refresh_token: rotation with
|
||||
// lineage-kill replay detection (T-08-REFRESH-REPLAY) and the D-17 expiry
|
||||
// sweep that also runs here (in addition to /register).
|
||||
package wristband
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/base64"
|
||||
"errors"
|
||||
"net/http"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// errInvalidClient is Token's internal signal that client authentication
|
||||
// failed (T-08-SECRET-TIMING). It never reaches a response body directly:
|
||||
// Token translates it into the exact 401 invalid_client body plus
|
||||
// WWW-Authenticate: Basic realm="OAuth" (D-04/D-06).
|
||||
var errInvalidClient = errors.New("wristband: invalid client")
|
||||
|
||||
// errInvalidGrant is Token's internal signal for every exact PHP
|
||||
// OAuthInvalidGrantException case (missing/wrong code, expired code, PKCE
|
||||
// mismatch, client/redirect/resource binding failure, replay). Token
|
||||
// translates it into the exact 400 invalid_grant body with no description.
|
||||
var errInvalidGrant = errors.New("wristband: invalid grant")
|
||||
|
||||
// tokenIssueResult is what a successful grant produces: the two raw secrets
|
||||
// plus the response's scope/offline_access ingredients.
|
||||
type tokenIssueResult struct {
|
||||
AccessToken string
|
||||
RefreshToken string
|
||||
Scopes []string
|
||||
OfflineAccess bool
|
||||
}
|
||||
|
||||
type tokenSuccessBody struct {
|
||||
AccessToken string `json:"access_token"`
|
||||
TokenType string `json:"token_type"`
|
||||
ExpiresIn int64 `json:"expires_in"`
|
||||
RefreshToken string `json:"refresh_token"`
|
||||
Scope string `json:"scope"`
|
||||
}
|
||||
|
||||
type tokenErrorBody struct {
|
||||
// Error is the sole field: PHP's rfcError() writes {"error": code} with
|
||||
// no error_description, unlike Register's richer error body.
|
||||
Error string `json:"error"`
|
||||
}
|
||||
|
||||
// Token handles POST /oauth/mcp/token (D-09: raw route, no middleware).
|
||||
//
|
||||
// D-02: a JSON request body is rejected before any form parsing, so query
|
||||
// parameters on a JSON call can never smuggle a grant through (Pitfall 4).
|
||||
// Otherwise every parameter comes from the merged r.Form (net/http's own
|
||||
// ParseForm precedence puts body values ahead of query values — verified
|
||||
// against the Go 1.27 stdlib source, not assumed). Basic credentials, when
|
||||
// present, always override client_id/client_secret form values.
|
||||
func (s *Server) Token(w http.ResponseWriter, r *http.Request) {
|
||||
if isJSONContentType(r.Header.Get("Content-Type")) {
|
||||
writeTokenError(w, http.StatusBadRequest, "invalid_request")
|
||||
return
|
||||
}
|
||||
if s.backend == nil {
|
||||
w.WriteHeader(http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
// Errors from ParseForm on a malformed urlencoded body are not
|
||||
// distinguished from "no fields at all": the grant_type-empty check
|
||||
// below already produces the exact invalid_request PHP would give for
|
||||
// an unparseable/empty request.
|
||||
_ = r.ParseForm()
|
||||
|
||||
grantType := r.Form.Get("grant_type")
|
||||
if grantType == "" {
|
||||
writeTokenError(w, http.StatusBadRequest, "invalid_request")
|
||||
return
|
||||
}
|
||||
if grantType != "authorization_code" && grantType != "refresh_token" {
|
||||
writeTokenError(w, http.StatusBadRequest, "unsupported_grant_type")
|
||||
return
|
||||
}
|
||||
|
||||
ctx := r.Context()
|
||||
|
||||
// D-17: wristband's expiry sweep also runs on /token (PHP has no sweep
|
||||
// at all here). It deletes only rows already past ExpiresAt; unexpired
|
||||
// rotated/revoked refresh rows and unexpired used codes stay as replay
|
||||
// evidence. A sweep failure is treated as an opaque 500 like any other
|
||||
// store failure -- it must never silently skip and must never leak a
|
||||
// house-shaped body onto this raw RFC endpoint.
|
||||
sweepAt := s.now()
|
||||
if err := s.backend.WithinTx(ctx, func(tx Tx) error {
|
||||
if err := tx.DeleteExpiredCodes(ctx, sweepAt); err != nil {
|
||||
return err
|
||||
}
|
||||
return tx.DeleteExpiredRefreshTokens(ctx, sweepAt)
|
||||
}); err != nil {
|
||||
w.WriteHeader(http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
|
||||
client, err := s.authenticateClient(ctx, r)
|
||||
if err != nil {
|
||||
if errors.Is(err, errInvalidClient) {
|
||||
w.Header().Set("WWW-Authenticate", `Basic realm="OAuth"`)
|
||||
writeTokenError(w, http.StatusUnauthorized, "invalid_client")
|
||||
return
|
||||
}
|
||||
w.WriteHeader(http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
|
||||
var issued tokenIssueResult
|
||||
if grantType == "authorization_code" {
|
||||
issued, err = s.exchangeAuthorizationCode(ctx, r, client)
|
||||
} else {
|
||||
issued, err = s.rotateRefreshToken(ctx, r, client)
|
||||
}
|
||||
if err != nil {
|
||||
if errors.Is(err, errInvalidGrant) {
|
||||
writeTokenError(w, http.StatusBadRequest, "invalid_grant")
|
||||
return
|
||||
}
|
||||
w.WriteHeader(http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
|
||||
scope := strings.Join(issued.Scopes, " ")
|
||||
if issued.OfflineAccess && !stringSliceContains(issued.Scopes, "offline_access") {
|
||||
scope = strings.TrimSpace(scope + " offline_access")
|
||||
}
|
||||
body := tokenSuccessBody{
|
||||
AccessToken: issued.AccessToken,
|
||||
TokenType: "Bearer",
|
||||
ExpiresIn: int64(s.opts.AccessTokenTTL.Seconds()),
|
||||
RefreshToken: issued.RefreshToken,
|
||||
Scope: scope,
|
||||
}
|
||||
writeExactJSON(w, http.StatusOK, body, map[string]string{
|
||||
"Cache-Control": "no-store, private",
|
||||
"Pragma": "no-cache",
|
||||
})
|
||||
}
|
||||
|
||||
// authenticateClient ports OAuthTokenController::authenticateClient. Basic
|
||||
// credentials override form credentials; a public client (auth method
|
||||
// "none") needs no secret at all, but a confidential client with a missing
|
||||
// or wrong secret is errInvalidClient (T-08-SECRET-TIMING: the secret
|
||||
// comparison is a fixed-length sha256-hex constant-time compare, never a
|
||||
// direct string compare of variable-length secrets).
|
||||
func (s *Server) authenticateClient(ctx context.Context, r *http.Request) (*ClientRecord, error) {
|
||||
clientID := r.Form.Get("client_id")
|
||||
secret := r.Form.Get("client_secret")
|
||||
|
||||
if authz := r.Header.Get("Authorization"); strings.HasPrefix(authz, "Basic ") {
|
||||
decoded, err := base64.StdEncoding.DecodeString(strings.TrimPrefix(authz, "Basic "))
|
||||
if err != nil {
|
||||
return nil, errInvalidClient
|
||||
}
|
||||
idx := strings.IndexByte(string(decoded), ':')
|
||||
if idx < 0 {
|
||||
return nil, errInvalidClient
|
||||
}
|
||||
clientID = string(decoded[:idx])
|
||||
secret = string(decoded[idx+1:])
|
||||
}
|
||||
if clientID == "" {
|
||||
return nil, errInvalidClient
|
||||
}
|
||||
|
||||
var client *ClientRecord
|
||||
err := s.backend.WithinTx(ctx, func(tx Tx) error {
|
||||
c, err := tx.ByClientID(ctx, clientID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
client = c
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if client == nil || client.RevokedAt != nil {
|
||||
return nil, errInvalidClient
|
||||
}
|
||||
if client.TokenEndpointAuthMethod == "none" {
|
||||
return client, nil
|
||||
}
|
||||
if secret == "" || client.ClientSecretHash == nil {
|
||||
return nil, errInvalidClient
|
||||
}
|
||||
if !constantEqual(*client.ClientSecretHash, sha256Hex(secret)) {
|
||||
return nil, errInvalidClient
|
||||
}
|
||||
return client, nil
|
||||
}
|
||||
|
||||
// verifyPkce ports OAuthCodeManager::verifyPkce: method must be S256, and
|
||||
// the comparison between the stored challenge and the verifier's S256
|
||||
// transform is constant-time (T-08-PKCE/D-04).
|
||||
func verifyPkce(verifier, challenge, method string) bool {
|
||||
if method != "S256" {
|
||||
return false
|
||||
}
|
||||
return constantEqual(challenge, s256Challenge(verifier))
|
||||
}
|
||||
|
||||
// exchangeAuthorizationCode ports OAuthCodeManager::exchangeCode inside one
|
||||
// WithinTx callback: lock the code row, validate every binding, consume it,
|
||||
// mint the inv_ access token, and create its refresh-token successor
|
||||
// atomically (D-07/T-08-CODE-REPLAY). A validation failure returns
|
||||
// errInvalidGrant before any mutation, so the surrounding transaction has
|
||||
// nothing to roll back; a second exchange attempt against an already-used
|
||||
// row always loses (single-use row lock via ByCodeHashForUpdate).
|
||||
func (s *Server) exchangeAuthorizationCode(ctx context.Context, r *http.Request, client *ClientRecord) (tokenIssueResult, error) {
|
||||
code := r.Form.Get("code")
|
||||
redirectURI := r.Form.Get("redirect_uri")
|
||||
verifier := r.Form.Get("code_verifier")
|
||||
resource := r.Form.Get("resource")
|
||||
|
||||
if code == "" || redirectURI == "" || verifier == "" {
|
||||
return tokenIssueResult{}, errInvalidGrant
|
||||
}
|
||||
|
||||
var result tokenIssueResult
|
||||
err := s.backend.WithinTx(ctx, func(tx Tx) error {
|
||||
rec, err := tx.ByCodeHashForUpdate(ctx, sha256Hex(code))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if rec == nil ||
|
||||
rec.UsedAt != nil ||
|
||||
!rec.ExpiresAt.After(s.now()) ||
|
||||
rec.UserID == nil ||
|
||||
rec.ClientID != client.ClientID ||
|
||||
rec.RedirectURI != redirectURI ||
|
||||
(resource != "" && rec.Resource != nil && resource != *rec.Resource) ||
|
||||
!verifyPkce(verifier, rec.CodeChallenge, rec.CodeChallengeMethod) {
|
||||
return errInvalidGrant
|
||||
}
|
||||
|
||||
if err := tx.MarkUsed(ctx, rec.ID); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
name := truncateRunes(client.ClientName, 120)
|
||||
expiresAt := s.now().Add(s.opts.AccessTokenTTL)
|
||||
minted, err := tx.Mint(ctx, *rec.UserID, name, rec.Scopes, expiresAt, rec.CollectionIDs, client.ClientID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
rawRefresh, err := s.randomBytes(32)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
accessTokenID := minted.ID
|
||||
refreshRec := &RefreshTokenRecord{
|
||||
TokenHash: sha256Hex(rawRefresh),
|
||||
APITokenID: &accessTokenID,
|
||||
ClientID: client.ClientID,
|
||||
UserID: *rec.UserID,
|
||||
Scopes: rec.Scopes,
|
||||
CollectionIDs: rec.CollectionIDs,
|
||||
ExpiresAt: s.now().Add(s.opts.RefreshTokenTTL),
|
||||
OfflineAccess: rec.OfflineAccess,
|
||||
}
|
||||
if err := tx.Create(ctx, refreshRec); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
result = tokenIssueResult{
|
||||
AccessToken: minted.Secret,
|
||||
RefreshToken: rawRefresh,
|
||||
Scopes: rec.Scopes,
|
||||
OfflineAccess: rec.OfflineAccess,
|
||||
}
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
return tokenIssueResult{}, err
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
|
||||
// rotateRefreshToken ports OAuthCodeManager::rotateRefresh inside one
|
||||
// WithinTx callback (D-04/D-05/D-07): the presented refresh secret is row-
|
||||
// locked, validated (not expired, not revoked, bound to the authenticated
|
||||
// client), and then either rotated (mint a new access token, revoke the old
|
||||
// one, create and link a successor refresh row) or -- if it was already
|
||||
// rotated once before -- treated as a replay: the entire lineage (every
|
||||
// rotated-to successor and each linked access token) is revoked and the
|
||||
// callback still returns nil so that revocation commits (T-08-REFRESH-
|
||||
// REPLAY). Token then maps the recorded "replayed" outcome to invalid_grant
|
||||
// outside the transaction, exactly mirroring PHP's own commit-then-throw
|
||||
// shape: the security-relevant kill must land even though the protocol
|
||||
// response is an error.
|
||||
func (s *Server) rotateRefreshToken(ctx context.Context, r *http.Request, client *ClientRecord) (tokenIssueResult, error) {
|
||||
raw := r.Form.Get("refresh_token")
|
||||
if raw == "" {
|
||||
return tokenIssueResult{}, errInvalidGrant
|
||||
}
|
||||
|
||||
var result tokenIssueResult
|
||||
var replayed bool
|
||||
err := s.backend.WithinTx(ctx, func(tx Tx) error {
|
||||
rec, err := tx.ByTokenHashForUpdate(ctx, sha256Hex(raw))
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if rec == nil ||
|
||||
!rec.ExpiresAt.After(s.now()) ||
|
||||
rec.RevokedAt != nil ||
|
||||
rec.ClientID != client.ClientID {
|
||||
return errInvalidGrant
|
||||
}
|
||||
|
||||
if rec.RotatedToID != nil {
|
||||
// T-08-REFRESH-REPLAY: a spent (already-rotated) refresh token
|
||||
// was presented again. Kill the whole lineage and commit that
|
||||
// kill; the caller maps replayed -> invalid_grant afterward.
|
||||
if err := tx.RevokeLineage(ctx, rec.ID); err != nil {
|
||||
return err
|
||||
}
|
||||
replayed = true
|
||||
return nil
|
||||
}
|
||||
|
||||
if rec.APITokenID != nil {
|
||||
if err := tx.Revoke(ctx, *rec.APITokenID); err != nil {
|
||||
return err
|
||||
}
|
||||
}
|
||||
|
||||
name := truncateRunes(client.ClientName, 120)
|
||||
expiresAt := s.now().Add(s.opts.AccessTokenTTL)
|
||||
minted, err := tx.Mint(ctx, rec.UserID, name, rec.Scopes, expiresAt, rec.CollectionIDs, client.ClientID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
rawNew, err := s.randomBytes(32)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
accessTokenID := minted.ID
|
||||
successor := &RefreshTokenRecord{
|
||||
TokenHash: sha256Hex(rawNew),
|
||||
APITokenID: &accessTokenID,
|
||||
ClientID: client.ClientID,
|
||||
UserID: rec.UserID,
|
||||
Scopes: rec.Scopes,
|
||||
CollectionIDs: rec.CollectionIDs,
|
||||
ExpiresAt: s.now().Add(s.opts.RefreshTokenTTL),
|
||||
OfflineAccess: rec.OfflineAccess,
|
||||
}
|
||||
if err := tx.Create(ctx, successor); err != nil {
|
||||
return err
|
||||
}
|
||||
if err := tx.MarkRotated(ctx, rec.ID, successor.ID); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
result = tokenIssueResult{
|
||||
AccessToken: minted.Secret,
|
||||
RefreshToken: rawNew,
|
||||
Scopes: rec.Scopes,
|
||||
OfflineAccess: rec.OfflineAccess,
|
||||
}
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
return tokenIssueResult{}, err
|
||||
}
|
||||
if replayed {
|
||||
return tokenIssueResult{}, errInvalidGrant
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
|
||||
// Revoke atomically kills an OAuth-issued access token and its entire
|
||||
// refresh-token lineage (08-06-PLAN.md D-08; PHP
|
||||
// ConnectedAppController::destroy + OAuthCodeManager::revokeChain). It is
|
||||
// the cascade-revoke seam the app's connected-app controller calls instead
|
||||
// of touching refresh rows directly: the access token is revoked
|
||||
// unconditionally, and if a refresh row is still linked to it, the whole
|
||||
// lineage it anchors is revoked too (a live connected-app token is always
|
||||
// the terminal row of its chain, so this also protects a stale predecessor
|
||||
// replay from ever reviving it).
|
||||
func (s *Server) Revoke(ctx context.Context, apiTokenID uint) error {
|
||||
return s.backend.WithinTx(ctx, func(tx Tx) error {
|
||||
if err := tx.Revoke(ctx, apiTokenID); err != nil {
|
||||
return err
|
||||
}
|
||||
rec, err := tx.ByAPITokenIDForUpdate(ctx, apiTokenID)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if rec == nil {
|
||||
return nil
|
||||
}
|
||||
return tx.RevokeLineage(ctx, rec.ID)
|
||||
})
|
||||
}
|
||||
|
||||
func writeTokenError(w http.ResponseWriter, status int, code string) {
|
||||
// PHP's rfcError() sets no explicit Cache-Control; Laravel's own
|
||||
// session-cookie default for an otherwise-unheadered JSON response is
|
||||
// "no-cache, private" (matches the live-recorded byte contract, same
|
||||
// default the house wire.WriteJSON convention already uses elsewhere).
|
||||
writeExactJSON(w, status, tokenErrorBody{Error: code}, map[string]string{"Cache-Control": "no-cache, private"})
|
||||
}
|
||||
983
modules/wristband/token_test.go
Normal file
983
modules/wristband/token_test.go
Normal file
@@ -0,0 +1,983 @@
|
||||
package wristband
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"net/url"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
const tokenTestRedirect = "https://chatgpt.com/connector/oauth/cb"
|
||||
|
||||
// insertTokenTestClient inserts an already-usable ClientRecord directly into
|
||||
// backend (bypassing CreateWithCap's cap/sweep policy, which this plan's
|
||||
// tests do not exercise) and returns it.
|
||||
func insertTokenTestClient(backend *memoryBackend, clientID, authMethod string, secretHash *string, ceiling []string) *ClientRecord {
|
||||
backend.mu.Lock()
|
||||
defer backend.mu.Unlock()
|
||||
backend.nextID++
|
||||
rec := &ClientRecord{
|
||||
ID: backend.nextID,
|
||||
ClientID: clientID,
|
||||
ClientSecretHash: secretHash,
|
||||
ClientName: "Test Client",
|
||||
RedirectURIs: []string{tokenTestRedirect},
|
||||
GrantTypes: []string{"authorization_code", "refresh_token"},
|
||||
TokenEndpointAuthMethod: authMethod,
|
||||
ScopeCeiling: ceiling,
|
||||
CreatedAt: time.Now(),
|
||||
}
|
||||
backend.clients = append(backend.clients, rec)
|
||||
return rec
|
||||
}
|
||||
|
||||
// insertTokenTestCode seeds an already-issued (post-consent) code row
|
||||
// directly into backend, matching the shape 08-05's consent flow will
|
||||
// produce via AuthCodeStore.MarkIssued: CodeHash set, RequestID nil, UserID
|
||||
// set. mutate, when non-nil, is applied to the record before it is stored so
|
||||
// individual tests can adjust ExpiresAt/UsedAt/ClientID/etc.
|
||||
func insertTokenTestCode(t *testing.T, backend *memoryBackend, clientID string, challenge string, mutate func(*AuthCodeRecord)) (rawCode string, rec *AuthCodeRecord) {
|
||||
t.Helper()
|
||||
raw, err := randomBase64URL(32)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
backend.mu.Lock()
|
||||
defer backend.mu.Unlock()
|
||||
backend.nextID++
|
||||
userID := uint(1)
|
||||
hash := sha256Hex(raw)
|
||||
rec = &AuthCodeRecord{
|
||||
ID: backend.nextID,
|
||||
CodeHash: &hash,
|
||||
ClientID: clientID,
|
||||
UserID: &userID,
|
||||
RedirectURI: tokenTestRedirect,
|
||||
Scopes: []string{"read", "write"},
|
||||
CodeChallenge: challenge,
|
||||
CodeChallengeMethod: "S256",
|
||||
ExpiresAt: time.Now().Add(5 * time.Minute),
|
||||
}
|
||||
if mutate != nil {
|
||||
mutate(rec)
|
||||
}
|
||||
backend.codes = append(backend.codes, rec)
|
||||
return raw, rec
|
||||
}
|
||||
|
||||
// tokenRequest builds a POST /oauth/mcp/token request from form (encoded as
|
||||
// the body) with an optional Authorization header, matching the D-02 body
|
||||
// parser every test in this file exercises.
|
||||
func tokenRequest(form url.Values, contentType string) *http.Request {
|
||||
if contentType == "" {
|
||||
contentType = "application/x-www-form-urlencoded"
|
||||
}
|
||||
req := httptest.NewRequest(http.MethodPost, "/oauth/mcp/token", strings.NewReader(form.Encode()))
|
||||
req.Header.Set("Content-Type", contentType)
|
||||
return req
|
||||
}
|
||||
|
||||
// TestPhase8RedCodeExchange is the Phase 8 Wave 4 RED anchor (08-04-PLAN.md
|
||||
// Task 1, D-02/D-04/D-05/D-07). It drives one valid S256 authorization-code
|
||||
// exchange through the real (in-memory-backed) Server.Token and asserts the
|
||||
// exact RFC 6749 success contract. It fails with the
|
||||
// PHASE8_RED:code-exchange sentinel while Token is the 501 stub;
|
||||
// scripts/check-phase8-red.sh verifies this failure is fail-closed.
|
||||
func TestPhase8RedCodeExchange(t *testing.T) {
|
||||
backend := newMemoryBackend()
|
||||
srv := newTestServer(backend)
|
||||
insertTokenTestClient(backend, "cli-red", "none", nil, nil)
|
||||
verifier, challenge := s256Pair(t)
|
||||
rawCode, _ := insertTokenTestCode(t, backend, "cli-red", challenge, nil)
|
||||
|
||||
req := tokenRequest(url.Values{
|
||||
"grant_type": {"authorization_code"},
|
||||
"code": {rawCode},
|
||||
"code_verifier": {verifier},
|
||||
"redirect_uri": {tokenTestRedirect},
|
||||
"client_id": {"cli-red"},
|
||||
}, "")
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Token(rec, req)
|
||||
|
||||
if rec.Code != http.StatusOK {
|
||||
t.Fatalf("PHASE8_RED:code-exchange: status = %d, want %d (body=%s)", rec.Code, http.StatusOK, rec.Body.String())
|
||||
}
|
||||
var got map[string]any
|
||||
if err := json.Unmarshal(rec.Body.Bytes(), &got); err != nil {
|
||||
t.Fatalf("PHASE8_RED:code-exchange: decode response: %v", err)
|
||||
}
|
||||
access, _ := got["access_token"].(string)
|
||||
refresh, _ := got["refresh_token"].(string)
|
||||
if access == "" || refresh == "" {
|
||||
t.Fatalf("PHASE8_RED:code-exchange: access_token/refresh_token empty in %v", got)
|
||||
}
|
||||
if got["token_type"] != "Bearer" {
|
||||
t.Fatalf("PHASE8_RED:code-exchange: token_type = %v, want Bearer", got["token_type"])
|
||||
}
|
||||
if got["scope"] != "read write" {
|
||||
t.Fatalf("PHASE8_RED:code-exchange: scope = %v, want %q", got["scope"], "read write")
|
||||
}
|
||||
if cc := rec.Header().Get("Cache-Control"); cc != "no-store, private" {
|
||||
t.Fatalf("PHASE8_RED:code-exchange: Cache-Control = %q, want \"no-store, private\"", cc)
|
||||
}
|
||||
}
|
||||
|
||||
// assertTokenError decodes rec as the exact PHP token error body ({"error":
|
||||
// code}, no error_description) and asserts status/code.
|
||||
func assertTokenError(t *testing.T, rec *httptest.ResponseRecorder, status int, code string) map[string]any {
|
||||
t.Helper()
|
||||
if rec.Code != status {
|
||||
t.Fatalf("status = %d, want %d (body=%s)", rec.Code, status, rec.Body.String())
|
||||
}
|
||||
var got map[string]any
|
||||
if err := json.Unmarshal(rec.Body.Bytes(), &got); err != nil {
|
||||
t.Fatalf("decode error body: %v (body=%s)", err, rec.Body.String())
|
||||
}
|
||||
if got["error"] != code {
|
||||
t.Fatalf("error = %v, want %q", got["error"], code)
|
||||
}
|
||||
if _, has := got["error_description"]; has {
|
||||
t.Fatalf("body = %s carries error_description, PHP token errors never do", rec.Body.String())
|
||||
}
|
||||
return got
|
||||
}
|
||||
|
||||
// newTokenExchangeFixture seeds one usable public client and one valid
|
||||
// pending-issued code bound to it, returning everything a caller needs to
|
||||
// build a successful exchange request (and to mutate before breaking it).
|
||||
func newTokenExchangeFixture(t *testing.T) (srv *Server, backend *memoryBackend, verifier string, rawCode string, code *AuthCodeRecord) {
|
||||
t.Helper()
|
||||
backend = newMemoryBackend()
|
||||
srv = newTestServer(backend)
|
||||
insertTokenTestClient(backend, "cli-tok", "none", nil, nil)
|
||||
var challenge string
|
||||
verifier, challenge = s256Pair(t)
|
||||
rawCode, code = insertTokenTestCode(t, backend, "cli-tok", challenge, nil)
|
||||
return
|
||||
}
|
||||
|
||||
func validExchangeForm(rawCode, verifier, clientID string) url.Values {
|
||||
return url.Values{
|
||||
"grant_type": {"authorization_code"},
|
||||
"code": {rawCode},
|
||||
"code_verifier": {verifier},
|
||||
"redirect_uri": {tokenTestRedirect},
|
||||
"client_id": {clientID},
|
||||
}
|
||||
}
|
||||
|
||||
// TestTokenRejectsJSONBodyEvenWithValidQueryParams proves D-02/Pitfall 4: a
|
||||
// JSON content type is rejected before ParseForm ever runs, so a valid
|
||||
// grant cannot be smuggled through the query string of a JSON-labeled
|
||||
// request.
|
||||
func TestTokenRejectsJSONBodyEvenWithValidQueryParams(t *testing.T) {
|
||||
srv, _, verifier, rawCode, _ := newTokenExchangeFixture(t)
|
||||
|
||||
q := validExchangeForm(rawCode, verifier, "cli-tok")
|
||||
req := httptest.NewRequest(http.MethodPost, "/oauth/mcp/token?"+q.Encode(), strings.NewReader(`{"grant_type":"authorization_code"}`))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Token(rec, req)
|
||||
|
||||
assertTokenError(t, rec, http.StatusBadRequest, "invalid_request")
|
||||
}
|
||||
|
||||
// TestTokenBodyOverQueryPrecedence proves D-02: when the same key appears in
|
||||
// both the form body and the query string, the body value wins (matching
|
||||
// net/http's own documented ParseForm precedence, verified against the
|
||||
// stdlib source for this plan).
|
||||
func TestTokenBodyOverQueryPrecedence(t *testing.T) {
|
||||
srv, _, verifier, rawCode, _ := newTokenExchangeFixture(t)
|
||||
|
||||
form := validExchangeForm(rawCode, verifier, "cli-tok")
|
||||
req := httptest.NewRequest(http.MethodPost, "/oauth/mcp/token?grant_type=refresh_token", strings.NewReader(form.Encode()))
|
||||
req.Header.Set("Content-Type", "application/x-www-form-urlencoded")
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Token(rec, req)
|
||||
|
||||
if rec.Code != http.StatusOK {
|
||||
t.Fatalf("status = %d, want %d (body=%s): body grant_type must win over query grant_type", rec.Code, http.StatusOK, rec.Body.String())
|
||||
}
|
||||
}
|
||||
|
||||
// TestTokenBasicCredentialsOverrideFormCredentials proves the confidential
|
||||
// client's Basic header wins over (wrong) form client_id/client_secret
|
||||
// values.
|
||||
func TestTokenBasicCredentialsOverrideFormCredentials(t *testing.T) {
|
||||
backend := newMemoryBackend()
|
||||
srv := newTestServer(backend)
|
||||
hash := sha256Hex("correct-secret")
|
||||
insertTokenTestClient(backend, "cli-basic", "client_secret_basic", &hash, nil)
|
||||
verifier, challenge := s256Pair(t)
|
||||
rawCode, _ := insertTokenTestCode(t, backend, "cli-basic", challenge, nil)
|
||||
|
||||
form := url.Values{
|
||||
"grant_type": {"authorization_code"},
|
||||
"code": {rawCode},
|
||||
"code_verifier": {verifier},
|
||||
"redirect_uri": {tokenTestRedirect},
|
||||
"client_id": {"cli-basic"},
|
||||
"client_secret": {"wrong-form-secret"},
|
||||
}
|
||||
req := tokenRequest(form, "")
|
||||
req.SetBasicAuth("cli-basic", "correct-secret")
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Token(rec, req)
|
||||
|
||||
if rec.Code != http.StatusOK {
|
||||
t.Fatalf("status = %d, want %d (body=%s): Basic header must override form client_secret", rec.Code, http.StatusOK, rec.Body.String())
|
||||
}
|
||||
}
|
||||
|
||||
func TestTokenMissingGrantTypeIsInvalidRequest(t *testing.T) {
|
||||
srv, _, _, _, _ := newTokenExchangeFixture(t)
|
||||
req := tokenRequest(url.Values{}, "")
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Token(rec, req)
|
||||
assertTokenError(t, rec, http.StatusBadRequest, "invalid_request")
|
||||
}
|
||||
|
||||
func TestTokenUnsupportedGrantTypeIsRejected(t *testing.T) {
|
||||
srv, _, _, _, _ := newTokenExchangeFixture(t)
|
||||
req := tokenRequest(url.Values{"grant_type": {"client_credentials"}}, "")
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Token(rec, req)
|
||||
assertTokenError(t, rec, http.StatusBadRequest, "unsupported_grant_type")
|
||||
}
|
||||
|
||||
func TestTokenUnknownClientIsInvalidClientWithBasicChallenge(t *testing.T) {
|
||||
srv, _, verifier, rawCode, _ := newTokenExchangeFixture(t)
|
||||
form := validExchangeForm(rawCode, verifier, "does-not-exist")
|
||||
req := tokenRequest(form, "")
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Token(rec, req)
|
||||
|
||||
assertTokenError(t, rec, http.StatusUnauthorized, "invalid_client")
|
||||
if wa := rec.Header().Get("WWW-Authenticate"); wa != `Basic realm="OAuth"` {
|
||||
t.Fatalf("WWW-Authenticate = %q, want %q", wa, `Basic realm="OAuth"`)
|
||||
}
|
||||
}
|
||||
|
||||
func TestTokenRevokedClientIsInvalidClient(t *testing.T) {
|
||||
backend := newMemoryBackend()
|
||||
srv := newTestServer(backend)
|
||||
client := insertTokenTestClient(backend, "cli-revoked", "none", nil, nil)
|
||||
now := time.Now()
|
||||
client.RevokedAt = &now
|
||||
verifier, challenge := s256Pair(t)
|
||||
rawCode, _ := insertTokenTestCode(t, backend, "cli-revoked", challenge, nil)
|
||||
|
||||
req := tokenRequest(validExchangeForm(rawCode, verifier, "cli-revoked"), "")
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Token(rec, req)
|
||||
assertTokenError(t, rec, http.StatusUnauthorized, "invalid_client")
|
||||
}
|
||||
|
||||
func TestTokenConfidentialClientMissingSecretIsInvalidClient(t *testing.T) {
|
||||
backend := newMemoryBackend()
|
||||
srv := newTestServer(backend)
|
||||
hash := sha256Hex("s3cret")
|
||||
insertTokenTestClient(backend, "cli-conf-missing", "client_secret_post", &hash, nil)
|
||||
verifier, challenge := s256Pair(t)
|
||||
rawCode, _ := insertTokenTestCode(t, backend, "cli-conf-missing", challenge, nil)
|
||||
|
||||
req := tokenRequest(validExchangeForm(rawCode, verifier, "cli-conf-missing"), "")
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Token(rec, req)
|
||||
assertTokenError(t, rec, http.StatusUnauthorized, "invalid_client")
|
||||
if wa := rec.Header().Get("WWW-Authenticate"); wa != `Basic realm="OAuth"` {
|
||||
t.Fatalf("WWW-Authenticate = %q, want %q", wa, `Basic realm="OAuth"`)
|
||||
}
|
||||
}
|
||||
|
||||
// TestTokenConfidentialClientWrongSecretIsInvalidClient is the plan's named
|
||||
// "invalid confidential client" case: exact status/body/no-newline,
|
||||
// Cache-Control absent (only success responses carry it), and the Basic
|
||||
// realm="OAuth" challenge (T-08-SECRET-TIMING: comparison goes through
|
||||
// constantEqual/sha256Hex, never a direct string compare).
|
||||
func TestTokenConfidentialClientWrongSecretIsInvalidClient(t *testing.T) {
|
||||
backend := newMemoryBackend()
|
||||
srv := newTestServer(backend)
|
||||
hash := sha256Hex("correct-secret")
|
||||
insertTokenTestClient(backend, "cli-conf-wrong", "client_secret_post", &hash, nil)
|
||||
verifier, challenge := s256Pair(t)
|
||||
rawCode, _ := insertTokenTestCode(t, backend, "cli-conf-wrong", challenge, nil)
|
||||
|
||||
form := validExchangeForm(rawCode, verifier, "cli-conf-wrong")
|
||||
form.Set("client_secret", "wrong-secret")
|
||||
req := tokenRequest(form, "")
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Token(rec, req)
|
||||
|
||||
if rec.Code != http.StatusUnauthorized {
|
||||
t.Fatalf("status = %d, want %d", rec.Code, http.StatusUnauthorized)
|
||||
}
|
||||
if body := rec.Body.String(); body != `{"error":"invalid_client"}` {
|
||||
t.Fatalf("body = %q, want exact %q", body, `{"error":"invalid_client"}`)
|
||||
}
|
||||
if strings.HasSuffix(rec.Body.String(), "\n") {
|
||||
t.Fatal("body has a trailing newline")
|
||||
}
|
||||
if wa := rec.Header().Get("WWW-Authenticate"); wa != `Basic realm="OAuth"` {
|
||||
t.Fatalf("WWW-Authenticate = %q, want %q", wa, `Basic realm="OAuth"`)
|
||||
}
|
||||
if cc := rec.Header().Get("Cache-Control"); cc != "no-cache, private" {
|
||||
t.Fatalf("Cache-Control = %q, want %q", cc, "no-cache, private")
|
||||
}
|
||||
}
|
||||
|
||||
func TestTokenPublicClientIgnoresSuppliedSecret(t *testing.T) {
|
||||
srv, _, verifier, rawCode, _ := newTokenExchangeFixture(t)
|
||||
form := validExchangeForm(rawCode, verifier, "cli-tok")
|
||||
form.Set("client_secret", "irrelevant-for-a-public-client")
|
||||
req := tokenRequest(form, "")
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Token(rec, req)
|
||||
if rec.Code != http.StatusOK {
|
||||
t.Fatalf("status = %d, want %d (body=%s)", rec.Code, http.StatusOK, rec.Body.String())
|
||||
}
|
||||
}
|
||||
|
||||
func TestTokenWrongVerifierIsInvalidGrant(t *testing.T) {
|
||||
srv, backend, _, rawCode, _ := newTokenExchangeFixture(t)
|
||||
form := validExchangeForm(rawCode, "wrong-verifier-entirely", "cli-tok")
|
||||
req := tokenRequest(form, "")
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Token(rec, req)
|
||||
assertTokenError(t, rec, http.StatusBadRequest, "invalid_grant")
|
||||
|
||||
backend.mu.Lock()
|
||||
defer backend.mu.Unlock()
|
||||
if len(backend.tokens) != 0 {
|
||||
t.Fatal("a wrong-verifier exchange must not mint an access token")
|
||||
}
|
||||
}
|
||||
|
||||
func TestTokenClientMismatchIsInvalidGrant(t *testing.T) {
|
||||
backend := newMemoryBackend()
|
||||
srv := newTestServer(backend)
|
||||
insertTokenTestClient(backend, "cli-owner", "none", nil, nil)
|
||||
insertTokenTestClient(backend, "cli-other", "none", nil, nil)
|
||||
verifier, challenge := s256Pair(t)
|
||||
rawCode, _ := insertTokenTestCode(t, backend, "cli-owner", challenge, nil)
|
||||
|
||||
req := tokenRequest(validExchangeForm(rawCode, verifier, "cli-other"), "")
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Token(rec, req)
|
||||
assertTokenError(t, rec, http.StatusBadRequest, "invalid_grant")
|
||||
}
|
||||
|
||||
func TestTokenRedirectURIMismatchIsInvalidGrant(t *testing.T) {
|
||||
srv, _, verifier, rawCode, _ := newTokenExchangeFixture(t)
|
||||
form := validExchangeForm(rawCode, verifier, "cli-tok")
|
||||
form.Set("redirect_uri", "https://evil.example.test/cb")
|
||||
req := tokenRequest(form, "")
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Token(rec, req)
|
||||
assertTokenError(t, rec, http.StatusBadRequest, "invalid_grant")
|
||||
}
|
||||
|
||||
func TestTokenResourceMismatchIsInvalidGrant(t *testing.T) {
|
||||
backend := newMemoryBackend()
|
||||
srv := newTestServer(backend)
|
||||
insertTokenTestClient(backend, "cli-res", "none", nil, nil)
|
||||
verifier, challenge := s256Pair(t)
|
||||
res := "https://mcp.plytarium.com/mcp"
|
||||
rawCode, _ := insertTokenTestCode(t, backend, "cli-res", challenge, func(rec *AuthCodeRecord) {
|
||||
rec.Resource = &res
|
||||
})
|
||||
|
||||
form := validExchangeForm(rawCode, verifier, "cli-res")
|
||||
form.Set("resource", "https://wrong.example.test/mcp")
|
||||
req := tokenRequest(form, "")
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Token(rec, req)
|
||||
assertTokenError(t, rec, http.StatusBadRequest, "invalid_grant")
|
||||
}
|
||||
|
||||
func TestTokenResourceOmittedIsAccepted(t *testing.T) {
|
||||
backend := newMemoryBackend()
|
||||
srv := newTestServer(backend)
|
||||
insertTokenTestClient(backend, "cli-res-omit", "none", nil, nil)
|
||||
verifier, challenge := s256Pair(t)
|
||||
res := "https://mcp.plytarium.com/mcp"
|
||||
rawCode, _ := insertTokenTestCode(t, backend, "cli-res-omit", challenge, func(rec *AuthCodeRecord) {
|
||||
rec.Resource = &res
|
||||
})
|
||||
|
||||
req := tokenRequest(validExchangeForm(rawCode, verifier, "cli-res-omit"), "")
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Token(rec, req)
|
||||
if rec.Code != http.StatusOK {
|
||||
t.Fatalf("status = %d, want %d (body=%s)", rec.Code, http.StatusOK, rec.Body.String())
|
||||
}
|
||||
}
|
||||
|
||||
func TestTokenExpiredCodeIsInvalidGrant(t *testing.T) {
|
||||
backend := newMemoryBackend()
|
||||
srv := newTestServer(backend)
|
||||
insertTokenTestClient(backend, "cli-expired", "none", nil, nil)
|
||||
verifier, challenge := s256Pair(t)
|
||||
rawCode, _ := insertTokenTestCode(t, backend, "cli-expired", challenge, func(rec *AuthCodeRecord) {
|
||||
rec.ExpiresAt = time.Now().Add(-1 * time.Second)
|
||||
})
|
||||
|
||||
req := tokenRequest(validExchangeForm(rawCode, verifier, "cli-expired"), "")
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Token(rec, req)
|
||||
assertTokenError(t, rec, http.StatusBadRequest, "invalid_grant")
|
||||
}
|
||||
|
||||
func TestTokenMissingRequiredFieldsAreInvalidGrant(t *testing.T) {
|
||||
cases := []struct {
|
||||
name string
|
||||
strip func(url.Values)
|
||||
}{
|
||||
{"missing-code", func(v url.Values) { v.Del("code") }},
|
||||
{"missing-redirect-uri", func(v url.Values) { v.Del("redirect_uri") }},
|
||||
{"missing-code-verifier", func(v url.Values) { v.Del("code_verifier") }},
|
||||
}
|
||||
for _, tc := range cases {
|
||||
t.Run(tc.name, func(t *testing.T) {
|
||||
srv, _, verifier, rawCode, _ := newTokenExchangeFixture(t)
|
||||
form := validExchangeForm(rawCode, verifier, "cli-tok")
|
||||
tc.strip(form)
|
||||
req := tokenRequest(form, "")
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Token(rec, req)
|
||||
assertTokenError(t, rec, http.StatusBadRequest, "invalid_grant")
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// TestTokenCodeSequentialReplayIsInvalidGrantSecondTime is the sequential
|
||||
// half of T-08-CODE-REPLAY: exchanging the same code twice succeeds exactly
|
||||
// once.
|
||||
func TestTokenCodeSequentialReplayIsInvalidGrantSecondTime(t *testing.T) {
|
||||
srv, _, verifier, rawCode, _ := newTokenExchangeFixture(t)
|
||||
form := validExchangeForm(rawCode, verifier, "cli-tok")
|
||||
|
||||
first := httptest.NewRecorder()
|
||||
srv.Token(first, tokenRequest(form, ""))
|
||||
if first.Code != http.StatusOK {
|
||||
t.Fatalf("first exchange status = %d, want %d (body=%s)", first.Code, http.StatusOK, first.Body.String())
|
||||
}
|
||||
|
||||
second := httptest.NewRecorder()
|
||||
srv.Token(second, tokenRequest(form, ""))
|
||||
assertTokenError(t, second, http.StatusBadRequest, "invalid_grant")
|
||||
}
|
||||
|
||||
// TestTokenCodeConcurrentReplayHasExactlyOneWinner is the concurrent half of
|
||||
// T-08-CODE-REPLAY (Pitfall 8): two synchronized goroutines racing to
|
||||
// exchange the same code must produce exactly one 200 and one invalid_grant,
|
||||
// never two successes. memoryBackend serializes the whole WithinTx closure
|
||||
// behind one mutex (08-PATTERNS.md), which is exactly the seam this test
|
||||
// exercises; the real-Postgres row-lock proof lives in fonoteka.go's
|
||||
// classes/auth package.
|
||||
func TestTokenCodeConcurrentReplayHasExactlyOneWinner(t *testing.T) {
|
||||
srv, _, verifier, rawCode, _ := newTokenExchangeFixture(t)
|
||||
form := validExchangeForm(rawCode, verifier, "cli-tok")
|
||||
|
||||
results := make([]int, 2)
|
||||
start := make(chan struct{})
|
||||
done := make(chan struct{})
|
||||
for i := range 2 {
|
||||
go func(i int) {
|
||||
<-start
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Token(rec, tokenRequest(form, ""))
|
||||
results[i] = rec.Code
|
||||
done <- struct{}{}
|
||||
}(i)
|
||||
}
|
||||
close(start)
|
||||
<-done
|
||||
<-done
|
||||
|
||||
successCount, grantErrCount := 0, 0
|
||||
for _, code := range results {
|
||||
switch code {
|
||||
case http.StatusOK:
|
||||
successCount++
|
||||
case http.StatusBadRequest:
|
||||
grantErrCount++
|
||||
default:
|
||||
t.Fatalf("unexpected status %d", code)
|
||||
}
|
||||
}
|
||||
if successCount != 1 || grantErrCount != 1 {
|
||||
t.Fatalf("successCount=%d grantErrCount=%d, want 1 and 1 (results=%v)", successCount, grantErrCount, results)
|
||||
}
|
||||
}
|
||||
|
||||
func TestTokenOfflineAccessAppendedToScope(t *testing.T) {
|
||||
backend := newMemoryBackend()
|
||||
srv := newTestServer(backend)
|
||||
insertTokenTestClient(backend, "cli-offline", "none", nil, nil)
|
||||
verifier, challenge := s256Pair(t)
|
||||
rawCode, _ := insertTokenTestCode(t, backend, "cli-offline", challenge, func(rec *AuthCodeRecord) {
|
||||
rec.OfflineAccess = true
|
||||
})
|
||||
|
||||
req := tokenRequest(validExchangeForm(rawCode, verifier, "cli-offline"), "")
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Token(rec, req)
|
||||
if rec.Code != http.StatusOK {
|
||||
t.Fatalf("status = %d, want %d (body=%s)", rec.Code, http.StatusOK, rec.Body.String())
|
||||
}
|
||||
var got map[string]any
|
||||
if err := json.Unmarshal(rec.Body.Bytes(), &got); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if got["scope"] != "read write offline_access" {
|
||||
t.Fatalf("scope = %v, want %q", got["scope"], "read write offline_access")
|
||||
}
|
||||
}
|
||||
|
||||
func TestTokenSuccessResponseHasNoEnvelopeAndNoTrailingNewline(t *testing.T) {
|
||||
srv, _, verifier, rawCode, _ := newTokenExchangeFixture(t)
|
||||
req := tokenRequest(validExchangeForm(rawCode, verifier, "cli-tok"), "")
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Token(rec, req)
|
||||
|
||||
if strings.HasSuffix(rec.Body.String(), "\n") {
|
||||
t.Fatal("body has a trailing newline")
|
||||
}
|
||||
var got map[string]any
|
||||
if err := json.Unmarshal(rec.Body.Bytes(), &got); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, hasData := got["data"]; hasData {
|
||||
t.Fatal("body has a house \"data\" envelope")
|
||||
}
|
||||
if got["expires_in"] != float64(3600) {
|
||||
t.Fatalf("expires_in = %v, want 3600", got["expires_in"])
|
||||
}
|
||||
if cc := rec.Header().Get("Cache-Control"); cc != "no-store, private" {
|
||||
t.Fatalf("Cache-Control = %q, want \"no-store, private\"", cc)
|
||||
}
|
||||
if p := rec.Header().Get("Pragma"); p != "no-cache" {
|
||||
t.Fatalf("Pragma = %q, want \"no-cache\"", p)
|
||||
}
|
||||
}
|
||||
|
||||
// TestTokenRefreshGrantDispatchIsAcceptedButNotYetImplemented proves Token's
|
||||
// own grant-type validity check accepts "refresh_token" exactly like PHP
|
||||
// does (it is not unsupported_grant_type): a syntactically well-formed but
|
||||
// never-issued refresh secret reaches rotateRefreshToken's real lookup and
|
||||
// is rejected as invalid_grant, not unsupported_grant_type.
|
||||
func TestTokenRefreshGrantDispatchIsAcceptedButNotYetImplemented(t *testing.T) {
|
||||
srv, _, _, _, _ := newTokenExchangeFixture(t)
|
||||
req := tokenRequest(url.Values{
|
||||
"grant_type": {"refresh_token"},
|
||||
"refresh_token": {"whatever"},
|
||||
"client_id": {"cli-tok"},
|
||||
}, "")
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Token(rec, req)
|
||||
assertTokenError(t, rec, http.StatusBadRequest, "invalid_grant")
|
||||
}
|
||||
|
||||
// TestPhase8RedLifecycleFramework is the Phase 8 Wave 6 RED anchor
|
||||
// (08-06-PLAN.md Task 1, D-04/D-17). It drives a full refresh-rotation and
|
||||
// replay lifecycle against the real (in-memory-backed) Server.Token:
|
||||
// exchange a code, rotate the resulting refresh token, replay the spent
|
||||
// original, and prove the whole lineage -- both access tokens and both
|
||||
// refresh rows -- ends up dead. It fails with the
|
||||
// PHASE8_RED:lifecycle-framework sentinel while rotateRefreshToken is
|
||||
// 08-04's invalid_grant placeholder (the rotate step below expects 200 but
|
||||
// gets 400); scripts/check-phase8-red.sh verifies this failure is
|
||||
// fail-closed.
|
||||
func TestPhase8RedLifecycleFramework(t *testing.T) {
|
||||
backend := newMemoryBackend()
|
||||
srv := newTestServer(backend)
|
||||
insertTokenTestClient(backend, "cli-lifecycle", "none", nil, nil)
|
||||
verifier, challenge := s256Pair(t)
|
||||
rawCode, _ := insertTokenTestCode(t, backend, "cli-lifecycle", challenge, nil)
|
||||
|
||||
first := httptest.NewRecorder()
|
||||
srv.Token(first, tokenRequest(validExchangeForm(rawCode, verifier, "cli-lifecycle"), ""))
|
||||
if first.Code != http.StatusOK {
|
||||
t.Fatalf("PHASE8_RED:lifecycle-framework: exchange status = %d, want %d (body=%s)", first.Code, http.StatusOK, first.Body.String())
|
||||
}
|
||||
firstAccess, firstRefresh := decodeTokenPair(t, first)
|
||||
|
||||
rotate := httptest.NewRecorder()
|
||||
srv.Token(rotate, refreshRequest(firstRefresh, "cli-lifecycle"))
|
||||
if rotate.Code != http.StatusOK {
|
||||
t.Fatalf("PHASE8_RED:lifecycle-framework: rotation status = %d, want %d (body=%s)", rotate.Code, http.StatusOK, rotate.Body.String())
|
||||
}
|
||||
secondAccess, secondRefresh := decodeTokenPair(t, rotate)
|
||||
if secondRefresh == firstRefresh {
|
||||
t.Fatalf("PHASE8_RED:lifecycle-framework: rotation must issue a new refresh secret")
|
||||
}
|
||||
|
||||
replay := httptest.NewRecorder()
|
||||
srv.Token(replay, refreshRequest(firstRefresh, "cli-lifecycle"))
|
||||
if replay.Code != http.StatusBadRequest {
|
||||
t.Fatalf("PHASE8_RED:lifecycle-framework: replay status = %d, want %d (body=%s)", replay.Code, http.StatusBadRequest, replay.Body.String())
|
||||
}
|
||||
|
||||
backend.mu.Lock()
|
||||
defer backend.mu.Unlock()
|
||||
if !refreshRowRevoked(backend, firstRefresh) {
|
||||
t.Fatal("PHASE8_RED:lifecycle-framework: original refresh row must be revoked after replay")
|
||||
}
|
||||
if !refreshRowRevoked(backend, secondRefresh) {
|
||||
t.Fatal("PHASE8_RED:lifecycle-framework: rotated successor refresh row must be revoked after replay of its predecessor")
|
||||
}
|
||||
if !accessTokenRevoked(backend, firstAccess) {
|
||||
t.Fatal("PHASE8_RED:lifecycle-framework: original access token must be revoked")
|
||||
}
|
||||
if !accessTokenRevoked(backend, secondAccess) {
|
||||
t.Fatal("PHASE8_RED:lifecycle-framework: rotated access token must be revoked after replay")
|
||||
}
|
||||
}
|
||||
|
||||
// decodeTokenPair extracts access_token/refresh_token from a successful
|
||||
// Token response body.
|
||||
func decodeTokenPair(t *testing.T, rec *httptest.ResponseRecorder) (access, refresh string) {
|
||||
t.Helper()
|
||||
var got map[string]any
|
||||
if err := json.Unmarshal(rec.Body.Bytes(), &got); err != nil {
|
||||
t.Fatalf("decode token body: %v (body=%s)", err, rec.Body.String())
|
||||
}
|
||||
access, _ = got["access_token"].(string)
|
||||
refresh, _ = got["refresh_token"].(string)
|
||||
if access == "" || refresh == "" {
|
||||
t.Fatalf("access_token/refresh_token empty in %v", got)
|
||||
}
|
||||
return access, refresh
|
||||
}
|
||||
|
||||
// refreshRequest builds a POST /oauth/mcp/token grant_type=refresh_token
|
||||
// request.
|
||||
func refreshRequest(rawRefresh, clientID string) *http.Request {
|
||||
return tokenRequest(url.Values{
|
||||
"grant_type": {"refresh_token"},
|
||||
"refresh_token": {rawRefresh},
|
||||
"client_id": {clientID},
|
||||
}, "")
|
||||
}
|
||||
|
||||
// refreshRowRevoked reports whether the refresh row matching rawRefresh has
|
||||
// a non-nil RevokedAt. The caller must already hold backend.mu.
|
||||
func refreshRowRevoked(backend *memoryBackend, rawRefresh string) bool {
|
||||
hash := sha256Hex(rawRefresh)
|
||||
for _, r := range backend.refresh {
|
||||
if r.TokenHash == hash {
|
||||
return r.RevokedAt != nil
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// accessTokenRevoked reports whether the IssuedToken matching rawAccess is
|
||||
// marked revoked in backend.revoked. The caller must already hold
|
||||
// backend.mu.
|
||||
func accessTokenRevoked(backend *memoryBackend, rawAccess string) bool {
|
||||
for _, tok := range backend.tokens {
|
||||
if tok.Secret == rawAccess {
|
||||
return backend.revoked[tok.ID]
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// TestRefreshRotationIssuesNewPairAndKeepsPredecessorAsEvidence proves D-04's
|
||||
// normal-path rotation: a fresh (not-yet-rotated) refresh token succeeds
|
||||
// exactly once, mints a new access/refresh pair with the same scopes, and
|
||||
// leaves the predecessor row retrievable (not deleted) with RotatedToID set
|
||||
// but RevokedAt nil -- a rotated-but-unreplayed row is not itself "revoked"
|
||||
// (D-17 evidence retention).
|
||||
func TestRefreshRotationIssuesNewPairAndKeepsPredecessorAsEvidence(t *testing.T) {
|
||||
backend := newMemoryBackend()
|
||||
srv := newTestServer(backend)
|
||||
insertTokenTestClient(backend, "cli-rotate", "none", nil, nil)
|
||||
verifier, challenge := s256Pair(t)
|
||||
rawCode, _ := insertTokenTestCode(t, backend, "cli-rotate", challenge, nil)
|
||||
|
||||
first := httptest.NewRecorder()
|
||||
srv.Token(first, tokenRequest(validExchangeForm(rawCode, verifier, "cli-rotate"), ""))
|
||||
firstAccess, firstRefresh := decodeTokenPair(t, first)
|
||||
|
||||
rotate := httptest.NewRecorder()
|
||||
srv.Token(rotate, refreshRequest(firstRefresh, "cli-rotate"))
|
||||
if rotate.Code != http.StatusOK {
|
||||
t.Fatalf("status = %d, want %d (body=%s)", rotate.Code, http.StatusOK, rotate.Body.String())
|
||||
}
|
||||
var got map[string]any
|
||||
if err := json.Unmarshal(rotate.Body.Bytes(), &got); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if got["scope"] != "read write" {
|
||||
t.Fatalf("scope = %v, want %q", got["scope"], "read write")
|
||||
}
|
||||
secondAccess, secondRefresh := decodeTokenPair(t, rotate)
|
||||
if secondAccess == firstAccess || secondRefresh == firstRefresh {
|
||||
t.Fatal("rotation must mint a brand new access/refresh pair")
|
||||
}
|
||||
|
||||
backend.mu.Lock()
|
||||
defer backend.mu.Unlock()
|
||||
hash := sha256Hex(firstRefresh)
|
||||
for _, r := range backend.refresh {
|
||||
if r.TokenHash == hash {
|
||||
if r.RevokedAt != nil {
|
||||
t.Fatal("a rotated-but-unreplayed predecessor must not itself be revoked")
|
||||
}
|
||||
if r.RotatedToID == nil {
|
||||
t.Fatal("predecessor must have RotatedToID set to its successor")
|
||||
}
|
||||
return
|
||||
}
|
||||
}
|
||||
t.Fatal("predecessor refresh row not found (must not be deleted)")
|
||||
}
|
||||
|
||||
// TestRefreshWrongClientIsInvalidGrant proves the client-binding check: a
|
||||
// refresh token issued to one client cannot be redeemed by another.
|
||||
func TestRefreshWrongClientIsInvalidGrant(t *testing.T) {
|
||||
backend := newMemoryBackend()
|
||||
srv := newTestServer(backend)
|
||||
insertTokenTestClient(backend, "cli-owner-r", "none", nil, nil)
|
||||
insertTokenTestClient(backend, "cli-other-r", "none", nil, nil)
|
||||
verifier, challenge := s256Pair(t)
|
||||
rawCode, _ := insertTokenTestCode(t, backend, "cli-owner-r", challenge, nil)
|
||||
|
||||
first := httptest.NewRecorder()
|
||||
srv.Token(first, tokenRequest(validExchangeForm(rawCode, verifier, "cli-owner-r"), ""))
|
||||
_, firstRefresh := decodeTokenPair(t, first)
|
||||
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Token(rec, refreshRequest(firstRefresh, "cli-other-r"))
|
||||
assertTokenError(t, rec, http.StatusBadRequest, "invalid_grant")
|
||||
}
|
||||
|
||||
// TestRefreshExpiredIsInvalidGrant proves an expired refresh row is
|
||||
// rejected even though it was never rotated or revoked.
|
||||
func TestRefreshExpiredIsInvalidGrant(t *testing.T) {
|
||||
backend := newMemoryBackend()
|
||||
srv := newTestServer(backend)
|
||||
insertTokenTestClient(backend, "cli-refresh-expired", "none", nil, nil)
|
||||
verifier, challenge := s256Pair(t)
|
||||
rawCode, _ := insertTokenTestCode(t, backend, "cli-refresh-expired", challenge, nil)
|
||||
|
||||
first := httptest.NewRecorder()
|
||||
srv.Token(first, tokenRequest(validExchangeForm(rawCode, verifier, "cli-refresh-expired"), ""))
|
||||
_, firstRefresh := decodeTokenPair(t, first)
|
||||
|
||||
backend.mu.Lock()
|
||||
hash := sha256Hex(firstRefresh)
|
||||
for _, r := range backend.refresh {
|
||||
if r.TokenHash == hash {
|
||||
r.ExpiresAt = time.Now().Add(-1 * time.Minute)
|
||||
}
|
||||
}
|
||||
backend.mu.Unlock()
|
||||
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Token(rec, refreshRequest(firstRefresh, "cli-refresh-expired"))
|
||||
assertTokenError(t, rec, http.StatusBadRequest, "invalid_grant")
|
||||
}
|
||||
|
||||
// TestRefreshUnknownTokenIsInvalidGrant proves a syntactically valid but
|
||||
// never-issued refresh secret is rejected exactly like PHP's missing-record
|
||||
// case, not distinguished from any other invalid_grant.
|
||||
func TestRefreshUnknownTokenIsInvalidGrant(t *testing.T) {
|
||||
srv, _, _, _, _ := newTokenExchangeFixture(t)
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Token(rec, refreshRequest("does-not-exist-at-all", "cli-tok"))
|
||||
assertTokenError(t, rec, http.StatusBadRequest, "invalid_grant")
|
||||
}
|
||||
|
||||
// TestRefreshMissingTokenIsInvalidGrant proves an empty refresh_token value
|
||||
// is rejected before any store lookup.
|
||||
func TestRefreshMissingTokenIsInvalidGrant(t *testing.T) {
|
||||
srv, _, _, _, _ := newTokenExchangeFixture(t)
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Token(rec, tokenRequest(url.Values{
|
||||
"grant_type": {"refresh_token"},
|
||||
"client_id": {"cli-tok"},
|
||||
}, ""))
|
||||
assertTokenError(t, rec, http.StatusBadRequest, "invalid_grant")
|
||||
}
|
||||
|
||||
// TestRefreshConcurrentReplayHasExactlyOneWinner is the concurrent half of
|
||||
// T-08-REFRESH-REPLAY: two synchronized goroutines racing to rotate the same
|
||||
// refresh token must produce exactly one 200 and one invalid_grant, never
|
||||
// two successors sharing one predecessor.
|
||||
func TestRefreshConcurrentReplayHasExactlyOneWinner(t *testing.T) {
|
||||
backend := newMemoryBackend()
|
||||
srv := newTestServer(backend)
|
||||
insertTokenTestClient(backend, "cli-refresh-race", "none", nil, nil)
|
||||
verifier, challenge := s256Pair(t)
|
||||
rawCode, _ := insertTokenTestCode(t, backend, "cli-refresh-race", challenge, nil)
|
||||
|
||||
first := httptest.NewRecorder()
|
||||
srv.Token(first, tokenRequest(validExchangeForm(rawCode, verifier, "cli-refresh-race"), ""))
|
||||
_, firstRefresh := decodeTokenPair(t, first)
|
||||
|
||||
results := make([]int, 2)
|
||||
start := make(chan struct{})
|
||||
done := make(chan struct{})
|
||||
for i := range 2 {
|
||||
go func(i int) {
|
||||
<-start
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Token(rec, refreshRequest(firstRefresh, "cli-refresh-race"))
|
||||
results[i] = rec.Code
|
||||
done <- struct{}{}
|
||||
}(i)
|
||||
}
|
||||
close(start)
|
||||
<-done
|
||||
<-done
|
||||
|
||||
successCount, grantErrCount := 0, 0
|
||||
for _, code := range results {
|
||||
switch code {
|
||||
case http.StatusOK:
|
||||
successCount++
|
||||
case http.StatusBadRequest:
|
||||
grantErrCount++
|
||||
default:
|
||||
t.Fatalf("unexpected status %d", code)
|
||||
}
|
||||
}
|
||||
if successCount != 1 || grantErrCount != 1 {
|
||||
t.Fatalf("successCount=%d grantErrCount=%d, want 1 and 1 (results=%v)", successCount, grantErrCount, results)
|
||||
}
|
||||
}
|
||||
|
||||
// TestTokenSweepDeletesExpiredRowsButKeepsUnexpiredEvidence proves D-17: an
|
||||
// expired pending/code row and an expired refresh row are gone after the
|
||||
// next /token call's sweep, while an unexpired-but-revoked refresh row (real
|
||||
// replay evidence) survives untouched.
|
||||
func TestTokenSweepDeletesExpiredRowsButKeepsUnexpiredEvidence(t *testing.T) {
|
||||
backend := newMemoryBackend()
|
||||
srv := newTestServer(backend)
|
||||
insertTokenTestClient(backend, "cli-sweep", "none", nil, nil)
|
||||
|
||||
_, expiredCode := insertTokenTestCode(t, backend, "cli-sweep", "unused-challenge", func(rec *AuthCodeRecord) {
|
||||
rec.ExpiresAt = time.Now().Add(-1 * time.Hour)
|
||||
})
|
||||
expiredCodeID := expiredCode.ID
|
||||
|
||||
backend.mu.Lock()
|
||||
backend.nextID++
|
||||
expiredRefreshID := backend.nextID
|
||||
backend.refresh = append(backend.refresh, &RefreshTokenRecord{
|
||||
ID: expiredRefreshID,
|
||||
TokenHash: sha256Hex("expired-refresh-secret"),
|
||||
ClientID: "cli-sweep",
|
||||
UserID: 1,
|
||||
Scopes: []string{"read"},
|
||||
ExpiresAt: time.Now().Add(-1 * time.Hour),
|
||||
})
|
||||
now := time.Now()
|
||||
backend.nextID++
|
||||
evidenceRefreshID := backend.nextID
|
||||
backend.refresh = append(backend.refresh, &RefreshTokenRecord{
|
||||
ID: evidenceRefreshID,
|
||||
TokenHash: sha256Hex("revoked-but-unexpired-refresh-secret"),
|
||||
ClientID: "cli-sweep",
|
||||
UserID: 1,
|
||||
Scopes: []string{"read"},
|
||||
ExpiresAt: time.Now().Add(24 * time.Hour),
|
||||
RevokedAt: &now,
|
||||
})
|
||||
backend.mu.Unlock()
|
||||
|
||||
// Any /token call runs the sweep; use a deliberately-broken grant so no
|
||||
// mutation beyond the sweep happens.
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Token(rec, tokenRequest(url.Values{
|
||||
"grant_type": {"refresh_token"},
|
||||
"refresh_token": {"does-not-exist"},
|
||||
"client_id": {"cli-sweep"},
|
||||
}, ""))
|
||||
|
||||
backend.mu.Lock()
|
||||
defer backend.mu.Unlock()
|
||||
for _, c := range backend.codes {
|
||||
if c.ID == expiredCodeID {
|
||||
t.Fatal("expired code row must be swept")
|
||||
}
|
||||
}
|
||||
for _, r := range backend.refresh {
|
||||
if r.ID == expiredRefreshID {
|
||||
t.Fatal("expired refresh row must be swept")
|
||||
}
|
||||
}
|
||||
found := false
|
||||
for _, r := range backend.refresh {
|
||||
if r.ID == evidenceRefreshID {
|
||||
found = true
|
||||
}
|
||||
}
|
||||
if !found {
|
||||
t.Fatal("unexpired revoked refresh row must survive the sweep as replay evidence")
|
||||
}
|
||||
}
|
||||
|
||||
// TestServerRevokeKillsAccessAndLineage proves Server.Revoke (the seam the
|
||||
// app's connected-app controller calls, 08-06-PLAN.md D-08): revoking a live
|
||||
// OAuth access token also revokes its linked refresh row.
|
||||
func TestServerRevokeKillsAccessAndLineage(t *testing.T) {
|
||||
backend := newMemoryBackend()
|
||||
srv := newTestServer(backend)
|
||||
insertTokenTestClient(backend, "cli-revoke", "none", nil, nil)
|
||||
verifier, challenge := s256Pair(t)
|
||||
rawCode, _ := insertTokenTestCode(t, backend, "cli-revoke", challenge, nil)
|
||||
|
||||
exchange := httptest.NewRecorder()
|
||||
srv.Token(exchange, tokenRequest(validExchangeForm(rawCode, verifier, "cli-revoke"), ""))
|
||||
access, refresh := decodeTokenPair(t, exchange)
|
||||
|
||||
backend.mu.Lock()
|
||||
var tokenID uint
|
||||
for _, tok := range backend.tokens {
|
||||
if tok.Secret == access {
|
||||
tokenID = tok.ID
|
||||
}
|
||||
}
|
||||
backend.mu.Unlock()
|
||||
if tokenID == 0 {
|
||||
t.Fatal("minted access token not found in backend")
|
||||
}
|
||||
|
||||
if err := srv.Revoke(t.Context(), tokenID); err != nil {
|
||||
t.Fatalf("Revoke: %v", err)
|
||||
}
|
||||
|
||||
backend.mu.Lock()
|
||||
if !backend.revoked[tokenID] {
|
||||
t.Fatal("access token must be revoked")
|
||||
}
|
||||
if !refreshRowRevoked(backend, refresh) {
|
||||
t.Fatal("linked refresh row must be revoked")
|
||||
}
|
||||
backend.mu.Unlock()
|
||||
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Token(rec, refreshRequest(refresh, "cli-revoke"))
|
||||
assertTokenError(t, rec, http.StatusBadRequest, "invalid_grant")
|
||||
}
|
||||
|
||||
func TestTokenBackendUnavailableIsOpaque500(t *testing.T) {
|
||||
opts := DefaultOptions()
|
||||
opts.Issuer = "https://plytarium.com"
|
||||
srv := NewServer(opts)
|
||||
req := tokenRequest(url.Values{"grant_type": {"authorization_code"}}, "")
|
||||
rec := httptest.NewRecorder()
|
||||
srv.Token(rec, req)
|
||||
if rec.Code != http.StatusInternalServerError {
|
||||
t.Fatalf("status = %d, want %d", rec.Code, http.StatusInternalServerError)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user