- Move remaining beach packages and embedded admin assets\n- Rewrite framework, example, build, and gate paths
621 lines
20 KiB
Go
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)
|
|
}
|
|
}
|