From d6fd62c2e722260520e3a40a409f421e13f08f54 Mon Sep 17 00:00:00 2001 From: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sun, 26 Jul 2026 00:30:56 +0000 Subject: [PATCH] feat(proxy): cap admin login lockout state Co-Authored-By: Ishaan Jaffer <155045088+ishaan-berri@users.noreply.github.com> --- litellm/proxy/auth/login_lockout.py | 103 ++++++++++ litellm/proxy/auth/login_utils.py | 30 +++ .../proxy/auth/test_login_utils.py | 192 +++++++++++++++++- 3 files changed, 322 insertions(+), 3 deletions(-) create mode 100644 litellm/proxy/auth/login_lockout.py diff --git a/litellm/proxy/auth/login_lockout.py b/litellm/proxy/auth/login_lockout.py new file mode 100644 index 00000000000..6ea6d65eea5 --- /dev/null +++ b/litellm/proxy/auth/login_lockout.py @@ -0,0 +1,103 @@ +"""In-memory login lockout policy for the Admin UI. + +The state is local to one process and is lost when that process restarts. It is +not shared between workers, so deployments requiring distributed enforcement +must use an external rate-limiting mechanism. The tracked identity count is +capped; during a flood of distinct usernames, an entry can be evicted early, so +the cap should remain well above the number of realistic users. +""" + +import math +import time +from collections.abc import Callable +from dataclasses import dataclass +from typing import Literal + + +_MAX_TRACKED_IDENTITIES = 4096 + + +@dataclass(frozen=True, slots=True) +class LoginBlock: + """A temporary block and the number of seconds until it expires.""" + + kind: Literal["cooldown", "lockout"] + remaining_seconds: int + + +class LoginLockout: + """Track failed login attempts and determine whether an identity is blocked.""" + + __slots__ = ( + "_failures", + "_time_fn", + "_window_seconds", + "_cooldown_threshold", + "_cooldown_seconds", + "_lockout_threshold", + "_lockout_seconds", + ) + + def __init__( + self, + *, + time_fn: Callable[[], float] = time.monotonic, + window_seconds: float = 900, + cooldown_threshold: int = 3, + cooldown_seconds: float = 60, + lockout_threshold: int = 5, + lockout_seconds: float = 900, + ) -> None: + self._failures: dict[str, tuple[float, ...]] = {} + self._time_fn = time_fn + self._window_seconds = window_seconds + self._cooldown_threshold = cooldown_threshold + self._cooldown_seconds = cooldown_seconds + self._lockout_threshold = lockout_threshold + self._lockout_seconds = lockout_seconds + + @staticmethod + def normalize_username(username: str) -> str: + return username.strip().lower() + + def check(self, username: str) -> LoginBlock | None: + key = self.normalize_username(username) + now = self._time_fn() + failures = self._prune(key, now) + if len(failures) >= self._lockout_threshold: + return self._block("lockout", failures[-1], now, self._lockout_seconds) + if len(failures) >= self._cooldown_threshold: + return self._block("cooldown", failures[-1], now, self._cooldown_seconds) + return None + + def record_failure(self, username: str) -> None: + key = self.normalize_username(username) + now = self._time_fn() + failures = self._prune(key, now) + self._failures[key] = (*failures, now) + if len(self._failures) > _MAX_TRACKED_IDENTITIES: + self._failures.pop(next(iter(self._failures))) + + def clear(self, username: str) -> None: + self._failures.pop(self.normalize_username(username), None) + + def _prune(self, key: str, now: float) -> tuple[float, ...]: + cutoff = now - self._window_seconds + failures = tuple(timestamp for timestamp in self._failures.get(key, ()) if timestamp > cutoff) + if failures: + self._failures[key] = failures + else: + self._failures.pop(key, None) + return failures + + @staticmethod + def _block( + kind: Literal["cooldown", "lockout"], + latest_failure: float, + now: float, + duration: float, + ) -> LoginBlock | None: + remaining_seconds = max(0, math.ceil(duration - (now - latest_failure))) + if remaining_seconds == 0: + return None + return LoginBlock(kind=kind, remaining_seconds=remaining_seconds) diff --git a/litellm/proxy/auth/login_utils.py b/litellm/proxy/auth/login_utils.py index f35d94c986e..f4bfe62b5cb 100644 --- a/litellm/proxy/auth/login_utils.py +++ b/litellm/proxy/auth/login_utils.py @@ -24,6 +24,7 @@ from litellm.proxy._types import ( UpdateUserRequest, UserAPIKeyAuth, ) +from litellm.proxy.auth.login_lockout import LoginBlock, LoginLockout from litellm.proxy.management_endpoints.internal_user_endpoints import user_update from litellm.proxy.management_endpoints.key_management_endpoints import ( generate_key_helper_fn, @@ -42,6 +43,26 @@ from litellm.secret_managers.main import get_secret_bool from litellm.types.proxy.ui_sso import ReturnedUITokenObject +_login_lockout = LoginLockout() + + +def _raise_login_block(block: LoginBlock) -> None: + if block.kind == "lockout": + wait = max(1, (block.remaining_seconds + 59) // 60) + message = f"Account temporarily locked due to too many failed login attempts. Try again in {wait} minutes." + else: + message = ( + "Login temporarily paused due to too many failed login attempts. " + f"Try again in {block.remaining_seconds} seconds." + ) + raise ProxyException( + message=message, + type=ProxyErrorTypes.auth_error, + param="invalid_credentials", + code=401, + ) + + async def _rehash_password_if_needed(user_id: str, password: str, stored: str) -> None: """Rehash legacy password (SHA256) to scrypt on successful login.""" if stored.startswith("scrypt:"): @@ -111,6 +132,8 @@ async def authenticate_user( password: str, master_key: Optional[str], prisma_client: Optional[PrismaClient], + *, + login_lockout: LoginLockout = _login_lockout, ) -> LoginResult: """ Authenticate a user and generate an API key for UI access. @@ -140,6 +163,9 @@ async def authenticate_user( ) ui_username, ui_password = get_ui_credentials(master_key) + block = login_lockout.check(username) + if block is not None: + _raise_login_block(block) # Check if we can find the `username` in the db. On the UI, users can enter username=their email _user_row: Optional[LiteLLM_UserTable] = None @@ -239,6 +265,7 @@ async def authenticate_user( key = ExperimentalUIJWTToken.get_experimental_ui_login_jwt_auth_token(user_info) + login_lockout.clear(username) return LoginResult( user_id=user_id, key=key, @@ -293,6 +320,7 @@ async def authenticate_user( key = response["token"] # type: ignore + login_lockout.clear(username) return LoginResult( user_id=user_id, key=key, @@ -301,6 +329,7 @@ async def authenticate_user( login_method="username_password", ) else: + login_lockout.record_failure(username) raise ProxyException( message=f"Invalid credentials used to access UI.\nNot valid credentials for {username}", type=ProxyErrorTypes.auth_error, @@ -308,6 +337,7 @@ async def authenticate_user( code=401, ) else: + login_lockout.record_failure(username) raise ProxyException( message="Invalid credentials used to access UI.\nCheck 'UI_USERNAME', 'UI_PASSWORD' in .env file", type=ProxyErrorTypes.auth_error, diff --git a/tests/test_litellm/proxy/auth/test_login_utils.py b/tests/test_litellm/proxy/auth/test_login_utils.py index c589014f276..7bd8202926e 100644 --- a/tests/test_litellm/proxy/auth/test_login_utils.py +++ b/tests/test_litellm/proxy/auth/test_login_utils.py @@ -23,6 +23,192 @@ from litellm.proxy.auth.login_utils import ( authenticate_user, get_ui_credentials, ) +from litellm.proxy.auth.login_lockout import LoginLockout, _MAX_TRACKED_IDENTITIES + + +class FakeClock: + def __init__(self) -> None: + self.value = 0.0 + + def __call__(self) -> float: + return self.value + + def advance(self, seconds: float) -> None: + self.value += seconds + + +async def _authenticate_admin( + *, + password: str, + login_lockout: LoginLockout, + prisma_client: None, +) -> LoginResult: + with patch.dict( + os.environ, + { + "UI_USERNAME": "admin", + "UI_PASSWORD": "correct-password", + "DATABASE_URL": "postgresql://test:test@localhost/test", + }, + ): + with patch( + "litellm.proxy.auth.login_utils.generate_key_helper_fn", + new_callable=AsyncMock, + ) as mock_generate_key: + mock_generate_key.return_value = {"token": "test-token"} + with patch( + "litellm.proxy.auth.login_utils.user_update", + new_callable=AsyncMock, + return_value=None, + ): + with patch("litellm.proxy.auth.login_utils.get_secret_bool", return_value=False): + return await authenticate_user( + username="admin", + password=password, + master_key="master-key", + prisma_client=prisma_client, + login_lockout=login_lockout, + ) + + +@pytest.mark.asyncio +async def test_authenticate_user_cooldown_blocks_correct_password_until_expiry(): + clock = FakeClock() + login_lockout = LoginLockout(time_fn=clock) + prisma_client = None + + for _ in range(3): + with pytest.raises(ProxyException) as exc_info: + await _authenticate_admin( + password="wrong-password", + login_lockout=login_lockout, + prisma_client=prisma_client, + ) + assert exc_info.value.param == "invalid_credentials" + + with pytest.raises(ProxyException) as blocked: + await _authenticate_admin( + password="correct-password", + login_lockout=login_lockout, + prisma_client=prisma_client, + ) + + assert blocked.value.code == "401" + assert "Try again in 60 seconds" in blocked.value.message + + clock.advance(60) + result = await _authenticate_admin( + password="correct-password", + login_lockout=login_lockout, + prisma_client=prisma_client, + ) + assert result.user_id == LITELLM_PROXY_ADMIN_NAME + + +@pytest.mark.asyncio +async def test_authenticate_user_lockout_blocks_correct_password_until_expiry(): + clock = FakeClock() + login_lockout = LoginLockout(time_fn=clock) + prisma_client = None + + for _ in range(3): + with pytest.raises(ProxyException): + await _authenticate_admin( + password="wrong-password", + login_lockout=login_lockout, + prisma_client=prisma_client, + ) + clock.advance(60) + with pytest.raises(ProxyException): + await _authenticate_admin( + password="wrong-password", + login_lockout=login_lockout, + prisma_client=prisma_client, + ) + clock.advance(60) + with pytest.raises(ProxyException): + await _authenticate_admin( + password="wrong-password", + login_lockout=login_lockout, + prisma_client=prisma_client, + ) + with pytest.raises(ProxyException) as lockout: + await _authenticate_admin( + password="correct-password", + login_lockout=login_lockout, + prisma_client=prisma_client, + ) + + assert "Account temporarily locked" in lockout.value.message + assert "Try again in 15 minutes" in lockout.value.message + + clock.advance(900) + result = await _authenticate_admin( + password="correct-password", + login_lockout=login_lockout, + prisma_client=prisma_client, + ) + assert result.user_id == LITELLM_PROXY_ADMIN_NAME + + +@pytest.mark.asyncio +async def test_authenticate_user_success_clears_failed_attempts(): + clock = FakeClock() + login_lockout = LoginLockout(time_fn=clock) + prisma_client = None + + for _ in range(3): + with pytest.raises(ProxyException): + await _authenticate_admin( + password="wrong-password", + login_lockout=login_lockout, + prisma_client=prisma_client, + ) + + clock.advance(60) + await _authenticate_admin( + password="correct-password", + login_lockout=login_lockout, + prisma_client=prisma_client, + ) + + with pytest.raises(ProxyException) as invalid: + await _authenticate_admin( + password="wrong-password", + login_lockout=login_lockout, + prisma_client=prisma_client, + ) + assert invalid.value.param == "invalid_credentials" + assert "temporarily" not in invalid.value.message + + +def test_login_lockout_prunes_old_failures_and_tracks_usernames_independently(): + clock = FakeClock() + login_lockout = LoginLockout(time_fn=clock) + + for _ in range(3): + login_lockout.record_failure(" Admin@Example.com ") + login_lockout.record_failure("other@example.com") + assert login_lockout.check("admin@example.com") is not None + assert login_lockout.check("OTHER@example.com") is None + + clock.advance(901) + assert login_lockout.check("admin@example.com") is None + assert login_lockout.check("other@example.com") is None + + +def test_login_lockout_caps_identities_with_oldest_entry_eviction(): + login_lockout = LoginLockout() + + for _ in range(3): + login_lockout.record_failure("oldest@example.com") + + for index in range(_MAX_TRACKED_IDENTITIES): + login_lockout.record_failure(f"user-{index}@example.com") + + assert len(login_lockout._failures) == _MAX_TRACKED_IDENTITIES + assert login_lockout.check("oldest@example.com") is None + assert login_lockout.check(f"user-{_MAX_TRACKED_IDENTITIES - 1}@example.com") is None def test_get_ui_credentials_prefers_explicit_password(): @@ -88,7 +274,7 @@ async def test_authenticate_user_admin_login_with_ui_credentials(): "litellm.proxy.auth.login_utils.user_update", new_callable=AsyncMock, return_value=None, - ) as mock_user_update: + ): with patch( "litellm.proxy.auth.login_utils.get_secret_bool", return_value=False, @@ -146,7 +332,7 @@ async def test_authenticate_user_admin_login_with_master_key_as_password(): "litellm.proxy.auth.login_utils.user_update", new_callable=AsyncMock, return_value=None, - ) as mock_user_update: + ): with patch( "litellm.proxy.auth.login_utils.get_secret_bool", return_value=False, @@ -387,7 +573,7 @@ async def test_authenticate_user_admin_login_with_non_ascii_characters(): "litellm.proxy.auth.login_utils.user_update", new_callable=AsyncMock, return_value=None, - ) as mock_user_update: + ): with patch( "litellm.proxy.auth.login_utils.get_secret_bool", return_value=False,