mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-07 08:26:10 +00:00
fix(spend-tracking): reverse-hash dirty spend keys in Postgres instead of paging token tables
This commit is contained in:
parent
f1f0294796
commit
ce95afe2bd
10 changed files with 189 additions and 346 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue