diff --git a/litellm/proxy/spend_tracking/key_metadata_recovery.py b/litellm/proxy/spend_tracking/key_metadata_recovery.py index 33339f16d52..78169283dff 100644 --- a/litellm/proxy/spend_tracking/key_metadata_recovery.py +++ b/litellm/proxy/spend_tracking/key_metadata_recovery.py @@ -1,7 +1,8 @@ -from collections.abc import Mapping, Sequence +from collections.abc import Awaitable, Callable, Mapping, Sequence from collections.abc import Set as AbstractSet -from typing import Final, Protocol +from typing import Final, Protocol, TypeVar +from prisma.errors import PrismaError from typing_extensions import TypedDict from litellm._logging import verbose_proxy_logger @@ -13,6 +14,8 @@ from litellm.repositories.verification_token_repository import ( VerificationTokenRepository, ) +_T = TypeVar("_T") + # Cap reverse-hash scans so a Usage page with orphaned double-hashed api_key # values cannot pull an unbounded VerificationToken table into memory. _MAX_DOUBLE_HASH_TOKEN_SCAN: Final = 10_000 @@ -56,6 +59,18 @@ class _TokenAliasRecord(Protocol): def user_id(self) -> str | None: ... +async def _db_or_empty( + load: Callable[[], Awaitable[_T]], + warning: str, + count: int, +) -> _T | None: + try: + return await load() + except PrismaError as e: + verbose_proxy_logger.warning(warning, count, e) + return None + + def _token_digest_metadata( records: Sequence[_TokenAliasRecord], wanted: AbstractSet[str], @@ -76,16 +91,12 @@ async def _reverse_hash_active_key_metadata( prisma_client: PrismaClient, wanted: AbstractSet[str], ) -> dict[str, KeyMetadataDict]: - try: - active_records: Final[Sequence[_TokenAliasRecord]] = await VerificationTokenRepository( - prisma_client - ).table.find_many(take=_MAX_DOUBLE_HASH_TOKEN_SCAN) - except Exception as e: - verbose_proxy_logger.warning( - "Failed reverse-hash recovery against active keys for %d missing keys: %s", - len(wanted), - e, - ) + active_records: Final = await _db_or_empty( + lambda: VerificationTokenRepository(prisma_client).table.find_many(take=_MAX_DOUBLE_HASH_TOKEN_SCAN), + "Failed reverse-hash recovery against active keys for %d missing keys: %s", + len(wanted), + ) + if active_records is None: return {} return _token_digest_metadata(active_records, wanted) @@ -94,19 +105,15 @@ async def _reverse_hash_deleted_key_metadata( prisma_client: PrismaClient, wanted: AbstractSet[str], ) -> dict[str, KeyMetadataDict]: - try: - deleted_records: Final[Sequence[_TokenAliasRecord]] = await DeletedVerificationTokenRepository( - prisma_client - ).table.find_many( + deleted_records: Final = await _db_or_empty( + lambda: DeletedVerificationTokenRepository(prisma_client).table.find_many( take=_MAX_DOUBLE_HASH_TOKEN_SCAN, order={"deleted_at": "desc"}, - ) - except Exception as e: - verbose_proxy_logger.warning( - "Failed reverse-hash recovery against deleted keys for %d missing keys: %s", - len(wanted), - e, - ) + ), + "Failed reverse-hash recovery against deleted keys for %d missing keys: %s", + len(wanted), + ) + if deleted_records is None: return {} return _token_digest_metadata(deleted_records, wanted) @@ -126,19 +133,14 @@ async def _spend_logs_key_metadata( prisma_client: PrismaClient, wanted: AbstractSet[str], ) -> dict[str, KeyMetadataDict]: - try: - spend_log_rows: Final = await prisma_client.db.query_raw( + spend_log_rows: Final = await _db_or_empty( + lambda: prisma_client.db.query_raw( _SPEND_LOGS_KEY_METADATA_SQL, list(wanted), - ) - except Exception as e: - verbose_proxy_logger.warning( - "Failed SpendLogs metadata recovery for %d missing keys: %s", - len(wanted), - e, - ) - return {} - + ), + "Failed SpendLogs metadata recovery for %d missing keys: %s", + len(wanted), + ) if not isinstance(spend_log_rows, list): return {} @@ -160,14 +162,12 @@ async def _emails_for_user_ids( ) -> Mapping[str, str]: if not user_ids: return {} - try: - users: Final = await UserRepository(prisma_client).table.find_many(where={"user_id": {"in": list(user_ids)}}) - except Exception as e: - verbose_proxy_logger.warning( - "Failed user_email recovery for %d user ids: %s", - len(user_ids), - e, - ) + users: Final = await _db_or_empty( + lambda: UserRepository(prisma_client).table.find_many(where={"user_id": {"in": list(user_ids)}}), + "Failed user_email recovery for %d user ids: %s", + len(user_ids), + ) + if users is None: return {} return { user.user_id: user.user_email 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 3d08b2909c4..6f0e912c5b1 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 @@ -2,6 +2,7 @@ from types import SimpleNamespace from unittest.mock import AsyncMock, MagicMock import pytest +from prisma.errors import PrismaError from litellm.proxy.spend_tracking.key_metadata_recovery import ( fill_missing_api_key_aliases, @@ -83,3 +84,30 @@ async def test_fill_missing_api_key_aliases_updates_null_alias_and_email_rows(): assert filled[0]["team_id"] == "team-9" assert filled[0]["user_email"] == "bob@example.com" assert filled[1]["api_key_alias"] == "named-key" + + +@pytest.mark.asyncio +async def test_recover_falls_back_to_spend_logs_when_token_scan_raises_prisma_error(): + token = "b" * 64 + double_hashed = hash_token(token) + mock_prisma = MagicMock() + mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(side_effect=PrismaError("db down")) + mock_prisma.db.litellm_deletedverificationtoken.find_many = AsyncMock(side_effect=PrismaError("db down")) + mock_prisma.db.litellm_usertable.find_many = AsyncMock(return_value=[]) + mock_prisma.db.query_raw = AsyncMock( + return_value=[ + { + "api_key": double_hashed, + "key_alias": "from-spend-logs", + "team_id": "team-sl", + "user_id": "carol", + "user_email": "carol@example.com", + } + ] + ) + + result = await recover_double_hashed_key_metadata(mock_prisma, {double_hashed}) + + assert result[double_hashed]["key_alias"] == "from-spend-logs" + assert result[double_hashed]["team_id"] == "team-sl" + assert result[double_hashed]["user_email"] == "carol@example.com"