from __future__ import annotations import unittest from types import SimpleNamespace from unittest.mock import MagicMock, patch from fastapi import HTTPException from redis.exceptions import RedisError from govoplan_access.backend.api.v1 import auth from govoplan_access.backend.security.login_throttle import ( AttemptBucket, InMemoryLoginAttemptStore, LoginThrottle, LoginThrottleDecision, ResilientLoginAttemptStore, ) from govoplan_core.api.v1.schemas import LoginRequest class LoginThrottleTests(unittest.TestCase): def test_identity_and_client_buckets_enforce_independent_limits(self) -> None: throttle = LoginThrottle( InMemoryLoginAttemptStore(), identity_limit=2, client_limit=3, window_seconds=60, ) context = { "normalized_email": "person@example.test", "tenant_slug": "tenant-a", "client_address": "192.0.2.4", } self.assertTrue(throttle.record_failure(**context).allowed) self.assertFalse(throttle.record_failure(**context).allowed) other_identity = {**context, "normalized_email": "other@example.test"} self.assertFalse(throttle.record_failure(**other_identity).allowed) def test_success_clears_identity_bucket_without_erasing_client_failures(self) -> None: throttle = LoginThrottle( InMemoryLoginAttemptStore(), identity_limit=2, client_limit=3, window_seconds=60, ) context = { "normalized_email": "person@example.test", "tenant_slug": "tenant-a", "client_address": "192.0.2.8", } self.assertTrue(throttle.record_failure(**context).allowed) throttle.record_success( normalized_email=context["normalized_email"], tenant_slug=context["tenant_slug"], ) self.assertTrue(throttle.check(**context).allowed) other_identity = {**context, "normalized_email": "other@example.test"} self.assertTrue(throttle.record_failure(**other_identity).allowed) self.assertFalse(throttle.record_failure(**other_identity).allowed) def test_rotating_untrusted_tenant_slug_does_not_bypass_identity_limit(self) -> None: throttle = LoginThrottle( InMemoryLoginAttemptStore(), identity_limit=2, client_limit=100, window_seconds=60, ) self.assertTrue( throttle.record_failure( normalized_email="person@example.test", tenant_slug="tenant-a", client_address="192.0.2.1", ).allowed ) self.assertFalse( throttle.record_failure( normalized_email="person@example.test", tenant_slug="made-up-tenant", client_address="198.51.100.9", ).allowed ) def test_bucket_keys_do_not_contain_identity_or_client_data(self) -> None: store = MagicMock() store.increment.return_value = AttemptBucket(count=1, retry_after_seconds=60) throttle = LoginThrottle( store, identity_limit=10, client_limit=100, window_seconds=60, ) throttle.record_failure( normalized_email="private.person@example.test", tenant_slug="private-tenant", client_address="192.0.2.9", ) keys = [call.args[0] for call in store.increment.call_args_list] self.assertEqual(len(keys), 2) for key in keys: self.assertNotIn("private", key) self.assertNotIn("example", key) self.assertNotIn("192.0.2.9", key) def test_redis_failure_uses_local_store_during_retry_window(self) -> None: primary = MagicMock() primary.read.side_effect = RedisError("not available") fallback = InMemoryLoginAttemptStore() store = ResilientLoginAttemptStore(primary, fallback, retry_seconds=60) with self.assertLogs( "govoplan_access.backend.security.login_throttle", level="WARNING", ) as captured: self.assertEqual(store.read("bucket"), AttemptBucket()) result = store.increment("bucket", window_seconds=60) self.assertEqual(result.count, 1) self.assertEqual(primary.read.call_count, 1) primary.increment.assert_not_called() self.assertIn("process-local fallback", captured.output[0]) def test_redis_recovery_does_not_erase_failures_counted_by_the_fallback(self) -> None: primary = MagicMock() primary.read.side_effect = [RedisError("not available"), AttemptBucket(count=1, retry_after_seconds=30)] primary.increment.return_value = AttemptBucket(count=2, retry_after_seconds=30) fallback = InMemoryLoginAttemptStore() store = ResilientLoginAttemptStore(primary, fallback, retry_seconds=60) with self.assertLogs("govoplan_access.backend.security.login_throttle", level="WARNING"): store.read("bucket") store.increment("bucket", window_seconds=60) store.increment("bucket", window_seconds=60) store._primary_unavailable_until = 0 # Simulate the next Redis retry window. recovered = store.read("bucket") incremented = store.increment("bucket", window_seconds=60) self.assertEqual(recovered.count, 2) self.assertEqual(incremented.count, 3) def test_in_memory_store_is_bounded(self) -> None: store = InMemoryLoginAttemptStore(max_entries=2) for key in ("one", "two", "three"): store.increment(key, window_seconds=60) active = sum(store.read(key).count > 0 for key in ("one", "two", "three")) self.assertEqual(active, 2) class LoginThrottleRouteTests(unittest.TestCase): def setUp(self) -> None: self.payload = LoginRequest( email="Person@Example.Test", password="attempt", tenant_slug="tenant-a", ) self.request = SimpleNamespace(client=SimpleNamespace(host="192.0.2.10")) def test_throttled_identity_gets_same_generic_failure_detail(self) -> None: throttle = MagicMock() throttle.check.return_value = LoginThrottleDecision(False, 42) with patch.object(auth, "_configured_login_throttle", return_value=throttle): with self.assertRaises(HTTPException) as raised: auth._resolve_throttled_login_user(MagicMock(), self.payload, self.request) # type: ignore[arg-type] self.assertEqual(raised.exception.status_code, 429) self.assertEqual(raised.exception.detail, "Invalid login") self.assertEqual(raised.exception.headers, {"Retry-After": "42"}) def test_failed_login_is_counted_and_threshold_response_stays_generic(self) -> None: throttle = MagicMock() throttle.check.return_value = LoginThrottleDecision(True) throttle.record_failure.return_value = LoginThrottleDecision(False, 60) generic_failure = HTTPException(status_code=401, detail="Invalid login") with ( patch.object(auth, "_configured_login_throttle", return_value=throttle), patch.object(auth, "_resolve_login_user", side_effect=generic_failure), ): with self.assertRaises(HTTPException) as raised: auth._resolve_throttled_login_user(MagicMock(), self.payload, self.request) # type: ignore[arg-type] self.assertEqual(raised.exception.status_code, 429) self.assertEqual(raised.exception.detail, "Invalid login") throttle.record_failure.assert_called_once_with( normalized_email="person@example.test", tenant_slug="tenant-a", client_address="192.0.2.10", ) def test_success_clears_only_the_identity_bucket(self) -> None: throttle = MagicMock() throttle.check.return_value = LoginThrottleDecision(True) resolved = (object(), object(), object()) with ( patch.object(auth, "_configured_login_throttle", return_value=throttle), patch.object(auth, "_resolve_login_user", return_value=resolved), ): result = auth._resolve_throttled_login_user(MagicMock(), self.payload, self.request) # type: ignore[arg-type] self.assertEqual(result, resolved) throttle.record_failure.assert_not_called() throttle.record_success.assert_called_once_with( normalized_email="person@example.test", tenant_slug="tenant-a", ) if __name__ == "__main__": unittest.main()