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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
3 changes: 3 additions & 0 deletions internal/api/handler_auth.go
Original file line number Diff line number Diff line change
Expand Up @@ -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())
}
Expand Down
16 changes: 16 additions & 0 deletions internal/api/handler_auth_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
3 changes: 3 additions & 0 deletions internal/auth/interfaces.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
6 changes: 4 additions & 2 deletions internal/auth/service.go
Original file line number Diff line number Diff line change
Expand Up @@ -7,6 +7,7 @@ import (
"errors"
"fmt"
"net/mail"
"slices"
"strings"
"sync"
"sync/atomic"
Expand Down Expand Up @@ -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 {
Expand Down Expand Up @@ -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
}
Expand Down
165 changes: 165 additions & 0 deletions internal/auth/service_credentials_db_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down Expand Up @@ -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)
})
}
7 changes: 3 additions & 4 deletions internal/auth/service_mfa_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down Expand Up @@ -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{
Expand Down
12 changes: 5 additions & 7 deletions internal/auth/service_password.go
Original file line number Diff line number Diff line change
Expand Up @@ -388,22 +388,20 @@ 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
}

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
Expand All @@ -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)
Expand Down
2 changes: 1 addition & 1 deletion internal/auth/service_password_callback_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down
Loading
Loading