From 75b64f98c280218bb0f1c0764402fffece26f60d Mon Sep 17 00:00:00 2001 From: Trey Date: Tue, 6 Oct 2026 10:10:20 -0700 Subject: [PATCH 1/4] Add compare-and-set DCR credential update DCRCredentialStore could only create-if-absent or overwrite-if-present, so a caller coordinating DCR re-registration across replicas could not make its overwrite safe against a concurrent writer: a lock holder that stalls past its lease would clobber a newer row on resume. The check and the write cannot be made atomic from outside the store because the lock and row keys hash to different Cluster slots and the row encoding is unexported. Implements changes for issue #6757: - Add UpdateDCRCredentialsIfUnchanged and ErrDCRCredentialsChanged - Memory: compare and write under the storage mutex - Redis: single-key WATCH/MULTI comparing decoded stored forms, so it is Cluster-safe and tolerant of the stored one-second time precision - TTL handling matches UpdateDCRCredentialsIfPresent - Regenerate mocks; add unit and Sentinel integration tests Co-Authored-By: Claude Opus 5.5 --- pkg/authserver/storage/memory.go | 42 ++++ pkg/authserver/storage/memory_test.go | 135 +++++++++++ pkg/authserver/storage/mocks/mock_storage.go | 15 ++ pkg/authserver/storage/redis.go | 127 +++++++++-- .../storage/redis_integration_test.go | 53 +++++ pkg/authserver/storage/redis_test.go | 215 ++++++++++++++++++ pkg/authserver/storage/types.go | 56 +++++ 7 files changed, 623 insertions(+), 20 deletions(-) diff --git a/pkg/authserver/storage/memory.go b/pkg/authserver/storage/memory.go index d1a17a526b..f116605b45 100644 --- a/pkg/authserver/storage/memory.go +++ b/pkg/authserver/storage/memory.go @@ -2058,6 +2058,48 @@ func (s *MemoryStorage) UpdateDCRCredentialsIfPresent(_ context.Context, creds * return cloneDCRCredentials(creds), nil } +// UpdateDCRCredentialsIfUnchanged replaces the entry at creds.Key with creds +// only when the stored entry still equals expected, returning ErrNotFound +// (wrapped) when no entry exists and ErrDCRCredentialsChanged when it differs. +// The compare and the write happen under s.mu, so no other writer can slip in +// between them. Presence and ClientSecretExpiresAt handling are identical to +// UpdateDCRCredentialsIfPresent. +func (s *MemoryStorage) UpdateDCRCredentialsIfUnchanged( + _ context.Context, creds, expected *DCRCredentials, +) (*DCRCredentials, error) { + if err := validateDCRCompareAndSet(creds, expected); err != nil { + return nil, err + } + + s.mu.Lock() + defer s.mu.Unlock() + + existing, ok := s.dcrCredentials[creds.Key] + if !ok { + return nil, notFoundRFC6749Error("DCR credentials not found") + } + if !dcrCredentialsEqual(existing, expected) { + return nil, ErrDCRCredentialsChanged + } + + s.dcrCredentials[creds.Key] = cloneDCRCredentials(creds) + return cloneDCRCredentials(creds), nil +} + +// dcrCredentialsEqual reports whether a and b hold the same persisted values. +// Time fields are compared with time.Time.Equal rather than ==, so two values +// for the same instant compare equal regardless of monotonic-clock reading or +// location. +func dcrCredentialsEqual(a, b *DCRCredentials) bool { + ac, bc := *a, *b + if !ac.CreatedAt.Equal(bc.CreatedAt) || !ac.ClientSecretExpiresAt.Equal(bc.ClientSecretExpiresAt) { + return false + } + ac.CreatedAt, bc.CreatedAt = time.Time{}, time.Time{} + ac.ClientSecretExpiresAt, bc.ClientSecretExpiresAt = time.Time{}, time.Time{} + return ac == bc +} + // GetDCRCredentials retrieves DCR credentials by key. // Returns a defensive copy; returns ErrNotFound (wrapped) on miss. func (s *MemoryStorage) GetDCRCredentials(_ context.Context, key DCRKey) (*DCRCredentials, error) { diff --git a/pkg/authserver/storage/memory_test.go b/pkg/authserver/storage/memory_test.go index e243adb2c6..91933385f2 100644 --- a/pkg/authserver/storage/memory_test.go +++ b/pkg/authserver/storage/memory_test.go @@ -3026,6 +3026,141 @@ func TestMemoryStorage_DCRCredentials_UpdateCopyIsolatesCaller(t *testing.T) { }) } +// TestMemoryStorage_DCRCredentials_UpdateIfUnchanged pins the +// UpdateDCRCredentialsIfUnchanged compare-and-set contract: a row that still +// equals expected is overwritten, a row that changed after expected was read +// is refused with ErrDCRCredentialsChanged, and an absent row is refused with +// ErrNotFound — in both refusal cases nothing is written. +func TestMemoryStorage_DCRCredentials_UpdateIfUnchanged(t *testing.T) { + t.Parallel() + + t.Run("unchanged row is overwritten", func(t *testing.T) { + withStorage(t, func(ctx context.Context, s *MemoryStorage) { + key := dcrFixtureKey() + _, err := s.StoreDCRCredentialsIfAbsent(ctx, dcrCASFixture(key, "original")) + require.NoError(t, err) + expected, err := s.GetDCRCredentials(ctx, key) + require.NoError(t, err) + + replacement := dcrCASFixture(key, "replacement") + got, err := s.UpdateDCRCredentialsIfUnchanged(ctx, replacement, expected) + require.NoError(t, err) + assert.Equal(t, *replacement, *got) + + reread, err := s.GetDCRCredentials(ctx, key) + require.NoError(t, err) + assert.Equal(t, *replacement, *reread) + }) + }) + + t.Run("changed row is refused and left intact", func(t *testing.T) { + withStorage(t, func(ctx context.Context, s *MemoryStorage) { + key := dcrFixtureKey() + _, err := s.StoreDCRCredentialsIfAbsent(ctx, dcrCASFixture(key, "original")) + require.NoError(t, err) + stale, err := s.GetDCRCredentials(ctx, key) + require.NoError(t, err) + + // A concurrent writer replaces the row after the stale read. + newer := dcrCASFixture(key, "newer") + _, err = s.UpdateDCRCredentialsIfPresent(ctx, newer) + require.NoError(t, err) + + _, err = s.UpdateDCRCredentialsIfUnchanged(ctx, dcrCASFixture(key, "stale-writer"), stale) + require.ErrorIs(t, err, ErrDCRCredentialsChanged) + + reread, err := s.GetDCRCredentials(ctx, key) + require.NoError(t, err) + assert.Equal(t, *newer, *reread, "a refused compare-and-set must not write") + }) + }) + + t.Run("absent row returns not-found and does not create", func(t *testing.T) { + withStorage(t, func(ctx context.Context, s *MemoryStorage) { + key := dcrFixtureKey() + _, err := s.UpdateDCRCredentialsIfUnchanged(ctx, + dcrCASFixture(key, "replacement"), dcrCASFixture(key, "original")) + requireNotFoundError(t, err) + + _, getErr := s.GetDCRCredentials(ctx, key) + requireNotFoundError(t, getErr) + }) + }) + + t.Run("equal instants in different locations compare equal", func(t *testing.T) { + withStorage(t, func(ctx context.Context, s *MemoryStorage) { + key := dcrFixtureKey() + stored := dcrCASFixture(key, "original") + stored.CreatedAt = time.Now() + _, err := s.StoreDCRCredentialsIfAbsent(ctx, stored) + require.NoError(t, err) + + expected := dcrCASFixture(key, "original") + expected.CreatedAt = stored.CreatedAt.In(time.FixedZone("elsewhere", 3600)) + _, err = s.UpdateDCRCredentialsIfUnchanged(ctx, dcrCASFixture(key, "replacement"), expected) + require.NoError(t, err) + }) + }) +} + +// TestMemoryStorage_DCRCredentials_UpdateIfUnchangedInvalidInput pins the +// input contract: creds runs the shared validateDCRCredentialsForStore gate, +// expected must be non-nil, and expected must address the same key as creds. +func TestMemoryStorage_DCRCredentials_UpdateIfUnchangedInvalidInput(t *testing.T) { + t.Parallel() + + key := dcrFixtureKey() + otherKey := key + otherKey.UpstreamID = "other-upstream" + + tests := []struct { + name string + creds *DCRCredentials + expected *DCRCredentials + }{ + {name: "nil creds", creds: nil, expected: dcrCASFixture(key, "original")}, + {name: "nil expected", creds: dcrCASFixture(key, "replacement"), expected: nil}, + {name: "mismatched key", creds: dcrCASFixture(key, "replacement"), expected: dcrCASFixture(otherKey, "original")}, + } + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + withStorage(t, func(ctx context.Context, s *MemoryStorage) { + _, err := s.UpdateDCRCredentialsIfUnchanged(ctx, tt.creds, tt.expected) + require.Error(t, err) + assert.ErrorIs(t, err, fosite.ErrInvalidRequest) + }) + }) + } +} + +// TestMemoryStorage_DCRCredentials_UpdateIfUnchangedConcurrent pins that the +// compare and the write are atomic: of N writers that all read the same row +// and race to replace it, exactly one wins and the rest are refused. +func TestMemoryStorage_DCRCredentials_UpdateIfUnchangedConcurrent(t *testing.T) { + withStorage(t, func(ctx context.Context, s *MemoryStorage) { + const goroutines = 8 + key := dcrFixtureKey() + _, err := s.StoreDCRCredentialsIfAbsent(ctx, dcrCASFixture(key, "original")) + require.NoError(t, err) + expected, err := s.GetDCRCredentials(ctx, key) + require.NoError(t, err) + + errs := make([]error, goroutines) + var wg sync.WaitGroup + for g := range goroutines { + wg.Add(1) + go func() { + defer wg.Done() + _, errs[g] = s.UpdateDCRCredentialsIfUnchanged(ctx, + dcrCASFixture(key, fmt.Sprintf("secret-%d", g)), expected) + }() + } + wg.Wait() + + requireSingleDCRCASWinner(t, errs) + }) +} + // TestMemoryStorage_DCRCredentials_StoreInvalidInputRejected pins the // fail-loud-on-invalid-input contract: nil creds, an unpopulated Key // (empty Issuer, UpstreamID, RedirectURI, or ScopesHash), and missing RFC 7591 diff --git a/pkg/authserver/storage/mocks/mock_storage.go b/pkg/authserver/storage/mocks/mock_storage.go index a36dcd147a..64356c9c5f 100644 --- a/pkg/authserver/storage/mocks/mock_storage.go +++ b/pkg/authserver/storage/mocks/mock_storage.go @@ -88,6 +88,21 @@ func (mr *MockDCRCredentialStoreMockRecorder) UpdateDCRCredentialsIfPresent(ctx, return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "UpdateDCRCredentialsIfPresent", reflect.TypeOf((*MockDCRCredentialStore)(nil).UpdateDCRCredentialsIfPresent), ctx, creds) } +// UpdateDCRCredentialsIfUnchanged mocks base method. +func (m *MockDCRCredentialStore) UpdateDCRCredentialsIfUnchanged(ctx context.Context, creds, expected *storage.DCRCredentials) (*storage.DCRCredentials, error) { + m.ctrl.T.Helper() + ret := m.ctrl.Call(m, "UpdateDCRCredentialsIfUnchanged", ctx, creds, expected) + ret0, _ := ret[0].(*storage.DCRCredentials) + ret1, _ := ret[1].(error) + return ret0, ret1 +} + +// UpdateDCRCredentialsIfUnchanged indicates an expected call of UpdateDCRCredentialsIfUnchanged. +func (mr *MockDCRCredentialStoreMockRecorder) UpdateDCRCredentialsIfUnchanged(ctx, creds, expected any) *gomock.Call { + mr.mock.ctrl.T.Helper() + return mr.mock.ctrl.RecordCallWithMethodType(mr.mock, "UpdateDCRCredentialsIfUnchanged", reflect.TypeOf((*MockDCRCredentialStore)(nil).UpdateDCRCredentialsIfUnchanged), ctx, creds, expected) +} + // MockPendingAuthorizationStorage is a mock of PendingAuthorizationStorage interface. type MockPendingAuthorizationStorage struct { ctrl *gomock.Controller diff --git a/pkg/authserver/storage/redis.go b/pkg/authserver/storage/redis.go index fefc4ec726..83b0f61c6f 100644 --- a/pkg/authserver/storage/redis.go +++ b/pkg/authserver/storage/redis.go @@ -1879,6 +1879,33 @@ type storedDCRCredentials struct { ClientSecretExpiresAt int64 `json:"client_secret_expires_at"` } +// newStoredDCRCredentials converts creds to its stored form, mirroring +// newStoredUpstreamTokens. Time fields are truncated to Unix seconds; a zero +// time.Time becomes the 0 "not set" sentinel. +func newStoredDCRCredentials(creds *DCRCredentials) storedDCRCredentials { + stored := storedDCRCredentials{ + KeyIssuer: creds.Key.Issuer, + KeyUpstreamID: creds.Key.UpstreamID, + KeyRedirectURI: creds.Key.RedirectURI, + KeyScopesHash: creds.Key.ScopesHash, + ProviderName: creds.ProviderName, + ClientID: creds.ClientID, + ClientSecret: creds.ClientSecret, + TokenEndpointAuthMethod: creds.TokenEndpointAuthMethod, + RegistrationAccessToken: creds.RegistrationAccessToken, + RegistrationClientURI: creds.RegistrationClientURI, + AuthorizationEndpoint: creds.AuthorizationEndpoint, + TokenEndpoint: creds.TokenEndpoint, + } + if !creds.CreatedAt.IsZero() { + stored.CreatedAt = creds.CreatedAt.Unix() + } + if !creds.ClientSecretExpiresAt.IsZero() { + stored.ClientSecretExpiresAt = creds.ClientSecretExpiresAt.Unix() + } + return stored +} + // toDCRCredentials decodes the stored form back into the public type, mirroring // storedUpstreamTokens.toUpstreamTokens. Zero epoch values become the zero // time.Time, preserving the "not set" sentinel. @@ -2024,26 +2051,7 @@ func (s *RedisStorage) StoreDCRCredentialsIfAbsent(ctx context.Context, creds *D // // See StoreDCRCredentialsIfAbsent's docstring for the past-expiry rationale. func marshalDCRCredentialsForStore(creds *DCRCredentials) ([]byte, time.Duration, error) { - stored := storedDCRCredentials{ - KeyIssuer: creds.Key.Issuer, - KeyUpstreamID: creds.Key.UpstreamID, - KeyRedirectURI: creds.Key.RedirectURI, - KeyScopesHash: creds.Key.ScopesHash, - ProviderName: creds.ProviderName, - ClientID: creds.ClientID, - ClientSecret: creds.ClientSecret, - TokenEndpointAuthMethod: creds.TokenEndpointAuthMethod, - RegistrationAccessToken: creds.RegistrationAccessToken, - RegistrationClientURI: creds.RegistrationClientURI, - AuthorizationEndpoint: creds.AuthorizationEndpoint, - TokenEndpoint: creds.TokenEndpoint, - } - if !creds.CreatedAt.IsZero() { - stored.CreatedAt = creds.CreatedAt.Unix() - } - if !creds.ClientSecretExpiresAt.IsZero() { - stored.ClientSecretExpiresAt = creds.ClientSecretExpiresAt.Unix() - } + stored := newStoredDCRCredentials(creds) data, err := json.Marshal(stored) //nolint:gosec // G117 - internal Redis storage serialization, not exposed to users if err != nil { @@ -2186,6 +2194,85 @@ func (s *RedisStorage) UpdateDCRCredentialsIfPresent(ctx context.Context, creds return cloneDCRCredentials(creds), nil } +// UpdateDCRCredentialsIfUnchanged replaces the row at creds.Key with creds +// only when the stored row still equals expected, returning ErrNotFound +// (wrapped) when no row exists and ErrDCRCredentialsChanged when it differs. +// +// The compare and the write run in a WATCH/MULTI transaction on the single +// row key — the same pattern StoreDCRCredentialsIfAbsent uses — so the +// operation is atomic against concurrent writers and valid in standalone, +// Sentinel and Cluster modes alike (one key means one hash slot). If another +// writer touches the key between the GET and EXEC, EXEC aborts with +// redis.TxFailedErr and the read-compare-write is retried, up to +// maxDCRClaimRetries; the retry re-reads the row and so normally resolves to +// ErrDCRCredentialsChanged. +// +// The comparison is between decoded stored forms (expected is normalised via +// newStoredDCRCredentials), not raw JSON bytes, so it is unaffected by field +// ordering or by a row written by an older encoder, and a value obtained from +// GetDCRCredentials always matches the row it was read from despite the +// one-second time precision of the stored form. +// +// Presence and TTL handling are identical to UpdateDCRCredentialsIfPresent: +// an expired-but-present row is updatable, and the rewritten row's TTL is +// derived from creds by marshalDCRCredentialsForStore. The write uses SET XX +// for parity with that method, although WATCH already guarantees the key +// still exists at EXEC. +func (s *RedisStorage) UpdateDCRCredentialsIfUnchanged( + ctx context.Context, creds, expected *DCRCredentials, +) (*DCRCredentials, error) { + if err := validateDCRCompareAndSet(creds, expected); err != nil { + return nil, err + } + + key := redisDCRKey(s.keyPrefix, creds.Key) + + data, ttl, err := marshalDCRCredentialsForStore(creds) + if err != nil { + return nil, err + } + want := newStoredDCRCredentials(expected) + + txFn := func(tx *redis.Tx) error { + existingData, getErr := tx.Get(ctx, key).Bytes() + if errors.Is(getErr, redis.Nil) { + return notFoundRFC6749Error("DCR credentials not found") + } + if getErr != nil { + return fmt.Errorf("failed to get existing dcr credentials: %w", getErr) + } + + var existing storedDCRCredentials + if unmarshalErr := json.Unmarshal(existingData, &existing); unmarshalErr != nil { + return fmt.Errorf("failed to unmarshal existing dcr credentials: %w", unmarshalErr) + } + if existing != want { + return ErrDCRCredentialsChanged + } + + _, pipeErr := tx.TxPipelined(ctx, func(pipe redis.Pipeliner) error { + pipe.SetArgs(ctx, key, data, redis.SetArgs{Mode: "XX", TTL: ttl}) + return nil + }) + return pipeErr + } + + var watchErr error + for attempt := 0; attempt < maxDCRClaimRetries; attempt++ { + watchErr = s.client.Watch(ctx, txFn, key) + if watchErr == nil { + return cloneDCRCredentials(creds), nil + } + if errors.Is(watchErr, ErrNotFound) || errors.Is(watchErr, ErrDCRCredentialsChanged) { + return nil, watchErr + } + if !errors.Is(watchErr, redis.TxFailedErr) { + return nil, fmt.Errorf("failed to compare-and-set dcr credentials: %w", watchErr) + } + } + return nil, fmt.Errorf("failed to compare-and-set dcr credentials after %d attempts: %w", maxDCRClaimRetries, watchErr) +} + // GetDCRCredentials retrieves the credentials previously persisted under key. // Returns ErrNotFound (wrapped) when no entry exists. The returned value is a // fresh struct decoded from JSON, which acts as a defensive copy. diff --git a/pkg/authserver/storage/redis_integration_test.go b/pkg/authserver/storage/redis_integration_test.go index 23712d6429..9c4066e78e 100644 --- a/pkg/authserver/storage/redis_integration_test.go +++ b/pkg/authserver/storage/redis_integration_test.go @@ -1600,6 +1600,59 @@ func TestIntegration_DCRCredentials_UpdateIfPresent(t *testing.T) { }) } +// TestIntegration_DCRCredentials_UpdateIfUnchanged pins the +// UpdateDCRCredentialsIfUnchanged WATCH/MULTI path against a real Redis +// Sentinel cluster: an unchanged row is overwritten with a TTL derived from the +// incoming creds, and a row replaced after the caller's read is refused and +// left intact. +func TestIntegration_DCRCredentials_UpdateIfUnchanged(t *testing.T) { + t.Parallel() + + t.Run("unchanged row is overwritten and TTL applied", func(t *testing.T) { + withIntegrationStorage(t, func(ctx context.Context, s *RedisStorage) { + key := dcrFixtureKey() + _, err := s.StoreDCRCredentialsIfAbsent(ctx, dcrCASFixture(key, "original")) + require.NoError(t, err) + expected, err := s.GetDCRCredentials(ctx, key) + require.NoError(t, err) + + replacement := dcrCASFixture(key, "replacement") + replacement.ClientSecretExpiresAt = time.Now().Add(24 * time.Hour).Truncate(time.Second) + _, err = s.UpdateDCRCredentialsIfUnchanged(ctx, replacement, expected) + require.NoError(t, err) + + reread, err := s.GetDCRCredentials(ctx, key) + require.NoError(t, err) + assert.Equal(t, "replacement", reread.ClientSecret) + + ttl, err := s.client.TTL(ctx, redisDCRKey(s.keyPrefix, key)).Result() + require.NoError(t, err) + assert.Greater(t, ttl, time.Duration(0)) + assert.LessOrEqual(t, ttl, 24*time.Hour) + }) + }) + + t.Run("changed row is refused and left intact", func(t *testing.T) { + withIntegrationStorage(t, func(ctx context.Context, s *RedisStorage) { + key := dcrFixtureKey() + _, err := s.StoreDCRCredentialsIfAbsent(ctx, dcrCASFixture(key, "original")) + require.NoError(t, err) + stale, err := s.GetDCRCredentials(ctx, key) + require.NoError(t, err) + + _, err = s.UpdateDCRCredentialsIfPresent(ctx, dcrCASFixture(key, "newer")) + require.NoError(t, err) + + _, err = s.UpdateDCRCredentialsIfUnchanged(ctx, dcrCASFixture(key, "stale-writer"), stale) + require.ErrorIs(t, err, ErrDCRCredentialsChanged) + + reread, err := s.GetDCRCredentials(ctx, key) + require.NoError(t, err) + assert.Equal(t, "newer", reread.ClientSecret) + }) + }) +} + // TestIntegration_DCRCredentials_TTL pins the RFC 7591 §3.2.1 TTL contract // against a real Redis Sentinel cluster: TTL command observes the expected // state for both the expiring and the never-expires cases. The unit-level diff --git a/pkg/authserver/storage/redis_test.go b/pkg/authserver/storage/redis_test.go index 9f0c716616..e856cc89cf 100644 --- a/pkg/authserver/storage/redis_test.go +++ b/pkg/authserver/storage/redis_test.go @@ -3825,6 +3825,36 @@ func dcrFixtureKey() DCRKey { } } +// dcrCASFixture returns a valid DCRCredentials at key whose ClientSecret is +// secret, for the UpdateDCRCredentialsIfUnchanged tests on every backend. +func dcrCASFixture(key DCRKey, secret string) *DCRCredentials { + return &DCRCredentials{ + Key: key, + ClientID: "client-" + secret, + ClientSecret: secret, + AuthorizationEndpoint: "https://idp.example.com/auth", + TokenEndpoint: "https://idp.example.com/token", + } +} + +// requireSingleDCRCASWinner asserts that of a set of racing +// UpdateDCRCredentialsIfUnchanged calls made against the same expected value, +// exactly one succeeded and every other one was refused with +// ErrDCRCredentialsChanged. +func requireSingleDCRCASWinner(t *testing.T, errs []error) { + t.Helper() + winners := 0 + for gid, e := range errs { + if e == nil { + winners++ + continue + } + require.ErrorIsf(t, e, ErrDCRCredentialsChanged, + "a losing writer must be refused as changed (goroutine %d)", gid) + } + require.Equal(t, 1, winners, "exactly one racing compare-and-set must win") +} + func TestRedisStorage_DCRCredentials_RoundTrip(t *testing.T) { withRedisStorage(t, func(ctx context.Context, s *RedisStorage, _ *miniredis.Miniredis) { // Truncate to second precision: time fields are stored as int64 unix seconds. @@ -4310,6 +4340,191 @@ func TestRedisStorage_DCRCredentials_UpdateConcurrent(t *testing.T) { }) } +// TestRedisStorage_DCRCredentials_UpdateIfUnchanged pins the +// UpdateDCRCredentialsIfUnchanged compare-and-set contract on the Redis +// backend, mirroring the memory backend's equivalent, plus the Redis-specific +// TTL and stored-precision behaviour. +func TestRedisStorage_DCRCredentials_UpdateIfUnchanged(t *testing.T) { + t.Parallel() + + t.Run("unchanged row is overwritten", func(t *testing.T) { + withRedisStorage(t, func(ctx context.Context, s *RedisStorage, _ *miniredis.Miniredis) { + key := dcrFixtureKey() + _, err := s.StoreDCRCredentialsIfAbsent(ctx, dcrCASFixture(key, "original")) + require.NoError(t, err) + expected, err := s.GetDCRCredentials(ctx, key) + require.NoError(t, err) + + replacement := dcrCASFixture(key, "replacement") + got, err := s.UpdateDCRCredentialsIfUnchanged(ctx, replacement, expected) + require.NoError(t, err) + assert.Equal(t, *replacement, *got) + + reread, err := s.GetDCRCredentials(ctx, key) + require.NoError(t, err) + assert.Equal(t, *replacement, *reread) + }) + }) + + t.Run("changed row is refused and left intact", func(t *testing.T) { + withRedisStorage(t, func(ctx context.Context, s *RedisStorage, _ *miniredis.Miniredis) { + key := dcrFixtureKey() + _, err := s.StoreDCRCredentialsIfAbsent(ctx, dcrCASFixture(key, "original")) + require.NoError(t, err) + stale, err := s.GetDCRCredentials(ctx, key) + require.NoError(t, err) + + // A concurrent writer replaces the row after the stale read. + newer := dcrCASFixture(key, "newer") + _, err = s.UpdateDCRCredentialsIfPresent(ctx, newer) + require.NoError(t, err) + + _, err = s.UpdateDCRCredentialsIfUnchanged(ctx, dcrCASFixture(key, "stale-writer"), stale) + require.ErrorIs(t, err, ErrDCRCredentialsChanged) + + reread, err := s.GetDCRCredentials(ctx, key) + require.NoError(t, err) + assert.Equal(t, *newer, *reread, "a refused compare-and-set must not write") + }) + }) + + t.Run("absent row returns not-found and does not create", func(t *testing.T) { + withRedisStorage(t, func(ctx context.Context, s *RedisStorage, _ *miniredis.Miniredis) { + key := dcrFixtureKey() + _, err := s.UpdateDCRCredentialsIfUnchanged(ctx, + dcrCASFixture(key, "replacement"), dcrCASFixture(key, "original")) + requireRedisNotFoundError(t, err) + + _, getErr := s.GetDCRCredentials(ctx, key) + requireRedisNotFoundError(t, getErr) + }) + }) + + t.Run("expected is compared at stored one-second precision", func(t *testing.T) { + withRedisStorage(t, func(ctx context.Context, s *RedisStorage, _ *miniredis.Miniredis) { + key := dcrFixtureKey() + // A sub-second CreatedAt is truncated on the way in, so the caller's + // own pre-store value must still match the stored row. + stored := dcrCASFixture(key, "original") + stored.CreatedAt = time.Unix(1_700_000_000, 123_456_789) + _, err := s.StoreDCRCredentialsIfAbsent(ctx, stored) + require.NoError(t, err) + + _, err = s.UpdateDCRCredentialsIfUnchanged(ctx, dcrCASFixture(key, "replacement"), stored) + require.NoError(t, err) + }) + }) + + t.Run("TTL follows UpdateDCRCredentialsIfPresent rules", func(t *testing.T) { + withRedisStorage(t, func(ctx context.Context, s *RedisStorage, mr *miniredis.Miniredis) { + key := dcrFixtureKey() + redisKey := redisDCRKey(s.keyPrefix, key) + _, err := s.StoreDCRCredentialsIfAbsent(ctx, dcrCASFixture(key, "original")) + require.NoError(t, err) + require.Equal(t, time.Duration(0), mr.TTL(redisKey)) + + // Future expiry sets a TTL. + expected, err := s.GetDCRCredentials(ctx, key) + require.NoError(t, err) + expiring := dcrCASFixture(key, "expiring") + expiring.ClientSecretExpiresAt = time.Now().Add(12 * time.Hour).Truncate(time.Second) + _, err = s.UpdateDCRCredentialsIfUnchanged(ctx, expiring, expected) + require.NoError(t, err) + ttl := mr.TTL(redisKey) + assert.Greater(t, ttl, time.Duration(0), "future expiry must set a positive TTL") + assert.LessOrEqual(t, ttl, 12*time.Hour) + + // Zero expiry clears it again. + expected, err = s.GetDCRCredentials(ctx, key) + require.NoError(t, err) + _, err = s.UpdateDCRCredentialsIfUnchanged(ctx, dcrCASFixture(key, "persistent"), expected) + require.NoError(t, err) + assert.Equal(t, time.Duration(0), mr.TTL(redisKey), "zero expiry must clear the TTL") + + // Past expiry uses the bounded pastExpiryDCRTTL. + expected, err = s.GetDCRCredentials(ctx, key) + require.NoError(t, err) + expired := dcrCASFixture(key, "expired") + expired.ClientSecretExpiresAt = time.Now().Add(-time.Hour).Truncate(time.Second) + _, err = s.UpdateDCRCredentialsIfUnchanged(ctx, expired, expected) + require.NoError(t, err) + assert.Equal(t, pastExpiryDCRTTL, mr.TTL(redisKey)) + }) + }) +} + +// TestRedisStorage_DCRCredentials_UpdateIfUnchangedInvalidInput pins that the +// Redis backend runs the shared validateDCRCompareAndSet gate. The full matrix +// is covered by TestMemoryStorage_DCRCredentials_UpdateIfUnchangedInvalidInput +// against the same function; this only confirms the wiring. +func TestRedisStorage_DCRCredentials_UpdateIfUnchangedInvalidInput(t *testing.T) { + withRedisStorage(t, func(ctx context.Context, s *RedisStorage, _ *miniredis.Miniredis) { + _, err := s.UpdateDCRCredentialsIfUnchanged(ctx, dcrCASFixture(dcrFixtureKey(), "replacement"), nil) + require.Error(t, err) + assert.ErrorIs(t, err, fosite.ErrInvalidRequest) + }) +} + +// TestRedisStorage_DCRCredentials_UpdateIfUnchangedConnectionFailure pins that +// a connection failure surfaces as a generic error, never as a spurious +// ErrNotFound or ErrDCRCredentialsChanged. +func TestRedisStorage_DCRCredentials_UpdateIfUnchangedConnectionFailure(t *testing.T) { + t.Parallel() + + mr := miniredis.RunT(t) + client := redis.NewClient(&redis.Options{Addr: mr.Addr()}) + s := NewRedisStorageWithClient(client, "test:auth:") + mr.Close() + + key := dcrFixtureKey() + _, err := s.UpdateDCRCredentialsIfUnchanged(context.Background(), + dcrCASFixture(key, "replacement"), dcrCASFixture(key, "original")) + require.Error(t, err) + assert.NotErrorIs(t, err, ErrNotFound) + assert.NotErrorIs(t, err, ErrDCRCredentialsChanged) +} + +// TestRedisStorage_DCRCredentials_UpdateIfUnchangedConcurrent pins that the +// WATCH/MULTI compare and write are atomic: of N writers that all read the +// same row and race to replace it, exactly one wins, every other one is +// refused as changed, and the stored row is the winner's. +func TestRedisStorage_DCRCredentials_UpdateIfUnchangedConcurrent(t *testing.T) { + withRedisStorage(t, func(ctx context.Context, s *RedisStorage, _ *miniredis.Miniredis) { + const goroutines = 8 + key := dcrFixtureKey() + _, err := s.StoreDCRCredentialsIfAbsent(ctx, dcrCASFixture(key, "original")) + require.NoError(t, err) + expected, err := s.GetDCRCredentials(ctx, key) + require.NoError(t, err) + + start := make(chan struct{}) + errs := make([]error, goroutines) + results := make([]*DCRCredentials, goroutines) + var wg sync.WaitGroup + for g := range goroutines { + wg.Add(1) + go func() { + defer wg.Done() + <-start + results[g], errs[g] = s.UpdateDCRCredentialsIfUnchanged(ctx, + dcrCASFixture(key, fmt.Sprintf("secret-%d", g)), expected) + }() + } + close(start) + wg.Wait() + + requireSingleDCRCASWinner(t, errs) + + got, err := s.GetDCRCredentials(ctx, key) + require.NoError(t, err) + for g, e := range errs { + if e == nil { + assert.Equal(t, *results[g], *got, "the stored row must be the winner's") + } + } + }) +} + // TestRedisStorage_DCRCredentials_FirstClaimWins_Concurrent proves the // create-if-absent race-safety property holds under genuine concurrency, not // just the sequential shape of TestRedisStorage_DCRCredentials_FirstClaimWins: diff --git a/pkg/authserver/storage/types.go b/pkg/authserver/storage/types.go index b6964c6470..1c32667f8a 100644 --- a/pkg/authserver/storage/types.go +++ b/pkg/authserver/storage/types.go @@ -79,6 +79,13 @@ var ( // CompareAndSwapUpstreamTokens for the full coordination contract). ErrConcurrentRefresh = errors.New("storage: upstream token row changed concurrently") + // ErrDCRCredentialsChanged is returned by + // DCRCredentialStore.UpdateDCRCredentialsIfUnchanged when the stored row + // no longer equals the caller's expected value: another writer replaced it + // after the caller read it. Nothing is written. A row that is absent + // altogether is reported as ErrNotFound instead. + ErrDCRCredentialsChanged = errors.New("storage: dcr credentials row changed concurrently") + // ErrInvalidState is returned when an operation requires an item to be in a // particular lifecycle state (e.g. a pending device request) but it is not. ErrInvalidState = errors.New("storage: item is not in the required state") @@ -306,6 +313,24 @@ func validateDCRCredentialsForStore(creds *DCRCredentials) error { return nil } +// validateDCRCompareAndSet enforces the input contract of +// DCRCredentialStore.UpdateDCRCredentialsIfUnchanged on top of +// validateDCRCredentialsForStore: expected must be present and must address +// the same row as creds, so a caller cannot gate a write to one key on the +// contents of another. +func validateDCRCompareAndSet(creds, expected *DCRCredentials) error { + if err := validateDCRCredentialsForStore(creds); err != nil { + return err + } + if expected == nil { + return fosite.ErrInvalidRequest.WithHint("expected dcr credentials cannot be nil") + } + if expected.Key != creds.Key { + return fosite.ErrInvalidRequest.WithHint("expected dcr credentials key must match creds key") + } + return nil +} + // DCRCredentials is the persisted form of an RFC 7591 Dynamic Client // Registration result. All fields are populated from the upstream's DCR // response. The RFC 7592 management fields (RegistrationAccessToken, @@ -473,6 +498,37 @@ type DCRCredentialStore interface { // from the incoming creds, so an update can extend, shorten, or clear the // row's TTL exactly as an initial store would. UpdateDCRCredentialsIfPresent(ctx context.Context, creds *DCRCredentials) (*DCRCredentials, error) + + // UpdateDCRCredentialsIfUnchanged replaces the record at creds.Key with + // creds, iff the record currently stored there still equals expected — + // the value the caller previously read via GetDCRCredentials. It is the + // compare-and-set sibling of UpdateDCRCredentialsIfPresent, for callers + // that coordinate re-registration across replicas and must not clobber a + // newer row written by a concurrent writer after their read (e.g. a lock + // holder that stalled past its lease before writing). + // + // Returns ErrNotFound (wrapped) if no record exists at the key, and + // ErrDCRCredentialsChanged if the stored record no longer equals + // expected. In both cases nothing is written. On success it returns the + // stored value (a defensive copy) and the returned *DCRCredentials is + // non-nil. + // + // expected must be non-nil and its Key must equal creds.Key. Equality is + // over every persisted field, compared at the backend's storage precision + // (the Redis backend persists times at one-second resolution), so a value + // returned by GetDCRCredentials always compares equal to the row it was + // read from. A decorator that transforms fields on the way in and out + // (encryption, compression, …) must pass the inner store the stored form + // of expected — what the inner GetDCRCredentials returned — not the + // decoded form it handed to its own caller. + // + // Presence is physical, exactly as for UpdateDCRCredentialsIfPresent, and + // the rewritten row's backend TTL follows the same ClientSecretExpiresAt + // rules. Implementations MUST perform the compare and the write + // atomically with respect to other writers of the same key, and MUST + // touch only that one key so the operation stays valid under Redis + // Cluster. + UpdateDCRCredentialsIfUnchanged(ctx context.Context, creds, expected *DCRCredentials) (*DCRCredentials, error) } // User represents a user account in the authorization server. From 06f784b426481646f44406c733c604711eabc290 Mon Sep 17 00:00:00 2001 From: Trey Date: Tue, 6 Oct 2026 10:12:21 -0700 Subject: [PATCH 2/4] Add start barrier to memory DCR CAS race test Co-Authored-By: Claude Opus 5.5 --- pkg/authserver/storage/memory_test.go | 3 +++ 1 file changed, 3 insertions(+) diff --git a/pkg/authserver/storage/memory_test.go b/pkg/authserver/storage/memory_test.go index 91933385f2..0c29ea07d9 100644 --- a/pkg/authserver/storage/memory_test.go +++ b/pkg/authserver/storage/memory_test.go @@ -3145,16 +3145,19 @@ func TestMemoryStorage_DCRCredentials_UpdateIfUnchangedConcurrent(t *testing.T) expected, err := s.GetDCRCredentials(ctx, key) require.NoError(t, err) + start := make(chan struct{}) errs := make([]error, goroutines) var wg sync.WaitGroup for g := range goroutines { wg.Add(1) go func() { defer wg.Done() + <-start _, errs[g] = s.UpdateDCRCredentialsIfUnchanged(ctx, dcrCASFixture(key, fmt.Sprintf("secret-%d", g)), expected) }() } + close(start) wg.Wait() requireSingleDCRCASWinner(t, errs) From 3d6f069563be8643f27cf99e2a9f68faac624983 Mon Sep 17 00:00:00 2001 From: Trey Date: Tue, 6 Oct 2026 10:23:40 -0700 Subject: [PATCH 3/4] Clarify DCR compare-and-set contract Addresses stacklok/toolhive#6758 review comments: - MEDIUM types.go (4198367657): document how a decorator with a non-reproducible transform supplies the stored form of expected - LOW redis.go (4198367706): document the retry-exhaustion outcome - LOW memory.go (4198367719): note new time.Time fields must be compared with Equal in dcrCredentialsEqual - LOW redis.go (4198367731): extend maxDCRClaimRetries comment to cover UpdateDCRCredentialsIfUnchanged Co-Authored-By: Claude Opus 5.5 --- pkg/authserver/storage/memory.go | 7 +++++++ pkg/authserver/storage/redis.go | 19 +++++++++++-------- pkg/authserver/storage/types.go | 28 ++++++++++++++++++++++++---- 3 files changed, 42 insertions(+), 12 deletions(-) diff --git a/pkg/authserver/storage/memory.go b/pkg/authserver/storage/memory.go index f116605b45..a87c41111a 100644 --- a/pkg/authserver/storage/memory.go +++ b/pkg/authserver/storage/memory.go @@ -2090,6 +2090,13 @@ func (s *MemoryStorage) UpdateDCRCredentialsIfUnchanged( // Time fields are compared with time.Time.Equal rather than ==, so two values // for the same instant compare equal regardless of monotonic-clock reading or // location. +// +// The final == covers every other field and only compiles while +// DCRCredentials stays comparable. A new time.Time field added to +// DCRCredentials MUST be compared with Equal and zeroed here alongside the +// existing two; == on it would reintroduce the monotonic-clock/location +// mismatch. TestDCRCredentialsEqual_TimeFieldsCompareByInstant fails if one +// is missed. func dcrCredentialsEqual(a, b *DCRCredentials) bool { ac, bc := *a, *b if !ac.CreatedAt.Equal(bc.CreatedAt) || !ac.ClientSecretExpiresAt.Equal(bc.ClientSecretExpiresAt) { diff --git a/pkg/authserver/storage/redis.go b/pkg/authserver/storage/redis.go index 83b0f61c6f..ac102c0484 100644 --- a/pkg/authserver/storage/redis.go +++ b/pkg/authserver/storage/redis.go @@ -39,12 +39,13 @@ const nullMarker = "null" // confuse this row with a healthy long-lived registration. const pastExpiryDCRTTL = time.Second -// maxDCRClaimRetries bounds StoreDCRCredentialsIfAbsent's WATCH/MULTI retry -// loop. go-redis does not retry Watch internally: a concurrent write to the -// watched key (another replica claiming, refreshing, or evicting the same -// row) aborts the pipelined EXEC with redis.TxFailedErr, which Watch returns -// to the caller unwrapped. Retrying a small, fixed number of times lets a -// real concurrent write during the exact race this method exists to close +// maxDCRClaimRetries bounds the WATCH/MULTI retry loops of +// StoreDCRCredentialsIfAbsent and UpdateDCRCredentialsIfUnchanged. go-redis +// does not retry Watch internally: a concurrent write to the watched key +// (another replica claiming, refreshing, updating, or evicting the same row) +// aborts the pipelined EXEC with redis.TxFailedErr, which Watch returns to +// the caller unwrapped. Retrying a small, fixed number of times lets a real +// concurrent write during the exact race these methods exist to close // resolve on its own rather than failing the caller with a spurious error. // Mirrors maxConfiguredClientReconcileRetries above for the same reason. const maxDCRClaimRetries = 3 @@ -2205,10 +2206,12 @@ func (s *RedisStorage) UpdateDCRCredentialsIfPresent(ctx context.Context, creds // writer touches the key between the GET and EXEC, EXEC aborts with // redis.TxFailedErr and the read-compare-write is retried, up to // maxDCRClaimRetries; the retry re-reads the row and so normally resolves to -// ErrDCRCredentialsChanged. +// ErrDCRCredentialsChanged. If every attempt aborts, a generic error is +// returned and nothing was written by this call. // // The comparison is between decoded stored forms (expected is normalised via -// newStoredDCRCredentials), not raw JSON bytes, so it is unaffected by field +// newStoredDCRCredentials, and storedDCRCredentials holds only strings and +// int64s, so == compares every field by value), not raw JSON bytes, so it is unaffected by field // ordering or by a row written by an older encoder, and a value obtained from // GetDCRCredentials always matches the row it was read from despite the // one-second time precision of the stored form. diff --git a/pkg/authserver/storage/types.go b/pkg/authserver/storage/types.go index 1c32667f8a..bd9a61da82 100644 --- a/pkg/authserver/storage/types.go +++ b/pkg/authserver/storage/types.go @@ -513,14 +513,34 @@ type DCRCredentialStore interface { // stored value (a defensive copy) and the returned *DCRCredentials is // non-nil. // + // Any other error leaves the row's state to be determined by re-reading + // it. This includes transport failures, where the write may or may not + // have been applied. It also includes a backend that bounds its internal + // retries (the Redis backend) giving up because other writers kept + // touching the key between its compare and its write; in that case this + // call wrote nothing, and the caller may re-read and try again. + // // expected must be non-nil and its Key must equal creds.Key. Equality is // over every persisted field, compared at the backend's storage precision // (the Redis backend persists times at one-second resolution), so a value // returned by GetDCRCredentials always compares equal to the row it was - // read from. A decorator that transforms fields on the way in and out - // (encryption, compression, …) must pass the inner store the stored form - // of expected — what the inner GetDCRCredentials returned — not the - // decoded form it handed to its own caller. + // read from. + // + // A decorator that transforms fields on the way in and out (encryption, + // compression, …) must pass the inner store the stored form of expected — + // what the inner GetDCRCredentials returned — not the decoded form it + // handed to its own caller. When the transform is not reproducible (e.g. + // encryption with a fresh nonce per seal), the decorator cannot derive + // that stored form from the caller's decoded expected. Instead it should + // re-read the inner row, decode it, return ErrDCRCredentialsChanged if + // the decoded value differs from the caller's expected, and otherwise + // call the inner UpdateDCRCredentialsIfUnchanged with the raw inner row + // it just read as expected. That stays atomic: if the row changes between + // the decorator's re-read and the inner write, the inner compare refuses + // it. A decorator that recovers a row it cannot decode at all already + // holds the raw inner row and passes it directly. A consumer outside the + // decorator that needs to compare-and-set a row it cannot decode would + // need an opaque version or etag instead; that is not provided today. // // Presence is physical, exactly as for UpdateDCRCredentialsIfPresent, and // the rewritten row's backend TTL follows the same ClientSecretExpiresAt From 170008591b6449c434d559bde1dcc0c78c75203b Mon Sep 17 00:00:00 2001 From: Trey Date: Tue, 6 Oct 2026 10:24:59 -0700 Subject: [PATCH 4/4] Cover remaining DCR compare-and-set paths Addresses stacklok/toolhive#6758 review comments: - MEDIUM redis_test.go (4198367681): test WATCH retry, retry exhaustion, corrupt stored row, and expired-but-present row - MEDIUM memory_test.go (4198367695): per-field comparison table run on both backends, with a reflection check that every field is listed - LOW memory.go (4198367719): reflection test that every time.Time field in dcrCredentialsEqual compares by instant Co-Authored-By: Claude Opus 5.5 --- pkg/authserver/storage/memory_test.go | 47 ++++++ pkg/authserver/storage/redis_test.go | 235 ++++++++++++++++++++++++++ 2 files changed, 282 insertions(+) diff --git a/pkg/authserver/storage/memory_test.go b/pkg/authserver/storage/memory_test.go index 0c29ea07d9..22c75c7590 100644 --- a/pkg/authserver/storage/memory_test.go +++ b/pkg/authserver/storage/memory_test.go @@ -23,6 +23,7 @@ import ( "errors" "fmt" "net/url" + "reflect" "sync" "sync/atomic" "testing" @@ -3103,6 +3104,52 @@ func TestMemoryStorage_DCRCredentials_UpdateIfUnchanged(t *testing.T) { }) } +// TestMemoryStorage_DCRCredentials_UpdateIfUnchangedPerField runs the shared +// per-field comparison test against the memory backend. +func TestMemoryStorage_DCRCredentials_UpdateIfUnchangedPerField(t *testing.T) { + t.Parallel() + runDCRCASPerFieldTest(t, func(t *testing.T) DCRCredentialStore { + t.Helper() + s := NewMemoryStorage() + t.Cleanup(func() { _ = s.Close() }) + return s + }) +} + +// TestDCRCredentialsEqual_TimeFieldsCompareByInstant pins that +// dcrCredentialsEqual compares every time.Time field of DCRCredentials by +// instant, not with ==. It discovers the fields by reflection, so a time +// field added to DCRCredentials without a matching Equal in the helper fails +// here instead of silently reintroducing monotonic-clock/location mismatches. +func TestDCRCredentialsEqual_TimeFieldsCompareByInstant(t *testing.T) { + t.Parallel() + + timeType := reflect.TypeFor[time.Time]() + typ := reflect.TypeFor[DCRCredentials]() + elsewhere := time.FixedZone("elsewhere", 3600) + + for i := range typ.NumField() { + field := typ.Field(i) + if field.Type != timeType { + continue + } + t.Run(field.Name, func(t *testing.T) { + t.Parallel() + now := time.Now() // carries a monotonic reading + a := dcrCASFixture(dcrFixtureKey(), "original") + b := cloneDCRCredentials(a) + reflect.ValueOf(a).Elem().Field(i).Set(reflect.ValueOf(now)) + reflect.ValueOf(b).Elem().Field(i).Set(reflect.ValueOf(now.Round(0).In(elsewhere))) + + assert.True(t, dcrCredentialsEqual(a, b), + "equal instants in %s must compare equal regardless of monotonic reading or location", field.Name) + + reflect.ValueOf(b).Elem().Field(i).Set(reflect.ValueOf(now.Add(time.Nanosecond))) + assert.False(t, dcrCredentialsEqual(a, b), "different instants in %s must not compare equal", field.Name) + }) + } +} + // TestMemoryStorage_DCRCredentials_UpdateIfUnchangedInvalidInput pins the // input contract: creds runs the shared validateDCRCredentialsForStore gate, // expected must be non-nil, and expected must address the same key as creds. diff --git a/pkg/authserver/storage/redis_test.go b/pkg/authserver/storage/redis_test.go index e856cc89cf..1c9335164d 100644 --- a/pkg/authserver/storage/redis_test.go +++ b/pkg/authserver/storage/redis_test.go @@ -17,6 +17,7 @@ import ( "fmt" "log/slog" "net/url" + "reflect" "strings" "sync" "sync/atomic" @@ -3837,6 +3838,81 @@ func dcrCASFixture(key DCRKey, secret string) *DCRCredentials { } } +// dcrCASFieldMutations mutates exactly one persisted field of a +// DCRCredentials each, keyed by field name. Key is excluded: a mismatched key +// is rejected as invalid input before any compare. Time fields move by a full +// second so the change survives the Redis backend's one-second precision. +var dcrCASFieldMutations = map[string]func(*DCRCredentials){ + "ProviderName": func(c *DCRCredentials) { c.ProviderName += "-changed" }, + "ClientID": func(c *DCRCredentials) { c.ClientID += "-changed" }, + "ClientSecret": func(c *DCRCredentials) { c.ClientSecret += "-changed" }, + "TokenEndpointAuthMethod": func(c *DCRCredentials) { c.TokenEndpointAuthMethod += "-changed" }, + "RegistrationAccessToken": func(c *DCRCredentials) { c.RegistrationAccessToken += "-changed" }, + "RegistrationClientURI": func(c *DCRCredentials) { c.RegistrationClientURI += "-changed" }, + "AuthorizationEndpoint": func(c *DCRCredentials) { c.AuthorizationEndpoint += "-changed" }, + "TokenEndpoint": func(c *DCRCredentials) { c.TokenEndpoint += "-changed" }, + "CreatedAt": func(c *DCRCredentials) { c.CreatedAt = c.CreatedAt.Add(time.Second) }, + "ClientSecretExpiresAt": func(c *DCRCredentials) { c.ClientSecretExpiresAt = c.ClientSecretExpiresAt.Add(time.Second) }, +} + +// fullDCRCASFixture returns a DCRCredentials at key with every field +// populated, so each mutation in dcrCASFieldMutations changes a non-zero value. +func fullDCRCASFixture(key DCRKey) *DCRCredentials { + c := dcrCASFixture(key, "original") + c.ProviderName = "provider" + c.TokenEndpointAuthMethod = "client_secret_basic" + c.RegistrationAccessToken = "rat" + c.RegistrationClientURI = "https://idp.example.com/register/client-original" + c.CreatedAt = time.Unix(1_700_000_000, 0) + c.ClientSecretExpiresAt = time.Now().Add(24 * time.Hour).Truncate(time.Second) + return c +} + +// runDCRCASPerFieldTest pins that every persisted field takes part in the +// UpdateDCRCredentialsIfUnchanged comparison: for each field, an expected that +// differs from the stored row in only that field is refused with +// ErrDCRCredentialsChanged and the row is left intact. newStore returns a +// fresh, empty store per subtest. +func runDCRCASPerFieldTest(t *testing.T, newStore func(t *testing.T) DCRCredentialStore) { + t.Helper() + + // Every DCRCredentials field except Key must have a mutation, so a field + // added later cannot silently escape the comparison. + typ := reflect.TypeFor[DCRCredentials]() + for i := range typ.NumField() { + name := typ.Field(i).Name + if name == "Key" { + continue + } + require.Containsf(t, dcrCASFieldMutations, name, + "DCRCredentials.%s has no entry in dcrCASFieldMutations", name) + } + + for name, mutate := range dcrCASFieldMutations { + t.Run(name, func(t *testing.T) { + t.Parallel() + ctx := context.Background() + s := newStore(t) + key := dcrFixtureKey() + stored := fullDCRCASFixture(key) + _, err := s.StoreDCRCredentialsIfAbsent(ctx, stored) + require.NoError(t, err) + + before, err := s.GetDCRCredentials(ctx, key) + require.NoError(t, err) + expected := cloneDCRCredentials(before) + mutate(expected) + + _, err = s.UpdateDCRCredentialsIfUnchanged(ctx, dcrCASFixture(key, "replacement"), expected) + require.ErrorIs(t, err, ErrDCRCredentialsChanged) + + after, err := s.GetDCRCredentials(ctx, key) + require.NoError(t, err) + assert.Equal(t, *before, *after, "a refused compare-and-set must not write") + }) + } +} + // requireSingleDCRCASWinner asserts that of a set of racing // UpdateDCRCredentialsIfUnchanged calls made against the same expected value, // exactly one succeeded and every other one was refused with @@ -4525,6 +4601,165 @@ func TestRedisStorage_DCRCredentials_UpdateIfUnchangedConcurrent(t *testing.T) { }) } +// TestRedisStorage_DCRCredentials_UpdateIfUnchangedPerField runs the shared +// per-field comparison test against the Redis backend. +func TestRedisStorage_DCRCredentials_UpdateIfUnchangedPerField(t *testing.T) { + t.Parallel() + runDCRCASPerFieldTest(t, func(t *testing.T) DCRCredentialStore { + t.Helper() + s, _ := newTestRedisStorage(t) + t.Cleanup(func() { _ = s.Close() }) + return s + }) +} + +// dcrCASInterferenceHook is a go-redis hook that, before each MULTI/EXEC +// pipeline, rewrites key with its current value through a separate client. +// The rewrite leaves the row's contents unchanged but invalidates the WATCH, +// so EXEC aborts with redis.TxFailedErr. It interferes with the first limit +// transactions only (limit < 0 means every one). +type dcrCASInterferenceHook struct { + other *redis.Client + key string + limit int64 + count atomic.Int64 +} + +func (*dcrCASInterferenceHook) DialHook(next redis.DialHook) redis.DialHook { return next } + +func (*dcrCASInterferenceHook) ProcessHook(next redis.ProcessHook) redis.ProcessHook { return next } + +func (h *dcrCASInterferenceHook) ProcessPipelineHook(next redis.ProcessPipelineHook) redis.ProcessPipelineHook { + return func(ctx context.Context, cmds []redis.Cmder) error { + if h.limit < 0 || h.count.Load() < h.limit { + h.count.Add(1) + val, err := h.other.Get(ctx, h.key).Result() + if err != nil { + return err + } + if err := h.other.SetArgs(ctx, h.key, val, redis.SetArgs{KeepTTL: true}).Err(); err != nil { + return err + } + } + return next(ctx, cmds) + } +} + +// newInterferedDCRCASStorage returns a Redis-backed store seeded with an +// "original" row at dcrFixtureKey, the value read back from it, and the +// interference hook installed on the store's client. +func newInterferedDCRCASStorage(t *testing.T, limit int64) (*RedisStorage, *DCRCredentials, *dcrCASInterferenceHook) { + t.Helper() + mr := miniredis.RunT(t) + client := redis.NewClient(&redis.Options{Addr: mr.Addr()}) + other := redis.NewClient(&redis.Options{Addr: mr.Addr()}) + s := NewRedisStorageWithClient(client, "test:auth:") + t.Cleanup(func() { + _ = s.Close() + _ = other.Close() + }) + + ctx := context.Background() + key := dcrFixtureKey() + _, err := s.StoreDCRCredentialsIfAbsent(ctx, dcrCASFixture(key, "original")) + require.NoError(t, err) + expected, err := s.GetDCRCredentials(ctx, key) + require.NoError(t, err) + + hook := &dcrCASInterferenceHook{other: other, key: redisDCRKey(s.keyPrefix, key), limit: limit} + client.AddHook(hook) + return s, expected, hook +} + +// TestRedisStorage_DCRCredentials_UpdateIfUnchangedRetry pins the +// WATCH/MULTI retry loop: an aborted EXEC is retried and the retry succeeds +// when the row still matches, while an EXEC aborted on every attempt returns a +// generic error after maxDCRClaimRetries attempts and writes nothing. +func TestRedisStorage_DCRCredentials_UpdateIfUnchangedRetry(t *testing.T) { + t.Parallel() + + t.Run("aborted EXEC is retried and succeeds", func(t *testing.T) { + t.Parallel() + s, expected, hook := newInterferedDCRCASStorage(t, 1) + ctx := context.Background() + + replacement := dcrCASFixture(expected.Key, "replacement") + _, err := s.UpdateDCRCredentialsIfUnchanged(ctx, replacement, expected) + require.NoError(t, err) + assert.Equal(t, int64(1), hook.count.Load(), "exactly one attempt must have been interfered with") + + got, err := s.GetDCRCredentials(ctx, expected.Key) + require.NoError(t, err) + assert.Equal(t, *replacement, *got) + }) + + t.Run("retry exhaustion returns a generic error and writes nothing", func(t *testing.T) { + t.Parallel() + s, expected, hook := newInterferedDCRCASStorage(t, -1) + ctx := context.Background() + + _, err := s.UpdateDCRCredentialsIfUnchanged(ctx, dcrCASFixture(expected.Key, "replacement"), expected) + require.Error(t, err) + assert.ErrorIs(t, err, redis.TxFailedErr) + assert.NotErrorIs(t, err, ErrDCRCredentialsChanged) + assert.NotErrorIs(t, err, ErrNotFound) + assert.Equal(t, int64(maxDCRClaimRetries), hook.count.Load()) + + got, err := s.GetDCRCredentials(ctx, expected.Key) + require.NoError(t, err) + assert.Equal(t, *expected, *got, "an exhausted compare-and-set must not write") + }) +} + +// TestRedisStorage_DCRCredentials_UpdateIfUnchangedStoredRowStates pins the +// two remaining stored-row branches: a corrupt row surfaces as a generic error +// and is left untouched, and a row whose ClientSecretExpiresAt has passed but +// whose key still exists is updatable (presence is physical, as for +// UpdateDCRCredentialsIfPresent). +func TestRedisStorage_DCRCredentials_UpdateIfUnchangedStoredRowStates(t *testing.T) { + t.Parallel() + + t.Run("corrupt row returns a generic error", func(t *testing.T) { + withRedisStorage(t, func(ctx context.Context, s *RedisStorage, mr *miniredis.Miniredis) { + key := dcrFixtureKey() + redisKey := redisDCRKey(s.keyPrefix, key) + require.NoError(t, mr.Set(redisKey, "not-json")) + + _, err := s.UpdateDCRCredentialsIfUnchanged(ctx, + dcrCASFixture(key, "replacement"), dcrCASFixture(key, "original")) + require.Error(t, err) + assert.NotErrorIs(t, err, ErrDCRCredentialsChanged) + assert.NotErrorIs(t, err, ErrNotFound) + + raw, getErr := mr.Get(redisKey) + require.NoError(t, getErr) + assert.Equal(t, "not-json", raw, "a corrupt row must not be overwritten") + }) + }) + + t.Run("expired but present row is updatable", func(t *testing.T) { + withRedisStorage(t, func(ctx context.Context, s *RedisStorage, mr *miniredis.Miniredis) { + key := dcrFixtureKey() + stored := dcrCASFixture(key, "original") + stored.ClientSecretExpiresAt = time.Now().Add(-time.Hour).Truncate(time.Second) + _, err := s.StoreDCRCredentialsIfAbsent(ctx, stored) + require.NoError(t, err) + // Keep the key present past its bounded pastExpiryDCRTTL. + mr.SetTTL(redisDCRKey(s.keyPrefix, key), time.Hour) + + expected, err := s.GetDCRCredentials(ctx, key) + require.NoError(t, err) + replacement := dcrCASFixture(key, "replacement") + _, err = s.UpdateDCRCredentialsIfUnchanged(ctx, replacement, expected) + require.NoError(t, err, "an expired-but-present row must be updatable") + + got, err := s.GetDCRCredentials(ctx, key) + require.NoError(t, err) + assert.Equal(t, *replacement, *got) + }) + }) +} + // TestRedisStorage_DCRCredentials_FirstClaimWins_Concurrent proves the // create-if-absent race-safety property holds under genuine concurrency, not // just the sequential shape of TestRedisStorage_DCRCredentials_FirstClaimWins: