Merge pull request #18557 from BerriAI/litellm_deactivate_other_ui_token

[Feature] Expire Previous UI Session Tokens on Login
This commit is contained in:
yuneng-jiang
2026-01-01 09:11:17 -08:00
committed by GitHub
2 changed files with 329 additions and 0 deletions
+62
View File
@@ -34,6 +34,59 @@ 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.
@@ -174,6 +227,10 @@ async def authenticate_user(
)
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",
**{
@@ -260,6 +317,11 @@ async def authenticate_user(
hash_password, _password
):
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,6 +6,7 @@ to login_utils.py for better reusability.
"""
import os
from datetime import datetime, timezone, timedelta
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
@@ -21,6 +22,7 @@ from litellm.proxy._types import (
from litellm.proxy.auth.login_utils import (
LoginResult,
authenticate_user,
expire_previous_ui_session_tokens,
get_ui_credentials,
)
@@ -282,3 +284,268 @@ async def test_authenticate_user_database_required_for_admin():
finally:
if original_db_url:
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)