From 69baac81e43c6ecf3275b59e93985f7a7ed936a1 Mon Sep 17 00:00:00 2001 From: Cristian Magherusan-Stanciu Date: Mon, 5 Oct 2026 02:37:19 +0200 Subject: [PATCH 1/2] fix(auth): scope profile, password and reset writes to their own columns UpdateUserProfile, ChangePassword and RequestPasswordReset read a whole user row, changed a few fields and wrote the whole row back with UpdateUser. An MFA enrollment that committed between the read and the write was erased, so an attacker holding a session and the password (or, for the reset request, only the email) could put a freshly enrolled account back to MFA-disabled. Profile and password changes now persist through UpdateUserCredentials, which writes only email, password hash, salt and history, and only while the stored email and hash still match what the service read. Reset issuance goes through SetPasswordResetToken, which writes only the token columns and requires the account to still be active under the same email with an unchanged reset expiry. A lost race returns ErrUserChanged (409 at the API) and skips session revocation and the password-change hook; a lost reset race answers like any ineligible account and sends no mail. Closes #474 Refs #227 --- internal/api/handler_auth.go | 6 + internal/api/handler_auth_test.go | 23 ++ internal/auth/errors.go | 4 + internal/auth/interfaces.go | 3 + internal/auth/service_api_test.go | 2 +- internal/auth/service_credentials_db_test.go | 276 ++++++++++++++++++ internal/auth/service_password.go | 27 +- .../auth/service_password_callback_test.go | 12 +- internal/auth/service_password_test.go | 16 +- internal/auth/service_test.go | 4 +- internal/auth/service_user.go | 4 +- internal/auth/service_user_test.go | 4 +- internal/auth/store_postgres_credentials.go | 42 +++ internal/auth/test_helpers.go | 8 + internal/mocks/stores.go | 8 + internal/server/health_test.go | 8 + 16 files changed, 418 insertions(+), 29 deletions(-) create mode 100644 internal/auth/service_credentials_db_test.go create mode 100644 internal/auth/store_postgres_credentials.go 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..afc68b54 --- /dev/null +++ b/internal/auth/service_credentials_db_test.go @@ -0,0 +1,276 @@ +//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-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 } From ec7ddd677bf73eff40f7a2140dd98f9db38faee9 Mon Sep 17 00:00:00 2001 From: Cristian Magherusan-Stanciu Date: Mon, 5 Oct 2026 08:08:43 +0200 Subject: [PATCH 2/2] test(auth): cover the reset token email predicate SetPasswordResetToken only stores a token while the row still has the email the reset request read. Replacing the predicate with $2::text IS NOT NULL left every subtest green. Add a subtest that changes the email from the read barrier and asserts no token is stored or mailed to the old address. --- internal/auth/service_credentials_db_test.go | 17 +++++++++++++++++ 1 file changed, 17 insertions(+) diff --git a/internal/auth/service_credentials_db_test.go b/internal/auth/service_credentials_db_test.go index afc68b54..ae5d495a 100644 --- a/internal/auth/service_credentials_db_test.go +++ b/internal/auth/service_credentials_db_test.go @@ -263,6 +263,23 @@ func TestIntegration_CredentialWritesRejectStaleCredentials(t *testing.T) { 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) {