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 <yucheng@berri.ai>
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-09-10 05:44:30 +00:00 • committed by GitHub
parent 13837d319d
commit dde19adde1
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
6 changed files with 154 additions and 54 deletions

View file

@ -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")))

View file

@ -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:

View file

@ -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)

View file

@ -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

View file

@ -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)

View file

@ -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