diff --git a/internal/api/handler_auth.go b/internal/api/handler_auth.go index da0b91ab..f18e9369 100644 --- a/internal/api/handler_auth.go +++ b/internal/api/handler_auth.go @@ -375,6 +375,9 @@ func (h *Handler) resetPassword(ctx context.Context, req *events.LambdaFunctionU if errors.Is(err, auth.ErrAccountDeactivated) { return nil, NewClientError(403, err.Error()) } + if errors.Is(err, auth.ErrUserChanged) { + return nil, NewClientError(409, err.Error()) + } if isResetPasswordClientError(err) { return nil, NewClientError(400, err.Error()) } diff --git a/internal/api/handler_auth_test.go b/internal/api/handler_auth_test.go index d742f147..4e90e60e 100644 --- a/internal/api/handler_auth_test.go +++ b/internal/api/handler_auth_test.go @@ -759,6 +759,22 @@ func TestHandler_resetPassword_AccountDeactivated(t *testing.T) { assert.Contains(t, ce.Error(), "deactivated") } +// A reset that lost a race to a concurrent account change is a 409, not a 500 (issue #493). +func TestHandler_resetPassword_UserChanged(t *testing.T) { + ctx := context.Background() + mockAuth := new(MockAuthService) + t.Cleanup(func() { mockAuth.AssertExpectations(t) }) + mockAuth.On("ConfirmPasswordReset", ctx, mock.Anything).Return(auth.ErrUserChanged) + + handler := &Handler{auth: mockAuth} + encoded := base64.StdEncoding.EncodeToString([]byte("SecureT3st@789")) + req := &events.LambdaFunctionURLRequest{Body: `{"token": "valid-token", "new_password": "` + encoded + `"}`} + _, err := handler.resetPassword(ctx, req) + ce, ok := IsClientError(err) + require.True(t, ok, "ErrUserChanged must be wrapped as a client error, not a 500") + assert.Equal(t, 409, ce.code) +} + // Issue #459: ConfirmPasswordReset errors must surface as a 4xx client // error with the original message preserved, so the frontend renders a // specific reason rather than the opaque "Failed to reset password" that diff --git a/internal/auth/interfaces.go b/internal/auth/interfaces.go index b83ccdc8..4b90874c 100644 --- a/internal/auth/interfaces.go +++ b/internal/auth/interfaces.go @@ -14,6 +14,9 @@ type StoreInterface interface { 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 + CompletePasswordReset(ctx context.Context, user *User, readResetToken, readPasswordHash string) error + ConsumePasswordResetToken(ctx context.Context, userID, readResetToken string) error + ConsumeMFARecoveryCode(ctx context.Context, userID string, readCodes, remaining []string) 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.go b/internal/auth/service.go index 5c10d6c8..8acc62e0 100644 --- a/internal/auth/service.go +++ b/internal/auth/service.go @@ -7,6 +7,7 @@ import ( "errors" "fmt" "net/mail" + "slices" "strings" "sync" "sync/atomic" @@ -215,7 +216,7 @@ func (s *Service) getUserAndValidateStatus(ctx context.Context, email string) (* // Accepts either a TOTP code OR a single-use recovery code as proof // of MFA. Consumed recovery codes are removed from the user row on // success — the success path persists the updated codes slice via -// UpdateUser before returning. A failed recovery-code attempt does +// ConsumeMFARecoveryCode before returning. A failed recovery-code attempt does // NOT consume anything (the consumeRecoveryCode call only mutates // the slice on a match). func (s *Service) verifyPasswordAndMFA(ctx context.Context, user *User, req LoginRequest) error { @@ -251,8 +252,9 @@ func (s *Service) verifyPasswordAndMFA(ctx context.Context, user *User, req Logi // TOTP miss — try a recovery code. consumeRecoveryCode mutates // the user's slice on a match; persist the slice so the // consumed code can't be reused. + readCodes := slices.Clone(user.MFARecoveryCodes) if s.consumeRecoveryCode(user, req.MFACode) { - if err := s.store.UpdateUser(ctx, user); err != nil { + if err := s.store.ConsumeMFARecoveryCode(ctx, user.ID, readCodes, user.MFARecoveryCodes); err != nil { logging.Warnf("Failed to persist recovery-code consumption for user %s: %v", user.ID, err) return ErrInvalidMFACode } diff --git a/internal/auth/service_credentials_db_test.go b/internal/auth/service_credentials_db_test.go index 86070835..a317202d 100644 --- a/internal/auth/service_credentials_db_test.go +++ b/internal/auth/service_credentials_db_test.go @@ -48,6 +48,14 @@ func (s *credentialReadBarrier) GetUserByEmail(ctx context.Context, email string return u, err } +func (s *credentialReadBarrier) GetUserByResetToken(ctx context.Context, token string) (*User, error) { + u, err := s.StoreInterface.GetUserByResetToken(ctx, token) + if err == nil { + s.pause(ctx, u) + } + return u, err +} + type recordingMailSink struct { resets []string err error @@ -293,3 +301,160 @@ func TestIntegration_CredentialWritesRejectStaleCredentials(t *testing.T) { assert.NotEmpty(t, f.stored().PasswordResetToken) }) } + +func (f *credentialRaceFixture) issueResetToken() string { + token := "reset-token-" + f.user.ID + _, err := f.store.db.Exec(f.t.Context(), `UPDATE users SET password_reset_token = $2, + password_reset_expiry = NOW() + interval '1 hour' WHERE id = $1`, f.user.ID, hashSessionToken(token)) + require.NoError(f.t, err) + return token +} + +func TestIntegration_ResetConfirmPreservesConcurrentMFA(t *testing.T) { + store := NewPostgresStore(setupAuthTestDB(t)) + ctx := t.Context() + + t.Run("confirm", func(t *testing.T) { + f := newCredentialRaceFixture(t, store, "reset-confirm@example.com") + token := f.issueResetToken() + enrolled := f.enrollOnRead() + require.NoError(t, f.svc.ConfirmPasswordReset(ctx, PasswordResetConfirm{Token: token, NewPassword: credentialRaceNew})) + stored := f.stored() + assert.True(t, f.svc.verifyPassword(credentialRaceNew, stored.PasswordHash)) + assertOnlyChanged(t, enrolled(), stored, func(w *User) { + w.PasswordHash, w.Salt, w.PasswordHistory = stored.PasswordHash, "", []string{f.user.PasswordHash} + w.PasswordResetToken, w.PasswordResetExpiry = "", nil + w.PasswordVersion++ + }) + _, err := f.svc.ValidateSession(ctx, f.session) + require.Error(t, err, "a reset must revoke existing sessions") + f.requireFactorEnforced(f.user.Email, credentialRaceNew, enrolled()) + err = f.svc.ConfirmPasswordReset(ctx, PasswordResetConfirm{Token: token, NewPassword: "AnotherPassword789!"}) + require.ErrorContains(t, err, "invalid or expired reset token") + }) + + t.Run("rejected-password-consumes-token", func(t *testing.T) { + f := newCredentialRaceFixture(t, store, "reset-confirm-weak@example.com") + token := f.issueResetToken() + enrolled := f.enrollOnRead() + require.Error(t, f.svc.ConfirmPasswordReset(ctx, PasswordResetConfirm{Token: token, NewPassword: "weak"})) + assertOnlyChanged(t, enrolled(), f.stored(), func(w *User) { w.PasswordResetToken, w.PasswordResetExpiry = "", nil }) + f.requireFactorEnforced(f.user.Email, credentialRacePassword, enrolled()) + }) +} + +func TestIntegration_ResetConfirmRejectsStaleRead(t *testing.T) { + store := NewPostgresStore(setupAuthTestDB(t)) + ctx := t.Context() + + t.Run("after-concurrent-confirm", func(t *testing.T) { + f := newCredentialRaceFixture(t, store, "reset-twice-confirm@example.com") + token := f.issueResetToken() + winner := f.onRead(func(ctx context.Context, other *Service) { + require.NoError(t, other.ConfirmPasswordReset(ctx, PasswordResetConfirm{Token: token, NewPassword: "WinnerPassword789!"})) + }) + err := f.svc.ConfirmPasswordReset(ctx, PasswordResetConfirm{Token: token, NewPassword: credentialRaceNew}) + require.ErrorIs(t, err, ErrUserChanged, "a reset token must be spent once") + assert.Equal(t, winner(), f.stored()) + }) + + t.Run("after-concurrent-rejected-confirm", func(t *testing.T) { + f := newCredentialRaceFixture(t, store, "reset-after-rejected@example.com") + token := f.issueResetToken() + winner := f.onRead(func(ctx context.Context, other *Service) { + require.Error(t, other.ConfirmPasswordReset(ctx, PasswordResetConfirm{Token: token, NewPassword: "weak"})) + }) + err := f.svc.ConfirmPasswordReset(ctx, PasswordResetConfirm{Token: token, NewPassword: credentialRaceNew}) + require.ErrorIs(t, err, ErrUserChanged, "a token consumed by a rejected attempt must not set a password") + assert.Equal(t, winner(), f.stored()) + }) + + t.Run("rejected-confirm-after-concurrent-reissue", func(t *testing.T) { + f := newCredentialRaceFixture(t, store, "reset-rejected-reissue@example.com") + token := f.issueResetToken() + winner := f.onRead(func(ctx context.Context, other *Service) { + _, 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, other.RequestPasswordReset(ctx, f.user.Email)) + }) + err := f.svc.ConfirmPasswordReset(ctx, PasswordResetConfirm{Token: token, NewPassword: "weak"}) + require.Error(t, err) + require.NotErrorIs(t, err, ErrUserChanged) + reissued := winner().PasswordResetToken + require.NotEmpty(t, reissued) + require.NotEqual(t, hashSessionToken(token), reissued, "the concurrent request must have replaced the token") + assert.Equal(t, winner(), f.stored(), "a rejected confirm must not clear a token it did not read") + }) + + t.Run("after-concurrent-password-change", func(t *testing.T) { + f := newCredentialRaceFixture(t, store, "reset-after-change@example.com") + token := f.issueResetToken() + winner := f.onRead(func(ctx context.Context, other *Service) { + require.NoError(t, other.ChangePassword(ctx, f.user.ID, ChangePasswordRequest{CurrentPassword: credentialRacePassword, NewPassword: "WinnerPassword789!"})) + }) + err := f.svc.ConfirmPasswordReset(ctx, PasswordResetConfirm{Token: token, NewPassword: credentialRaceNew}) + require.ErrorIs(t, err, ErrUserChanged) + assert.Equal(t, winner(), f.stored()) + require.NoError(t, f.svc.ConfirmPasswordReset(ctx, PasswordResetConfirm{Token: token, NewPassword: credentialRaceNew}), + "a lost race must leave the token usable") + assert.Equal(t, []string{winner().PasswordHash, f.user.PasswordHash}, f.stored().PasswordHistory) + }) + + t.Run("after-concurrent-deactivation", func(t *testing.T) { + f := newCredentialRaceFixture(t, store, "reset-after-deactivation@example.com") + token := f.issueResetToken() + 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)) + }) + err := f.svc.ConfirmPasswordReset(ctx, PasswordResetConfirm{Token: token, NewPassword: credentialRaceNew}) + require.ErrorIs(t, err, ErrUserChanged) + assert.Equal(t, winner(), f.stored()) + require.ErrorIs(t, f.svc.ConfirmPasswordReset(ctx, PasswordResetConfirm{Token: token, NewPassword: credentialRaceNew}), ErrAccountDeactivated) + assert.Empty(t, f.stored().PasswordResetToken) + }) +} + +func (f *credentialRaceFixture) enrollMFA() []string { + setup, err := f.svc.MFASetup(f.t.Context(), f.user.ID, credentialRacePassword) + require.NoError(f.t, err) + codes, err := f.svc.MFAEnable(f.t.Context(), f.user.ID, generateTOTP(setup.Secret, time.Now().Unix()/30)) + require.NoError(f.t, err) + return codes +} + +func TestIntegration_RecoveryCodeLoginRejectsStaleRead(t *testing.T) { + store := NewPostgresStore(setupAuthTestDB(t)) + ctx := t.Context() + + t.Run("after-concurrent-use-of-same-code", func(t *testing.T) { + f := newCredentialRaceFixture(t, store, "recovery-twice@example.com") + req := LoginRequest{Email: f.user.Email, Password: credentialRacePassword, MFACode: f.enrollMFA()[0]} + winner := f.onRead(func(ctx context.Context, other *Service) { + _, err := other.Login(ctx, req) + require.NoError(t, err) + }) + _, err := f.svc.Login(ctx, req) + require.ErrorIs(t, err, ErrInvalidMFACode, "a recovery code must be spent once") + assert.Equal(t, winner(), f.stored()) + }) + + t.Run("after-concurrent-password-change", func(t *testing.T) { + f := newCredentialRaceFixture(t, store, "recovery-after-change@example.com") + codes := f.enrollMFA() + winner := f.onRead(func(ctx context.Context, other *Service) { + require.NoError(t, other.ChangePassword(ctx, f.user.ID, ChangePasswordRequest{CurrentPassword: credentialRacePassword, NewPassword: credentialRaceNew})) + }) + // Login checks the password it read, so it still succeeds here (a separate, pre-existing gap); + // this test pins only that the code write leaves the new password in place. + _, err := f.svc.Login(ctx, LoginRequest{Email: f.user.Email, Password: credentialRacePassword, MFACode: codes[0]}) + require.NoError(t, err) + stored := f.stored() + assert.Len(t, stored.MFARecoveryCodes, len(codes)-1) + assertOnlyChanged(t, winner(), stored, func(w *User) { w.MFARecoveryCodes = stored.MFARecoveryCodes }) + f.requireFactorEnforced(f.user.Email, credentialRaceNew, stored) + }) +} diff --git a/internal/auth/service_mfa_test.go b/internal/auth/service_mfa_test.go index b4b33705..67168c8e 100644 --- a/internal/auth/service_mfa_test.go +++ b/internal/auth/service_mfa_test.go @@ -450,7 +450,7 @@ func TestLogin_WithMFA_RecoveryCode_ConsumedOnce(t *testing.T) { user.MFARecoveryCodes = []string{hash} mockStore.On("GetUserByEmail", ctx, user.Email).Return(user, nil) - mockStore.On("UpdateUser", ctx, mock.AnythingOfType("*auth.User")).Return(nil) + mockStore.On("ConsumeMFARecoveryCode", ctx, user.ID, []string{hash}, []string{}).Return(nil).Once() mockStore.On("RecordSuccessfulLogin", ctx, user.ID).Return(nil) mockStore.On("RecordFailedLogin", ctx, user.ID).Return(nil) mockStore.On("CreateSession", ctx, mock.AnythingOfType("*auth.Session")).Return(nil) @@ -483,9 +483,8 @@ func TestLogin_WithMFA_RecoveryCode_PersistenceFailure(t *testing.T) { snapshot.MFARecoveryCodes = append([]string(nil), user.MFARecoveryCodes...) store.On("GetUserByEmail", ctx, user.Email).Return(&snapshot, nil).Once() } - store.On("UpdateUser", ctx, mock.MatchedBy(func(updated *User) bool { - return updated.ID == user.ID && len(updated.MFARecoveryCodes) == 0 - })).Return(errors.New("consumption write failed")).Twice() + store.On("ConsumeMFARecoveryCode", ctx, user.ID, []string{hash}, []string{}). + Return(errors.New("consumption write failed")).Twice() for range 2 { response, loginErr := service.Login(ctx, LoginRequest{ diff --git a/internal/auth/service_password.go b/internal/auth/service_password.go index c49dddb3..5a092608 100644 --- a/internal/auth/service_password.go +++ b/internal/auth/service_password.go @@ -388,14 +388,12 @@ func (s *Service) ConfirmPasswordReset(ctx context.Context, req PasswordResetCon return err } - // Invalidate the reset token before processing to ensure one-time use - user.PasswordResetToken = "" - user.PasswordResetExpiry = nil + readResetToken, readPasswordHash := user.PasswordResetToken, user.PasswordHash // An admin-deactivated account must not reactivate itself through a reset // (A03-006). The token is still consumed so the link cannot be replayed. if user.DeactivatedAt != nil { - if updateErr := s.store.UpdateUser(ctx, user); updateErr != nil { + if updateErr := s.store.ConsumePasswordResetToken(ctx, user.ID, readResetToken); updateErr != nil { logging.Warnf("Failed to invalidate reset token for deactivated user %s: %v", user.ID, updateErr) } return ErrAccountDeactivated @@ -403,7 +401,7 @@ func (s *Service) ConfirmPasswordReset(ctx context.Context, req PasswordResetCon if err := s.processPasswordReset(user, req.NewPassword); err != nil { // Token is consumed even on validation failure (one-time use) - if updateErr := s.store.UpdateUser(ctx, user); updateErr != nil { + if updateErr := s.store.ConsumePasswordResetToken(ctx, user.ID, readResetToken); updateErr != nil { logging.Warnf("Failed to invalidate reset token after password validation failure: %v", updateErr) } return err @@ -415,12 +413,12 @@ func (s *Service) ConfirmPasswordReset(ctx context.Context, req PasswordResetCon user.Active = true } - if err := s.store.UpdateUser(ctx, user); err != nil { + if err := s.store.CompletePasswordReset(ctx, user, readResetToken, readPasswordHash); err != nil { return err } // See the matching comment in ChangePassword: invalidate only after the - // new password is persisted, or a failed UpdateUser leaves the caller + // new password is persisted, or a failed write leaves the caller // with revoked credentials for a password that never actually changed. s.invalidateUserCredentialsBestEffort(ctx, user.ID, "password reset") s.notifyPasswordChange(ctx, user.ID, req.NewPassword) diff --git a/internal/auth/service_password_callback_test.go b/internal/auth/service_password_callback_test.go index 8aa8d10c..3fa5d329 100644 --- a/internal/auth/service_password_callback_test.go +++ b/internal/auth/service_password_callback_test.go @@ -133,7 +133,7 @@ func TestService_OnPasswordChange_ConfirmPasswordReset(t *testing.T) { mockStore.On("GetUserByResetToken", ctx, mock.AnythingOfType("string")).Return(testUser, nil).Once() mockStore.On("DeleteUserSessions", ctx, "user-456").Return(nil).Once() mockStore.On("ListAPIKeysByUser", ctx, "user-456").Return([]*UserAPIKey{}, nil).Once() - mockStore.On("UpdateUser", ctx, mock.AnythingOfType("*auth.User")).Return(nil).Once() + mockStore.On("CompletePasswordReset", ctx, mock.AnythingOfType("*auth.User"), hashSessionToken("valid-token"), "").Return(nil).Once() err := service.ConfirmPasswordReset(ctx, PasswordResetConfirm{ Token: "valid-token", diff --git a/internal/auth/service_password_test.go b/internal/auth/service_password_test.go index 60fd9d02..52ae211a 100644 --- a/internal/auth/service_password_test.go +++ b/internal/auth/service_password_test.go @@ -410,8 +410,7 @@ func TestService_ConfirmPasswordReset(t *testing.T) { mockStore.On("GetUserByResetToken", ctx, mock.AnythingOfType("string")).Return(testUser, nil).Once() mockStore.On("DeleteUserSessions", ctx, "user-123").Return(nil).Once() mockStore.On("ListAPIKeysByUser", ctx, "user-123").Return([]*UserAPIKey{}, nil).Once() - // UpdateUser is called once: password change + token invalidation in single call - mockStore.On("UpdateUser", ctx, mock.AnythingOfType("*auth.User")).Return(nil).Once() + mockStore.On("CompletePasswordReset", ctx, mock.AnythingOfType("*auth.User"), hashSessionToken("valid-reset-token"), "").Return(nil).Once() req := PasswordResetConfirm{ Token: "valid-reset-token", @@ -461,7 +460,7 @@ func TestService_ConfirmPasswordReset(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("CompletePasswordReset", ctx, mock.AnythingOfType("*auth.User"), hashSessionToken("valid-reset-token"), "").Return(nil).Once() err = service.ConfirmPasswordReset(ctx, PasswordResetConfirm{ Token: "valid-reset-token", @@ -543,7 +542,7 @@ func TestService_ConfirmPasswordReset(t *testing.T) { // Token is hashed before lookup mockStore.On("GetUserByResetToken", ctx, mock.AnythingOfType("string")).Return(testUser, nil).Once() // Token is invalidated even on password validation failure (one-time use) - mockStore.On("UpdateUser", ctx, mock.AnythingOfType("*auth.User")).Return(nil).Once() + mockStore.On("ConsumePasswordResetToken", ctx, "user-123", hashSessionToken("valid-reset-token")).Return(nil).Once() req := PasswordResetConfirm{ Token: "valid-reset-token", @@ -576,7 +575,7 @@ func TestService_ConfirmPasswordReset(t *testing.T) { mockStore.On("GetUserByResetToken", ctx, mock.AnythingOfType("string")).Return(testUser, nil).Once() // Token is invalidated even on password validation failure (one-time use) - mockStore.On("UpdateUser", ctx, mock.AnythingOfType("*auth.User")).Return(nil).Once() + mockStore.On("ConsumePasswordResetToken", ctx, "user-123", hashSessionToken("valid-reset-token")).Return(nil).Once() // Try to reuse a password from history req := PasswordResetConfirm{ @@ -615,7 +614,7 @@ func TestService_ConfirmPasswordReset(t *testing.T) { mockStore.On("GetUserByResetToken", ctx, mock.AnythingOfType("string")).Return(testUser, nil).Once() // Token is invalidated even on validation failure (one-time use). - mockStore.On("UpdateUser", ctx, mock.AnythingOfType("*auth.User")).Return(nil).Once() + mockStore.On("ConsumePasswordResetToken", ctx, "user-123", hashSessionToken("valid-reset-token")).Return(nil).Once() req := PasswordResetConfirm{ Token: "valid-reset-token", @@ -688,12 +687,8 @@ func TestService_ConfirmPasswordReset(t *testing.T) { mockStore.On("GetUserByResetToken", ctx, mock.AnythingOfType("string")).Return(deactivatedUser, nil).Once() // The token is still consumed (one-time use) even though the - // reactivation itself is refused. PasswordHash must be untouched: - // processPasswordReset (which would set the NEW password) must - // never run for a deactivated account. - mockStore.On("UpdateUser", ctx, mock.MatchedBy(func(u *User) bool { - return u.PasswordResetToken == "" && u.PasswordResetExpiry == nil && u.PasswordHash == originalHash - })).Return(nil).Once() + // reactivation itself is refused, and no password write happens. + mockStore.On("ConsumePasswordResetToken", ctx, "user-789", hashSessionToken("valid-reset-token")).Return(nil).Once() req := PasswordResetConfirm{ Token: "valid-reset-token", @@ -730,7 +725,8 @@ func TestService_ConfirmPasswordReset(t *testing.T) { mockStore.On("GetUserByResetToken", ctx, mock.AnythingOfType("string")).Return(invitedUser, nil).Once() mockStore.On("DeleteUserSessions", ctx, "user-790").Return(nil).Once() mockStore.On("ListAPIKeysByUser", ctx, "user-790").Return([]*UserAPIKey{}, nil).Once() - mockStore.On("UpdateUser", ctx, mock.AnythingOfType("*auth.User")).Return(nil).Once() + mockStore.On("CompletePasswordReset", ctx, mock.MatchedBy(func(u *User) bool { return u.Active }), + hashSessionToken("valid-invite-token"), "").Return(nil).Once() req := PasswordResetConfirm{ Token: "valid-invite-token", diff --git a/internal/auth/service_test.go b/internal/auth/service_test.go index 785cb274..84632342 100644 --- a/internal/auth/service_test.go +++ b/internal/auth/service_test.go @@ -702,8 +702,7 @@ func TestService_ErrorPaths(t *testing.T) { mockStore.On("GetUserByResetToken", ctx, mock.AnythingOfType("string")).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() - // UpdateUser is called once: password change + token invalidation in single call - mockStore.On("UpdateUser", ctx, mock.AnythingOfType("*auth.User")).Return(nil).Once() + mockStore.On("CompletePasswordReset", ctx, mock.AnythingOfType("*auth.User"), hashSessionToken("valid-reset-token"), "").Return(nil).Once() req := PasswordResetConfirm{ Token: "valid-reset-token", diff --git a/internal/auth/store_postgres_credentials.go b/internal/auth/store_postgres_credentials.go index 7d70db38..234b027d 100644 --- a/internal/auth/store_postgres_credentials.go +++ b/internal/auth/store_postgres_credentials.go @@ -23,6 +23,56 @@ func (s *PostgresStore) UpdateUserCredentials(ctx context.Context, user *User, r return nil } +// CompletePasswordReset writes the new password, activates the account and +// consumes the reset token, only while the row still holds the token and hash +// the caller read and the account is not deactivated (issue #493). +func (s *PostgresStore) CompletePasswordReset(ctx context.Context, user *User, readResetToken, readPasswordHash string) error { + result, err := s.db.Exec(ctx, ` + UPDATE users SET password_hash = $4, salt = $5, password_history = $6, active = $7, + password_reset_token = NULL, password_reset_expiry = NULL, updated_at = NOW() + WHERE id = $1 AND password_reset_token = $2 AND password_hash = $3 AND deactivated_at IS NULL + `, user.ID, readResetToken, readPasswordHash, user.PasswordHash, user.Salt, user.PasswordHistory, user.Active) + if err != nil { + return fmt.Errorf("failed to complete password reset: %w", err) + } + if result.RowsAffected() == 0 { + return ErrUserChanged + } + return nil +} + +// ConsumePasswordResetToken clears the reset token columns only while the row +// still holds the token the caller read. +func (s *PostgresStore) ConsumePasswordResetToken(ctx context.Context, userID, readResetToken string) error { + result, err := s.db.Exec(ctx, ` + UPDATE users SET password_reset_token = NULL, password_reset_expiry = NULL, updated_at = NOW() + WHERE id = $1 AND password_reset_token = $2 + `, userID, readResetToken) + if err != nil { + return fmt.Errorf("failed to consume reset token: %w", err) + } + if result.RowsAffected() == 0 { + return ErrUserChanged + } + return nil +} + +// ConsumeMFARecoveryCode writes only the recovery codes, and only while the row +// still holds the codes the caller read, so each code is spent once (issue #493). +func (s *PostgresStore) ConsumeMFARecoveryCode(ctx context.Context, userID string, readCodes, remaining []string) error { + result, err := s.db.Exec(ctx, ` + UPDATE users SET mfa_recovery_codes = $3, updated_at = NOW() + WHERE id = $1 AND mfa_recovery_codes = $2 + `, userID, readCodes, remaining) + if err != nil { + return fmt.Errorf("failed to consume recovery code: %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. diff --git a/internal/auth/test_helpers.go b/internal/auth/test_helpers.go index 6819761a..d06b2ef9 100644 --- a/internal/auth/test_helpers.go +++ b/internal/auth/test_helpers.go @@ -58,6 +58,18 @@ func (m *MockStore) SetPasswordResetToken(ctx context.Context, user *User, readE return m.Called(ctx, user, readExpiry).Error(0) } +func (m *MockStore) CompletePasswordReset(ctx context.Context, user *User, readResetToken, readPasswordHash string) error { + return m.Called(ctx, user, readResetToken, readPasswordHash).Error(0) +} + +func (m *MockStore) ConsumePasswordResetToken(ctx context.Context, userID, readResetToken string) error { + return m.Called(ctx, userID, readResetToken).Error(0) +} + +func (m *MockStore) ConsumeMFARecoveryCode(ctx context.Context, userID string, readCodes, remaining []string) error { + return m.Called(ctx, userID, readCodes, remaining).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 c3ad9ea7..e2ee4b80 100644 --- a/internal/mocks/stores.go +++ b/internal/mocks/stores.go @@ -805,6 +805,18 @@ func (m *MockAuthStore) SetPasswordResetToken(ctx context.Context, user *auth.Us return m.Called(ctx, user, readExpiry).Error(0) } +func (m *MockAuthStore) CompletePasswordReset(ctx context.Context, user *auth.User, readResetToken, readPasswordHash string) error { + return m.Called(ctx, user, readResetToken, readPasswordHash).Error(0) +} + +func (m *MockAuthStore) ConsumePasswordResetToken(ctx context.Context, userID, readResetToken string) error { + return m.Called(ctx, userID, readResetToken).Error(0) +} + +func (m *MockAuthStore) ConsumeMFARecoveryCode(ctx context.Context, userID string, readCodes, remaining []string) error { + return m.Called(ctx, userID, readCodes, remaining).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 95b3f213..67192af7 100644 --- a/internal/server/health_test.go +++ b/internal/server/health_test.go @@ -40,6 +40,18 @@ func (m *mockAuthStoreForHealth) SetPasswordResetToken(context.Context, *auth.Us return nil } +func (m *mockAuthStoreForHealth) CompletePasswordReset(context.Context, *auth.User, string, string) error { + return nil +} + +func (m *mockAuthStoreForHealth) ConsumePasswordResetToken(context.Context, string, string) error { + return nil +} + +func (m *mockAuthStoreForHealth) ConsumeMFARecoveryCode(context.Context, string, []string, []string) error { + return nil +} + func (m *mockAuthStoreForHealth) RecordFailedLogin(context.Context, string) error { return nil }