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
6 changes: 6 additions & 0 deletions internal/api/handler_auth.go
Original file line number Diff line number Diff line change
Expand Up @@ -430,6 +430,7 @@ func (h *Handler) updateProfile(ctx context.Context, req *events.LambdaFunctionU
// - ErrEmailInUse -> 409 with a neutral message that does NOT confirm whether
// another account holds the address; prevents account enumeration via the
// profile-update path (issue #929).
// - ErrUserChanged -> 409: the account changed after it was read (issue #474).
// - All other errors pass through unchanged for handleRequestError to render
// as 500.
func mapProfileUpdateError(err error) error {
Expand All @@ -438,6 +439,8 @@ func mapProfileUpdateError(err error) error {
return NewClientError(401, err.Error())
case errors.Is(err, auth.ErrEmailInUse):
return NewClientError(409, "Unable to update email")
case errors.Is(err, auth.ErrUserChanged):
return NewClientError(409, auth.ErrUserChanged.Error())
}
return err
}
Expand Down Expand Up @@ -505,6 +508,9 @@ func (h *Handler) changePassword(ctx context.Context, req *events.LambdaFunction
}

err = h.auth.ChangePasswordAPI(ctx, session.UserID, currentPassword, newPassword)
if errors.Is(err, auth.ErrUserChanged) {
return nil, NewClientError(409, err.Error())
}
if err != nil {
return nil, err
}
Expand Down
23 changes: 23 additions & 0 deletions internal/api/handler_auth_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -1084,6 +1084,29 @@ func TestHandler_updateProfile_DuplicateEmail(t *testing.T) {
assert.NotContains(t, ce.message, "already in use", "response must not confirm another account's existence")
}

func TestHandler_credentialWritesMapConcurrentChangeToConflict(t *testing.T) {
ctx := context.Background()
userID := "12345678-1234-1234-1234-123456789abc"
mockAuth := new(MockAuthService)
t.Cleanup(func() { mockAuth.AssertExpectations(t) })
mockAuth.On("ValidateSession", ctx, "test-token").Return(&Session{UserID: userID}, nil)
mockAuth.On("UpdateUserProfile", mock.Anything, userID, "new@example.com", "oldpass", "").Return(fmt.Errorf("failed to update user: %w", auth.ErrUserChanged))
mockAuth.On("ChangePasswordAPI", ctx, userID, "oldpass", "newpass").Return(auth.ErrUserChanged)
handler := &Handler{auth: mockAuth}
old, next := base64.StdEncoding.EncodeToString([]byte("oldpass")), base64.StdEncoding.EncodeToString([]byte("newpass"))
headers := map[string]string{"Authorization": "Bearer test-token"}

_, profileErr := handler.updateProfile(ctx, &events.LambdaFunctionURLRequest{Headers: headers,
Body: `{"email": "new@example.com", "current_password": "` + old + `"}`})
_, passwordErr := handler.changePassword(ctx, &events.LambdaFunctionURLRequest{Headers: headers,
Body: `{"current_password": "` + old + `", "new_password": "` + next + `"}`})
for _, err := range []error{profileErr, passwordErr} {
ce, ok := IsClientError(err)
require.True(t, ok, "a concurrent account change must be a 409, got %v", err)
assert.Equal(t, 409, ce.code)
}
}

// TestHandler_resetPassword_DecodesBase64 verifies issue #356: the
// resetPassword handler must base64-decode new_password before forwarding to
// the service, matching the pattern used by login / change-password /
Expand Down
4 changes: 4 additions & 0 deletions internal/auth/errors.go
Original file line number Diff line number Diff line change
Expand Up @@ -91,6 +91,10 @@ var (
// reset token mailed to the account.
ErrAccountDeactivated = errors.New("account is deactivated")

// ErrUserChanged: a field-scoped user write found the row changed since the
// caller read it (issue #474). Mapped to 409.
ErrUserChanged = errors.New("account changed concurrently; reload and retry")

// MFA login-gate sentinels — used by the login API handler to map
// to machine-readable response codes (mfa_required /
// invalid_mfa_code) so the frontend can branch on the error class
Expand Down
3 changes: 3 additions & 0 deletions internal/auth/interfaces.go
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@ package auth

import (
"context"
"time"
)

// StoreInterface defines the methods required for auth storage.
Expand All @@ -11,6 +12,8 @@ type StoreInterface interface {
GetUserByEmail(ctx context.Context, email string) (*User, error)
CreateUser(ctx context.Context, user *User) error
UpdateUser(ctx context.Context, user *User) error
UpdateUserCredentials(ctx context.Context, user *User, readEmail, readPasswordHash string) error
SetPasswordResetToken(ctx context.Context, user *User, readExpiry *time.Time) error
RecordFailedLogin(ctx context.Context, userID string) error
RecordSuccessfulLogin(ctx context.Context, userID string) error
DeleteUser(ctx context.Context, userID string) error
Expand Down
2 changes: 1 addition & 1 deletion internal/auth/service_api_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -419,7 +419,7 @@ func TestService_ChangePasswordAPI(t *testing.T) {
mockStore.On("GetUserByID", ctx, "user-123").Return(testUser, nil).Once()
mockStore.On("DeleteUserSessions", ctx, "user-123").Return(nil).Once()
mockStore.On("ListAPIKeysByUser", ctx, "user-123").Return([]*UserAPIKey{}, nil).Once()
mockStore.On("UpdateUser", ctx, mock.AnythingOfType("*auth.User")).Return(nil).Once()
mockStore.On("UpdateUserCredentials", ctx, mock.AnythingOfType("*auth.User"), mock.Anything, mock.Anything).Return(nil).Once()

err := service.ChangePasswordAPI(ctx, "user-123", "OldPassword123", "SecureTest@456")
require.NoError(t, err)
Expand Down
293 changes: 293 additions & 0 deletions internal/auth/service_credentials_db_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,293 @@
//go:build integration

package auth

import (
"context"
"errors"
"testing"
"time"

"github.com/stretchr/testify/assert"
"github.com/stretchr/testify/require"
)

const (
credentialRacePassword = "OriginalPassword123!"
credentialRaceNew = "ReplacedPassword456!"
)

// credentialReadBarrier runs afterRead once, between the service's user read
// and its write, to replay the stale-snapshot schedules from issue #474.
type credentialReadBarrier struct {
StoreInterface
afterRead func(context.Context, *User)
}

func (s *credentialReadBarrier) pause(ctx context.Context, u *User) {
if s.afterRead != nil {
f := s.afterRead
s.afterRead = nil
f(ctx, u)
}
}

func (s *credentialReadBarrier) GetUserByID(ctx context.Context, id string) (*User, error) {
u, err := s.StoreInterface.GetUserByID(ctx, id)
if err == nil {
s.pause(ctx, u)
}
return u, err
}

func (s *credentialReadBarrier) GetUserByEmail(ctx context.Context, email string) (*User, error) {
u, err := s.StoreInterface.GetUserByEmail(ctx, email)
if err == nil {
s.pause(ctx, u)
}
return u, err
}

type recordingMailSink struct {
resets []string
err error
}

func (m *recordingMailSink) SendPasswordResetEmail(_ context.Context, email, _ string) error {
m.resets = append(m.resets, email)
return m.err
}
func (m *recordingMailSink) SendWelcomeEmail(context.Context, string, string, string) error {
return nil
}
func (m *recordingMailSink) SendUserInviteEmail(context.Context, string, string) error { return nil }

type credentialRaceFixture struct {
t *testing.T
store *PostgresStore
barrier *credentialReadBarrier
mail *recordingMailSink
svc *Service
user *User
session string
}

func newCredentialRaceFixture(t *testing.T, store *PostgresStore, email string) *credentialRaceFixture {
f := &credentialRaceFixture{t: t, store: store, barrier: &credentialReadBarrier{StoreInterface: store}, mail: &recordingMailSink{}}
f.svc = f.newService(f.barrier, f.mail)
hash, err := f.svc.hashPassword(credentialRacePassword)
require.NoError(t, err)
f.user = &User{Email: email, PasswordHash: hash, Active: true, GroupIDs: []string{DefaultPurchaserGroupID}}
require.NoError(t, store.CreateUser(t.Context(), f.user))
login, err := f.svc.Login(t.Context(), LoginRequest{Email: email, Password: credentialRacePassword})
require.NoError(t, err)
f.session = login.Token
return f
}

func (f *credentialRaceFixture) newService(store StoreInterface, mail EmailSenderInterface) *Service {
svc := NewService(ServiceConfig{Store: store, EmailSender: mail, DashboardURL: "https://dashboard.example.com"})
svc.bcryptCostOverride = 4
return svc
}

// onRead schedules a concurrent action by a second, unpaused service and
// returns the row as it stood once that action committed.
func (f *credentialRaceFixture) onRead(action func(ctx context.Context, other *Service)) func() *User {
var after *User
f.barrier.afterRead = func(ctx context.Context, _ *User) {
action(ctx, f.newService(f.store, &recordingMailSink{}))
var err error
after, err = f.store.GetUserByID(ctx, f.user.ID)
require.NoError(f.t, err)
}
return func() *User {
require.NotNil(f.t, after, "the concurrent action never ran")
return after
}
}

// enrollOnRead lets the victim complete MFA enrollment after the stale read.
func (f *credentialRaceFixture) enrollOnRead() func() *User {
return f.onRead(func(ctx context.Context, victim *Service) {
setup, err := victim.MFASetup(ctx, f.user.ID, credentialRacePassword)
require.NoError(f.t, err)
codes, err := victim.MFAEnable(ctx, f.user.ID, generateTOTP(setup.Secret, time.Now().Unix()/30))
require.NoError(f.t, err)
require.NotEmpty(f.t, codes)
})
}

func (f *credentialRaceFixture) stored() *User {
u, err := f.store.GetUserByID(f.t.Context(), f.user.ID)
require.NoError(f.t, err)
return u
}

// requireFactorEnforced proves the victim's factor survived: login without it
// is refused and login with it succeeds.
func (f *credentialRaceFixture) requireFactorEnforced(email, password string, enrolled *User) {
stored := f.stored()
require.True(f.t, stored.MFAEnabled, "stale write erased the concurrent MFA enrollment")
assert.Equal(f.t, enrolled.MFASecret, stored.MFASecret)
assert.Equal(f.t, enrolled.MFARecoveryCodes, stored.MFARecoveryCodes)
_, err := f.svc.Login(f.t.Context(), LoginRequest{Email: email, Password: password})
require.ErrorIs(f.t, err, ErrMFARequired)
_, err = f.svc.Login(f.t.Context(), LoginRequest{Email: email, Password: password,
MFACode: generateTOTP(enrolled.MFASecret, time.Now().Unix()/30)})
require.NoError(f.t, err)
}

// assertOnlyChanged asserts every column other than the named ones still
// matches the post-enrollment row.
func assertOnlyChanged(t *testing.T, enrolled, stored *User, mutate func(want *User)) {
want := *enrolled
mutate(&want)
want.UpdatedAt, want.LastLoginAt = stored.UpdatedAt, stored.LastLoginAt
assert.Equal(t, &want, stored)
}

func TestIntegration_CredentialWritesPreserveConcurrentMFA(t *testing.T) {
store := NewPostgresStore(setupAuthTestDB(t))
ctx := t.Context()

t.Run("profile-email-only", func(t *testing.T) {
f := newCredentialRaceFixture(t, store, "profile-email@example.com")
enrolled := f.enrollOnRead()
require.NoError(t, f.svc.UpdateUserProfile(ctx, f.user.ID, "profile-email-new@example.com", credentialRacePassword, ""))
assertOnlyChanged(t, enrolled(), f.stored(), func(w *User) { w.Email = "profile-email-new@example.com" })
f.requireFactorEnforced("profile-email-new@example.com", credentialRacePassword, enrolled())
})

t.Run("profile-email-and-password", func(t *testing.T) {
f := newCredentialRaceFixture(t, store, "profile-both@example.com")
enrolled := f.enrollOnRead()
require.NoError(t, f.svc.UpdateUserProfile(ctx, f.user.ID, "profile-both-new@example.com", credentialRacePassword, credentialRaceNew))
stored := f.stored()
assert.True(t, f.svc.verifyPassword(credentialRaceNew, stored.PasswordHash))
assert.Equal(t, []string{f.user.PasswordHash}, stored.PasswordHistory)
assertOnlyChanged(t, enrolled(), stored, func(w *User) {
w.Email, w.PasswordHash, w.Salt, w.PasswordHistory = "profile-both-new@example.com", stored.PasswordHash, "", stored.PasswordHistory
})
_, err := f.svc.ValidateSession(ctx, f.session)
require.Error(t, err, "a password change must revoke existing sessions")
f.requireFactorEnforced("profile-both-new@example.com", credentialRaceNew, enrolled())
_, err = f.svc.Login(ctx, LoginRequest{Email: "profile-both-new@example.com", Password: credentialRacePassword})
require.EqualError(t, err, genericLoginError)
})

t.Run("change-password", func(t *testing.T) {
f := newCredentialRaceFixture(t, store, "change-password@example.com")
enrolled := f.enrollOnRead()
require.NoError(t, f.svc.ChangePassword(ctx, f.user.ID, ChangePasswordRequest{CurrentPassword: credentialRacePassword, NewPassword: credentialRaceNew}))
stored := f.stored()
assert.True(t, f.svc.verifyPassword(credentialRaceNew, stored.PasswordHash))
assert.Equal(t, []string{f.user.PasswordHash}, stored.PasswordHistory)
assertOnlyChanged(t, enrolled(), stored, func(w *User) {
w.PasswordHash, w.Salt, w.PasswordHistory = stored.PasswordHash, "", stored.PasswordHistory
})
_, err := f.svc.ValidateSession(ctx, f.session)
require.Error(t, err, "a password change must revoke existing sessions")
f.requireFactorEnforced(f.user.Email, credentialRaceNew, enrolled())
_, err = f.svc.Login(ctx, LoginRequest{Email: f.user.Email, Password: credentialRacePassword})
require.EqualError(t, err, genericLoginError)
})

t.Run("reset-request-with-failing-delivery", func(t *testing.T) {
f := newCredentialRaceFixture(t, store, "reset-request@example.com")
f.mail.err = errors.New("synthetic delivery failure")
enrolled := f.enrollOnRead()
require.NoError(t, f.svc.RequestPasswordReset(ctx, f.user.Email))
assert.Equal(t, []string{f.user.Email}, f.mail.resets)
stored := f.stored()
require.NotEmpty(t, stored.PasswordResetToken)
require.NotNil(t, stored.PasswordResetExpiry)
assertOnlyChanged(t, enrolled(), stored, func(w *User) {
w.PasswordResetToken, w.PasswordResetExpiry = stored.PasswordResetToken, stored.PasswordResetExpiry
})
f.requireFactorEnforced(f.user.Email, credentialRacePassword, enrolled())

_, err := store.db.Exec(ctx, "UPDATE users SET password_reset_expiry = NOW() - interval '1 hour' WHERE id = $1", f.user.ID)
require.NoError(t, err)
require.NoError(t, f.svc.RequestPasswordReset(ctx, f.user.Email))
assert.Len(t, f.mail.resets, 2, "a reissue over an expired token must match the stored expiry")
assert.NotEqual(t, stored.PasswordResetToken, f.stored().PasswordResetToken)
})
}

func TestIntegration_CredentialWritesRejectStaleCredentials(t *testing.T) {
store := NewPostgresStore(setupAuthTestDB(t))
ctx := t.Context()

t.Run("change-password-after-concurrent-change", func(t *testing.T) {
f := newCredentialRaceFixture(t, store, "stale-change@example.com")
winner := f.onRead(func(ctx context.Context, other *Service) {
require.NoError(t, other.ChangePassword(ctx, f.user.ID, ChangePasswordRequest{CurrentPassword: credentialRacePassword, NewPassword: "WinnerPassword789!"}))
login, err := other.Login(ctx, LoginRequest{Email: f.user.Email, Password: "WinnerPassword789!"})
require.NoError(t, err)
f.session = login.Token
})
err := f.svc.ChangePassword(ctx, f.user.ID, ChangePasswordRequest{CurrentPassword: credentialRacePassword, NewPassword: credentialRaceNew})
require.ErrorIs(t, err, ErrUserChanged)
assert.Equal(t, winner(), f.stored())
_, err = f.svc.ValidateSession(ctx, f.session)
require.NoError(t, err, "a rejected write must not revoke sessions")
})

t.Run("profile-after-concurrent-email-change", func(t *testing.T) {
f := newCredentialRaceFixture(t, store, "stale-profile@example.com")
winner := f.onRead(func(ctx context.Context, _ *Service) {
u, err := store.GetUserByID(ctx, f.user.ID)
require.NoError(t, err)
u.Email = "stale-profile-admin@example.com"
require.NoError(t, store.UpdateUser(ctx, u))
})
err := f.svc.UpdateUserProfile(ctx, f.user.ID, "stale-profile-self@example.com", credentialRacePassword, credentialRaceNew)
require.ErrorIs(t, err, ErrUserChanged)
assert.Equal(t, winner(), f.stored())
_, err = f.svc.ValidateSession(ctx, f.session)
require.NoError(t, err, "a rejected write must not revoke sessions")
})

t.Run("reset-after-concurrent-deactivation", func(t *testing.T) {
f := newCredentialRaceFixture(t, store, "stale-reset-deactivated@example.com")
winner := f.onRead(func(ctx context.Context, _ *Service) {
u, err := store.GetUserByID(ctx, f.user.ID)
require.NoError(t, err)
now := time.Now()
u.Active, u.DeactivatedAt = false, &now
require.NoError(t, store.UpdateUser(ctx, u))
})
require.NoError(t, f.svc.RequestPasswordReset(ctx, f.user.Email))
assert.Empty(t, f.mail.resets)
assert.Equal(t, winner(), f.stored())
})

t.Run("reset-after-concurrent-email-change", func(t *testing.T) {
f := newCredentialRaceFixture(t, store, "stale-reset-email@example.com")
winner := f.onRead(func(ctx context.Context, _ *Service) {
u, err := store.GetUserByID(ctx, f.user.ID)
require.NoError(t, err)
u.Email = "stale-reset-email-new@example.com"
require.NoError(t, store.UpdateUser(ctx, u))
})
require.NoError(t, f.svc.RequestPasswordReset(ctx, f.user.Email))
assert.Empty(t, f.mail.resets, "a token must not be mailed to the address the user just left")
stored := f.stored()
assert.Equal(t, winner(), stored)
assert.Equal(t, "stale-reset-email-new@example.com", stored.Email)
assert.Empty(t, stored.PasswordResetToken)
assert.Nil(t, stored.PasswordResetExpiry)
})

t.Run("reset-after-concurrent-reset", func(t *testing.T) {
f := newCredentialRaceFixture(t, store, "stale-reset-twice@example.com")
winner := f.onRead(func(ctx context.Context, other *Service) {
require.NoError(t, other.RequestPasswordReset(ctx, f.user.Email))
})
require.NoError(t, f.svc.RequestPasswordReset(ctx, f.user.Email))
assert.Empty(t, f.mail.resets, "the losing request must not mail a token it did not store")
assert.Equal(t, winner(), f.stored())
assert.NotEmpty(t, f.stored().PasswordResetToken)
})
}
Loading
Loading