diff --git a/litellm/__init__.py b/litellm/__init__.py index a74a79635f..112d58d49d 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -351,7 +351,7 @@ default_team_settings: Optional[List] = None max_user_budget: Optional[float] = None default_max_internal_user_budget: Optional[float] = None max_internal_user_budget: Optional[float] = None -max_ui_session_budget: Optional[float] = 10 # $10 USD budgets for UI Chat sessions +max_ui_session_budget: Optional[float] = 0.25 # $0.25 USD budgets for UI Chat sessions internal_user_budget_duration: Optional[str] = None tag_budget_config: Optional[Dict[str, "BudgetConfig"]] = None max_end_user_budget: Optional[float] = None diff --git a/litellm/proxy/auth/login_utils.py b/litellm/proxy/auth/login_utils.py index 939cfefadc..4df773dec2 100644 --- a/litellm/proxy/auth/login_utils.py +++ b/litellm/proxy/auth/login_utils.py @@ -34,59 +34,6 @@ from litellm.secret_managers.main import get_secret_bool from litellm.types.proxy.ui_sso import ReturnedUITokenObject -async def expire_previous_ui_session_tokens( - user_id: str, prisma_client: Optional[PrismaClient] -) -> None: - """ - Expire (block) all other valid UI session tokens for a user. - - This prevents accumulation of multiple valid UI session tokens that - are supposed to be short-lived test keys. Only affects keys with - team_id = "litellm-dashboard" and that haven't expired yet. - - Args: - user_id: The user ID whose previous UI session tokens should be expired - prisma_client: Database client for performing the update - """ - if prisma_client is None: - return - - try: - from datetime import datetime, timezone - - current_time = datetime.now(timezone.utc) - - # Find all unblocked AND non-expired UI session tokens for this user - ui_session_tokens = await prisma_client.db.litellm_verificationtoken.find_many( - where={ - "user_id": user_id, - "team_id": "litellm-dashboard", - "OR": [ - {"blocked": None}, # Tokens that have never been blocked (null) - {"blocked": False}, # Tokens explicitly set to not blocked - ], - "expires": {"gt": current_time}, # Only get tokens that haven't expired - } - ) - - if not ui_session_tokens: - return - - # Block all the found tokens - tokens_to_block = [token.token for token in ui_session_tokens if token.token] - - if tokens_to_block: - await prisma_client.db.litellm_verificationtoken.update_many( - where={"token": {"in": tokens_to_block}}, - data={"blocked": True} - ) - - except Exception: - # Silently fail - don't block login if cleanup fails - # This is a best-effort operation - pass - - def get_ui_credentials(master_key: Optional[str]) -> tuple[str, str]: """ Get UI username and password from environment variables or master key. @@ -227,10 +174,6 @@ async def authenticate_user( # noqa: PLR0915 ) if os.getenv("DATABASE_URL") is not None: - # Expire any previous UI session tokens for this user - await expire_previous_ui_session_tokens( - user_id=key_user_id, prisma_client=prisma_client - ) response = await generate_key_helper_fn( request_type="key", **{ @@ -317,11 +260,6 @@ async def authenticate_user( # noqa: PLR0915 password.encode("utf-8"), _password.encode("utf-8") ) or secrets.compare_digest(hash_password.encode("utf-8"), _password.encode("utf-8")): if os.getenv("DATABASE_URL") is not None: - # Expire any previous UI session tokens for this user - await expire_previous_ui_session_tokens( - user_id=user_id, prisma_client=prisma_client - ) - response = await generate_key_helper_fn( request_type="key", **{ # type: ignore diff --git a/tests/test_litellm/proxy/auth/test_login_utils.py b/tests/test_litellm/proxy/auth/test_login_utils.py index a0e29e0610..6d2a85522f 100644 --- a/tests/test_litellm/proxy/auth/test_login_utils.py +++ b/tests/test_litellm/proxy/auth/test_login_utils.py @@ -6,7 +6,6 @@ to login_utils.py for better reusability. """ import os -from datetime import datetime, timezone, timedelta from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -22,7 +21,6 @@ from litellm.proxy._types import ( from litellm.proxy.auth.login_utils import ( LoginResult, authenticate_user, - expire_previous_ui_session_tokens, get_ui_credentials, ) @@ -288,31 +286,26 @@ async def test_authenticate_user_email_case_insensitive_login(): }, ): with patch( - "litellm.proxy.auth.login_utils.expire_previous_ui_session_tokens", + "litellm.proxy.auth.login_utils.generate_key_helper_fn", new_callable=AsyncMock, - return_value=None, - ): - with patch( - "litellm.proxy.auth.login_utils.generate_key_helper_fn", - new_callable=AsyncMock, - ) as mock_generate_key: - mock_generate_key.side_effect = [ - {"token": "token-1"}, - {"token": "token-2"}, - ] + ) as mock_generate_key: + mock_generate_key.side_effect = [ + {"token": "token-1"}, + {"token": "token-2"}, + ] - result_mixed = await authenticate_user( - username=login_email_mixed_case, - password=correct_password, - master_key=master_key, - prisma_client=mock_prisma_client, - ) - result_lower = await authenticate_user( - username=stored_email, - password=correct_password, - master_key=master_key, - prisma_client=mock_prisma_client, - ) + result_mixed = await authenticate_user( + username=login_email_mixed_case, + password=correct_password, + master_key=master_key, + prisma_client=mock_prisma_client, + ) + result_lower = await authenticate_user( + username=stored_email, + password=correct_password, + master_key=master_key, + prisma_client=mock_prisma_client, + ) assert result_mixed.user_id == result_lower.user_id == "test-user-123" assert result_mixed.user_email == result_lower.user_email == stored_email @@ -363,271 +356,6 @@ async def test_authenticate_user_database_required_for_admin(): os.environ["DATABASE_URL"] = original_db_url -@pytest.mark.asyncio -async def test_expire_previous_ui_session_tokens_none_prisma_client(): - """Test that function returns early when prisma_client is None""" - await expire_previous_ui_session_tokens("test-user", None) - # Should not raise any exception - - -@pytest.mark.asyncio -async def test_expire_previous_ui_session_tokens_only_litellm_dashboard_team(): - """Test that only tokens with team_id='litellm-dashboard' are expired""" - user_id = "test-user" - current_time = datetime.now(timezone.utc) - - # Create mock tokens with proper attributes - token1 = MagicMock() - token1.token = "token1" - token1.user_id = user_id - token1.team_id = "litellm-dashboard" - token1.blocked = None - token1.expires = current_time + timedelta(hours=1) - - token2 = MagicMock() - token2.token = "token2" - token2.user_id = user_id - token2.team_id = "other-team" - token2.blocked = None - token2.expires = current_time + timedelta(hours=1) - - def mock_find_many(**kwargs): - """Mock find_many that filters tokens based on query criteria""" - where_clause = kwargs.get("where", {}) - filtered_tokens = [] - - for token in [token1, token2]: - # Check user_id match - if token.user_id != where_clause.get("user_id"): - continue - # Check team_id match - if token.team_id != where_clause.get("team_id"): - continue - # Check blocked condition (None or False) - if token.blocked is not None and token.blocked is not False: - continue - # Check expires > current_time - if token.expires <= where_clause.get("expires", {}).get("gt"): - continue - filtered_tokens.append(token) - - return filtered_tokens - - mock_prisma_client = MagicMock() - mock_prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(side_effect=mock_find_many) - mock_prisma_client.db.litellm_verificationtoken.update_many = AsyncMock() - - await expire_previous_ui_session_tokens(user_id, mock_prisma_client) - - # Should only call update_many with the litellm-dashboard token - mock_prisma_client.db.litellm_verificationtoken.update_many.assert_called_once_with( - where={"token": {"in": ["token1"]}}, - data={"blocked": True} - ) - - -@pytest.mark.asyncio -async def test_expire_previous_ui_session_tokens_blocks_null_and_false(): - """Test that tokens with blocked=None and blocked=False are both processed""" - user_id = "test-user" - current_time = datetime.now(timezone.utc) - - # Create mock tokens with proper attributes - token1 = MagicMock() - token1.token = "token1" - token1.user_id = user_id - token1.team_id = "litellm-dashboard" - token1.blocked = None - token1.expires = current_time + timedelta(hours=1) - - token2 = MagicMock() - token2.token = "token2" - token2.user_id = user_id - token2.team_id = "litellm-dashboard" - token2.blocked = False - token2.expires = current_time + timedelta(hours=1) - - token3 = MagicMock() - token3.token = "token3" - token3.user_id = user_id - token3.team_id = "litellm-dashboard" - token3.blocked = True # This should be ignored - token3.expires = current_time + timedelta(hours=1) - - def mock_find_many(**kwargs): - """Mock find_many that filters tokens based on query criteria""" - where_clause = kwargs.get("where", {}) - filtered_tokens = [] - - for token in [token1, token2, token3]: - # Check user_id match - if token.user_id != where_clause.get("user_id"): - continue - # Check team_id match - if token.team_id != where_clause.get("team_id"): - continue - # Check blocked condition (None or False) - if token.blocked is not None and token.blocked is not False: - continue - # Check expires > current_time - if token.expires <= where_clause.get("expires", {}).get("gt"): - continue - filtered_tokens.append(token) - - return filtered_tokens - - mock_prisma_client = MagicMock() - mock_prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(side_effect=mock_find_many) - mock_prisma_client.db.litellm_verificationtoken.update_many = AsyncMock() - - await expire_previous_ui_session_tokens(user_id, mock_prisma_client) - - # Should only block token1 and token2 (not token3 which is already blocked) - mock_prisma_client.db.litellm_verificationtoken.update_many.assert_called_once_with( - where={"token": {"in": ["token1", "token2"]}}, - data={"blocked": True} - ) - - -@pytest.mark.asyncio -async def test_expire_previous_ui_session_tokens_only_non_expired(): - """Test that only non-expired tokens are processed""" - user_id = "test-user" - current_time = datetime.now(timezone.utc) - - # Create mock tokens with proper attributes - token1 = MagicMock() - token1.token = "token1" - token1.user_id = user_id - token1.team_id = "litellm-dashboard" - token1.blocked = None - token1.expires = current_time + timedelta(hours=1) # Not expired - - token2 = MagicMock() - token2.token = "token2" - token2.user_id = user_id - token2.team_id = "litellm-dashboard" - token2.blocked = None - token2.expires = current_time - timedelta(hours=1) # Already expired - - def mock_find_many(**kwargs): - """Mock find_many that filters tokens based on query criteria""" - where_clause = kwargs.get("where", {}) - filtered_tokens = [] - - for token in [token1, token2]: - # Check user_id match - if token.user_id != where_clause.get("user_id"): - continue - # Check team_id match - if token.team_id != where_clause.get("team_id"): - continue - # Check blocked condition (None or False) - if token.blocked is not None and token.blocked is not False: - continue - # Check expires > current_time - if token.expires <= where_clause.get("expires", {}).get("gt"): - continue - filtered_tokens.append(token) - - return filtered_tokens - - mock_prisma_client = MagicMock() - mock_prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(side_effect=mock_find_many) - mock_prisma_client.db.litellm_verificationtoken.update_many = AsyncMock() - - await expire_previous_ui_session_tokens(user_id, mock_prisma_client) - - # Should only block the non-expired token - mock_prisma_client.db.litellm_verificationtoken.update_many.assert_called_once_with( - where={"token": {"in": ["token1"]}}, - data={"blocked": True} - ) - - -@pytest.mark.asyncio -async def test_expire_previous_ui_session_tokens_no_tokens_found(): - """Test behavior when no valid tokens are found""" - user_id = "test-user" - - mock_prisma_client = MagicMock() - mock_prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) - mock_prisma_client.db.litellm_verificationtoken.update_many = AsyncMock() - - await expire_previous_ui_session_tokens(user_id, mock_prisma_client) - - # Should not call update_many when no tokens found - mock_prisma_client.db.litellm_verificationtoken.update_many.assert_not_called() - - -@pytest.mark.asyncio -async def test_expire_previous_ui_session_tokens_filters_none_token(): - """Test that tokens with None token value are filtered out""" - user_id = "test-user" - current_time = datetime.now(timezone.utc) - - # Create mock tokens with proper attributes - token1 = MagicMock() - token1.token = "token1" - token1.user_id = user_id - token1.team_id = "litellm-dashboard" - token1.blocked = None - token1.expires = current_time + timedelta(hours=1) - - token2 = MagicMock() - token2.token = None # This should be filtered out in the token collection step - token2.user_id = user_id - token2.team_id = "litellm-dashboard" - token2.blocked = None - token2.expires = current_time + timedelta(hours=1) - - def mock_find_many(**kwargs): - """Mock find_many that filters tokens based on query criteria""" - where_clause = kwargs.get("where", {}) - filtered_tokens = [] - - for token in [token1, token2]: - # Check user_id match - if token.user_id != where_clause.get("user_id"): - continue - # Check team_id match - if token.team_id != where_clause.get("team_id"): - continue - # Check blocked condition (None or False) - if token.blocked is not None and token.blocked is not False: - continue - # Check expires > current_time - if token.expires <= where_clause.get("expires", {}).get("gt"): - continue - filtered_tokens.append(token) - - return filtered_tokens - - mock_prisma_client = MagicMock() - mock_prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(side_effect=mock_find_many) - mock_prisma_client.db.litellm_verificationtoken.update_many = AsyncMock() - - await expire_previous_ui_session_tokens(user_id, mock_prisma_client) - - # Should only block token1 (token with None value should be filtered out) - mock_prisma_client.db.litellm_verificationtoken.update_many.assert_called_once_with( - where={"token": {"in": ["token1"]}}, - data={"blocked": True} - ) - - -@pytest.mark.asyncio -async def test_expire_previous_ui_session_tokens_exception_handling(): - """Test that exceptions during token expiry are silently handled""" - user_id = "test-user" - - mock_prisma_client = MagicMock() - mock_prisma_client.db.litellm_verificationtoken.find_many = AsyncMock(side_effect=Exception("Database error")) - - # Should not raise exception despite database error - await expire_previous_ui_session_tokens(user_id, mock_prisma_client) - - @pytest.mark.asyncio async def test_authenticate_user_admin_login_with_non_ascii_characters(): """Test admin login with non-ASCII characters in password (issue #19559)""" @@ -701,6 +429,80 @@ def test_authenticate_user_non_ascii_direct_comparison(): assert result is False +@pytest.mark.asyncio +async def test_authenticate_user_multiple_logins_generate_unique_tokens(): + """Test that multiple logins for the same user each generate unique tokens. + + This test verifies that users can have multiple concurrent UI sessions. + Previous UI session tokens should NOT be expired/blocked when a new session is created. + """ + master_key = "sk-1234" + ui_username = "admin" + ui_password = "sk-1234" + + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_usertable.find_first = AsyncMock(return_value=None) + + with patch.dict( + os.environ, + { + "UI_USERNAME": ui_username, + "UI_PASSWORD": ui_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: + # Each login should generate a unique token + mock_generate_key.side_effect = [ + {"token": "session-token-1", "user_id": LITELLM_PROXY_ADMIN_NAME}, + {"token": "session-token-2", "user_id": LITELLM_PROXY_ADMIN_NAME}, + {"token": "session-token-3", "user_id": LITELLM_PROXY_ADMIN_NAME}, + ] + + 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, + ): + # Simulate multiple logins from the same user + result1 = await authenticate_user( + username=ui_username, + password=ui_password, + master_key=master_key, + prisma_client=mock_prisma_client, + ) + result2 = await authenticate_user( + username=ui_username, + password=ui_password, + master_key=master_key, + prisma_client=mock_prisma_client, + ) + result3 = await authenticate_user( + username=ui_username, + password=ui_password, + master_key=master_key, + prisma_client=mock_prisma_client, + ) + + # Each login should return a unique token + assert result1.key == "session-token-1" + assert result2.key == "session-token-2" + assert result3.key == "session-token-3" + + # All tokens should be different (concurrent sessions allowed) + assert len({result1.key, result2.key, result3.key}) == 3 + + # generate_key_helper_fn should be called 3 times (once per login) + assert mock_generate_key.call_count == 3 + + @pytest.mark.asyncio async def test_authenticate_user_database_login_with_non_ascii_password(): """Test database user login with non-ASCII characters in password (issue #19559)""" @@ -736,23 +538,18 @@ async def test_authenticate_user_database_login_with_non_ascii_password(): }, ): with patch( - "litellm.proxy.auth.login_utils.expire_previous_ui_session_tokens", + "litellm.proxy.auth.login_utils.generate_key_helper_fn", new_callable=AsyncMock, - return_value=None, - ): - with patch( - "litellm.proxy.auth.login_utils.generate_key_helper_fn", - new_callable=AsyncMock, - ) as mock_generate_key: - mock_generate_key.return_value = {"token": "token-123"} + ) as mock_generate_key: + mock_generate_key.return_value = {"token": "token-123"} - result = await authenticate_user( - username=user_email, - password=password_with_special_char, - master_key=master_key, - prisma_client=mock_prisma_client, - ) + result = await authenticate_user( + username=user_email, + password=password_with_special_char, + master_key=master_key, + prisma_client=mock_prisma_client, + ) - assert isinstance(result, LoginResult) - assert result.user_id == "test-user-123" - assert result.user_email == user_email + assert isinstance(result, LoginResult) + assert result.user_id == "test-user-123" + assert result.user_email == user_email