From dde19adde171c3262388e589b6a3f8c4b352433c Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Thu, 10 Sep 2026 05:44:30 +0000 Subject: [PATCH] fix(proxy): bound concurrent key and spend-counter DB lookups to stop prisma pool thrash (#40387) * fix(proxy): bound concurrent key and spend-counter DB lookups to stop prisma pool thrash A cache-miss burst fanned every get_key_object DB fallback and SpendCounterReseed point lookup into the prisma query-engine httpx pool at once. httpcore's request assignment is O(queued x connections) per event, so the event loop spent most of its time in pool bookkeeping and the logging worker's 20s wait_for tripped. Callers now wait on a small shared semaphore (PROXY_DB_LOOKUP_MAX_CONCURRENCY, default 25) instead of queueing inside httpcore Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(proxy): type the in-flight counting prisma fake Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor(proxy): drop module docstring from db_lookup_gate Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(proxy): type the in-flight counting table fake Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: yucheng Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- litellm/constants.py | 1 + litellm/proxy/auth/auth_checks.py | 62 ++++++++++--------- litellm/proxy/db/db_lookup_gate.py | 23 +++++++ litellm/proxy/db/spend_counter_reseed.py | 50 ++++++++------- .../proxy/auth/test_auth_checks.py | 41 +++++++++++- .../proxy/db/test_spend_counter_reseed.py | 31 ++++++++++ 6 files changed, 154 insertions(+), 54 deletions(-) create mode 100644 litellm/proxy/db/db_lookup_gate.py diff --git a/litellm/constants.py b/litellm/constants.py index d545890dd59..0a7ef363e92 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -1682,6 +1682,7 @@ SPEND_LOG_QUEUE_POLL_INTERVAL: Final = float(os.getenv("SPEND_LOG_QUEUE_POLL_INT RESPONSES_SESSION_LOOKUP_MAX_ATTEMPTS: Final = max(1, int(os.getenv("RESPONSES_SESSION_LOOKUP_MAX_ATTEMPTS", "3"))) RESPONSES_SESSION_LOOKUP_RETRY_INTERVAL: Final = float(os.getenv("RESPONSES_SESSION_LOOKUP_RETRY_INTERVAL", "0.2")) SPEND_COUNTER_RESEED_LOCKS_MAX_SIZE: Final = int(os.getenv("SPEND_COUNTER_RESEED_LOCKS_MAX_SIZE", 10000)) +PROXY_DB_LOOKUP_MAX_CONCURRENCY: Final = max(1, int(os.getenv("PROXY_DB_LOOKUP_MAX_CONCURRENCY", "25"))) DEFAULT_CRON_JOB_LOCK_TTL_SECONDS: Final = int(os.getenv("DEFAULT_CRON_JOB_LOCK_TTL_SECONDS", 60)) # 1 minute PROXY_BUDGET_RESCHEDULER_MIN_TIME: Final = int(os.getenv("PROXY_BUDGET_RESCHEDULER_MIN_TIME", 597)) RESET_BUDGET_JOB_BATCH_SIZE: Final = max(1, int(os.getenv("RESET_BUDGET_JOB_BATCH_SIZE", "500"))) diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index a71a1993064..1efc9611fe6 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -94,6 +94,7 @@ from litellm.proxy.common_utils.user_api_key_cache import ( team_membership_auth_cache_key, team_membership_reservation_cache_key, ) +from litellm.proxy.db.db_lookup_gate import db_lookup_gate from litellm.proxy.db.exception_handler import PrismaDBExceptionHandler from litellm.proxy.guardrails.tool_name_extraction import ( TOOL_CAPABLE_CALL_TYPES, @@ -3450,36 +3451,37 @@ async def _fetch_key_object_from_db_with_reconnect( """ 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_transport_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 + async with db_lookup_gate.current(): + 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_transport_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 def jwt_key_mapping_cache_key(jwt_claim_name: str, jwt_claim_value: str) -> str: diff --git a/litellm/proxy/db/db_lookup_gate.py b/litellm/proxy/db/db_lookup_gate.py new file mode 100644 index 00000000000..2fd427687bd --- /dev/null +++ b/litellm/proxy/db/db_lookup_gate.py @@ -0,0 +1,23 @@ +import asyncio +from typing import Final + +from litellm.constants import PROXY_DB_LOOKUP_MAX_CONCURRENCY + + +class LoopBoundSemaphore: + __slots__ = ("_loop", "_semaphore", "_value") + + def __init__(self, value: int) -> None: + self._value: Final = value + self._loop: asyncio.AbstractEventLoop | None = None + self._semaphore: asyncio.Semaphore | None = None + + def current(self) -> asyncio.Semaphore: + loop: Final = asyncio.get_running_loop() + if self._semaphore is None or self._loop is not loop: + self._semaphore = asyncio.Semaphore(self._value) + self._loop = loop + return self._semaphore + + +db_lookup_gate: Final = LoopBoundSemaphore(PROXY_DB_LOOKUP_MAX_CONCURRENCY) diff --git a/litellm/proxy/db/spend_counter_reseed.py b/litellm/proxy/db/spend_counter_reseed.py index a38b8a47dbd..e35b1c8c82b 100644 --- a/litellm/proxy/db/spend_counter_reseed.py +++ b/litellm/proxy/db/spend_counter_reseed.py @@ -23,6 +23,7 @@ from litellm._logging import verbose_proxy_logger from litellm.constants import SPEND_COUNTER_RESEED_LOCKS_MAX_SIZE from litellm.litellm_core_utils.duration_parser import duration_in_seconds from litellm.proxy._types import Litellm_EntityType +from litellm.proxy.db.db_lookup_gate import db_lookup_gate from litellm.repositories.organization_repository import OrganizationRepository from litellm.repositories.table_repositories import ( BudgetWindowSpendRepository, @@ -121,30 +122,33 @@ class SpendCounterReseed: if SpendCounterReseed._is_key_or_team_window_counter(counter_key): return None try: - if counter_key.startswith("spend:key:"): - token: Final = counter_key[len("spend:key:") :] - row = await VerificationTokenRepository(prisma_client).table.find_unique(where={"token": token}) - elif counter_key.startswith("spend:team_member:"): - suffix: Final = counter_key[len("spend:team_member:") :] - if ":" not in suffix: + async with db_lookup_gate.current(): + if counter_key.startswith("spend:key:"): + token: Final = counter_key[len("spend:key:") :] + row = await VerificationTokenRepository(prisma_client).table.find_unique(where={"token": token}) + elif counter_key.startswith("spend:team_member:"): + suffix: Final = counter_key[len("spend:team_member:") :] + if ":" not in suffix: + return None + user_id, team_id = suffix.rsplit(":", 1) + row = await TeamMembershipRepository(prisma_client).table.find_unique( + where={"user_id_team_id": {"user_id": user_id, "team_id": team_id}} + ) + elif counter_key.startswith("spend:team:"): + team_id = counter_key[len("spend:team:") :] + row = await TeamRepository(prisma_client).table.find_unique(where={"team_id": team_id}) + elif counter_key.startswith("spend:user:"): + user_id = counter_key[len("spend:user:") :] + row = await UserRepository(prisma_client).table.find_unique(where={"user_id": user_id}) + elif counter_key.startswith(END_USER_COUNTER_PREFIX) or counter_key.startswith("spend:tag:"): + return None + elif counter_key.startswith("spend:org:"): + org_id: Final = counter_key[len("spend:org:") :] + row = await OrganizationRepository(prisma_client).table.find_unique( + where={"organization_id": org_id} + ) + else: return None - user_id, team_id = suffix.rsplit(":", 1) - row = await TeamMembershipRepository(prisma_client).table.find_unique( - where={"user_id_team_id": {"user_id": user_id, "team_id": team_id}} - ) - elif counter_key.startswith("spend:team:"): - team_id = counter_key[len("spend:team:") :] - row = await TeamRepository(prisma_client).table.find_unique(where={"team_id": team_id}) - elif counter_key.startswith("spend:user:"): - user_id = counter_key[len("spend:user:") :] - row = await UserRepository(prisma_client).table.find_unique(where={"user_id": user_id}) - elif counter_key.startswith(END_USER_COUNTER_PREFIX) or counter_key.startswith("spend:tag:"): - return None - elif counter_key.startswith("spend:org:"): - org_id: Final = counter_key[len("spend:org:") :] - row = await OrganizationRepository(prisma_client).table.find_unique(where={"organization_id": org_id}) - else: - return None except Exception: verbose_proxy_logger.exception("SpendCounterReseed.from_db: failed for %s", counter_key) return None diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py index 5bfef2b6445..8777e24e209 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -1,7 +1,7 @@ import asyncio import json from types import SimpleNamespace -from typing import TYPE_CHECKING, Optional +from typing import TYPE_CHECKING, Final, Optional from unittest.mock import AsyncMock, MagicMock, patch if TYPE_CHECKING: @@ -37,6 +37,7 @@ from litellm.proxy.auth.auth_checks import ( _can_object_call_vector_stores, _check_end_user_budget, _check_team_member_budget, + _fetch_key_object_from_db_with_reconnect, _get_fuzzy_user_object, _get_team_db_check, _log_budget_lookup_failure, @@ -55,6 +56,7 @@ from litellm.caching.redis_cache import RedisCache from litellm.constants import ( DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL, END_USER_RESTRICTED_REGISTRY_MAX_SIZE, + PROXY_DB_LOOKUP_MAX_CONCURRENCY, REGISTRY_ERROR_NEGATIVE_CACHE_TTL, TAG_REGISTRY_MAX_SIZE, ) @@ -564,6 +566,43 @@ async def test_get_key_object_should_raise_if_reconnect_fails_on_db_connection_e assert mock_prisma_client.get_data.await_count == 1 +class _InFlightCountingPrisma: + def __init__(self) -> None: + self.in_flight = 0 + self.max_in_flight = 0 + + async def get_data( + self, token: str, table_name: str, parent_otel_span: None, proxy_logging_obj: None + ) -> UserAPIKeyAuth: + self.in_flight += 1 + self.max_in_flight = max(self.max_in_flight, self.in_flight) + await asyncio.sleep(0.001) + self.in_flight -= 1 + return UserAPIKeyAuth(token=token) + + +@pytest.mark.asyncio +async def test_fetch_key_object_from_db_bounds_in_flight_prisma_requests(): + prisma: Final = _InFlightCountingPrisma() + burst: Final = PROXY_DB_LOOKUP_MAX_CONCURRENCY * 5 + + results: Final = await asyncio.gather( + *( + _fetch_key_object_from_db_with_reconnect( + hashed_token=f"hashed-token-{i}", + prisma_client=prisma, # pyright: ignore[reportArgumentType] # fake stands in for PrismaClient + parent_otel_span=None, + proxy_logging_obj=None, + ) + for i in range(burst) + ) + ) + + assert len(results) == burst + assert {r.token for r in results if r is not None} == {f"hashed-token-{i}" for i in range(burst)} + assert prisma.max_in_flight == PROXY_DB_LOOKUP_MAX_CONCURRENCY + + def _fake_redis_cache(): fake_redis = MagicMock() fake_redis.async_get_cache = AsyncMock(return_value=None) diff --git a/tests/test_litellm/proxy/db/test_spend_counter_reseed.py b/tests/test_litellm/proxy/db/test_spend_counter_reseed.py index 8cb3fc665eb..0361d5cfe8f 100644 --- a/tests/test_litellm/proxy/db/test_spend_counter_reseed.py +++ b/tests/test_litellm/proxy/db/test_spend_counter_reseed.py @@ -7,6 +7,7 @@ allowed to run: only when the row is missing or belongs to an older window. from __future__ import annotations +import asyncio from datetime import datetime, timedelta, timezone from types import SimpleNamespace from typing import Final @@ -14,6 +15,7 @@ from typing import Final import pytest from litellm.caching.dual_cache import DualCache +from litellm.constants import PROXY_DB_LOOKUP_MAX_CONCURRENCY from litellm.proxy.db.spend_counter_reseed import SpendCounterReseed WINDOW_START = datetime(2026, 8, 1, tzinfo=timezone.utc) @@ -42,6 +44,19 @@ class _FakeSpendLogsTable: return [{by[0]: where.get(by[0]), "_sum": {"spend": self._total}}] +class _InFlightCountingTable: + def __init__(self) -> None: + self.in_flight = 0 + self.max_in_flight = 0 + + async def find_unique(self, where: dict[str, str]) -> SimpleNamespace: + self.in_flight += 1 + self.max_in_flight = max(self.max_in_flight, self.in_flight) + await asyncio.sleep(0.001) + self.in_flight -= 1 + return SimpleNamespace(token=where["token"], spend=1.0) + + class _FakePrismaClient: def __init__( self, @@ -55,6 +70,7 @@ class _FakePrismaClient: litellm_budgetwindowspend=_FakeFindUniqueTable(row=row, error=error), litellm_spendlogs=_FakeSpendLogsTable(total=spend_logs_total), litellm_endusertable=_FakeFindUniqueTable(row=end_user_row, error=end_user_error), + litellm_verificationtoken=_InFlightCountingTable(), ) @@ -306,6 +322,21 @@ async def test_end_user_from_db_returns_none_without_a_row_a_client_or_on_db_err ) +@pytest.mark.asyncio +async def test_from_db_bounds_in_flight_prisma_requests_across_counter_keys(): + """Per-counter singleflight only collapses duplicates of one key. A cold-cache burst + over many distinct keys must still not flood the prisma engine HTTP pool (LIT-6435).""" + prisma: Final = _FakePrismaClient() + burst: Final = PROXY_DB_LOOKUP_MAX_CONCURRENCY * 5 + + results: Final = await asyncio.gather( + *(SpendCounterReseed.from_db(prisma_client=prisma, counter_key=f"spend:key:hashed-{i}") for i in range(burst)) + ) + + assert results == [1.0] * burst + assert prisma.db.litellm_verificationtoken.max_in_flight == PROXY_DB_LOOKUP_MAX_CONCURRENCY + + @pytest.mark.asyncio async def test_from_db_still_never_reads_the_end_user_row(): """A cold end-user counter keeps seeding from the cached end-user object the auth