fix(spend): page reverse-hash recovery past the first 10k keys

Historical dirty spend on large installs was still unlabeled when the
matching token sat past the first page. Keep scanning until the digest
matches or the table ends.

Co-authored-by: Mateo Wang <mateo-berri@users.noreply.github.com>
This commit is contained in:
Cursor Agent 2026-09-03 15:44:18 +00:00
parent 4db370e851
commit d426b99f56
No known key found for this signature in database
2 changed files with 142 additions and 26 deletions

View file

@ -16,7 +16,7 @@ from litellm.repositories.verification_token_repository import (
_T = TypeVar("_T")
_MAX_DOUBLE_HASH_TOKEN_SCAN: Final = 10_000
_TOKEN_SCAN_PAGE: Final = 10_000
_SPEND_LOGS_KEY_METADATA_SQL: Final = """
SELECT DISTINCT ON (api_key)
@ -108,46 +108,86 @@ def _token_digest_metadata(
)
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]:
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 _EMPTY_KEY_METADATA
return _token_digest_metadata(active_records, wanted)
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]:
deleted_records: Final = await _db_or_empty(
lambda: DeletedVerificationTokenRepository(prisma_client).table.find_many(
take=_MAX_DOUBLE_HASH_TOKEN_SCAN,
order={"deleted_at": "desc"}, # mutable-ok: Prisma find_many order= is a dict
),
"Failed reverse-hash recovery against deleted keys for %d missing keys: %s",
len(wanted),
)
if deleted_records is None:
return _EMPTY_KEY_METADATA
return _token_digest_metadata(deleted_records, wanted)
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,
wanted: AbstractSet[str],
*,
page_size: int,
) -> Mapping[str, KeyMetadataDict]:
from_active: Final = await _reverse_hash_active_key_metadata(prisma_client, wanted)
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)
from_deleted: Final = await _reverse_hash_deleted_key_metadata(prisma_client, still_wanted, page_size=page_size)
return MappingProxyType({**from_active, **from_deleted})
@ -228,21 +268,24 @@ 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
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. Prefer a bounded reverse-hash against active/deleted
tokens; fall back to SpendLogs metadata. Emails come from SpendLogs when
present, otherwise from UserTable via the recovered key's user_id.
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.
"""
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)
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

View file

@ -111,3 +111,76 @@ async def test_recover_falls_back_to_spend_logs_when_token_scan_raises_prisma_er
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=[])
result = await recover_double_hashed_key_metadata(mock_prisma, {double_hashed}, token_scan_page_size=2)
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,
]
mock_prisma.db.query_raw.assert_not_called()