mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
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:
parent
13837d319d
commit
dde19adde1
6 changed files with 154 additions and 54 deletions
|
|
@ -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")))
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
23
litellm/proxy/db/db_lookup_gate.py
Normal file
23
litellm/proxy/db/db_lookup_gate.py
Normal 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)
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue