Skip to content
Open
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
7 changes: 2 additions & 5 deletions internal/auth/service.go
Original file line number Diff line number Diff line change
Expand Up @@ -210,7 +210,7 @@ func (s *Service) getUserAndValidateStatus(ctx context.Context, email string) (*
// the request didn't carry a code, so the API handler can map to a
// machine-readable response (`{"error":"mfa_required"}`) rather than
// a generic 401. Returns ErrInvalidMFACode (sentinel) when the code
// was provided but didn't match TOTP or any stored recovery code.
// was provided but didn't match TOTP or a persistently consumed recovery code.
//
// 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
Expand Down Expand Up @@ -254,10 +254,7 @@ func (s *Service) verifyPasswordAndMFA(ctx context.Context, user *User, req Logi
if s.consumeRecoveryCode(user, req.MFACode) {
if err := s.store.UpdateUser(ctx, user); err != nil {
logging.Warnf("Failed to persist recovery-code consumption for user %s: %v", user.ID, err)
// The recovery code already verified; still allow
// login but warn — repeated use of the same code on
// the next login will fail because the slice is
// stale, which is the safe failure mode.
return ErrInvalidMFACode
}
return nil
}
Expand Down
35 changes: 35 additions & 0 deletions internal/auth/service_mfa_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -465,6 +465,41 @@ func TestLogin_WithMFA_RecoveryCode_ConsumedOnce(t *testing.T) {
assert.Empty(t, user.MFARecoveryCodes, "consumed recovery code must be removed from the slice")
}

func TestLogin_WithMFA_RecoveryCode_PersistenceFailure(t *testing.T) {
t.Parallel()
ctx := t.Context()
store := new(MockStore)
service := createTestService(store, new(MockEmailSender))
code := "ABCD-2345"
hash, err := service.hashRecoveryCode(code)
require.NoError(t, err)
user := createTestUser(t, "SecurePass@123")
user.MFAEnabled = true
user.MFASecret = "JBSWY3DPEHPK3PXP"
user.MFARecoveryCodes = []string{hash}

for range 2 {
snapshot := *user
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()

for range 2 {
response, loginErr := service.Login(ctx, LoginRequest{
Email: user.Email, Password: "SecurePass@123", MFACode: code,
})
assert.ErrorIs(t, loginErr, ErrInvalidMFACode)
assert.True(t, response == nil, "failed consumption must not return a token")
}
store.AssertNotCalled(t, "CreateSession", mock.Anything, mock.Anything)
store.AssertNotCalled(t, "RecordSuccessfulLogin", mock.Anything, mock.Anything)
store.AssertNotCalled(t, "RecordFailedLogin", mock.Anything, mock.Anything)
store.AssertExpectations(t)
}

// ---------------------------------------------------------------
// Sentinel-identity tests (issue #512).
//
Expand Down
72 changes: 72 additions & 0 deletions internal/auth/service_recovery_db_test.go
Original file line number Diff line number Diff line change
@@ -0,0 +1,72 @@
//go:build integration

package auth

import (
"context"
"testing"
"time"

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

func TestIntegration_LoginRecoveryPersistenceFailure(t *testing.T) {
db := setupAuthTestDB(t)
store := NewPostgresStore(db)
service := NewService(ServiceConfig{Store: store})
service.bcryptCostOverride = 4
ctx := t.Context()
password, err := service.hashPassword("SyntheticPassword123!")
require.NoError(t, err)
code := "ABCD-2345"
codeHash, err := service.hashRecoveryCode(code)
require.NoError(t, err)
user := &User{
Email: "recovery-persistence@example.com", PasswordHash: password, Active: true,
GroupIDs: []string{DefaultPurchaserGroupID}, FailedLoginAttempts: 2,
MFAEnabled: true, MFASecret: "JBSWY3DPEHPK3PXP", MFARecoveryCodes: []string{codeHash},
}
require.NoError(t, store.CreateUser(ctx, user))
_, err = db.Exec(ctx, `ALTER TABLE users ADD CONSTRAINT test_recovery_consumption_failure
CHECK (email <> 'recovery-persistence@example.com' OR cardinality(mfa_recovery_codes) > 0)`)
require.NoError(t, err)
t.Cleanup(func() {
cleanupCtx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
defer cancel()
_, cleanupErr := db.Exec(cleanupCtx, "ALTER TABLE users DROP CONSTRAINT IF EXISTS test_recovery_consumption_failure")
assert.NoError(t, cleanupErr)
})
request := LoginRequest{Email: user.Email, Password: "SyntheticPassword123!", MFACode: code}
for range 2 {
response, loginErr := service.Login(ctx, request)
assert.ErrorIs(t, loginErr, ErrInvalidMFACode)
assert.True(t, response == nil, "failed consumption must not return a token")
stored, readErr := store.GetUserByID(ctx, user.ID)
require.NoError(t, readErr)
assert.Equal(t, []string{codeHash}, stored.MFARecoveryCodes)
assert.Equal(t, 2, stored.FailedLoginAttempts)
assert.Nil(t, stored.LastLoginAt)
var sessions int
require.NoError(t, db.QueryRow(ctx, "SELECT count(*) FROM sessions WHERE user_id=$1", user.ID).Scan(&sessions))
assert.Zero(t, sessions)
}
_, err = db.Exec(ctx, "ALTER TABLE users DROP CONSTRAINT test_recovery_consumption_failure")
require.NoError(t, err)
response, err := service.Login(ctx, request)
require.NoError(t, err)
require.NotNil(t, response)
_, err = service.ValidateSession(ctx, response.Token)
require.NoError(t, err)
stored, err := store.GetUserByID(ctx, user.ID)
require.NoError(t, err)
assert.Empty(t, stored.MFARecoveryCodes)
assert.Zero(t, stored.FailedLoginAttempts)
assert.NotNil(t, stored.LastLoginAt)
response, err = service.Login(ctx, request)
assert.ErrorIs(t, err, ErrInvalidMFACode)
assert.True(t, response == nil, "consumed code must not return another token")
var sessions int
require.NoError(t, db.QueryRow(ctx, "SELECT count(*) FROM sessions WHERE user_id=$1", user.ID).Scan(&sessions))
assert.Equal(t, 1, sessions)
}
Loading