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