Files
summercms/modules/wristband/registration_test.go
Jakub Zych 5e50b166ef 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
2026-09-28 02:21:02 +02:00

621 lines
20 KiB
Go

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