mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-06 08:16:43 +00:00
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:
parent
4db370e851
commit
d426b99f56
2 changed files with 142 additions and 26 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue