mirror of
https://github.com/tiennm99/litellm.git
synced 2026-08-08 20:22:16 +00:00
fix(proxy): self-heal Prisma connection for auth and runtime (#21706)
* fix(proxy): add prisma reconnect primitive and db watchdog * fix(proxy): start and stop prisma watchdog in lifecycle * fix(auth): retry key lookup once after prisma reconnect * test(proxy): add prisma self-heal watchdog coverage * test(auth): cover reconnect-once behavior for key lookup * refactor(auth): extract db reconnect helper and remove inline import * fix(proxy): apply reconnect cooldown after attempt and add auth timeout path * fix(auth): bound reconnect latency on key lookup path * test(auth): assert reconnect timeout argument in key lookup * test(proxy): verify reconnect cooldown timestamp set after attempt * fix(proxy): harden prisma reconnect cycle semantics * test(proxy): cover watchdog reconnect + timeout budget * fix(proxy): bound watchdog probe and reconnect paths * test(proxy): cover watchdog timeout and probe behavior * fix(proxy): narrow prisma db connection error classification * fix(proxy): add auth reconnect lock timeout budget * fix(auth): pass lock timeout for db reconnect retries * test(proxy): cover narrow prisma connection error detection * test(proxy): add reconnect lock-timeout behavior coverage * test(auth): assert reconnect lock timeout argument * fix(proxy): avoid lock leak race in reconnect lock timeout path * test(proxy): cover reconnect lock-timeout race cleanup
This commit is contained in:
@@ -41,11 +41,11 @@ from litellm.proxy._types import (
|
||||
LiteLLM_ObjectPermissionTable,
|
||||
LiteLLM_OrganizationMembershipTable,
|
||||
LiteLLM_OrganizationTable,
|
||||
LiteLLM_ProjectTableCachedObj,
|
||||
LiteLLM_TagTable,
|
||||
LiteLLM_TeamMembership,
|
||||
LiteLLM_TeamTable,
|
||||
LiteLLM_TeamTableCachedObj,
|
||||
LiteLLM_ProjectTableCachedObj,
|
||||
LiteLLM_UserTable,
|
||||
LiteLLMRoutes,
|
||||
LitellmUserRoles,
|
||||
@@ -57,6 +57,7 @@ from litellm.proxy._types import (
|
||||
UserAPIKeyAuth,
|
||||
)
|
||||
from litellm.proxy.auth.route_checks import RouteChecks
|
||||
from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler
|
||||
from litellm.proxy.route_llm_request import route_request
|
||||
from litellm.proxy.utils import PrismaClient, ProxyLogging, log_db_metrics
|
||||
from litellm.router import Router
|
||||
@@ -1982,6 +1983,51 @@ class ExperimentalUIJWTToken:
|
||||
)
|
||||
|
||||
|
||||
async def _fetch_key_object_from_db_with_reconnect(
|
||||
hashed_token: str,
|
||||
prisma_client: PrismaClient,
|
||||
parent_otel_span: Optional[Span],
|
||||
proxy_logging_obj: Optional[ProxyLogging],
|
||||
) -> Optional[BaseModel]:
|
||||
"""
|
||||
Fetch key object from DB and retry once if a DB connection error can be healed.
|
||||
"""
|
||||
try:
|
||||
return await prisma_client.get_data(
|
||||
token=hashed_token,
|
||||
table_name="combined_view",
|
||||
parent_otel_span=parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
except Exception as e:
|
||||
if PrismaDBExceptionHandler.is_database_connection_error(e):
|
||||
did_reconnect = False
|
||||
if hasattr(prisma_client, "attempt_db_reconnect"):
|
||||
auth_reconnect_timeout = getattr(
|
||||
prisma_client, "_db_auth_reconnect_timeout_seconds", 2.0
|
||||
)
|
||||
if not isinstance(auth_reconnect_timeout, (int, float)):
|
||||
auth_reconnect_timeout = 2.0
|
||||
auth_reconnect_lock_timeout = getattr(
|
||||
prisma_client, "_db_auth_reconnect_lock_timeout_seconds", 0.1
|
||||
)
|
||||
if not isinstance(auth_reconnect_lock_timeout, (int, float)):
|
||||
auth_reconnect_lock_timeout = 0.1
|
||||
did_reconnect = await prisma_client.attempt_db_reconnect(
|
||||
reason="auth_get_key_object_lookup_failure",
|
||||
timeout_seconds=auth_reconnect_timeout,
|
||||
lock_timeout_seconds=auth_reconnect_lock_timeout,
|
||||
)
|
||||
if did_reconnect:
|
||||
return await prisma_client.get_data(
|
||||
token=hashed_token,
|
||||
table_name="combined_view",
|
||||
parent_otel_span=parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
raise
|
||||
|
||||
|
||||
@log_db_metrics
|
||||
async def get_key_object(
|
||||
hashed_token: str,
|
||||
@@ -2020,11 +2066,13 @@ async def get_key_object(
|
||||
)
|
||||
|
||||
# else, check db
|
||||
_valid_token: Optional[BaseModel] = await prisma_client.get_data(
|
||||
token=hashed_token,
|
||||
table_name="combined_view",
|
||||
parent_otel_span=parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
_valid_token: Optional[BaseModel] = (
|
||||
await _fetch_key_object_from_db_with_reconnect(
|
||||
hashed_token=hashed_token,
|
||||
prisma_client=prisma_client,
|
||||
parent_otel_span=parent_otel_span,
|
||||
proxy_logging_obj=proxy_logging_obj,
|
||||
)
|
||||
)
|
||||
|
||||
if _valid_token is None:
|
||||
|
||||
@@ -38,8 +38,30 @@ class PrismaDBExceptionHandler:
|
||||
|
||||
if isinstance(e, DB_CONNECTION_ERROR_TYPES):
|
||||
return True
|
||||
if isinstance(e, prisma.errors.PrismaError):
|
||||
if isinstance(
|
||||
e, (prisma.errors.ClientNotConnectedError, prisma.errors.HTTPClientClosedError)
|
||||
):
|
||||
return True
|
||||
if isinstance(e, prisma.errors.PrismaError):
|
||||
error_message = str(e).lower()
|
||||
# Treat generic PrismaError as connection error only when its text
|
||||
# clearly indicates transport/connectivity failure.
|
||||
connection_keywords = (
|
||||
"can't reach database server",
|
||||
"cannot reach database server",
|
||||
"can't connect",
|
||||
"cannot connect",
|
||||
"connection error",
|
||||
"connection closed",
|
||||
"timed out",
|
||||
"timeout",
|
||||
"connection refused",
|
||||
"network is unreachable",
|
||||
"no route to host",
|
||||
"broken pipe",
|
||||
)
|
||||
if any(keyword in error_message for keyword in connection_keywords):
|
||||
return True
|
||||
if isinstance(e, ProxyException) and e.type == ProxyErrorTypes.no_db_connection:
|
||||
return True
|
||||
return False
|
||||
|
||||
@@ -388,10 +388,10 @@ from litellm.proxy.management_endpoints.model_management_endpoints import (
|
||||
from litellm.proxy.management_endpoints.organization_endpoints import (
|
||||
router as organization_router,
|
||||
)
|
||||
from litellm.proxy.management_endpoints.policy_endpoints import router as policy_router
|
||||
from litellm.proxy.management_endpoints.project_endpoints import (
|
||||
router as project_router,
|
||||
)
|
||||
from litellm.proxy.management_endpoints.policy_endpoints import router as policy_router
|
||||
from litellm.proxy.management_endpoints.router_settings_endpoints import (
|
||||
router as router_settings_router,
|
||||
)
|
||||
@@ -902,6 +902,15 @@ async def proxy_startup_event(app: FastAPI): # noqa: PLR0915
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.error(f"Error stopping token refresh task: {e}")
|
||||
|
||||
# Shutdown event - stop Prisma DB health watchdog task
|
||||
if prisma_client is not None and hasattr(
|
||||
prisma_client, "stop_db_health_watchdog_task"
|
||||
):
|
||||
try:
|
||||
await prisma_client.stop_db_health_watchdog_task()
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.error(f"Error stopping DB health watchdog task: {e}")
|
||||
|
||||
await proxy_shutdown_event() # type: ignore[reportGeneralTypeIssues]
|
||||
|
||||
|
||||
@@ -5829,6 +5838,9 @@ class ProxyStartupEvent:
|
||||
is not True
|
||||
):
|
||||
await prisma_client.health_check()
|
||||
|
||||
if hasattr(prisma_client, "start_db_health_watchdog_task"):
|
||||
await prisma_client.start_db_health_watchdog_task()
|
||||
return prisma_client
|
||||
except Exception as e:
|
||||
PrismaDBExceptionHandler.handle_db_exception(e)
|
||||
|
||||
@@ -102,6 +102,7 @@ from litellm.proxy.db.create_views import (
|
||||
should_create_missing_views,
|
||||
)
|
||||
from litellm.proxy.db.db_spend_update_writer import DBSpendUpdateWriter
|
||||
from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler
|
||||
from litellm.proxy.db.log_db_metrics import log_db_metrics
|
||||
from litellm.proxy.db.prisma_client import PrismaWrapper
|
||||
from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import (
|
||||
@@ -2266,6 +2267,32 @@ class PrismaClient:
|
||||
else False
|
||||
),
|
||||
) # Client to connect to Prisma db
|
||||
self._db_reconnect_lock = asyncio.Lock()
|
||||
self._db_health_watchdog_task: Optional[asyncio.Task] = None
|
||||
self._db_last_reconnect_attempt_ts: float = 0.0
|
||||
self._db_reconnect_cooldown_seconds: int = max(
|
||||
1, int(os.getenv("PRISMA_RECONNECT_COOLDOWN_SECONDS", "15"))
|
||||
)
|
||||
self._db_health_watchdog_interval_seconds: int = max(
|
||||
5, int(os.getenv("PRISMA_HEALTH_WATCHDOG_INTERVAL_SECONDS", "30"))
|
||||
)
|
||||
self._db_health_watchdog_enabled: bool = (
|
||||
str_to_bool(os.getenv("PRISMA_HEALTH_WATCHDOG_ENABLED", "true")) is True
|
||||
)
|
||||
self._db_health_watchdog_probe_timeout_seconds: float = max(
|
||||
0.5,
|
||||
float(os.getenv("PRISMA_HEALTH_WATCHDOG_PROBE_TIMEOUT_SECONDS", "5.0")),
|
||||
)
|
||||
self._db_watchdog_reconnect_timeout_seconds: float = max(
|
||||
1.0, float(os.getenv("PRISMA_WATCHDOG_RECONNECT_TIMEOUT_SECONDS", "30.0"))
|
||||
)
|
||||
self._db_auth_reconnect_timeout_seconds: float = max(
|
||||
0.5, float(os.getenv("PRISMA_AUTH_RECONNECT_TIMEOUT_SECONDS", "2.0"))
|
||||
)
|
||||
self._db_auth_reconnect_lock_timeout_seconds: float = max(
|
||||
0.0,
|
||||
float(os.getenv("PRISMA_AUTH_RECONNECT_LOCK_TIMEOUT_SECONDS", "0.1")),
|
||||
)
|
||||
verbose_proxy_logger.debug("Success - Created Prisma Client")
|
||||
|
||||
def get_request_status(
|
||||
@@ -3533,6 +3560,207 @@ class PrismaClient:
|
||||
)
|
||||
raise e
|
||||
|
||||
async def _run_reconnect_cycle(
|
||||
self, timeout_seconds: Optional[float] = None
|
||||
) -> None:
|
||||
"""
|
||||
Run a reconnect cycle with direct db operations and a single overall timeout
|
||||
budget to avoid long retries on hot paths (e.g. auth).
|
||||
"""
|
||||
async def _do_direct_reconnect() -> None:
|
||||
try:
|
||||
await self.db.disconnect()
|
||||
except Exception as disconnect_err:
|
||||
verbose_proxy_logger.debug(
|
||||
"Prisma DB disconnect before reconnect failed (ignored): %s",
|
||||
disconnect_err,
|
||||
)
|
||||
|
||||
await self.db.connect()
|
||||
await self.db.query_raw("SELECT 1")
|
||||
|
||||
effective_timeout = (
|
||||
timeout_seconds
|
||||
if timeout_seconds is not None
|
||||
else self._db_watchdog_reconnect_timeout_seconds
|
||||
)
|
||||
await asyncio.wait_for(_do_direct_reconnect(), timeout=effective_timeout)
|
||||
|
||||
async def attempt_db_reconnect(
|
||||
self,
|
||||
reason: str,
|
||||
force: bool = False,
|
||||
timeout_seconds: Optional[float] = None,
|
||||
lock_timeout_seconds: Optional[float] = None,
|
||||
) -> bool:
|
||||
"""
|
||||
Attempt to reconnect the Prisma client in a singleflight manner.
|
||||
|
||||
Returns:
|
||||
bool: True if reconnection succeeded, else False.
|
||||
"""
|
||||
now = time.time()
|
||||
if (
|
||||
force is False
|
||||
and now - self._db_last_reconnect_attempt_ts
|
||||
< self._db_reconnect_cooldown_seconds
|
||||
):
|
||||
verbose_proxy_logger.debug(
|
||||
"Skipping DB reconnect attempt due to cooldown. reason=%s",
|
||||
reason,
|
||||
)
|
||||
return False
|
||||
|
||||
async def _attempt_reconnect_inside_lock() -> bool:
|
||||
now = time.time()
|
||||
if (
|
||||
force is False
|
||||
and now - self._db_last_reconnect_attempt_ts
|
||||
< self._db_reconnect_cooldown_seconds
|
||||
):
|
||||
verbose_proxy_logger.debug(
|
||||
"Skipping DB reconnect attempt inside lock due to cooldown. reason=%s",
|
||||
reason,
|
||||
)
|
||||
return False
|
||||
|
||||
verbose_proxy_logger.warning(
|
||||
"Attempting Prisma DB reconnect. reason=%s", reason
|
||||
)
|
||||
|
||||
reconnect_succeeded = False
|
||||
try:
|
||||
await self._run_reconnect_cycle(timeout_seconds=timeout_seconds)
|
||||
reconnect_succeeded = True
|
||||
verbose_proxy_logger.info(
|
||||
"Prisma DB reconnect succeeded. reason=%s", reason
|
||||
)
|
||||
except Exception as reconnect_err:
|
||||
verbose_proxy_logger.error(
|
||||
"Prisma DB reconnect failed. reason=%s error=%s",
|
||||
reason,
|
||||
reconnect_err,
|
||||
)
|
||||
finally:
|
||||
# Start cooldown after reconnect attempt has completed.
|
||||
self._db_last_reconnect_attempt_ts = time.time()
|
||||
|
||||
return reconnect_succeeded
|
||||
|
||||
if lock_timeout_seconds is None:
|
||||
async with self._db_reconnect_lock:
|
||||
return await _attempt_reconnect_inside_lock()
|
||||
|
||||
lock_acquired_by_timeout_task = False
|
||||
|
||||
async def _acquire_reconnect_lock() -> bool:
|
||||
nonlocal lock_acquired_by_timeout_task
|
||||
await self._db_reconnect_lock.acquire()
|
||||
lock_acquired_by_timeout_task = True
|
||||
return True
|
||||
|
||||
acquire_task = asyncio.create_task(_acquire_reconnect_lock())
|
||||
done, _pending = await asyncio.wait(
|
||||
{acquire_task},
|
||||
timeout=lock_timeout_seconds,
|
||||
return_when=asyncio.FIRST_COMPLETED,
|
||||
)
|
||||
if acquire_task not in done:
|
||||
acquire_task.cancel()
|
||||
try:
|
||||
await acquire_task
|
||||
except asyncio.CancelledError:
|
||||
pass
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
# Defensive cleanup for timeout/cancel race on Python 3.9-3.11.
|
||||
if lock_acquired_by_timeout_task:
|
||||
try:
|
||||
self._db_reconnect_lock.release()
|
||||
except RuntimeError:
|
||||
pass
|
||||
verbose_proxy_logger.debug(
|
||||
"Skipping DB reconnect attempt due to lock acquisition timeout. reason=%s timeout=%ss",
|
||||
reason,
|
||||
lock_timeout_seconds,
|
||||
)
|
||||
return False
|
||||
|
||||
try:
|
||||
acquire_task.result()
|
||||
except Exception as lock_acquire_err:
|
||||
verbose_proxy_logger.debug(
|
||||
"Skipping DB reconnect attempt due to lock acquisition error. reason=%s error=%s",
|
||||
reason,
|
||||
lock_acquire_err,
|
||||
)
|
||||
return False
|
||||
|
||||
try:
|
||||
return await _attempt_reconnect_inside_lock()
|
||||
finally:
|
||||
self._db_reconnect_lock.release()
|
||||
|
||||
async def start_db_health_watchdog_task(self) -> None:
|
||||
"""
|
||||
Start a background task that probes DB health and attempts reconnect on failure.
|
||||
"""
|
||||
if self._db_health_watchdog_enabled is not True:
|
||||
verbose_proxy_logger.debug(
|
||||
"Prisma DB health watchdog disabled via PRISMA_HEALTH_WATCHDOG_ENABLED"
|
||||
)
|
||||
return
|
||||
if self._db_health_watchdog_task is not None:
|
||||
return
|
||||
self._db_health_watchdog_task = asyncio.create_task(
|
||||
self._db_health_watchdog_loop()
|
||||
)
|
||||
verbose_proxy_logger.info(
|
||||
"Started Prisma DB health watchdog (interval=%ss, reconnect_cooldown=%ss, probe_timeout=%ss, reconnect_timeout=%ss)",
|
||||
self._db_health_watchdog_interval_seconds,
|
||||
self._db_reconnect_cooldown_seconds,
|
||||
self._db_health_watchdog_probe_timeout_seconds,
|
||||
self._db_watchdog_reconnect_timeout_seconds,
|
||||
)
|
||||
|
||||
async def stop_db_health_watchdog_task(self) -> None:
|
||||
"""
|
||||
Stop DB health watchdog task gracefully.
|
||||
"""
|
||||
if self._db_health_watchdog_task is None:
|
||||
return
|
||||
self._db_health_watchdog_task.cancel()
|
||||
try:
|
||||
await self._db_health_watchdog_task
|
||||
except asyncio.CancelledError:
|
||||
pass
|
||||
self._db_health_watchdog_task = None
|
||||
verbose_proxy_logger.info("Stopped Prisma DB health watchdog")
|
||||
|
||||
async def _db_health_watchdog_loop(self) -> None:
|
||||
while True:
|
||||
try:
|
||||
await asyncio.sleep(self._db_health_watchdog_interval_seconds)
|
||||
await asyncio.wait_for(
|
||||
self.db.query_raw("SELECT 1"),
|
||||
timeout=self._db_health_watchdog_probe_timeout_seconds,
|
||||
)
|
||||
except asyncio.CancelledError:
|
||||
break
|
||||
except Exception as e:
|
||||
if isinstance(
|
||||
e, asyncio.TimeoutError
|
||||
) or PrismaDBExceptionHandler.is_database_connection_error(e):
|
||||
await self.attempt_db_reconnect(
|
||||
reason="db_health_watchdog_connection_error",
|
||||
timeout_seconds=self._db_watchdog_reconnect_timeout_seconds,
|
||||
)
|
||||
else:
|
||||
verbose_proxy_logger.debug(
|
||||
"Prisma DB health watchdog observed non-DB error: %s", e
|
||||
)
|
||||
|
||||
@backoff.on_exception(
|
||||
backoff.expo,
|
||||
Exception,
|
||||
|
||||
@@ -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,63 @@ 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_with(
|
||||
reason="auth_get_key_object_lookup_failure",
|
||||
timeout_seconds=2.0,
|
||||
lock_timeout_seconds=0.1,
|
||||
)
|
||||
|
||||
|
||||
@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_with(
|
||||
reason="auth_get_key_object_lookup_failure",
|
||||
timeout_seconds=2.0,
|
||||
lock_timeout_seconds=0.1,
|
||||
)
|
||||
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 +240,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")
|
||||
|
||||
|
||||
@@ -31,10 +31,28 @@ from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler
|
||||
|
||||
|
||||
# Test is_database_connection_error method
|
||||
@pytest.mark.parametrize(
|
||||
"prisma_error",
|
||||
[
|
||||
HTTPClientClosedError(),
|
||||
ClientNotConnectedError(),
|
||||
PrismaError("can't reach database server"),
|
||||
PrismaError("connection refused"),
|
||||
PrismaError("timed out while connecting"),
|
||||
],
|
||||
)
|
||||
def test_is_database_connection_error_prisma_connection_errors(prisma_error):
|
||||
"""
|
||||
Test that only Prisma connection-related errors are considered DB connection errors.
|
||||
"""
|
||||
assert PrismaDBExceptionHandler.is_database_connection_error(prisma_error) == True
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"prisma_error",
|
||||
[
|
||||
PrismaError(),
|
||||
PrismaError("validation failed on query"),
|
||||
DataError(data={"user_facing_error": {"meta": {"table": "test_table"}}}),
|
||||
UniqueViolationError(
|
||||
data={"user_facing_error": {"meta": {"table": "test_table"}}}
|
||||
@@ -52,15 +70,10 @@ from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler
|
||||
RecordNotFoundError(
|
||||
data={"user_facing_error": {"meta": {"table": "test_table"}}}
|
||||
),
|
||||
HTTPClientClosedError(),
|
||||
ClientNotConnectedError(),
|
||||
],
|
||||
)
|
||||
def test_is_database_connection_error_prisma_errors(prisma_error):
|
||||
"""
|
||||
Test that all Prisma errors are considered database connection errors
|
||||
"""
|
||||
assert PrismaDBExceptionHandler.is_database_connection_error(prisma_error) == True
|
||||
def test_is_database_connection_error_non_connection_prisma_errors(prisma_error):
|
||||
assert PrismaDBExceptionHandler.is_database_connection_error(prisma_error) == False
|
||||
|
||||
|
||||
def test_is_database_connection_generic_errors():
|
||||
|
||||
@@ -0,0 +1,276 @@
|
||||
import asyncio
|
||||
import os
|
||||
import sys
|
||||
import time
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
|
||||
import pytest
|
||||
|
||||
sys.path.insert(
|
||||
0, os.path.abspath("../../../..")
|
||||
) # Adds the parent directory to the system path
|
||||
|
||||
from litellm.proxy.utils import PrismaClient, ProxyLogging
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def mock_prisma_binary():
|
||||
"""Mock prisma.Prisma to avoid requiring generated Prisma binaries for unit tests."""
|
||||
mock_module = MagicMock()
|
||||
with patch.dict(sys.modules, {"prisma": mock_module}):
|
||||
yield
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_proxy_logging():
|
||||
proxy_logging = AsyncMock(spec=ProxyLogging)
|
||||
proxy_logging.failure_handler = AsyncMock()
|
||||
return proxy_logging
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_attempt_db_reconnect_should_succeed(mock_proxy_logging):
|
||||
client = PrismaClient(database_url="mock://test", proxy_logging_obj=mock_proxy_logging)
|
||||
client.db.disconnect = AsyncMock(return_value=None)
|
||||
client.db.connect = AsyncMock(return_value=None)
|
||||
client.db.query_raw = AsyncMock(return_value=[{"result": 1}])
|
||||
|
||||
result = await client.attempt_db_reconnect(
|
||||
reason="unit_test_reconnect_success",
|
||||
force=True,
|
||||
)
|
||||
|
||||
assert result is True
|
||||
client.db.disconnect.assert_awaited_once()
|
||||
client.db.connect.assert_awaited_once()
|
||||
client.db.query_raw.assert_awaited_once_with("SELECT 1")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_attempt_db_reconnect_should_skip_when_in_cooldown(mock_proxy_logging):
|
||||
client = PrismaClient(database_url="mock://test", proxy_logging_obj=mock_proxy_logging)
|
||||
client.db.disconnect = AsyncMock(return_value=None)
|
||||
client.db.connect = AsyncMock(return_value=None)
|
||||
client.db.query_raw = AsyncMock(return_value=[{"result": 1}])
|
||||
client._db_reconnect_cooldown_seconds = 120
|
||||
client._db_last_reconnect_attempt_ts = time.time()
|
||||
|
||||
result = await client.attempt_db_reconnect(
|
||||
reason="unit_test_reconnect_cooldown",
|
||||
force=False,
|
||||
)
|
||||
|
||||
assert result is False
|
||||
client.db.disconnect.assert_not_called()
|
||||
client.db.connect.assert_not_called()
|
||||
client.db.query_raw.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_attempt_db_reconnect_should_skip_when_lock_timeout_expires(
|
||||
mock_proxy_logging,
|
||||
):
|
||||
client = PrismaClient(database_url="mock://test", proxy_logging_obj=mock_proxy_logging)
|
||||
client.db.disconnect = AsyncMock(return_value=None)
|
||||
client.db.connect = AsyncMock(return_value=None)
|
||||
client.db.query_raw = AsyncMock(return_value=[{"result": 1}])
|
||||
|
||||
await client._db_reconnect_lock.acquire()
|
||||
try:
|
||||
result = await client.attempt_db_reconnect(
|
||||
reason="unit_test_reconnect_lock_timeout",
|
||||
force=True,
|
||||
timeout_seconds=0.1,
|
||||
lock_timeout_seconds=0.01,
|
||||
)
|
||||
finally:
|
||||
client._db_reconnect_lock.release()
|
||||
|
||||
assert result is False
|
||||
client.db.disconnect.assert_not_called()
|
||||
client.db.connect.assert_not_called()
|
||||
client.db.query_raw.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_attempt_db_reconnect_should_not_leak_lock_on_timeout_race(
|
||||
mock_proxy_logging,
|
||||
):
|
||||
client = PrismaClient(database_url="mock://test", proxy_logging_obj=mock_proxy_logging)
|
||||
client.db.disconnect = AsyncMock(return_value=None)
|
||||
client.db.connect = AsyncMock(return_value=None)
|
||||
client.db.query_raw = AsyncMock(return_value=[{"result": 1}])
|
||||
|
||||
async def _fake_wait(tasks, timeout=None, return_when=None):
|
||||
# Let the acquire task run first, then emulate a timeout response
|
||||
# from asyncio.wait to exercise timeout-race cleanup.
|
||||
await asyncio.sleep(0)
|
||||
return set(), set(tasks)
|
||||
|
||||
with patch("litellm.proxy.utils.asyncio.wait", side_effect=_fake_wait):
|
||||
result = await client.attempt_db_reconnect(
|
||||
reason="unit_test_reconnect_lock_timeout_race",
|
||||
force=True,
|
||||
timeout_seconds=0.1,
|
||||
lock_timeout_seconds=0.01,
|
||||
)
|
||||
|
||||
assert result is False
|
||||
assert client._db_reconnect_lock.locked() is False
|
||||
client.db.disconnect.assert_not_called()
|
||||
client.db.connect.assert_not_called()
|
||||
client.db.query_raw.assert_not_called()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_attempt_db_reconnect_should_set_cooldown_after_attempt(mock_proxy_logging):
|
||||
client = PrismaClient(database_url="mock://test", proxy_logging_obj=mock_proxy_logging)
|
||||
client._db_last_reconnect_attempt_ts = 0.0
|
||||
client._db_reconnect_cooldown_seconds = 10
|
||||
client.db.disconnect = AsyncMock(return_value=None)
|
||||
client.db.connect = AsyncMock(return_value=None)
|
||||
client.db.query_raw = AsyncMock(return_value=[{"result": 1}])
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.utils.time.time", side_effect=[100.0, 101.0, 150.0, 200.0]
|
||||
):
|
||||
result = await client.attempt_db_reconnect(
|
||||
reason="unit_test_cooldown_timestamp_after_attempt",
|
||||
timeout_seconds=0.1,
|
||||
)
|
||||
|
||||
assert result is True
|
||||
assert client._db_last_reconnect_attempt_ts == 200.0
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_run_reconnect_cycle_watchdog_should_use_direct_db_ops(mock_proxy_logging):
|
||||
client = PrismaClient(database_url="mock://test", proxy_logging_obj=mock_proxy_logging)
|
||||
client.disconnect = AsyncMock(side_effect=AssertionError("wrapper disconnect used"))
|
||||
client.connect = AsyncMock(side_effect=AssertionError("wrapper connect used"))
|
||||
client.db.disconnect = AsyncMock(return_value=None)
|
||||
client.db.connect = AsyncMock(return_value=None)
|
||||
client.db.query_raw = AsyncMock(return_value=[{"result": 1}])
|
||||
|
||||
await client._run_reconnect_cycle(timeout_seconds=None)
|
||||
|
||||
client.db.disconnect.assert_awaited_once()
|
||||
client.db.connect.assert_awaited_once()
|
||||
client.db.query_raw.assert_awaited_once_with("SELECT 1")
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_run_reconnect_cycle_watchdog_should_use_default_timeout_budget(
|
||||
mock_proxy_logging,
|
||||
):
|
||||
client = PrismaClient(database_url="mock://test", proxy_logging_obj=mock_proxy_logging)
|
||||
client._db_watchdog_reconnect_timeout_seconds = 0.1
|
||||
client.db.disconnect = AsyncMock(return_value=None)
|
||||
|
||||
async def _slow_connect():
|
||||
await asyncio.sleep(0.08)
|
||||
|
||||
async def _slow_query(_query: str):
|
||||
await asyncio.sleep(0.08)
|
||||
return [{"result": 1}]
|
||||
|
||||
client.db.connect = AsyncMock(side_effect=_slow_connect)
|
||||
client.db.query_raw = AsyncMock(side_effect=_slow_query)
|
||||
|
||||
with pytest.raises(asyncio.TimeoutError):
|
||||
await client._run_reconnect_cycle(timeout_seconds=None)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_run_reconnect_cycle_timeout_should_use_single_overall_budget(
|
||||
mock_proxy_logging,
|
||||
):
|
||||
client = PrismaClient(database_url="mock://test", proxy_logging_obj=mock_proxy_logging)
|
||||
client.db.disconnect = AsyncMock(return_value=None)
|
||||
|
||||
async def _slow_connect():
|
||||
await asyncio.sleep(0.08)
|
||||
|
||||
async def _slow_query(_query: str):
|
||||
await asyncio.sleep(0.08)
|
||||
return [{"result": 1}]
|
||||
|
||||
client.db.connect = AsyncMock(side_effect=_slow_connect)
|
||||
client.db.query_raw = AsyncMock(side_effect=_slow_query)
|
||||
|
||||
with pytest.raises(asyncio.TimeoutError):
|
||||
await client._run_reconnect_cycle(timeout_seconds=0.1)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_db_health_watchdog_should_trigger_reconnect_on_db_error(mock_proxy_logging):
|
||||
client = PrismaClient(database_url="mock://test", proxy_logging_obj=mock_proxy_logging)
|
||||
client.db.query_raw = AsyncMock(side_effect=Exception("db connection dropped"))
|
||||
client.attempt_db_reconnect = AsyncMock(return_value=True)
|
||||
client._db_health_watchdog_interval_seconds = 1
|
||||
client._db_watchdog_reconnect_timeout_seconds = 7.0
|
||||
client._db_health_watchdog_probe_timeout_seconds = 0.2
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.utils.asyncio.sleep",
|
||||
AsyncMock(side_effect=[None, asyncio.CancelledError()]),
|
||||
), patch(
|
||||
"litellm.proxy.db.exception_handler.PrismaDBExceptionHandler.is_database_connection_error",
|
||||
return_value=True,
|
||||
):
|
||||
await client._db_health_watchdog_loop()
|
||||
|
||||
client.attempt_db_reconnect.assert_awaited_once_with(
|
||||
reason="db_health_watchdog_connection_error",
|
||||
timeout_seconds=7.0,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_db_health_watchdog_should_trigger_reconnect_on_probe_timeout(
|
||||
mock_proxy_logging,
|
||||
):
|
||||
client = PrismaClient(database_url="mock://test", proxy_logging_obj=mock_proxy_logging)
|
||||
client.db.query_raw = AsyncMock(side_effect=asyncio.TimeoutError())
|
||||
client.attempt_db_reconnect = AsyncMock(return_value=True)
|
||||
client._db_health_watchdog_interval_seconds = 1
|
||||
client._db_watchdog_reconnect_timeout_seconds = 9.0
|
||||
client._db_health_watchdog_probe_timeout_seconds = 0.2
|
||||
|
||||
with patch(
|
||||
"litellm.proxy.utils.asyncio.sleep",
|
||||
AsyncMock(side_effect=[None, asyncio.CancelledError()]),
|
||||
), patch(
|
||||
"litellm.proxy.db.exception_handler.PrismaDBExceptionHandler.is_database_connection_error",
|
||||
return_value=False,
|
||||
):
|
||||
await client._db_health_watchdog_loop()
|
||||
|
||||
client.attempt_db_reconnect.assert_awaited_once_with(
|
||||
reason="db_health_watchdog_connection_error",
|
||||
timeout_seconds=9.0,
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_db_health_watchdog_start_stop_lifecycle(mock_proxy_logging):
|
||||
client = PrismaClient(database_url="mock://test", proxy_logging_obj=mock_proxy_logging)
|
||||
client._db_health_watchdog_enabled = True
|
||||
client._db_health_watchdog_interval_seconds = 3600
|
||||
|
||||
loop = asyncio.get_running_loop()
|
||||
dummy_task = loop.create_task(asyncio.sleep(3600))
|
||||
|
||||
def _fake_create_task(coro):
|
||||
# create_task is patched in this test, so explicitly close the incoming coroutine
|
||||
# to avoid "coroutine was never awaited" warnings.
|
||||
coro.close()
|
||||
return dummy_task
|
||||
|
||||
with patch("litellm.proxy.utils.asyncio.create_task", side_effect=_fake_create_task):
|
||||
await client.start_db_health_watchdog_task()
|
||||
assert client._db_health_watchdog_task is dummy_task
|
||||
|
||||
await client.stop_db_health_watchdog_task()
|
||||
assert client._db_health_watchdog_task is None
|
||||
assert dummy_task.cancelled() is True
|
||||
Reference in New Issue
Block a user