From ce95afe2bdfb47f4325c59fea89b8c8f8fb0a5a6 Mon Sep 17 00:00:00 2001 From: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Date: Thu, 3 Sep 2026 17:58:17 -0700 Subject: [PATCH] fix(spend-tracking): reverse-hash dirty spend keys in Postgres instead of paging token tables --- litellm/integrations/cloudzero/database.py | 2 - litellm/integrations/focus/database.py | 2 - litellm/proxy/_types.py | 1 - .../spend_tracking/key_metadata_recovery.py | 232 ++++------------- .../spend_tracking/spend_tracking_utils.py | 1 - .../integrations/cloudzero/test_cloudzero.py | 2 +- .../test_common_daily_activity.py | 58 ++--- .../test_key_metadata_recovery.py | 234 ++++++++---------- .../test_spend_management_endpoints.py | 1 - .../test_spend_tracking_utils.py | 2 - 10 files changed, 189 insertions(+), 346 deletions(-) diff --git a/litellm/integrations/cloudzero/database.py b/litellm/integrations/cloudzero/database.py index e630bd85114..4adf725fd0f 100644 --- a/litellm/integrations/cloudzero/database.py +++ b/litellm/integrations/cloudzero/database.py @@ -105,8 +105,6 @@ class LiteLLMDatabase: if isinstance(db_response, list) else [] ) - # v1.99 double-hashed DailyUserSpend.api_key values miss the - # VerificationToken join above; recover alias/team for those rows. recovered_rows: Final = await fill_missing_api_key_aliases(client, usage_rows) return pl.DataFrame(tuple(recovered_rows), infer_schema_length=None) except Exception as e: diff --git a/litellm/integrations/focus/database.py b/litellm/integrations/focus/database.py index 02b1e9e944b..f214aa02b5b 100644 --- a/litellm/integrations/focus/database.py +++ b/litellm/integrations/focus/database.py @@ -107,8 +107,6 @@ class FocusLiteLLMDatabase: if isinstance(db_response, list) else [] ) - # v1.99 double-hashed DailyUserSpend.api_key values miss the - # VerificationToken join above; recover alias/team for those rows. recovered_rows: Final = await fill_missing_api_key_aliases(client, usage_rows) return pl.DataFrame(tuple(recovered_rows), infer_schema_length=None) except Exception as exc: diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 9b4a55d3510..5d5a25e7cd6 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -3659,7 +3659,6 @@ class SpendLogsMetadata(TypedDict): user_api_key_project_alias: str | None user_api_key_org_id: str | None user_api_key_user_id: str | None - user_api_key_user_email: ReadOnly[str | None] user_api_key_team_alias: str | None spend_logs_metadata: dict | None # special param to log k,v pairs to spendlogs for a call requester_ip_address: str | None diff --git a/litellm/proxy/spend_tracking/key_metadata_recovery.py b/litellm/proxy/spend_tracking/key_metadata_recovery.py index 895f97e1a05..d524b158c9c 100644 --- a/litellm/proxy/spend_tracking/key_metadata_recovery.py +++ b/litellm/proxy/spend_tracking/key_metadata_recovery.py @@ -1,38 +1,30 @@ from collections.abc import Awaitable, Callable, Mapping, Sequence from collections.abc import Set as AbstractSet from types import MappingProxyType -from typing import Final, Protocol, TypeVar +from typing import Final, TypeVar +from pydantic import BaseModel, TypeAdapter from typing_extensions import ReadOnly, TypedDict from litellm._logging import verbose_proxy_logger from litellm.litellm_core_utils.litellm_logging import is_valid_sha256_hash -from litellm.proxy.utils import PrismaClient, hash_token -from litellm.repositories.table_repositories import DeletedVerificationTokenRepository +from litellm.proxy.utils import PrismaClient from litellm.repositories.user_repository import UserRepository -from litellm.repositories.verification_token_repository import ( - VerificationTokenRepository, -) _T = TypeVar("_T") -_TOKEN_SCAN_PAGE: Final = 10_000 +_ACTIVE_TOKEN_DIGEST_SQL: Final = """ +SELECT encode(sha256(convert_to(token, 'UTF8')), 'hex') AS digest, key_alias, team_id, user_id +FROM "LiteLLM_VerificationToken" +WHERE encode(sha256(convert_to(token, 'UTF8')), 'hex') = ANY($1::text[]) +""" -_SPEND_LOGS_KEY_METADATA_SQL: Final = """ -SELECT DISTINCT ON (api_key) - api_key, - metadata->>'user_api_key_alias' AS key_alias, - metadata->>'user_api_key_team_id' AS team_id, - metadata->>'user_api_key_user_id' AS user_id, - metadata->>'user_api_key_user_email' AS user_email -FROM "LiteLLM_SpendLogs" -WHERE api_key = ANY($1::text[]) - AND ( - NULLIF(metadata->>'user_api_key_alias', '') IS NOT NULL - OR NULLIF(metadata->>'user_api_key_team_id', '') IS NOT NULL - OR NULLIF(metadata->>'user_api_key_user_email', '') IS NOT NULL - ) -ORDER BY api_key, "startTime" DESC NULLS LAST +_DELETED_TOKEN_DIGEST_SQL: Final = """ +SELECT DISTINCT ON (token) + encode(sha256(convert_to(token, 'UTF8')), 'hex') AS digest, key_alias, team_id, user_id +FROM "LiteLLM_DeletedVerificationToken" +WHERE encode(sha256(convert_to(token, 'UTF8')), 'hex') = ANY($1::text[]) +ORDER BY token, deleted_at DESC """ @@ -43,24 +35,18 @@ class KeyMetadataDict(TypedDict, total=False): user_email: ReadOnly[str | None] +class _TokenDigestRow(BaseModel): + digest: str + key_alias: str | None = None + team_id: str | None = None + user_id: str | None = None + + +_TOKEN_DIGEST_ROWS: Final = TypeAdapter(tuple[_TokenDigestRow, ...]) _EMPTY_KEY_METADATA: Final[Mapping[str, KeyMetadataDict]] = MappingProxyType({}) _EMPTY_EMAILS: Final[Mapping[str, str]] = MappingProxyType({}) -class _TokenAliasRecord(Protocol): - @property - def token(self) -> str: ... - - @property - def key_alias(self) -> str | None: ... - - @property - def team_id(self) -> str | None: ... - - @property - def user_id(self) -> str | None: ... - - async def _db_or_empty( load: Callable[[], Awaitable[_T]], warning: str, @@ -75,142 +61,25 @@ async def _db_or_empty( return None -def _record_metadata(record: _TokenAliasRecord) -> KeyMetadataDict: - meta: Final[KeyMetadataDict] = { - "key_alias": record.key_alias, - "team_id": record.team_id, - "user_id": getattr(record, "user_id", None), - } - return meta - - -def _spend_log_row_metadata(row: Mapping[str, object]) -> KeyMetadataDict: - meta: Final[KeyMetadataDict] = { - "key_alias": row.get("key_alias") if isinstance(row.get("key_alias"), str) else None, - "team_id": row.get("team_id") if isinstance(row.get("team_id"), str) else None, - "user_id": row.get("user_id") if isinstance(row.get("user_id"), str) else None, - "user_email": row.get("user_email") if isinstance(row.get("user_email"), str) else None, - } - return meta - - -def _token_digest_metadata( - records: Sequence[_TokenAliasRecord], - wanted: AbstractSet[str], -) -> Mapping[str, KeyMetadataDict]: - return MappingProxyType( - { - digested: _record_metadata(record) - for record in records - for digested in (hash_token(record.token),) - if digested in wanted - } - ) - - -async def _paginate_token_digest_metadata( - load_page: Callable[[int], Awaitable[Sequence[_TokenAliasRecord] | None]], - wanted: AbstractSet[str], - *, - page_size: int, - skip: int = 0, - accumulated: Mapping[str, KeyMetadataDict] = _EMPTY_KEY_METADATA, -) -> Mapping[str, KeyMetadataDict]: - if not wanted: - return accumulated - records: Final = await load_page(skip) - if records is None: - return accumulated - page_hits: Final = _token_digest_metadata(records, wanted) - combined: Final[Mapping[str, KeyMetadataDict]] = ( - MappingProxyType({**accumulated, **page_hits}) if page_hits else accumulated - ) - still_wanted: Final = wanted - frozenset(page_hits) - if not still_wanted or len(records) < page_size: - return combined - return await _paginate_token_digest_metadata( - load_page, - still_wanted, - page_size=page_size, - skip=skip + page_size, - accumulated=combined, - ) - - -async def _reverse_hash_active_key_metadata( - prisma_client: PrismaClient, - wanted: AbstractSet[str], - *, - page_size: int, -) -> Mapping[str, KeyMetadataDict]: - async def load_page(skip: int) -> Sequence[_TokenAliasRecord] | None: - return await _db_or_empty( - lambda: VerificationTokenRepository(prisma_client).table.find_many( - take=page_size, - skip=skip, - order={"token": "asc"}, # mutable-ok: Prisma find_many order= is a dict - ), - "Failed reverse-hash recovery against active keys for %d missing keys: %s", - len(wanted), - ) - - return await _paginate_token_digest_metadata(load_page, wanted, page_size=page_size) - - -async def _reverse_hash_deleted_key_metadata( - prisma_client: PrismaClient, - wanted: AbstractSet[str], - *, - page_size: int, -) -> Mapping[str, KeyMetadataDict]: - async def load_page(skip: int) -> Sequence[_TokenAliasRecord] | None: - return await _db_or_empty( - lambda: DeletedVerificationTokenRepository(prisma_client).table.find_many( - take=page_size, - skip=skip, - order=[{"deleted_at": "desc"}, {"id": "asc"}], # mutable-ok: Prisma find_many order= is a dict - ), - "Failed reverse-hash recovery against deleted keys for %d missing keys: %s", - len(wanted), - ) - - return await _paginate_token_digest_metadata(load_page, wanted, page_size=page_size) - - async def _reverse_hash_key_metadata( prisma_client: PrismaClient, + sql: str, wanted: AbstractSet[str], *, - page_size: int, + warning: str, ) -> Mapping[str, KeyMetadataDict]: - from_active: Final = await _reverse_hash_active_key_metadata(prisma_client, wanted, page_size=page_size) - still_wanted: Final = wanted - frozenset(from_active) - if not still_wanted: - return from_active - from_deleted: Final = await _reverse_hash_deleted_key_metadata(prisma_client, still_wanted, page_size=page_size) - return MappingProxyType({**from_active, **from_deleted}) - - -async def _spend_logs_key_metadata( - prisma_client: PrismaClient, - wanted: AbstractSet[str], -) -> Mapping[str, KeyMetadataDict]: - spend_log_rows: Final = await _db_or_empty( - lambda: prisma_client.db.query_raw( - _SPEND_LOGS_KEY_METADATA_SQL, - tuple(wanted), - ), - "Failed SpendLogs metadata recovery for %d missing keys: %s", + rows: Final = await _db_or_empty( + lambda: prisma_client.db.query_raw(sql, sorted(wanted)), + warning, len(wanted), ) - if not isinstance(spend_log_rows, list): + if rows is None: return _EMPTY_KEY_METADATA - return MappingProxyType( { - row["api_key"]: _spend_log_row_metadata(row) - for row in spend_log_rows - if isinstance(row, dict) and isinstance(row.get("api_key"), str) and row["api_key"] in wanted + row.digest: KeyMetadataDict(key_alias=row.key_alias, team_id=row.team_id, user_id=row.user_id) + for row in _TOKEN_DIGEST_ROWS.validate_python(rows) + if row.digest in wanted } ) @@ -223,7 +92,7 @@ async def _emails_for_user_ids( return _EMPTY_EMAILS users: Final = await _db_or_empty( lambda: UserRepository(prisma_client).table.find_many( - where={"user_id": {"in": tuple(user_ids)}}, # mutable-ok: Prisma find_many where= is a dict + where={"user_id": {"in": list(user_ids)}}, # mutable-ok: Prisma find_many where= is a dict ), "Failed user_email recovery for %d user ids: %s", len(user_ids), @@ -268,31 +137,35 @@ async def attach_user_emails( async def recover_double_hashed_key_metadata( prisma_client: PrismaClient, missing_keys: AbstractSet[str], - *, - token_scan_page_size: int = _TOKEN_SCAN_PAGE, ) -> Mapping[str, KeyMetadataDict]: """ - Recover key_alias/team_id/user_email for DailyUserSpend.api_key values that + Recover key_alias/team_id/user_id for DailyUserSpend.api_key values that were double-hashed by the v1.99 spend-log provenance gate. Those rows store hash(VerificationToken.token) instead of the token, so the - exact join misses. Page through active then deleted tokens until every - wanted digest is found or the table ends; fall back to SpendLogs metadata. - Emails come from SpendLogs when present, otherwise from UserTable via the - recovered key's user_id. + exact join misses. Postgres hashes the token column itself, one pass over + active keys and one over deleted keys, so no key row crosses the wire. """ sha_missing: Final = frozenset(key for key in missing_keys if is_valid_sha256_hash(key)) if not sha_missing: return _EMPTY_KEY_METADATA - from_tokens: Final = await _reverse_hash_key_metadata(prisma_client, sha_missing, page_size=token_scan_page_size) - still_missing: Final = sha_missing - frozenset(from_tokens) - recovered: Final = ( - from_tokens - if not still_missing - else MappingProxyType({**from_tokens, **(await _spend_logs_key_metadata(prisma_client, still_missing))}) + from_active: Final = await _reverse_hash_key_metadata( + prisma_client, + _ACTIVE_TOKEN_DIGEST_SQL, + sha_missing, + warning="Failed reverse-hash recovery against active keys for %d missing keys: %s", ) - return await attach_user_emails(prisma_client, recovered) + still_missing: Final = sha_missing - frozenset(from_active) + if not still_missing: + return from_active + from_deleted: Final = await _reverse_hash_key_metadata( + prisma_client, + _DELETED_TOKEN_DIGEST_SQL, + still_missing, + warning="Failed reverse-hash recovery against deleted keys for %d missing keys: %s", + ) + return MappingProxyType({**from_active, **from_deleted}) def _row_with_recovered_fields( @@ -345,7 +218,10 @@ async def fill_missing_api_key_aliases( if not missing_keys: return tuple(rows) - recovered: Final = await recover_double_hashed_key_metadata(prisma_client, missing_keys) + recovered: Final = await attach_user_emails( + prisma_client, + await recover_double_hashed_key_metadata(prisma_client, missing_keys), + ) if not recovered: return tuple(rows) diff --git a/litellm/proxy/spend_tracking/spend_tracking_utils.py b/litellm/proxy/spend_tracking/spend_tracking_utils.py index 4c7333143c5..a37c3ba4405 100644 --- a/litellm/proxy/spend_tracking/spend_tracking_utils.py +++ b/litellm/proxy/spend_tracking/spend_tracking_utils.py @@ -139,7 +139,6 @@ def _get_spend_logs_metadata( user_api_key_project_alias=None, user_api_key_org_id=None, user_api_key_user_id=None, - user_api_key_user_email=None, user_api_key_team_alias=None, spend_logs_metadata=None, requester_ip_address=None, diff --git a/tests/test_litellm/integrations/cloudzero/test_cloudzero.py b/tests/test_litellm/integrations/cloudzero/test_cloudzero.py index c543156eedd..1b6e8ca513c 100644 --- a/tests/test_litellm/integrations/cloudzero/test_cloudzero.py +++ b/tests/test_litellm/integrations/cloudzero/test_cloudzero.py @@ -74,7 +74,7 @@ class TestCloudZeroHourlyExport: fake_db = MagicMock() async def query_raw_mock(query: str, *params): - if "LiteLLM_SpendLogs" in query: + if "sha256(" in query: return [] start_time_utc = params[0] if len(params) > 0 else None end_time_utc = params[1] if len(params) > 1 else None diff --git a/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py b/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py index b9c0d953086..37a54c4901a 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py +++ b/tests/test_litellm/proxy/management_endpoints/test_common_daily_activity.py @@ -459,33 +459,23 @@ async def test_get_api_key_metadata_recovers_double_hashed_key_via_reverse_hash( """ v1.99 spend logging re-hashed already-hashed api_key values when provenance was missing. Usage joins DailyUserSpend.api_key to VerificationToken.token, so those - rows looked like key-hash-... with a null alias. Reverse-hash recovery must map - hash(token) back to the key's alias for historical dirty spend. + rows looked like key-hash-... with a null alias. Recovery asks Postgres for the + key whose hashed token matches the dirty value and maps it back to its alias. """ from litellm.proxy.utils import hash_token - token = "a" * 64 - double_hashed = hash_token(token) + double_hashed = hash_token("a" * 64) mock_prisma = MagicMock() - - mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock( - side_effect=[ - [], # exact join miss - [ - SimpleNamespace( - token=token, - key_alias="batch-worker", - team_id="team-1", - user_id="alice", - ) - ], - ] - ) + mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) mock_prisma.db.litellm_deletedverificationtoken.find_many = AsyncMock(return_value=[]) mock_prisma.db.litellm_usertable.find_many = AsyncMock( return_value=[SimpleNamespace(user_id="alice", user_email="alice@example.com")] ) - mock_prisma.db.query_raw = AsyncMock(return_value=[]) + mock_prisma.db.query_raw = AsyncMock( + return_value=[ + {"digest": double_hashed, "key_alias": "batch-worker", "team_id": "team-1", "user_id": "alice"} + ] + ) result = await get_api_key_metadata( prisma_client=mock_prisma, @@ -495,37 +485,37 @@ async def test_get_api_key_metadata_recovers_double_hashed_key_via_reverse_hash( assert result[double_hashed]["key_alias"] == "batch-worker" assert result[double_hashed]["team_id"] == "team-1" assert result[double_hashed]["user_email"] == "alice@example.com" - mock_prisma.db.query_raw.assert_not_called() + ((digest_sql, digests),) = [call.args for call in mock_prisma.db.query_raw.call_args_list] + assert '"LiteLLM_VerificationToken"' in digest_sql + assert digests == [double_hashed] @pytest.mark.asyncio -async def test_get_api_key_metadata_recovers_double_hashed_key_via_spend_logs(): - """When the token tables cannot reverse-hash the dirty key, use SpendLogs metadata.""" +async def test_get_api_key_metadata_permanent_miss_never_pages_tokens_or_reads_spend_logs(): + """A dirty key no table can explain costs two digest lookups, never a token page walk or a SpendLogs scan.""" from litellm.proxy.utils import hash_token double_hashed = hash_token("b" * 64) mock_prisma = MagicMock() mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) mock_prisma.db.litellm_deletedverificationtoken.find_many = AsyncMock(return_value=[]) - mock_prisma.db.query_raw = AsyncMock( - return_value=[ - { - "api_key": double_hashed, - "key_alias": "from-spend-log", - "team_id": "team-spend", - } - ] - ) mock_prisma.db.litellm_usertable.find_many = AsyncMock(return_value=[]) + mock_prisma.db.query_raw = AsyncMock(return_value=[]) result = await get_api_key_metadata( prisma_client=mock_prisma, api_keys={double_hashed}, ) - assert result[double_hashed]["key_alias"] == "from-spend-log" - assert result[double_hashed]["team_id"] == "team-spend" - mock_prisma.db.query_raw.assert_called_once() + assert double_hashed not in result + issued_sql = [call.args[0] for call in mock_prisma.db.query_raw.call_args_list] + assert len(issued_sql) == 2 + assert not any("LiteLLM_SpendLogs" in sql for sql in issued_sql) + token_lookups = ( + mock_prisma.db.litellm_verificationtoken.find_many.call_args_list + + mock_prisma.db.litellm_deletedverificationtoken.find_many.call_args_list + ) + assert all("take" not in call.kwargs and "skip" not in call.kwargs for call in token_lookups) def test_key_metadata_includes_recovered_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 704de049fc8..43f7d20cf13 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 @@ -1,3 +1,4 @@ +from collections.abc import Sequence from types import SimpleNamespace from unittest.mock import AsyncMock, MagicMock @@ -11,55 +12,126 @@ from litellm.proxy.spend_tracking.key_metadata_recovery import ( from litellm.proxy.utils import hash_token +def _digest_row(digest: str, key_alias: str, team_id: str | None, user_id: str | None) -> dict[str, str | None]: + return {"digest": digest, "key_alias": key_alias, "team_id": team_id, "user_id": user_id} + + +def _query_raw_by_table( + active_rows: Sequence[dict[str, str | None]], + deleted_rows: Sequence[dict[str, str | None]], +) -> AsyncMock: + async def query_raw(sql: str, *params: object) -> list[dict[str, str | None]]: + if '"LiteLLM_VerificationToken"' in sql: + return list(active_rows) + if '"LiteLLM_DeletedVerificationToken"' in sql: + return list(deleted_rows) + raise AssertionError(f"unexpected query: {sql}") + + return AsyncMock(side_effect=query_raw) + + @pytest.mark.asyncio -async def test_recover_double_hashed_key_metadata_via_reverse_hash(): - token = "a" * 64 - double_hashed = hash_token(token) +async def test_recover_double_hashed_key_metadata_via_active_token_digest(): + double_hashed = hash_token("a" * 64) mock_prisma = MagicMock() - mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock( - return_value=[ - SimpleNamespace( - token=token, - key_alias="batch-worker", - team_id="team-1", - user_id="alice", - ) - ] + mock_prisma.db.query_raw = _query_raw_by_table( + active_rows=[_digest_row(double_hashed, "batch-worker", "team-1", "alice")], + deleted_rows=[], ) - mock_prisma.db.litellm_deletedverificationtoken.find_many = AsyncMock(return_value=[]) - mock_prisma.db.litellm_usertable.find_many = AsyncMock( - return_value=[SimpleNamespace(user_id="alice", user_email="alice@example.com")] - ) - mock_prisma.db.query_raw = AsyncMock(return_value=[]) result = await recover_double_hashed_key_metadata(mock_prisma, {double_hashed}) assert result[double_hashed]["key_alias"] == "batch-worker" assert result[double_hashed]["team_id"] == "team-1" - assert result[double_hashed]["user_email"] == "alice@example.com" + assert result[double_hashed]["user_id"] == "alice" + ((_, digests),) = [call.args for call in mock_prisma.db.query_raw.call_args_list] + assert digests == [double_hashed] + + +@pytest.mark.asyncio +async def test_recover_double_hashed_key_metadata_falls_back_to_deleted_tokens(): + double_hashed = hash_token("y" * 64) + mock_prisma = MagicMock() + mock_prisma.db.query_raw = _query_raw_by_table( + active_rows=[], + deleted_rows=[_digest_row(double_hashed, "deleted-key", "team-del", "erin")], + ) + + result = await recover_double_hashed_key_metadata(mock_prisma, {double_hashed}) + + assert result[double_hashed]["key_alias"] == "deleted-key" + assert result[double_hashed]["team_id"] == "team-del" + assert result[double_hashed]["user_id"] == "erin" + assert [call.args[1] for call in mock_prisma.db.query_raw.call_args_list] == [[double_hashed], [double_hashed]] + + +@pytest.mark.asyncio +async def test_recover_only_asks_deleted_tokens_for_digests_active_keys_missed(): + found_active = hash_token("1" * 64) + found_deleted = hash_token("2" * 64) + mock_prisma = MagicMock() + mock_prisma.db.query_raw = _query_raw_by_table( + active_rows=[_digest_row(found_active, "active-key", None, None)], + deleted_rows=[_digest_row(found_deleted, "deleted-key", None, None)], + ) + + result = await recover_double_hashed_key_metadata(mock_prisma, {found_active, found_deleted}) + + assert result[found_active]["key_alias"] == "active-key" + assert result[found_deleted]["key_alias"] == "deleted-key" + assert [call.args[1] for call in mock_prisma.db.query_raw.call_args_list] == [ + sorted((found_active, found_deleted)), + [found_deleted], + ] + + +@pytest.mark.asyncio +async def test_recover_permanent_miss_costs_two_digest_lookups_and_no_table_walk(): + double_hashed = hash_token("b" * 64) + mock_prisma = MagicMock() + mock_prisma.db.query_raw = _query_raw_by_table(active_rows=[], deleted_rows=[]) + + result = await recover_double_hashed_key_metadata(mock_prisma, {double_hashed}) + + assert result == {} + assert len(mock_prisma.db.query_raw.call_args_list) == 2 + mock_prisma.db.litellm_verificationtoken.find_many.assert_not_called() + mock_prisma.db.litellm_deletedverificationtoken.find_many.assert_not_called() + + +@pytest.mark.asyncio +async def test_recover_skips_keys_that_are_not_sha256_digests(): + mock_prisma = MagicMock() + mock_prisma.db.query_raw = AsyncMock(return_value=[]) + + result = await recover_double_hashed_key_metadata(mock_prisma, {"sk-plain-key", "key-hash-short"}) + + assert result == {} mock_prisma.db.query_raw.assert_not_called() @pytest.mark.asyncio -async def test_fill_missing_api_key_aliases_updates_null_alias_and_email_rows(): - token = "c" * 64 - double_hashed = hash_token(token) +async def test_recover_returns_empty_when_digest_lookup_raises_prisma_error(): + double_hashed = hash_token("c" * 64) mock_prisma = MagicMock() - mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock( - return_value=[ - SimpleNamespace( - token=token, - key_alias="recovered-alias", - team_id="team-9", - user_id="bob", - ) - ] + mock_prisma.db.query_raw = AsyncMock(side_effect=PrismaError("db down")) + + result = await recover_double_hashed_key_metadata(mock_prisma, {double_hashed}) + + assert result == {} + + +@pytest.mark.asyncio +async def test_fill_missing_api_key_aliases_updates_null_alias_and_email_rows(): + double_hashed = hash_token("d" * 64) + mock_prisma = MagicMock() + mock_prisma.db.query_raw = _query_raw_by_table( + active_rows=[_digest_row(double_hashed, "recovered-alias", "team-9", "bob")], + deleted_rows=[], ) - mock_prisma.db.litellm_deletedverificationtoken.find_many = AsyncMock(return_value=[]) mock_prisma.db.litellm_usertable.find_many = AsyncMock( return_value=[SimpleNamespace(user_id="bob", user_email="bob@example.com")] ) - mock_prisma.db.query_raw = AsyncMock(return_value=[]) rows = ( { @@ -83,104 +155,18 @@ async def test_fill_missing_api_key_aliases_updates_null_alias_and_email_rows(): assert filled[0]["api_key_alias"] == "recovered-alias" assert filled[0]["team_id"] == "team-9" assert filled[0]["user_email"] == "bob@example.com" + assert filled[0]["spend"] == 12.5 assert filled[1]["api_key_alias"] == "named-key" + assert mock_prisma.db.litellm_usertable.find_many.call_args.kwargs["where"] == {"user_id": {"in": ["bob"]}} @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) +async def test_fill_missing_api_key_aliases_leaves_rows_untouched_when_nothing_is_missing(): 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" - - -@pytest.mark.asyncio -async def test_recover_double_hashed_key_metadata_scans_past_first_page(): - token = "z" * 64 - double_hashed = hash_token(token) - decoys = ( - SimpleNamespace(token="1" * 64, key_alias="decoy-1", team_id=None, user_id=None), - SimpleNamespace(token="2" * 64, key_alias="decoy-2", team_id=None, user_id=None), - ) - match = SimpleNamespace(token=token, key_alias="late-key", team_id="team-late", user_id="dana") - - async def find_many(*, take: int | None = None, skip: int | None = None, order: object = None): - if skip == 0: - return list(decoys) - if skip == 2: - return [match] - return [] - - mock_prisma = MagicMock() - mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(side_effect=find_many) - mock_prisma.db.litellm_deletedverificationtoken.find_many = AsyncMock(return_value=[]) - mock_prisma.db.litellm_usertable.find_many = AsyncMock( - return_value=[SimpleNamespace(user_id="dana", user_email="dana@example.com")] - ) mock_prisma.db.query_raw = AsyncMock(return_value=[]) + rows = ({"api_key": hash_token("e" * 64), "api_key_alias": "named", "user_email": "x@example.com"},) - result = await recover_double_hashed_key_metadata(mock_prisma, {double_hashed}, token_scan_page_size=2) + filled = await fill_missing_api_key_aliases(mock_prisma, rows) - assert result[double_hashed]["key_alias"] == "late-key" - assert result[double_hashed]["team_id"] == "team-late" - assert result[double_hashed]["user_email"] == "dana@example.com" - assert [call.kwargs["skip"] for call in mock_prisma.db.litellm_verificationtoken.find_many.call_args_list] == [0, 2] - mock_prisma.db.litellm_deletedverificationtoken.find_many.assert_not_called() - mock_prisma.db.query_raw.assert_not_called() - - -@pytest.mark.asyncio -async def test_recover_double_hashed_key_metadata_pages_deleted_tokens(): - token = "y" * 64 - double_hashed = hash_token(token) - decoys = ( - SimpleNamespace(token="3" * 64, key_alias="deleted-decoy-1", team_id=None, user_id=None), - SimpleNamespace(token="4" * 64, key_alias="deleted-decoy-2", team_id=None, user_id=None), - ) - match = SimpleNamespace(token=token, key_alias="deleted-late-key", team_id="team-del", user_id="erin") - - async def find_deleted(*, take: int | None = None, skip: int | None = None, order: object = None): - if skip == 0: - return list(decoys) - if skip == 2: - return [match] - return [] - - mock_prisma = MagicMock() - mock_prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[]) - mock_prisma.db.litellm_deletedverificationtoken.find_many = AsyncMock(side_effect=find_deleted) - mock_prisma.db.litellm_usertable.find_many = AsyncMock( - return_value=[SimpleNamespace(user_id="erin", user_email="erin@example.com")] - ) - mock_prisma.db.query_raw = AsyncMock(return_value=[]) - - result = await recover_double_hashed_key_metadata(mock_prisma, {double_hashed}, token_scan_page_size=2) - - assert result[double_hashed]["key_alias"] == "deleted-late-key" - assert result[double_hashed]["user_email"] == "erin@example.com" - assert [ - call.kwargs["skip"] for call in mock_prisma.db.litellm_deletedverificationtoken.find_many.call_args_list - ] == [ - 0, - 2, - ] + assert filled == rows mock_prisma.db.query_raw.assert_not_called() diff --git a/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py b/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py index dfdad545003..73a29afd9b9 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py +++ b/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py @@ -511,7 +511,6 @@ ignored_keys = [ "metadata.user_api_key_project_alias", "metadata.user_api_key_org_id", "metadata.user_api_key_user_id", - "metadata.user_api_key_user_email", "metadata.user_api_key_team_alias", "metadata.spend_logs_metadata", "metadata.requester_ip_address", diff --git a/tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py b/tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py index 3cfc16eedec..fb0cdc175b1 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py +++ b/tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py @@ -2729,7 +2729,6 @@ def test_get_logging_payload_batch_attribution_keeps_verification_token_hash(): "user_api_key_hash": token_hash, "user_api_key_alias": "batch-creator", "user_api_key_user_id": "alice", - "user_api_key_user_email": "alice@example.com", "user_api_key_team_id": "team-1", } }, @@ -2746,7 +2745,6 @@ def test_get_logging_payload_batch_attribution_keeps_verification_token_hash(): parsed_meta = json.loads(payload["metadata"]) assert parsed_meta["user_api_key"] == token_hash assert parsed_meta["user_api_key_alias"] == "batch-creator" - assert parsed_meta["user_api_key_user_email"] == "alice@example.com" def test_get_spend_logs_metadata_provenance_bypass_requires_hash_match():