diff --git a/litellm/proxy/spend_tracking/key_metadata_recovery.py b/litellm/proxy/spend_tracking/key_metadata_recovery.py index 6299f576c37..796f6ffd995 100644 --- a/litellm/proxy/spend_tracking/key_metadata_recovery.py +++ b/litellm/proxy/spend_tracking/key_metadata_recovery.py @@ -120,11 +120,9 @@ async def _db_or_empty( warning: str, count: int, ) -> _T | None: - from prisma.errors import PrismaError - try: return await load() - except PrismaError as e: + except Exception as e: verbose_proxy_logger.warning(warning, count, e) return None diff --git a/tests/integration/_support/proxy.py b/tests/integration/_support/proxy.py index a444b93757d..193254b9e7f 100644 --- a/tests/integration/_support/proxy.py +++ b/tests/integration/_support/proxy.py @@ -1,4 +1,4 @@ -"""Run the normal single-process CLI with the existing behavior-suite test entitlement.""" +"""Run the normal CLI with the behavior-suite test entitlement, patched at import so spawned workers inherit it.""" import signal import sys @@ -7,6 +7,10 @@ from unittest.mock import patch from litellm import run_server +patch( # test-quality-ok: route entitlement only; license validation is outside these HTTP/DB contracts + "litellm.proxy.auth.litellm_license.LicenseCheck.is_premium", return_value=True +).start() + def _exit_on_reraised_term(signum: int, frame: FrameType | None) -> None: sys.exit(0) @@ -14,10 +18,7 @@ def _exit_on_reraised_term(signum: int, frame: FrameType | None) -> None: def main() -> None: signal.signal(signal.SIGTERM, _exit_on_reraised_term) - with patch( # test-quality-ok: route entitlement only; license validation is outside these HTTP/DB contracts - "litellm.proxy.auth.litellm_license.LicenseCheck.is_premium", return_value=True - ): - run_server() + run_server() if __name__ == "__main__": diff --git a/tests/integration/spend/test_daily_activity_key_metadata_query_timeout.py b/tests/integration/spend/test_daily_activity_key_metadata_query_timeout.py index 9e684ebc317..d037909545f 100644 --- a/tests/integration/spend/test_daily_activity_key_metadata_query_timeout.py +++ b/tests/integration/spend/test_daily_activity_key_metadata_query_timeout.py @@ -331,12 +331,16 @@ def _spender(proxy: Gateway, database_url: str, label: str) -> Spender: return Spender(alias, user_id, user_email, digest) -def _read(pinned: Pinned, digest: str) -> httpx.Response: - response: Final = pinned.request( +def _fetch(pinned: Pinned, digest: str) -> httpx.Response: + return pinned.request( "GET", "/user/daily/activity/aggregated", params={"start_date": _day(-1), "end_date": _day(1), "api_key": digest}, ) + + +def _read(pinned: Pinned, digest: str) -> httpx.Response: + response: Final = _fetch(pinned, digest) assert response.status_code == 200, response.text return response @@ -358,6 +362,14 @@ def _named(pinned: Pinned, spender: Spender) -> httpx.Response: ) +def _named_once_reconnected(pinned: Pinned, spender: Spender) -> httpx.Response: + return eventually( + lambda: _fetch(pinned, spender.digest), + lambda response: response.status_code == 200 and _aliases(response, spender.digest) == (spender.alias,), + seconds=MISS_TTL_BOUND, + ) + + @contextmanager def _locked_spend_logs(database_url: str) -> Generator[None]: with psycopg.connect(database_url) as connection: @@ -884,7 +896,7 @@ def test_usage_page_survives_a_dropped_database_connection_during_alias_recovery dropped: Final = _read(pinned, spender.digest) assert relay.tripped.is_set(), dropped.text assert _aliases(dropped, spender.digest) == (None,), dropped.text - recovered: Final = _named(pinned, spender) + recovered: Final = _named_once_reconnected(pinned, spender) assert _metadata(recovered, spender.digest) == ( _KeyMetadata(key_alias=spender.alias, user_id=spender.user_id, user_email=spender.user_email), ), recovered.text diff --git a/tests/test_litellm/proxy/spend_tracking/test_key_metadata_recovery.py b/tests/test_litellm/proxy/spend_tracking/test_key_metadata_recovery.py index d96156a9d10..c52d11e4bb0 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_key_metadata_recovery.py +++ b/tests/test_litellm/proxy/spend_tracking/test_key_metadata_recovery.py @@ -471,6 +471,34 @@ async def test_recover_key_metadata_from_spend_logs_never_caches_repeated_query_ assert result[digest]["key_alias"] == "back-online" +@pytest.mark.asyncio +async def test_recover_key_metadata_from_spend_logs_treats_a_dropped_connection_error_as_a_short_lived_miss(): + digest: Final = hash_token("cli-session-dropped-connection") + window: Final = (datetime(2026, 9, 7), datetime(2026, 9, 10)) + cache: Final = InMemoryCache(default_ttl=SPEND_LOG_KEY_METADATA_CACHE_TTL) + mock_prisma: Final = MagicMock() + _spend_log_transaction( + mock_prisma, + AsyncMock( + side_effect=[ + AttributeError("'NoneType' object has no attribute 'get'"), + [_spend_log_row(digest, "back-online", None, None)], + ] + ), + ) + started: Final = time.time() + + assert await recover_key_metadata_from_spend_logs(mock_prisma, {digest}, window, cache=cache) == {} + + miss_key: Final = next(key for key in cache.ttl_dict if digest in key and not key.endswith(":missed-before")) + assert cache.ttl_dict[miss_key] - started <= SPEND_LOG_KEY_METADATA_MISS_CACHE_TTL + 1 + assert f"{miss_key}:missed-before" not in cache.ttl_dict + cache.ttl_dict[miss_key] = time.time() - 1 + + result: Final = await recover_key_metadata_from_spend_logs(mock_prisma, {digest}, window, cache=cache) + + assert result[digest]["key_alias"] == "back-online" + @pytest.mark.asyncio async def test_recover_key_metadata_from_spend_logs_drops_the_owner_of_a_digest_shared_by_several_users(): shared_ui_digest = hash_token("ui-token")