diff --git a/litellm/constants.py b/litellm/constants.py index 084371e1dc0..4cabe444a1a 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -1659,6 +1659,7 @@ LITELLM_SETTINGS_SAFE_DB_OVERRIDES: Final = [ ] SPECIAL_LITELLM_AUTH_TOKEN: Final = ["ui-token"] DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL = int(os.getenv("DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL", 60)) +PRISMA_KEY_LOOKUP_CONCURRENCY: Final = get_env_int_in_range("PRISMA_KEY_LOOKUP_CONCURRENCY", 16, 1, 1000) DEFAULT_ACCESS_GROUP_CACHE_TTL: Final = int(os.getenv("DEFAULT_ACCESS_GROUP_CACHE_TTL", 600)) # Short TTL for negative MCP access-group existence lookups. Keeps unauthenticated # callers from forcing a DB query per request for unknown names, while bounding diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 66bbda1ca4e..d348771761e 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -3263,36 +3263,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 prisma_client.key_lookup_semaphore: + 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 @log_db_metrics diff --git a/litellm/proxy/hooks/proxy_track_cost_callback.py b/litellm/proxy/hooks/proxy_track_cost_callback.py index 6abfca1d3a0..8c61b5293e0 100644 --- a/litellm/proxy/hooks/proxy_track_cost_callback.py +++ b/litellm/proxy/hooks/proxy_track_cost_callback.py @@ -238,7 +238,10 @@ class _ProxyDBLogger(CustomLogger): if metadata.get("user_api_key") and not metadata.get("user_api_key_user_id"): metadata = await _ProxyDBLogger._enrich_failure_metadata_with_key_info( # rebind-ok: enriched metadata replaces the original metadata=metadata, - resolve_missing_key_identity=str(kwargs.get("call_type")) not in _CAPTURED_IDENTITY_CALL_TYPES, + resolve_missing_key_identity=( + "user_api_key_user_id" not in metadata + and str(kwargs.get("call_type")) not in _CAPTURED_IDENTITY_CALL_TYPES + ), ) _write_spend_metadata_to_kwargs(kwargs=kwargs, metadata=metadata) budget_reservation: Final = _get_budget_reservation_from_metadata(metadata=metadata) diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 2c571b4027b..e2d784081d2 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -26,6 +26,7 @@ from litellm.constants import ( DEFAULT_MODEL_CREATED_AT_TIME, LITELLM_LOGGING_NO_UPSTREAM_LLM_CALL, MAX_TEAM_LIST_LIMIT, + PRISMA_KEY_LOOKUP_CONCURRENCY, SPEND_LOG_QUEUE_MAX_BYTES, SPEND_LOG_WRITE_BATCH_MAX_BYTES, SPEND_LOG_WRITE_BATCH_MAX_ROWS, @@ -3493,6 +3494,7 @@ class PrismaClient: else: self.db = writer_wrapper # Client to connect to Prisma db self._db_reconnect_lock = asyncio.Lock() + self.key_lookup_semaphore: Final = asyncio.Semaphore(PRISMA_KEY_LOOKUP_CONCURRENCY) self._db_health_watchdog_task: asyncio.Task | None = 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"))) diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py index 3dea89ed67b..5aed260a97f 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: @@ -565,6 +565,60 @@ async def test_get_key_object_should_reconnect_once_on_db_connection_error(): ) +@pytest.mark.asyncio +async def test_key_lookup_limit_is_shared_by_auth_and_logging_and_survives_cancellation() -> None: + from litellm.proxy.auth.resolvers.store import IdentityStore + + started: Final = asyncio.Event() + semaphore: Final = asyncio.Semaphore(1) + + async def lookup(token: str, **_kwargs: object) -> UserAPIKeyAuth: + if token == "first-key": + started.set() + await asyncio.Event().wait() + return UserAPIKeyAuth(token=token) + + database: Final = MagicMock(key_lookup_semaphore=semaphore, get_data=AsyncMock(side_effect=lookup)) + cache: Final = UserApiKeyCache() + first: Final = asyncio.create_task( + get_key_object(hashed_token="first-key", prisma_client=database, user_api_key_cache=cache) + ) + second: Final = asyncio.create_task(IdentityStore(prisma_client=database, cache=cache).resolve("second-key")) + cancelled_waiter: Final = asyncio.create_task( + get_key_object(hashed_token="cancelled-key", prisma_client=database, user_api_key_cache=cache) + ) + tasks: Final = (first, second, cancelled_waiter) + try: + await asyncio.wait_for(started.wait(), timeout=5) + await asyncio.sleep(0) + assert database.get_data.await_count == 1 + + cancelled_waiter.cancel() + with pytest.raises(asyncio.CancelledError): + await cancelled_waiter + first.cancel() + with pytest.raises(asyncio.CancelledError): + await first + + principal: Final = await asyncio.wait_for(second, timeout=5) + assert principal.source_key is not None + assert principal.source_key.token == "second-key" + final_key: Final = await get_key_object( + hashed_token="after-cancellation", prisma_client=database, user_api_key_cache=cache + ) + assert final_key.token == "after-cancellation" + assert tuple(call.kwargs["token"] for call in database.get_data.await_args_list) == ( + "first-key", + "second-key", + "after-cancellation", + ) + assert not semaphore.locked() + finally: + for task in tasks: + task.cancel() + await asyncio.gather(*tasks, return_exceptions=True) + + @pytest.mark.asyncio async def test_get_key_object_should_raise_if_reconnect_fails_on_db_connection_error(): mock_prisma_client = MagicMock() diff --git a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py index ca517474a5c..58085c0e193 100644 --- a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py +++ b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py @@ -3,6 +3,7 @@ import pytest from datetime import datetime +from typing import Final from unittest.mock import AsyncMock, MagicMock, patch from litellm.proxy._types import UserAPIKeyAuth @@ -1319,8 +1320,8 @@ async def test_enrich_failure_metadata_ignores_flag_when_alias_present(): async def test_track_cost_callback_reads_key_only_for_in_request_logs(call_type, expect_key_read): """ The batch cost row is logged long after the batch was created, so it keeps the - identity persisted at create time. Every other call type still backfills from - the key. + identity persisted at create time. In-request logs can backfill identity + that was not captured during authentication. """ logger = _ProxyDBLogger() @@ -1339,7 +1340,6 @@ async def test_track_cost_callback_reads_key_only_for_in_request_logs(call_type, "metadata": { "user_api_key": "hashed_key", "user_api_key_alias": None, - "user_api_key_user_id": None, "user_api_key_team_id": None, "user_api_key_org_id": None, } @@ -1377,7 +1377,7 @@ async def test_track_cost_callback_reads_key_only_for_in_request_logs(call_type, assert written["user_api_key_user_id"] == "user-assigned-later" assert written["user_api_key_team_id"] == "team-assigned-later" else: - assert written["user_api_key_user_id"] is None + assert written.get("user_api_key_user_id") is None assert written["user_api_key_team_id"] is None assert written["user_api_key_org_id"] is None @@ -1875,3 +1875,54 @@ async def test_track_cost_callback_logs_unauthenticated_pass_through_request( assert mock_proxy_logging.db_spend_update_writer.update_database.await_count == ( 1 if expect_spend_log else 0 ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("metadata_key", ("metadata", "litellm_metadata")) +async def test_track_cost_callback_preserves_authenticated_null_identity(metadata_key: str) -> None: + from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup + + identity: Final = UserAPIKeyAuth(api_key="authenticated-key") + metadata: Final = { + **LiteLLMProxyRequestSetup.get_sanitized_user_information_from_key(identity), + "user_api_key": identity.api_key, + } + kwargs: Final = { + "call_type": "acompletion", + "model": "test-model", + "litellm_params": {metadata_key: metadata}, + "response_cost": 0.25, + } + with ( + patch( # test-quality-ok: [TQ008] observes the legacy callback's imported lookup boundary + "litellm.proxy.hooks.proxy_track_cost_callback.get_key_object", + new_callable=AsyncMock, + return_value=UserAPIKeyAuth(api_key=identity.api_key, user_id="reassigned-user"), + ) as lookup, + patch( # test-quality-ok: [TQ008] callback imports the counter singleton at call time + "litellm.proxy.proxy_server.increment_spend_counters", new_callable=AsyncMock + ) as increment, + patch( # test-quality-ok: [TQ008] callback imports the cache writer at call time + "litellm.proxy.proxy_server.update_cache", new_callable=AsyncMock + ), + patch( # test-quality-ok: [TQ008] callback imports the logging singleton at call time + "litellm.proxy.proxy_server.proxy_logging_obj" + ) as logging, + ): + logging.db_spend_update_writer.update_database = AsyncMock() + logging.slack_alerting_instance.customer_spend_alert = AsyncMock() + await _ProxyDBLogger()._PROXY_track_cost_callback( + kwargs=kwargs, + completion_response={"id": "authenticated-completion"}, + start_time=datetime.now(), + end_time=datetime.now(), + ) + + lookup.assert_not_awaited() + increment.assert_awaited_once() + logging.db_spend_update_writer.update_database.assert_awaited_once() + for call in (increment.await_args, logging.db_spend_update_writer.update_database.await_args): + assert call.kwargs["user_id"] is None + assert call.kwargs["team_id"] is None + assert call.kwargs["org_id"] is None + assert call.kwargs["response_cost"] == 0.25