fix(spend-tracking): reverse-hash dirty spend keys in Postgres instead of paging token tables

This commit is contained in:
mateo-berri 2026-09-03 17:58:17 -07:00
parent f1f0294796
commit ce95afe2bd
10 changed files with 189 additions and 346 deletions

View file

@ -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:

View file

@ -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:

View file

@ -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

View file

@ -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)

View file

@ -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,

View file

@ -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

View file

@ -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():

View file

@ -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()

View file

@ -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",

View file

@ -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():