mirror of
https://github.com/tiennm99/litellm.git
synced 2026-08-02 12:21:10 +00:00
remove key blocking
This commit is contained in:
+1
-1
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user