diff --git a/internal/api/handler_auth.go b/internal/api/handler_auth.go index a244ee3d..da0b91ab 100644 --- a/internal/api/handler_auth.go +++ b/internal/api/handler_auth.go @@ -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 { @@ -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 } @@ -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 } diff --git a/internal/api/handler_auth_test.go b/internal/api/handler_auth_test.go index c8590471..d742f147 100644 --- a/internal/api/handler_auth_test.go +++ b/internal/api/handler_auth_test.go @@ -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 / diff --git a/internal/auth/errors.go b/internal/auth/errors.go index f9af0e63..4f7c5a62 100644 --- a/internal/auth/errors.go +++ b/internal/auth/errors.go @@ -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 diff --git a/internal/auth/interfaces.go b/internal/auth/interfaces.go index bf526005..b83ccdc8 100644 --- a/internal/auth/interfaces.go +++ b/internal/auth/interfaces.go @@ -2,6 +2,7 @@ package auth import ( "context" + "time" ) // StoreInterface defines the methods required for auth storage. @@ -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 diff --git a/internal/auth/service_api_test.go b/internal/auth/service_api_test.go index 04e47c03..89eee24f 100644 --- a/internal/auth/service_api_test.go +++ b/internal/auth/service_api_test.go @@ -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) diff --git a/internal/auth/service_credentials_db_test.go b/internal/auth/service_credentials_db_test.go new file mode 100644 index 00000000..ae5d495a --- /dev/null +++ b/internal/auth/service_credentials_db_test.go @@ -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) + }) +} diff --git a/internal/auth/service_password.go b/internal/auth/service_password.go index 1fb74862..c49dddb3 100644 --- a/internal/auth/service_password.go +++ b/internal/auth/service_password.go @@ -263,6 +263,7 @@ func (s *Service) ChangePassword(ctx context.Context, userID string, req ChangeP return fmt.Errorf("failed to hash password: %w", err) } + readPasswordHash := user.PasswordHash // Update password history (add current password to history) user.PasswordHistory = addToPasswordHistory(user.PasswordHash, user.PasswordHistory) @@ -270,7 +271,7 @@ func (s *Service) ChangePassword(ctx context.Context, userID string, req ChangeP user.Salt = "" // Not used anymore user.PasswordHash = passwordHash - if err := s.store.UpdateUser(ctx, user); err != nil { + if err := s.store.UpdateUserCredentials(ctx, user, user.Email, readPasswordHash); err != nil { return err } @@ -342,13 +343,26 @@ func (s *Service) RequestPasswordReset(ctx context.Context, email string) error // Set expiry based on configured duration expiry := time.Now().Add(PasswordResetExpiry) + readExpiry := user.PasswordResetExpiry user.PasswordResetToken = tokenHash user.PasswordResetExpiry = &expiry - if err := s.store.UpdateUser(ctx, user); err != nil { - return fmt.Errorf("failed to save reset token: %w", err) + if err := s.store.SetPasswordResetToken(ctx, user, readExpiry); err != nil { + if errors.Is(err, ErrUserChanged) { + // Deactivated, re-addressed or reset concurrently: answer as for an ineligible account. + logging.Debugf("Password reset skipped for concurrently changed account: %s", redactEmail(email)) + return nil + } + return err } + s.sendPasswordResetEmail(ctx, user.Email, token) + return nil +} + +// sendPasswordResetEmail is best-effort so RequestPasswordReset's response +// never reveals whether an email was sent. +func (s *Service) sendPasswordResetEmail(ctx context.Context, email, token string) { // Skip the email entirely if dashboardURL is unconfigured — a broken // relative link in an inbox is worse than no email; the operator's // startup-time WARN already names the missing env var. Don't return an @@ -357,17 +371,14 @@ func (s *Service) RequestPasswordReset(ctx context.Context, email string) error // exists, and that includes whether or not a send happened). Issue #355. if s.dashboardURL == "" { logging.Errorf("RequestPasswordReset: skipping send — DashboardURL empty would produce a broken relative link (set DASHBOARD_URL).") - return nil + return } // Send reset email (use unhashed token in URL) resetURL := fmt.Sprintf("%s/reset-password?token=%s", s.dashboardURL, token) - if err := s.emailSender.SendPasswordResetEmail(ctx, user.Email, resetURL); err != nil { + if err := s.emailSender.SendPasswordResetEmail(ctx, email, resetURL); err != nil { logging.Errorf("Failed to send password reset email: %v", err) - // Don't return error to prevent email enumeration } - - return nil } // ConfirmPasswordReset completes a password reset. diff --git a/internal/auth/service_password_callback_test.go b/internal/auth/service_password_callback_test.go index 7e76c10b..8aa8d10c 100644 --- a/internal/auth/service_password_callback_test.go +++ b/internal/auth/service_password_callback_test.go @@ -35,7 +35,7 @@ func TestService_OnPasswordChange_ChangePassword(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.ChangePassword(ctx, "user-123", ChangePasswordRequest{ CurrentPassword: "OldSecure123!", @@ -57,7 +57,7 @@ func TestService_OnPasswordChange_ChangePassword(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.ChangePassword(ctx, "user-123", ChangePasswordRequest{ CurrentPassword: "OldSecure123!", @@ -89,14 +89,14 @@ func TestService_OnPasswordChange_ChangePassword(t *testing.T) { // credentials for a password that never actually changed), so a // failed UpdateUser must never reach DeleteUserSessions/ListAPIKeysByUser. mockStore.On("GetUserByID", ctx, "user-123").Return(testUser, nil).Once() - mockStore.On("UpdateUser", ctx, mock.AnythingOfType("*auth.User")).Return(assert.AnError).Once() + mockStore.On("UpdateUserCredentials", ctx, mock.AnythingOfType("*auth.User"), mock.Anything, mock.Anything).Return(assert.AnError).Once() err := service.ChangePassword(ctx, "user-123", ChangePasswordRequest{ CurrentPassword: "OldSecure123!", NewPassword: "NewSecure@456", }) require.Error(t, err) - assert.False(t, callbackCalled, "callback should not be called when UpdateUser fails") + assert.False(t, callbackCalled, "callback should not be called when the credential write fails") mockStore.AssertExpectations(t) }) } @@ -168,7 +168,7 @@ func TestService_OnPasswordChange_UpdateUserProfile(t *testing.T) { testUser := createTestUser(t, "OldSecure123!") mockStore.On("GetUserByID", ctx, "user-123").Return(testUser, 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() mockStore.On("DeleteUserSessions", ctx, "user-123").Return(nil).Once() mockStore.On("ListAPIKeysByUser", ctx, "user-123").Return([]*UserAPIKey{}, nil).Once() @@ -198,7 +198,7 @@ func TestService_OnPasswordChange_UpdateUserProfile(t *testing.T) { mockStore.On("GetUserByID", ctx, "user-123").Return(testUser, nil).Once() mockStore.On("GetUserByEmail", ctx, "new@example.com").Return(nil, 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.UpdateUserProfile(ctx, "user-123", "new@example.com", "OldSecure123!", "") require.NoError(t, err) diff --git a/internal/auth/service_password_test.go b/internal/auth/service_password_test.go index 426e0196..60fd9d02 100644 --- a/internal/auth/service_password_test.go +++ b/internal/auth/service_password_test.go @@ -26,7 +26,7 @@ func TestService_ChangePassword(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() req := ChangePasswordRequest{ CurrentPassword: "OldSecure123!", @@ -140,7 +140,7 @@ func TestService_ChangePassword(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.MatchedBy(func(u *User) bool { + mockStore.On("UpdateUserCredentials", ctx, mock.MatchedBy(func(u *User) bool { // Verify password history includes old password and maintains limit // Should have: original current password (newly added to history) + 2 existing = 3 total return len(u.PasswordHistory) == 3 && @@ -148,7 +148,7 @@ func TestService_ChangePassword(t *testing.T) { u.PasswordHistory[0] == originalHash && // Original current password should be first in history u.PasswordHistory[1] == hash1 && // Previous history items should follow u.PasswordHistory[2] == hash2 - })).Return(nil).Once() + }), mock.Anything, mock.Anything).Return(nil).Once() req := ChangePasswordRequest{ CurrentPassword: "CurrentS3cur3!", @@ -228,7 +228,7 @@ func TestService_ChangePassword(t *testing.T) { mockStore.On("DeleteUserSessions", ctx, "user-123").Return(nil).Once() mockStore.On("ListAPIKeysByUser", ctx, "user-123").Return([]*UserAPIKey{apiKeyRecord}, nil).Once() mockStore.On("UpdateAPIKey", ctx, apiKeyRecord).Return(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.ChangePassword(ctx, "user-123", ChangePasswordRequest{ CurrentPassword: "OldSecure123!", @@ -259,7 +259,7 @@ func TestService_RequestPasswordReset(t *testing.T) { testUser := createTestUser(t, "SecureS3cur3@123") mockStore.On("GetUserByEmail", ctx, "test@example.com").Return(testUser, nil).Once() - mockStore.On("UpdateUser", ctx, mock.AnythingOfType("*auth.User")).Return(nil).Once() + mockStore.On("SetPasswordResetToken", ctx, mock.AnythingOfType("*auth.User"), mock.Anything).Return(nil).Once() mockEmail.On("SendPasswordResetEmail", ctx, "test@example.com", mock.AnythingOfType("string")).Return(nil).Once() err := service.RequestPasswordReset(ctx, "test@example.com") @@ -296,7 +296,7 @@ func TestService_RequestPasswordReset(t *testing.T) { mockStore.AssertExpectations(t) }) - t.Run("return error when UpdateUser fails", func(t *testing.T) { + t.Run("return error when SetPasswordResetToken fails", func(t *testing.T) { mockStore := new(MockStore) mockEmail := new(MockEmailSender) service := createTestService(mockStore, mockEmail) @@ -304,7 +304,7 @@ func TestService_RequestPasswordReset(t *testing.T) { testUser := createTestUser(t, "SecureS3cur3@123") mockStore.On("GetUserByEmail", ctx, "test@example.com").Return(testUser, nil).Once() - mockStore.On("UpdateUser", ctx, mock.AnythingOfType("*auth.User")).Return(assert.AnError).Once() + mockStore.On("SetPasswordResetToken", ctx, mock.AnythingOfType("*auth.User"), mock.Anything).Return(assert.AnError).Once() err := service.RequestPasswordReset(ctx, "test@example.com") assert.Error(t, err) @@ -320,7 +320,7 @@ func TestService_RequestPasswordReset(t *testing.T) { testUser := createTestUser(t, "SecureS3cur3@123") mockStore.On("GetUserByEmail", ctx, "test@example.com").Return(testUser, nil).Once() - mockStore.On("UpdateUser", ctx, mock.AnythingOfType("*auth.User")).Return(nil).Once() + mockStore.On("SetPasswordResetToken", ctx, mock.AnythingOfType("*auth.User"), mock.Anything).Return(nil).Once() mockEmail.On("SendPasswordResetEmail", ctx, "test@example.com", mock.AnythingOfType("string")).Return(assert.AnError).Once() // Should not return error to prevent email enumeration diff --git a/internal/auth/service_test.go b/internal/auth/service_test.go index bf7a5cd8..785cb274 100644 --- a/internal/auth/service_test.go +++ b/internal/auth/service_test.go @@ -625,7 +625,7 @@ func TestService_ErrorPaths(t *testing.T) { testUser := createTestUser(t, "SecurePass@123") mockStore.On("GetUserByEmail", ctx, "test@example.com").Return(testUser, nil).Once() - mockStore.On("UpdateUser", ctx, mock.AnythingOfType("*auth.User")).Return(nil).Once() + mockStore.On("SetPasswordResetToken", ctx, mock.AnythingOfType("*auth.User"), mock.Anything).Return(nil).Once() mockEmail.On("SendPasswordResetEmail", ctx, "test@example.com", mock.AnythingOfType("string")).Return(fmt.Errorf("email error")).Once() // Should not return error to prevent email enumeration @@ -671,7 +671,7 @@ func TestService_ErrorPaths(t *testing.T) { mockStore.On("GetUserByID", ctx, "user-123").Return(testUser, nil).Once() mockStore.On("DeleteUserSessions", ctx, "user-123").Return(fmt.Errorf("session error")).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() req := ChangePasswordRequest{ CurrentPassword: "OldPassword123", diff --git a/internal/auth/service_user.go b/internal/auth/service_user.go index 146ddcaa..459e37da 100644 --- a/internal/auth/service_user.go +++ b/internal/auth/service_user.go @@ -765,6 +765,7 @@ func (s *Service) UpdateUserProfile(ctx context.Context, userID, email, currentP if !s.verifyPassword(currentPassword, user.PasswordHash) { return ErrCurrentPasswordIncorrect } + readEmail, readPasswordHash := user.Email, user.PasswordHash err = s.updateUserEmail(ctx, user, email) if err != nil { @@ -776,8 +777,7 @@ func (s *Service) UpdateUserProfile(ctx context.Context, userID, email, currentP return err } - user.UpdatedAt = time.Now() - if err := s.store.UpdateUser(ctx, user); err != nil { + if err := s.store.UpdateUserCredentials(ctx, user, readEmail, readPasswordHash); err != nil { return fmt.Errorf("failed to update user: %w", err) } diff --git a/internal/auth/service_user_test.go b/internal/auth/service_user_test.go index fc30e661..df42df4b 100644 --- a/internal/auth/service_user_test.go +++ b/internal/auth/service_user_test.go @@ -1209,7 +1209,7 @@ func TestService_UpdateUserProfile(t *testing.T) { mockStore.On("GetUserByID", ctx, "user-123").Return(testUser, nil).Once() mockStore.On("GetUserByEmail", ctx, "new@example.com").Return(nil, 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() mockStore.On("DeleteUserSessions", ctx, "user-123").Return(nil).Once() mockStore.On("ListAPIKeysByUser", ctx, "user-123").Return([]*UserAPIKey{}, nil).Once() @@ -1347,7 +1347,7 @@ func TestService_UpdateUserProfile(t *testing.T) { mockStore.On("GetUserByID", ctx, "user-123").Return(testUser, nil).Once() mockStore.On("GetUserByEmail", ctx, "new@example.com").Return(nil, 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.UpdateUserProfile(ctx, "user-123", "new@example.com", "OldPassword123", "") require.NoError(t, err) diff --git a/internal/auth/store_postgres_credentials.go b/internal/auth/store_postgres_credentials.go new file mode 100644 index 00000000..7d70db38 --- /dev/null +++ b/internal/auth/store_postgres_credentials.go @@ -0,0 +1,42 @@ +package auth + +import ( + "context" + "fmt" + "time" +) + +// UpdateUserCredentials writes only the email and password columns, and only +// while the row still holds the email and hash the caller read (issue #474). +func (s *PostgresStore) UpdateUserCredentials(ctx context.Context, user *User, readEmail, readPasswordHash string) error { + result, err := s.db.Exec(ctx, ` + UPDATE users SET email = $4, password_hash = $5, salt = $6, + password_history = $7, updated_at = NOW() + WHERE id = $1 AND email = $2 AND password_hash = $3 + `, user.ID, readEmail, readPasswordHash, user.Email, user.PasswordHash, user.Salt, user.PasswordHistory) + if err != nil { + return fmt.Errorf("failed to update user credentials: %w", err) + } + if result.RowsAffected() == 0 { + return ErrUserChanged + } + return nil +} + +// SetPasswordResetToken writes only the reset token columns, and only while the +// account is still active under the same email with the reset expiry the caller +// read, so a concurrent deactivation or reset issuance wins. +func (s *PostgresStore) SetPasswordResetToken(ctx context.Context, user *User, readExpiry *time.Time) error { + result, err := s.db.Exec(ctx, ` + UPDATE users SET password_reset_token = $3, password_reset_expiry = $4, updated_at = NOW() + WHERE id = $1 AND email = $2 AND deactivated_at IS NULL + AND password_reset_expiry IS NOT DISTINCT FROM $5 + `, user.ID, user.Email, user.PasswordResetToken, user.PasswordResetExpiry, readExpiry) + if err != nil { + return fmt.Errorf("failed to save reset token: %w", err) + } + if result.RowsAffected() == 0 { + return ErrUserChanged + } + return nil +} diff --git a/internal/auth/test_helpers.go b/internal/auth/test_helpers.go index edaa9e58..6819761a 100644 --- a/internal/auth/test_helpers.go +++ b/internal/auth/test_helpers.go @@ -50,6 +50,14 @@ func (m *MockStore) UpdateUser(ctx context.Context, user *User) error { return args.Error(0) } +func (m *MockStore) UpdateUserCredentials(ctx context.Context, user *User, readEmail, readPasswordHash string) error { + return m.Called(ctx, user, readEmail, readPasswordHash).Error(0) +} + +func (m *MockStore) SetPasswordResetToken(ctx context.Context, user *User, readExpiry *time.Time) error { + return m.Called(ctx, user, readExpiry).Error(0) +} + func (m *MockStore) RecordFailedLogin(ctx context.Context, userID string) error { return m.Called(ctx, userID).Error(0) } diff --git a/internal/mocks/stores.go b/internal/mocks/stores.go index 6c9e6d35..c3ad9ea7 100644 --- a/internal/mocks/stores.go +++ b/internal/mocks/stores.go @@ -797,6 +797,14 @@ func (m *MockAuthStore) UpdateUser(ctx context.Context, user *auth.User) error { return args.Error(0) } +func (m *MockAuthStore) UpdateUserCredentials(ctx context.Context, user *auth.User, readEmail, readPasswordHash string) error { + return m.Called(ctx, user, readEmail, readPasswordHash).Error(0) +} + +func (m *MockAuthStore) SetPasswordResetToken(ctx context.Context, user *auth.User, readExpiry *time.Time) error { + return m.Called(ctx, user, readExpiry).Error(0) +} + func (m *MockAuthStore) RecordFailedLogin(ctx context.Context, userID string) error { return m.Called(ctx, userID).Error(0) } diff --git a/internal/server/health_test.go b/internal/server/health_test.go index 17c52378..95b3f213 100644 --- a/internal/server/health_test.go +++ b/internal/server/health_test.go @@ -32,6 +32,14 @@ func (m *mockAuthStoreForHealth) UpdateUser(ctx context.Context, user *auth.User return nil } +func (m *mockAuthStoreForHealth) UpdateUserCredentials(context.Context, *auth.User, string, string) error { + return nil +} + +func (m *mockAuthStoreForHealth) SetPasswordResetToken(context.Context, *auth.User, *time.Time) error { + return nil +} + func (m *mockAuthStoreForHealth) RecordFailedLogin(context.Context, string) error { return nil }