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 }