Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
49 changes: 49 additions & 0 deletions pkg/authserver/storage/memory.go
Original file line number Diff line number Diff line change
Expand Up @@ -2058,6 +2058,55 @@ 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.
//
// 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
Comment thread
tgrunnagle marked this conversation as resolved.
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) {
Expand Down
185 changes: 185 additions & 0 deletions pkg/authserver/storage/memory_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -23,6 +23,7 @@ import (
"errors"
"fmt"
"net/url"
"reflect"
"sync"
"sync/atomic"
"testing"
Expand Down Expand Up @@ -3026,6 +3027,190 @@ 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) {
Comment thread
tgrunnagle marked this conversation as resolved.
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_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.
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)

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

// 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
Expand Down
15 changes: 15 additions & 0 deletions pkg/authserver/storage/mocks/mock_storage.go

Some generated files are not rendered by default. Learn more about how customized files appear on GitHub.

Loading
Loading