mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-05 02:41:56 +00:00
test(auth): cover reconnect-once behavior for key lookup
This commit is contained in:
parent
4dee7dc0b0
commit
b560164921
1 changed files with 55 additions and 2 deletions
|
|
@ -10,6 +10,7 @@ sys.path.insert(
|
|||
|
||||
from datetime import datetime, timedelta
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
|
|
@ -33,6 +34,7 @@ from litellm.proxy.auth.auth_checks import (
|
|||
_log_budget_lookup_failure,
|
||||
_virtual_key_max_budget_alert_check,
|
||||
_virtual_key_soft_budget_check,
|
||||
get_key_object,
|
||||
get_user_object,
|
||||
vector_store_access_check,
|
||||
)
|
||||
|
|
@ -50,9 +52,10 @@ def set_salt_key(monkeypatch):
|
|||
def reset_constants_module():
|
||||
"""Reset constants module to ensure clean state before each test"""
|
||||
import importlib
|
||||
|
||||
from litellm import constants
|
||||
from litellm.proxy.auth import auth_checks
|
||||
|
||||
|
||||
# Reload modules before test
|
||||
importlib.reload(constants)
|
||||
importlib.reload(auth_checks)
|
||||
|
|
@ -151,6 +154,55 @@ def test_get_key_object_from_ui_hash_key_invalid():
|
|||
assert key_object is None
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_key_object_should_reconnect_once_on_db_connection_error():
|
||||
mock_prisma_client = MagicMock()
|
||||
mock_prisma_client.get_data = AsyncMock(
|
||||
side_effect=[
|
||||
httpx.ConnectError("db connection reset"),
|
||||
UserAPIKeyAuth(token="hashed-token-1"),
|
||||
]
|
||||
)
|
||||
mock_prisma_client.attempt_db_reconnect = AsyncMock(return_value=True)
|
||||
|
||||
mock_cache = MagicMock()
|
||||
mock_cache.async_get_cache = AsyncMock(return_value=None)
|
||||
mock_cache.async_set_cache = AsyncMock()
|
||||
|
||||
key_obj = await get_key_object(
|
||||
hashed_token="hashed-token-1",
|
||||
prisma_client=mock_prisma_client,
|
||||
user_api_key_cache=mock_cache,
|
||||
)
|
||||
|
||||
assert key_obj.token == "hashed-token-1"
|
||||
assert mock_prisma_client.get_data.await_count == 2
|
||||
mock_prisma_client.attempt_db_reconnect.assert_awaited_once()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_key_object_should_raise_if_reconnect_fails_on_db_connection_error():
|
||||
mock_prisma_client = MagicMock()
|
||||
mock_prisma_client.get_data = AsyncMock(
|
||||
side_effect=httpx.ConnectError("db not reachable after outage")
|
||||
)
|
||||
mock_prisma_client.attempt_db_reconnect = AsyncMock(return_value=False)
|
||||
|
||||
mock_cache = MagicMock()
|
||||
mock_cache.async_get_cache = AsyncMock(return_value=None)
|
||||
mock_cache.async_set_cache = AsyncMock()
|
||||
|
||||
with pytest.raises(Exception, match="db not reachable after outage"):
|
||||
await get_key_object(
|
||||
hashed_token="hashed-token-2",
|
||||
prisma_client=mock_prisma_client,
|
||||
user_api_key_cache=mock_cache,
|
||||
)
|
||||
|
||||
mock_prisma_client.attempt_db_reconnect.assert_awaited_once()
|
||||
assert mock_prisma_client.get_data.await_count == 1
|
||||
|
||||
|
||||
def test_get_cli_jwt_auth_token_default_expiration(valid_sso_user_defined_values):
|
||||
"""Test generating CLI JWT token with default 24-hour expiration"""
|
||||
token = ExperimentalUIJWTToken.get_cli_jwt_auth_token(valid_sso_user_defined_values)
|
||||
|
|
@ -180,9 +232,10 @@ def test_get_cli_jwt_auth_token_custom_expiration(
|
|||
):
|
||||
"""Test generating CLI JWT token with custom expiration via environment variable"""
|
||||
import importlib
|
||||
|
||||
from litellm import constants
|
||||
from litellm.proxy.auth import auth_checks
|
||||
|
||||
|
||||
# Set custom expiration to 48 hours
|
||||
monkeypatch.setenv("LITELLM_CLI_JWT_EXPIRATION_HOURS", "48")
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue