From f1685d14d67faec8c98c84d6e72d01aab276cb29 Mon Sep 17 00:00:00 2001 From: Cristian Magherusan-Stanciu Date: Wed, 30 Sep 2026 03:51:28 +0200 Subject: [PATCH] fix(auth): persist login bookkeeping atomically Update only login counters, lockout and timestamps so an in-flight login cannot restore stale account security fields. Increment failed attempts in PostgreSQL to preserve concurrent failures. Verify service routing, concurrent security updates, mixed login ordering, lockout expiry and persistence errors against local PostgreSQL. --- internal/auth/interfaces.go | 2 + internal/auth/service.go | 11 +- internal/auth/service_lockout_test.go | 433 +++++---------------- internal/auth/service_login_db_test.go | 346 ++++++++++++++++ internal/auth/service_mfa_test.go | 4 +- internal/auth/service_test.go | 16 +- internal/auth/service_user.go | 13 +- internal/auth/store_postgres_login.go | 39 ++ internal/auth/store_postgres_login_test.go | 37 ++ internal/auth/test_helpers.go | 8 + internal/mocks/stores.go | 8 + internal/server/adapter_test.go | 2 +- internal/server/health_test.go | 8 + 13 files changed, 556 insertions(+), 371 deletions(-) create mode 100644 internal/auth/service_login_db_test.go create mode 100644 internal/auth/store_postgres_login.go create mode 100644 internal/auth/store_postgres_login_test.go diff --git a/internal/auth/interfaces.go b/internal/auth/interfaces.go index 52fd6085..bf526005 100644 --- a/internal/auth/interfaces.go +++ b/internal/auth/interfaces.go @@ -11,6 +11,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 + RecordFailedLogin(ctx context.Context, userID string) error + RecordSuccessfulLogin(ctx context.Context, userID string) error DeleteUser(ctx context.Context, userID string) error ListUsers(ctx context.Context) ([]User, error) GetUserByResetToken(ctx context.Context, token string) (*User, error) diff --git a/internal/auth/service.go b/internal/auth/service.go index aeefa30d..aad14e33 100644 --- a/internal/auth/service.go +++ b/internal/auth/service.go @@ -275,15 +275,8 @@ func (s *Service) completeSuccessfulLogin(ctx context.Context, user *User) (*Log return nil, fmt.Errorf("failed to create session: %w", err) } - now := time.Now() - user.LastLoginAt = &now - user.FailedLoginAttempts = 0 - user.LockedUntil = nil - // Deliberately do not return this error: the session was successfully created and the - // token already issued. Failing here would leave the caller with no token despite a - // valid login. The consequence is that LastLoginAt / FailedLoginAttempts may be stale - // in the store until the next successful login, which is an acceptable trade-off. - if err := s.store.UpdateUser(ctx, user); err != nil { + // The session already exists; bookkeeping failure must not hide its token. + if err := s.store.RecordSuccessfulLogin(ctx, user.ID); err != nil { logging.Warnf("Failed to update login info for user %s: %v", user.ID, err) } diff --git a/internal/auth/service_lockout_test.go b/internal/auth/service_lockout_test.go index 1b6f3c37..5efc859b 100644 --- a/internal/auth/service_lockout_test.go +++ b/internal/auth/service_lockout_test.go @@ -11,352 +11,107 @@ import ( "github.com/stretchr/testify/require" ) -// TestLogin_AccountLockout_BeforePasswordCheck verifies lockout check happens before password verification. -func TestLogin_AccountLockout_BeforePasswordCheck(t *testing.T) { - ctx := context.Background() - mockStore := new(MockStore) - mockEmail := new(MockEmailSender) - service := createTestService(mockStore, mockEmail) - - // Create a locked user with correct password - testUser := createTestUser(t, "CorrectPassword123") - lockUntil := time.Now().Add(10 * time.Minute) - testUser.LockedUntil = &lockUntil - testUser.FailedLoginAttempts = MaxFailedLoginAttempts - - mockStore.On("GetUserByEmail", ctx, "test@example.com").Return(testUser, nil).Once() - // UpdateUser should NOT be called since we fail on lockout check before password verification - // No password verification happens when account is locked - - req := LoginRequest{ - Email: "test@example.com", - Password: "CorrectPassword123", // Even with correct password - } - - resp, err := service.Login(ctx, req) - assert.Error(t, err) - assert.Nil(t, resp) - assert.Contains(t, err.Error(), "Check your email address and password and try again") // Generic error to prevent user enumeration - - mockStore.AssertExpectations(t) - // Verify UpdateUser was NOT called - lockout check happens first - mockStore.AssertNotCalled(t, "UpdateUser", ctx, mock.Anything) -} - -// TestLogin_AccountLockout_FailedAttempts verifies lockout occurs after max failed attempts. -func TestLogin_AccountLockout_FailedAttempts(t *testing.T) { - ctx := context.Background() - mockStore := new(MockStore) - mockEmail := new(MockEmailSender) - service := createTestService(mockStore, mockEmail) - - testUser := createTestUser(t, "CorrectPassword123") - testUser.FailedLoginAttempts = 4 // One more attempt will trigger lockout - - // Track calls to UpdateUser - var updatedUser *User - mockStore.On("GetUserByEmail", ctx, "test@example.com").Return(testUser, nil).Once() - mockStore.On("UpdateUser", ctx, mock.AnythingOfType("*auth.User")). - Run(func(args mock.Arguments) { - updatedUser = args.Get(1).(*User) - }). - Return(nil).Once() - - req := LoginRequest{ - Email: "test@example.com", - Password: "WrongPassword", // Wrong password triggers failed attempt - } - - resp, err := service.Login(ctx, req) - assert.Error(t, err) - assert.Nil(t, resp) - - // Verify user was locked - require.NotNil(t, updatedUser) - assert.Equal(t, MaxFailedLoginAttempts, updatedUser.FailedLoginAttempts) - assert.NotNil(t, updatedUser.LockedUntil) - assert.True(t, updatedUser.LockedUntil.After(time.Now())) - - mockStore.AssertExpectations(t) -} - -// TestLogin_AccountLockout_Duration verifies lockout duration is correct. -func TestLogin_AccountLockout_Duration(t *testing.T) { - ctx := context.Background() - mockStore := new(MockStore) - mockEmail := new(MockEmailSender) - service := createTestService(mockStore, mockEmail) - - testUser := createTestUser(t, "CorrectPassword123") - testUser.FailedLoginAttempts = MaxFailedLoginAttempts - 1 - - var updatedUser *User - mockStore.On("GetUserByEmail", ctx, "test@example.com").Return(testUser, nil).Once() - mockStore.On("UpdateUser", ctx, mock.AnythingOfType("*auth.User")). - Run(func(args mock.Arguments) { - updatedUser = args.Get(1).(*User) - }). - Return(nil).Once() - - req := LoginRequest{ - Email: "test@example.com", - Password: "WrongPassword", - } - - _, err := service.Login(ctx, req) - assert.Error(t, err) - - // Verify lockout duration is AccountLockoutDuration (15 minutes) - require.NotNil(t, updatedUser) - require.NotNil(t, updatedUser.LockedUntil) - - expectedLockout := time.Now().Add(AccountLockoutDuration) - // Allow 1 second tolerance for test execution time - assert.WithinDuration(t, expectedLockout, *updatedUser.LockedUntil, time.Second) - - mockStore.AssertExpectations(t) -} - -// TestLogin_AccountLockout_ExpiredLock verifies expired lockouts allow login. -func TestLogin_AccountLockout_ExpiredLock(t *testing.T) { - ctx := context.Background() - mockStore := new(MockStore) - mockEmail := new(MockEmailSender) - service := createTestService(mockStore, mockEmail) - - testUser := createTestUser(t, "CorrectPassword123") - // Lockout expired 1 minute ago - lockUntil := time.Now().Add(-1 * time.Minute) - testUser.LockedUntil = &lockUntil - testUser.FailedLoginAttempts = MaxFailedLoginAttempts - - mockStore.On("GetUserByEmail", ctx, "test@example.com").Return(testUser, nil).Once() - mockStore.On("CreateSession", ctx, mock.AnythingOfType("*auth.Session")).Return(nil).Once() - mockStore.On("UpdateUser", ctx, mock.AnythingOfType("*auth.User")).Return(nil).Once() - - req := LoginRequest{ - Email: "test@example.com", - Password: "CorrectPassword123", - } - - resp, err := service.Login(ctx, req) - require.NoError(t, err) - assert.NotNil(t, resp) - assert.NotEmpty(t, resp.Token) - - mockStore.AssertExpectations(t) -} - -// TestLogin_AccountLockout_ResetOnSuccess verifies successful login resets failed attempts. -func TestLogin_AccountLockout_ResetOnSuccess(t *testing.T) { - ctx := context.Background() - mockStore := new(MockStore) - mockEmail := new(MockEmailSender) - service := createTestService(mockStore, mockEmail) - - testUser := createTestUser(t, "CorrectPassword123") - testUser.FailedLoginAttempts = 3 // Some failed attempts, but not locked - - var updatedUser *User - mockStore.On("GetUserByEmail", ctx, "test@example.com").Return(testUser, nil).Once() - mockStore.On("CreateSession", ctx, mock.AnythingOfType("*auth.Session")).Return(nil).Once() - mockStore.On("UpdateUser", ctx, mock.AnythingOfType("*auth.User")). - Run(func(args mock.Arguments) { - updatedUser = args.Get(1).(*User) - }). - Return(nil).Once() - - req := LoginRequest{ - Email: "test@example.com", - Password: "CorrectPassword123", - } - - resp, err := service.Login(ctx, req) - require.NoError(t, err) - assert.NotNil(t, resp) - - // Verify failed attempts were reset - require.NotNil(t, updatedUser) - assert.Equal(t, 0, updatedUser.FailedLoginAttempts) - assert.Nil(t, updatedUser.LockedUntil) - - mockStore.AssertExpectations(t) -} - -// TestLogin_AccountLockout_IncrementalFailures verifies each failure increments counter. -func TestLogin_AccountLockout_IncrementalFailures(t *testing.T) { - ctx := context.Background() - - for attempt := 0; attempt < MaxFailedLoginAttempts; attempt++ { - t.Run(fmt.Sprintf("Attempt_%d", attempt+1), func(t *testing.T) { - mockStore := new(MockStore) - mockEmail := new(MockEmailSender) - service := createTestService(mockStore, mockEmail) - - testUser := createTestUser(t, "CorrectPassword123") - testUser.FailedLoginAttempts = attempt - - var updatedUser *User - mockStore.On("GetUserByEmail", ctx, "test@example.com").Return(testUser, nil).Once() - mockStore.On("UpdateUser", ctx, mock.AnythingOfType("*auth.User")). - Run(func(args mock.Arguments) { - updatedUser = args.Get(1).(*User) - }). - Return(nil).Once() - - req := LoginRequest{ - Email: "test@example.com", - Password: "WrongPassword", +func TestLogin_AccountLockout_Bookkeeping(t *testing.T) { + for _, tc := range []struct { + name string + attempts int + lockOffset time.Duration + password string + mfa bool + method string + }{ + {name: "locked", attempts: 5, lockOffset: time.Minute, password: "CorrectPassword123"}, + {name: "below threshold", attempts: 2, password: "wrong", method: "RecordFailedLogin"}, + {name: "threshold", attempts: 4, password: "wrong", method: "RecordFailedLogin"}, + {name: "expired failure", attempts: 5, lockOffset: -time.Minute, password: "wrong", method: "RecordFailedLogin"}, + {name: "expired success", attempts: 5, lockOffset: -time.Minute, password: "CorrectPassword123", method: "RecordSuccessfulLogin"}, + {name: "reset on success", attempts: 3, password: "CorrectPassword123", method: "RecordSuccessfulLogin"}, + {name: "MFA failure", attempts: 4, password: "CorrectPassword123", mfa: true, method: "RecordFailedLogin"}, + } { + t.Run(tc.name, func(t *testing.T) { + ctx := context.Background() + store := new(MockStore) + service := createTestService(store, new(MockEmailSender)) + user := createTestUser(t, "CorrectPassword123") + user.FailedLoginAttempts = tc.attempts + if tc.lockOffset != 0 { + until := time.Now().Add(tc.lockOffset) + user.LockedUntil = &until } - - _, err := service.Login(ctx, req) - assert.Error(t, err) - - // Verify attempt counter incremented - require.NotNil(t, updatedUser) - assert.Equal(t, attempt+1, updatedUser.FailedLoginAttempts) - - // Verify lockout only happens at MaxFailedLoginAttempts - if attempt+1 >= MaxFailedLoginAttempts { - assert.NotNil(t, updatedUser.LockedUntil) + user.MFAEnabled = tc.mfa + user.MFASecret = "JBSWY3DPEHPK3PXP" + store.On("GetUserByEmail", ctx, user.Email).Return(user, nil).Once() + if tc.method != "" { + store.On(tc.method, ctx, user.ID).Return(nil).Once() + } + if tc.method == "RecordSuccessfulLogin" { + store.On("CreateSession", ctx, mock.AnythingOfType("*auth.Session")).Return(nil).Once() + } + response, err := service.Login(ctx, LoginRequest{Email: user.Email, Password: tc.password, MFACode: "invalid"}) + if tc.method == "RecordSuccessfulLogin" { + require.NoError(t, err) + require.NotEmpty(t, response.Token) } else { - assert.Nil(t, updatedUser.LockedUntil) + require.Error(t, err) + assert.Nil(t, response) + if tc.mfa { + assert.ErrorIs(t, err, ErrInvalidMFACode) + } else { + assert.EqualError(t, err, genericLoginError) + } } - - mockStore.AssertExpectations(t) + store.AssertNotCalled(t, "UpdateUser", mock.Anything, mock.Anything) + store.AssertExpectations(t) }) } } -// TestLogin_AccountLockout_MFAFailure verifies MFA failures count toward lockout. -func TestLogin_AccountLockout_MFAFailure(t *testing.T) { - ctx := context.Background() - mockStore := new(MockStore) - mockEmail := new(MockEmailSender) - service := createTestService(mockStore, mockEmail) - - s := newTestService() - hash, _ := s.hashPassword("CorrectPassword123") - - testUser := &User{ - ID: "user-123", - Email: "test@example.com", - PasswordHash: hash, - Active: true, - MFAEnabled: true, - MFASecret: "JBSWY3DPEHPK3PXP", - FailedLoginAttempts: MaxFailedLoginAttempts - 1, - } - - var updatedUser *User - mockStore.On("GetUserByEmail", ctx, "test@example.com").Return(testUser, nil).Once() - mockStore.On("UpdateUser", ctx, mock.AnythingOfType("*auth.User")). - Run(func(args mock.Arguments) { - updatedUser = args.Get(1).(*User) - }). - Return(nil).Once() - - req := LoginRequest{ - Email: "test@example.com", - Password: "CorrectPassword123", // Correct password - MFACode: "000000", // Wrong MFA code - } - - _, err := service.Login(ctx, req) - assert.Error(t, err) - assert.ErrorIs(t, err, ErrInvalidMFACode) - - // Verify MFA failure incremented counter and locked account - require.NotNil(t, updatedUser) - assert.Equal(t, MaxFailedLoginAttempts, updatedUser.FailedLoginAttempts) - assert.NotNil(t, updatedUser.LockedUntil) - - mockStore.AssertExpectations(t) -} - -// TestLogin_AccountLockout_GenericErrorMessage verifies no information leakage. -func TestLogin_AccountLockout_GenericErrorMessage(t *testing.T) { - ctx := context.Background() - mockStore := new(MockStore) - mockEmail := new(MockEmailSender) - service := createTestService(mockStore, mockEmail) - - testUser := createTestUser(t, "CorrectPassword123") - lockUntil := time.Now().Add(10 * time.Minute) - testUser.LockedUntil = &lockUntil - - mockStore.On("GetUserByEmail", ctx, "test@example.com").Return(testUser, nil).Once() - - req := LoginRequest{ - Email: "test@example.com", - Password: "CorrectPassword123", +func TestLogin_BookkeepingErrors(t *testing.T) { + for _, tc := range []struct { + name, password, method string + sessionError bool + }{ + {name: "failure write", password: "wrong", method: "RecordFailedLogin"}, + {name: "success write", password: "CorrectPassword123", method: "RecordSuccessfulLogin"}, + {name: "session creation", password: "CorrectPassword123", sessionError: true}, + } { + t.Run(tc.name, func(t *testing.T) { + ctx := context.Background() + store := new(MockStore) + service := createTestService(store, new(MockEmailSender)) + user := createTestUser(t, "CorrectPassword123") + store.On("GetUserByEmail", ctx, user.Email).Return(user, nil).Once() + if tc.method != "" { + store.On(tc.method, ctx, user.ID).Return(assert.AnError).Once() + } + var issued *Session + if tc.password == "CorrectPassword123" { + var sessionErr error + if tc.sessionError { + sessionErr = assert.AnError + } + store.On("CreateSession", ctx, mock.AnythingOfType("*auth.Session")).Run(func(args mock.Arguments) { issued = args.Get(1).(*Session) }).Return(sessionErr).Once() + } + response, err := service.Login(ctx, LoginRequest{Email: user.Email, Password: tc.password}) + if tc.method == "RecordSuccessfulLogin" { + require.NoError(t, err) + require.NotNil(t, response) + require.NotNil(t, issued) + store.On("GetSession", ctx, hashSessionToken(response.Token)).Return(issued, nil).Once() + store.On("GetUserByID", ctx, user.ID).Return(user, nil).Once() + _, err = service.ValidateSession(ctx, response.Token) + require.NoError(t, err) + } else { + require.Error(t, err) + assert.Nil(t, response) + if tc.method == "RecordFailedLogin" { + assert.EqualError(t, err, genericLoginError) + } + } + store.AssertNotCalled(t, "UpdateUser", mock.Anything, mock.Anything) + store.AssertExpectations(t) + }) } - - _, err := service.Login(ctx, req) - assert.Error(t, err) - - // Error message should be generic to prevent user enumeration - // Should NOT reveal that account is locked - assert.Equal(t, "Check your email address and password and try again", err.Error()) - - mockStore.AssertExpectations(t) -} - -// TestRecordFailedLogin verifies recordFailedLogin function behavior. -func TestRecordFailedLogin(t *testing.T) { - ctx := context.Background() - - t.Run("increments counter below threshold", func(t *testing.T) { - mockStore := new(MockStore) - mockEmail := new(MockEmailSender) - service := createTestService(mockStore, mockEmail) - - testUser := createTestUser(t, "password") - testUser.FailedLoginAttempts = 2 - - mockStore.On("UpdateUser", ctx, mock.AnythingOfType("*auth.User")).Return(nil).Once() - - service.recordFailedLogin(ctx, testUser) - - assert.Equal(t, 3, testUser.FailedLoginAttempts) - assert.Nil(t, testUser.LockedUntil) - - mockStore.AssertExpectations(t) - }) - - t.Run("locks account at threshold", func(t *testing.T) { - mockStore := new(MockStore) - mockEmail := new(MockEmailSender) - service := createTestService(mockStore, mockEmail) - - testUser := createTestUser(t, "password") - testUser.FailedLoginAttempts = MaxFailedLoginAttempts - 1 - - mockStore.On("UpdateUser", ctx, mock.AnythingOfType("*auth.User")).Return(nil).Once() - - service.recordFailedLogin(ctx, testUser) - - assert.Equal(t, MaxFailedLoginAttempts, testUser.FailedLoginAttempts) - assert.NotNil(t, testUser.LockedUntil) - assert.True(t, testUser.LockedUntil.After(time.Now())) - - mockStore.AssertExpectations(t) - }) - - t.Run("handles update error gracefully", func(t *testing.T) { - mockStore := new(MockStore) - mockEmail := new(MockEmailSender) - service := createTestService(mockStore, mockEmail) - - testUser := createTestUser(t, "password") - - mockStore.On("UpdateUser", ctx, mock.AnythingOfType("*auth.User")).Return(assert.AnError).Once() - - // Should not panic even if update fails - service.recordFailedLogin(ctx, testUser) - - mockStore.AssertExpectations(t) - }) } // TestLogin_OWASPEnumerationInvariant guards against regression that re-introduces @@ -435,7 +190,7 @@ func TestLogin_OWASPEnumerationInvariant(t *testing.T) { user := sc.getUser(t) mockStore.On("GetUserByEmail", ctx, "test@example.com").Return(user, sc.storeError).Once() - mockStore.On("UpdateUser", ctx, mock.AnythingOfType("*auth.User")).Return(nil).Maybe() + mockStore.On("RecordFailedLogin", ctx, mock.AnythingOfType("string")).Return(nil).Maybe() req := LoginRequest{ Email: "test@example.com", diff --git a/internal/auth/service_login_db_test.go b/internal/auth/service_login_db_test.go new file mode 100644 index 00000000..0774bbf2 --- /dev/null +++ b/internal/auth/service_login_db_test.go @@ -0,0 +1,346 @@ +//go:build integration + +package auth + +import ( + "context" + "fmt" + "sync" + "testing" + "time" + + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/require" +) + +type loginReadBarrierStore struct { + StoreInterface + afterRead func(context.Context, *User) error +} + +func (s *loginReadBarrierStore) GetUserByEmail(ctx context.Context, email string) (*User, error) { + u, err := s.StoreInterface.GetUserByEmail(ctx, email) + if err == nil { + err = s.afterRead(ctx, u) + } + return u, err +} + +func TestIntegration_LoginPreservesSecurityChanges(t *testing.T) { + store := NewPostgresStore(setupAuthTestDB(t)) + ctx := t.Context() + for _, scenario := range []string{"wrong-password", "wrong-mfa", "password-success", "totp-success"} { + t.Run(scenario, func(t *testing.T) { + service := NewService(ServiceConfig{Store: store}) + service.bcryptCostOverride = 4 + hash, err := service.hashPassword("OriginalPassword123!") + require.NoError(t, err) + user := &User{Email: scenario + "@example.com", PasswordHash: hash, Active: true, + GroupIDs: []string{DefaultPurchaserGroupID}, MFASecret: "JBSWY3DPEHPK3PXP", + MFAEnabled: scenario == "wrong-mfa" || scenario == "totp-success"} + require.NoError(t, store.CreateUser(ctx, user)) + var changed *User + service.store = &loginReadBarrierStore{StoreInterface: store, afterRead: func(ctx context.Context, snapshot *User) error { + current, readErr := store.GetUserByID(ctx, snapshot.ID) + if readErr != nil { + return readErr + } + stamp := time.Now().UTC().Truncate(time.Microsecond) + current.PasswordHash, current.Salt = "new-password-hash", "new-salt" + current.Email = "changed-" + current.Email + current.MFAEnabled, current.MFASecret = !current.MFAEnabled, "changed-secret" + current.MFAPendingSecret, current.MFAPendingSecretExpiresAt = "pending-secret", &stamp + current.MFARecoveryCodes = []string{"new-code"} + current.GroupIDs = []string{DefaultAdminGroupID} + current.Active, current.DeactivatedAt = false, &stamp + current.PasswordResetToken, current.PasswordResetExpiry = "new-reset-token", &stamp + current.PasswordHistory = []string{"new-history"} + if updateErr := store.UpdateUser(ctx, current); updateErr != nil { + return updateErr + } + changed, readErr = store.GetUserByID(ctx, current.ID) + return readErr + }} + request := LoginRequest{Email: user.Email, Password: "OriginalPassword123!"} + success := scenario == "password-success" || scenario == "totp-success" + if scenario == "wrong-password" { + request.Password = "incorrect" + } + if scenario == "wrong-mfa" { + request.MFACode = "not-a-code" + } + if scenario == "totp-success" { + request.MFACode = generateTOTP(user.MFASecret, time.Now().Unix()/30) + } + response, loginErr := service.Login(ctx, request) + if success { + require.NoError(t, loginErr) + require.NotEmpty(t, response.Token) + } else { + require.Error(t, loginErr) + require.Nil(t, response) + } + require.NotNil(t, changed) + stored, err := store.GetUserByID(ctx, user.ID) + require.NoError(t, err) + if success { + assert.Zero(t, stored.FailedLoginAttempts) + require.NotNil(t, stored.LastLoginAt) + } else { + assert.Equal(t, 1, stored.FailedLoginAttempts) + assert.Nil(t, stored.LastLoginAt) + } + stored.UpdatedAt, stored.LastLoginAt = changed.UpdatedAt, changed.LastLoginAt + stored.FailedLoginAttempts, stored.LockedUntil = changed.FailedLoginAttempts, changed.LockedUntil + assert.Equal(t, changed, stored, "login must preserve every non-bookkeeping column") + }) + } +} + +func TestIntegration_LoginConcurrentFailures(t *testing.T) { + store := NewPostgresStore(setupAuthTestDB(t)) + ctx, cancel := context.WithTimeout(t.Context(), 10*time.Second) + var workers sync.WaitGroup + defer func() { cancel(); workers.Wait() }() + service := NewService(ServiceConfig{Store: store}) + service.bcryptCostOverride = 4 + hash, err := service.hashPassword("OriginalPassword123!") + require.NoError(t, err) + user := &User{Email: "parallel-login@example.com", PasswordHash: hash, Active: true, GroupIDs: []string{DefaultPurchaserGroupID}} + require.NoError(t, store.CreateUser(ctx, user)) + arrived, release := make(chan struct{}, MaxFailedLoginAttempts), make(chan struct{}) + service.store = &loginReadBarrierStore{StoreInterface: store, afterRead: func(ctx context.Context, _ *User) error { + arrived <- struct{}{} + select { + case <-release: + return nil + case <-ctx.Done(): + return ctx.Err() + } + }} + results := make(chan error, MaxFailedLoginAttempts) + before := time.Now() + for range MaxFailedLoginAttempts { + workers.Add(1) + go func() { + defer workers.Done() + response, loginErr := service.Login(ctx, LoginRequest{Email: user.Email, Password: "incorrect"}) + if response != nil { + results <- fmt.Errorf("failed login returned a session") + return + } + results <- loginErr + }() + } + for range MaxFailedLoginAttempts { + select { + case <-arrived: + case <-ctx.Done(): + t.Fatal(ctx.Err()) + } + } + close(release) + for range MaxFailedLoginAttempts { + require.EqualError(t, <-results, genericLoginError) + } + stored, err := store.GetUserByID(t.Context(), user.ID) + require.NoError(t, err) + assert.Equal(t, MaxFailedLoginAttempts, stored.FailedLoginAttempts) + require.NotNil(t, stored.LockedUntil) + assert.WithinRange(t, *stored.LockedUntil, before.Add(AccountLockoutDuration), time.Now().Add(AccountLockoutDuration)) + service.store = store + response, err := service.Login(t.Context(), LoginRequest{Email: user.Email, Password: "OriginalPassword123!"}) + require.EqualError(t, err, genericLoginError) + assert.Nil(t, response) +} + +type loginWriteOrderStore struct { + StoreInterface + firstDone chan error + proceed chan struct{} + successFirst bool +} + +func TestIntegration_LoginSequentialFailures(t *testing.T) { + store := NewPostgresStore(setupAuthTestDB(t)) + service := NewService(ServiceConfig{Store: store}) + service.bcryptCostOverride = 4 + hash, err := service.hashPassword("OriginalPassword123!") + require.NoError(t, err) + user := &User{Email: "sequential-login@example.com", PasswordHash: hash, Active: true, GroupIDs: []string{DefaultPurchaserGroupID}} + require.NoError(t, store.CreateUser(t.Context(), user)) + for attempt := 1; attempt <= MaxFailedLoginAttempts; attempt++ { + response, loginErr := service.Login(t.Context(), LoginRequest{Email: user.Email, Password: "incorrect"}) + require.EqualError(t, loginErr, genericLoginError) + assert.Nil(t, response) + stored, readErr := store.GetUserByID(t.Context(), user.ID) + require.NoError(t, readErr) + assert.Equal(t, attempt, stored.FailedLoginAttempts) + assert.Equal(t, attempt == MaxFailedLoginAttempts, stored.LockedUntil != nil) + } +} + +func (s *loginWriteOrderStore) RecordFailedLogin(ctx context.Context, id string) error { + if s.successFirst { + select { + case <-s.proceed: + case <-ctx.Done(): + return ctx.Err() + } + } + err := s.StoreInterface.RecordFailedLogin(ctx, id) + if !s.successFirst { + s.firstDone <- err + select { + case <-s.proceed: + case <-ctx.Done(): + return ctx.Err() + } + } + return err +} + +func (s *loginWriteOrderStore) RecordSuccessfulLogin(ctx context.Context, id string) error { + if !s.successFirst { + select { + case <-s.proceed: + case <-ctx.Done(): + return ctx.Err() + } + } + err := s.StoreInterface.RecordSuccessfulLogin(ctx, id) + if s.successFirst { + s.firstDone <- err + select { + case <-s.proceed: + case <-ctx.Done(): + return ctx.Err() + } + } + return err +} + +func TestIntegration_LoginMixedOrdering(t *testing.T) { + store := NewPostgresStore(setupAuthTestDB(t)) + for _, initial := range []int{0, MaxFailedLoginAttempts - 1} { + for _, successFirst := range []bool{false, true} { + t.Run(fmt.Sprintf("initial-%d-success-first-%v", initial, successFirst), func(t *testing.T) { + ctx, cancel := context.WithTimeout(t.Context(), 10*time.Second) + var workers sync.WaitGroup + defer func() { cancel(); workers.Wait() }() + service := NewService(ServiceConfig{Store: store}) + service.bcryptCostOverride = 4 + hash, err := service.hashPassword("OriginalPassword123!") + require.NoError(t, err) + user := &User{Email: fmt.Sprintf("mixed-%d-%v@example.com", initial, successFirst), PasswordHash: hash, Active: true, + GroupIDs: []string{DefaultPurchaserGroupID}, FailedLoginAttempts: initial} + require.NoError(t, store.CreateUser(ctx, user)) + ordered := &loginWriteOrderStore{StoreInterface: store, successFirst: successFirst, firstDone: make(chan error, 1), proceed: make(chan struct{})} + arrived, release := make(chan struct{}, 2), make(chan struct{}) + service.store = &loginReadBarrierStore{StoreInterface: ordered, afterRead: func(ctx context.Context, _ *User) error { + arrived <- struct{}{} + select { + case <-release: + return nil + case <-ctx.Done(): + return ctx.Err() + } + }} + type result struct { + response *LoginResponse + err error + } + successResult, failureResult := make(chan result, 1), make(chan result, 1) + workers.Add(2) + go func() { + defer workers.Done() + r, e := service.Login(ctx, LoginRequest{Email: user.Email, Password: "OriginalPassword123!"}) + successResult <- result{r, e} + }() + go func() { + defer workers.Done() + r, e := service.Login(ctx, LoginRequest{Email: user.Email, Password: "wrong"}) + failureResult <- result{r, e} + }() + for range 2 { + select { + case <-arrived: + case <-ctx.Done(): + t.Fatal(ctx.Err()) + } + } + close(release) + select { + case firstErr := <-ordered.firstDone: + require.NoError(t, firstErr) + case <-ctx.Done(): + t.Fatal(ctx.Err()) + } + intermediate, err := store.GetUserByID(ctx, user.ID) + require.NoError(t, err) + if successFirst { + assert.Zero(t, intermediate.FailedLoginAttempts) + assert.Nil(t, intermediate.LockedUntil) + } else { + assert.Equal(t, initial+1, intermediate.FailedLoginAttempts) + assert.Equal(t, initial+1 >= MaxFailedLoginAttempts, intermediate.LockedUntil != nil) + } + close(ordered.proceed) + success, failure := <-successResult, <-failureResult + require.NoError(t, success.err) + require.NotNil(t, success.response) + _, err = service.ValidateSession(ctx, success.response.Token) + require.NoError(t, err) + require.EqualError(t, failure.err, genericLoginError) + assert.Nil(t, failure.response) + stored, err := store.GetUserByID(ctx, user.ID) + require.NoError(t, err) + expected := 0 + if successFirst { + expected = 1 + } + assert.Equal(t, expected, stored.FailedLoginAttempts) + assert.Nil(t, stored.LockedUntil) + require.NotNil(t, stored.LastLoginAt) + }) + } + } +} + +func TestIntegration_LoginLockoutExpiry(t *testing.T) { + store := NewPostgresStore(setupAuthTestDB(t)) + service := NewService(ServiceConfig{Store: store}) + service.bcryptCostOverride = 4 + hash, err := service.hashPassword("OriginalPassword123!") + require.NoError(t, err) + for _, success := range []bool{false, true} { + t.Run(fmt.Sprintf("success-%v", success), func(t *testing.T) { + expired := time.Now().Add(-time.Minute) + user := &User{Email: fmt.Sprintf("expired-%v@example.com", success), PasswordHash: hash, Active: true, + GroupIDs: []string{DefaultPurchaserGroupID}, FailedLoginAttempts: MaxFailedLoginAttempts, LockedUntil: &expired} + require.NoError(t, store.CreateUser(t.Context(), user)) + password := "wrong" + if success { + password = "OriginalPassword123!" + } + before := time.Now() + response, loginErr := service.Login(t.Context(), LoginRequest{Email: user.Email, Password: password}) + stored, err := store.GetUserByID(t.Context(), user.ID) + require.NoError(t, err) + if success { + require.NoError(t, loginErr) + require.NotNil(t, response) + assert.Zero(t, stored.FailedLoginAttempts) + assert.Nil(t, stored.LockedUntil) + require.NotNil(t, stored.LastLoginAt) + assert.WithinRange(t, *stored.LastLoginAt, before, time.Now()) + } else { + require.EqualError(t, loginErr, genericLoginError) + assert.Nil(t, response) + assert.Equal(t, MaxFailedLoginAttempts+1, stored.FailedLoginAttempts) + require.NotNil(t, stored.LockedUntil) + assert.WithinRange(t, *stored.LockedUntil, before.Add(AccountLockoutDuration), time.Now().Add(AccountLockoutDuration)) + } + }) + } +} diff --git a/internal/auth/service_mfa_test.go b/internal/auth/service_mfa_test.go index 24f1778c..1a7b8d3e 100644 --- a/internal/auth/service_mfa_test.go +++ b/internal/auth/service_mfa_test.go @@ -450,9 +450,9 @@ func TestLogin_WithMFA_RecoveryCode_ConsumedOnce(t *testing.T) { user.MFARecoveryCodes = []string{hash} mockStore.On("GetUserByEmail", ctx, user.Email).Return(user, nil) - // recovery-code consumption persists user, then completeSuccessfulLogin - // persists again. Both must succeed. mockStore.On("UpdateUser", ctx, mock.AnythingOfType("*auth.User")).Return(nil) + 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) resp, err := service.Login(ctx, LoginRequest{ diff --git a/internal/auth/service_test.go b/internal/auth/service_test.go index ea4f7f0a..bf7a5cd8 100644 --- a/internal/auth/service_test.go +++ b/internal/auth/service_test.go @@ -23,7 +23,7 @@ func TestService_Login(t *testing.T) { mockStore.On("GetUserByEmail", ctx, "test@example.com").Return(testUser, nil).Once() mockStore.On("CreateSession", ctx, mock.AnythingOfType("*auth.Session")).Return(nil).Once() - mockStore.On("UpdateUser", ctx, mock.AnythingOfType("*auth.User")).Return(nil).Once() + mockStore.On("RecordSuccessfulLogin", ctx, mock.AnythingOfType("string")).Return(nil).Once() req := LoginRequest{ Email: "test@example.com", @@ -67,7 +67,7 @@ func TestService_Login(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).Maybe() + mockStore.On("RecordFailedLogin", ctx, mock.AnythingOfType("string")).Return(nil).Maybe() req := LoginRequest{ Email: "test@example.com", @@ -91,7 +91,7 @@ func TestService_Login(t *testing.T) { testUser.Active = false mockStore.On("GetUserByEmail", ctx, "test@example.com").Return(testUser, nil).Once() - mockStore.On("UpdateUser", ctx, mock.AnythingOfType("*auth.User")).Return(nil).Maybe() + mockStore.On("RecordFailedLogin", ctx, mock.AnythingOfType("string")).Return(nil).Maybe() req := LoginRequest{ Email: "test@example.com", @@ -422,7 +422,7 @@ func TestLogin_WithMFA(t *testing.T) { mockStore.On("GetUserByEmail", ctx, "mfa@example.com").Return(user, nil) mockStore.On("CreateSession", ctx, mock.AnythingOfType("*auth.Session")).Return(nil) - mockStore.On("UpdateUser", ctx, mock.AnythingOfType("*auth.User")).Return(nil) + mockStore.On("RecordSuccessfulLogin", ctx, mock.AnythingOfType("string")).Return(nil) req := LoginRequest{ Email: "mfa@example.com", @@ -457,7 +457,7 @@ func TestLogin_WithMFA_InvalidCode(t *testing.T) { mockStore.On("GetUserByEmail", ctx, "mfa@example.com").Return(user, nil) // Add mock for failed login recording due to invalid MFA code - mockStore.On("UpdateUser", ctx, mock.AnythingOfType("*auth.User")).Return(nil).Maybe() + mockStore.On("RecordFailedLogin", ctx, mock.AnythingOfType("string")).Return(nil).Maybe() req := LoginRequest{ Email: "mfa@example.com", @@ -556,7 +556,7 @@ func TestLogin_WithMFA_NoSecret(t *testing.T) { t.Cleanup(func() { mockStore.AssertExpectations(t) }) mockStore.On("GetUserByEmail", ctx, "mfa@example.com").Return(user, nil).Once() - mockStore.On("UpdateUser", ctx, mock.AnythingOfType("*auth.User")).Return(nil).Maybe() + mockStore.On("RecordFailedLogin", ctx, mock.AnythingOfType("string")).Return(nil).Maybe() req := LoginRequest{ Email: "mfa@example.com", @@ -726,7 +726,7 @@ func TestService_ErrorPaths(t *testing.T) { mockStore.On("GetUserByEmail", ctx, "test@example.com").Return(testUser, nil).Once() mockStore.On("CreateSession", ctx, mock.AnythingOfType("*auth.Session")).Return(nil).Once() - mockStore.On("UpdateUser", ctx, mock.AnythingOfType("*auth.User")).Return(fmt.Errorf("update error")).Once() + mockStore.On("RecordSuccessfulLogin", ctx, testUser.ID).Return(fmt.Errorf("update error")).Once() req := LoginRequest{ Email: "test@example.com", @@ -839,7 +839,7 @@ func TestService_Login_RFC5322DisplayName(t *testing.T) { // Store should be called with the bare address only, not the display-name form mockStore.On("GetUserByEmail", ctx, "test@example.com").Return(testUser, nil).Once() mockStore.On("CreateSession", ctx, mock.AnythingOfType("*auth.Session")).Return(nil).Once() - mockStore.On("UpdateUser", ctx, mock.AnythingOfType("*auth.User")).Return(nil).Once() + mockStore.On("RecordSuccessfulLogin", ctx, testUser.ID).Return(nil).Once() req := LoginRequest{ Email: `"Attacker" `, diff --git a/internal/auth/service_user.go b/internal/auth/service_user.go index 50efb088..146ddcaa 100644 --- a/internal/auth/service_user.go +++ b/internal/auth/service_user.go @@ -849,18 +849,7 @@ func (s *Service) ListUsers(ctx context.Context) ([]User, error) { // recordFailedLogin increments failed login attempts and locks the account if necessary. func (s *Service) recordFailedLogin(ctx context.Context, user *User) { - user.FailedLoginAttempts++ - now := time.Now() - user.UpdatedAt = now - - if user.FailedLoginAttempts >= MaxFailedLoginAttempts { - lockUntil := now.Add(AccountLockoutDuration) - user.LockedUntil = &lockUntil - logging.Warnf("Account locked due to %d failed login attempts: id=%s (locked until %v)", - user.FailedLoginAttempts, user.ID, lockUntil) - } - - if err := s.store.UpdateUser(ctx, user); err != nil { + if err := s.store.RecordFailedLogin(ctx, user.ID); err != nil { logging.Errorf("Failed to record failed login attempt for user %s: %v", user.ID, err) } } diff --git a/internal/auth/store_postgres_login.go b/internal/auth/store_postgres_login.go new file mode 100644 index 00000000..5dc42ac0 --- /dev/null +++ b/internal/auth/store_postgres_login.go @@ -0,0 +1,39 @@ +package auth + +import ( + "context" + "fmt" +) + +func (s *PostgresStore) RecordFailedLogin(ctx context.Context, userID string) error { + result, err := s.db.Exec(ctx, ` + UPDATE users SET + failed_login_attempts = failed_login_attempts + 1, + locked_until = CASE WHEN failed_login_attempts + 1 >= $2 + THEN NOW() + make_interval(secs => $3) ELSE locked_until END, + updated_at = NOW() + WHERE id = $1 + `, userID, MaxFailedLoginAttempts, AccountLockoutDuration.Seconds()) + if err != nil { + return fmt.Errorf("failed to record failed login: %w", err) + } + if result.RowsAffected() == 0 { + return fmt.Errorf("user not found: %s", userID) + } + return nil +} + +func (s *PostgresStore) RecordSuccessfulLogin(ctx context.Context, userID string) error { + result, err := s.db.Exec(ctx, ` + UPDATE users SET last_login_at = NOW(), failed_login_attempts = 0, + locked_until = NULL, updated_at = NOW() + WHERE id = $1 + `, userID) + if err != nil { + return fmt.Errorf("failed to record successful login: %w", err) + } + if result.RowsAffected() == 0 { + return fmt.Errorf("user not found: %s", userID) + } + return nil +} diff --git a/internal/auth/store_postgres_login_test.go b/internal/auth/store_postgres_login_test.go new file mode 100644 index 00000000..3b64176c --- /dev/null +++ b/internal/auth/store_postgres_login_test.go @@ -0,0 +1,37 @@ +package auth + +import ( + "context" + "testing" + + "github.com/jackc/pgx/v5/pgconn" + "github.com/stretchr/testify/assert" + "github.com/stretchr/testify/mock" + "github.com/stretchr/testify/require" +) + +func TestPostgresStore_LoginBookkeepingErrors(t *testing.T) { + for _, success := range []bool{false, true} { + for _, databaseError := range []bool{false, true} { + db := new(MockDBConnection) + store := NewPostgresStore(db) + var execErr error + if databaseError { + execErr = assert.AnError + } + db.On("Exec", mock.Anything, mock.AnythingOfType("string"), mock.Anything). + Return(pgconn.NewCommandTag("UPDATE 0"), execErr).Once() + operation := store.RecordFailedLogin + if success { + operation = store.RecordSuccessfulLogin + } + err := operation(context.Background(), "missing-user") + if databaseError { + require.ErrorIs(t, err, assert.AnError) + } else { + require.EqualError(t, err, "user not found: missing-user") + } + db.AssertExpectations(t) + } + } +} diff --git a/internal/auth/test_helpers.go b/internal/auth/test_helpers.go index f670eaf1..edaa9e58 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) RecordFailedLogin(ctx context.Context, userID string) error { + return m.Called(ctx, userID).Error(0) +} + +func (m *MockStore) RecordSuccessfulLogin(ctx context.Context, userID string) error { + return m.Called(ctx, userID).Error(0) +} + func (m *MockStore) DeleteUser(ctx context.Context, userID string) error { args := m.Called(ctx, userID) return args.Error(0) diff --git a/internal/mocks/stores.go b/internal/mocks/stores.go index 70a3a9cd..e7de8e0f 100644 --- a/internal/mocks/stores.go +++ b/internal/mocks/stores.go @@ -792,6 +792,14 @@ func (m *MockAuthStore) UpdateUser(ctx context.Context, user *auth.User) error { return args.Error(0) } +func (m *MockAuthStore) RecordFailedLogin(ctx context.Context, userID string) error { + return m.Called(ctx, userID).Error(0) +} + +func (m *MockAuthStore) RecordSuccessfulLogin(ctx context.Context, userID string) error { + return m.Called(ctx, userID).Error(0) +} + // DeleteUser mocks the DeleteUser operation. func (m *MockAuthStore) DeleteUser(ctx context.Context, userID string) error { args := m.Called(ctx, userID) diff --git a/internal/server/adapter_test.go b/internal/server/adapter_test.go index eb690ce9..b7a9d2b2 100644 --- a/internal/server/adapter_test.go +++ b/internal/server/adapter_test.go @@ -555,7 +555,7 @@ func TestSetupAdminThenLogin_RoundTrip(t *testing.T) { }). Return(true, nil).Once() mockStore.On("CreateSession", ctx, mock.AnythingOfType("*auth.Session")).Return(nil).Twice() - mockStore.On("UpdateUser", ctx, mock.AnythingOfType("*auth.User")).Return(nil).Maybe() + mockStore.On("RecordSuccessfulLogin", ctx, mock.AnythingOfType("string")).Return(nil).Once() setupReq := &events.LambdaFunctionURLRequest{ RequestContext: events.LambdaFunctionURLRequestContext{ diff --git a/internal/server/health_test.go b/internal/server/health_test.go index 2bfc97dd..17c52378 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) RecordFailedLogin(context.Context, string) error { + return nil +} + +func (m *mockAuthStoreForHealth) RecordSuccessfulLogin(context.Context, string) error { + return nil +} + func (m *mockAuthStoreForHealth) DeleteUser(ctx context.Context, userID string) error { return nil }