Merge pull request #40275 from BerriAI/litellm_lit6852_spend_attribution

fix(spend-tracking): recover key alias for session tokens from spend logs
This commit is contained in:
Mateo Wang 2026-09-08 19:05:43 -07:00 committed by GitHub
commit f8e456d105
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
5 changed files with 723 additions and 17 deletions

View file

@ -1767,6 +1767,10 @@ LITELLM_SETTINGS_SAFE_DB_OVERRIDES: Final = [
SPECIAL_LITELLM_AUTH_TOKEN: Final = ["ui-token"]
DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL = int(os.getenv("DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL", 60))
DEFAULT_ACCESS_GROUP_CACHE_TTL: Final = int(os.getenv("DEFAULT_ACCESS_GROUP_CACHE_TTL", 600))
SPEND_LOG_KEY_METADATA_CACHE_TTL: Final = 600
SPEND_LOG_KEY_METADATA_MISS_CACHE_TTL: Final = 30
SPEND_LOG_KEY_METADATA_CACHE_MAX_ITEMS: Final = 10000
SPEND_LOG_KEY_METADATA_QUERY_TIMEOUT_MS: Final = 5000
# Short TTL for negative MCP access-group existence lookups. Keeps unauthenticated
# callers from forcing a DB query per request for unknown names, while bounding
# staleness so a transient DB error (which surfaces as an empty list) cannot

View file

@ -14,6 +14,7 @@ from litellm.proxy._types import CommonProxyErrors
from litellm.proxy.spend_tracking.key_metadata_recovery import (
attach_user_emails,
recover_double_hashed_key_metadata,
recover_key_metadata_from_spend_logs,
)
from litellm.proxy.spend_tracking.ptu_feature_flag import is_ptu_cost_attribution_enabled
from litellm.proxy.utils import PrismaClient
@ -433,9 +434,29 @@ def update_breakdown_metrics(
return breakdown
def _spend_logs_window(dates: AbstractSet[str | None]) -> tuple[datetime, datetime] | None:
parsed: Final = sorted(day for day in (_parse_spend_date(raw) for raw in dates) if day is not None)
if not parsed:
return None
return (parsed[0] - timedelta(days=1), parsed[-1] + timedelta(days=2))
def _parse_spend_date(raw: str | None) -> datetime | None:
if not isinstance(raw, str):
return None
try:
return datetime.fromisoformat(raw)
except ValueError:
return None
_EMPTY_KEY_METADATA: Final[Mapping[str, _KeyMetadataDict]] = MappingProxyType({})
async def get_api_key_metadata(
prisma_client: PrismaClient,
api_keys: AbstractSet[str],
spend_logs_window: tuple[datetime, datetime] | None = None,
) -> Mapping[str, _KeyMetadataDict]:
"""Get api key metadata, falling back to deleted keys table for keys not found in active table.
@ -481,11 +502,17 @@ async def get_api_key_metadata(
)
still_missing: Final = api_keys - frozenset(result)
combined: Final = (
result
if not still_missing
else MappingProxyType({**result, **(await recover_double_hashed_key_metadata(prisma_client, still_missing))})
from_reverse_hash: Final = (
await recover_double_hashed_key_metadata(prisma_client, still_missing) if still_missing else _EMPTY_KEY_METADATA
)
after_token_recovery: Final = MappingProxyType({**result, **from_reverse_hash})
unresolved: Final = api_keys - frozenset(after_token_recovery)
from_spend_logs: Final = (
await recover_key_metadata_from_spend_logs(prisma_client, unresolved, spend_logs_window)
if unresolved and spend_logs_window is not None
else _EMPTY_KEY_METADATA
)
combined: Final = MappingProxyType({**after_token_recovery, **from_spend_logs})
return await attach_user_emails(prisma_client, combined)
@ -898,7 +925,9 @@ async def _aggregate_spend_records(
api_key_metadata: dict[str, _KeyMetadataDict] = {}
if api_keys:
api_key_metadata = await get_api_key_metadata(prisma_client, api_keys)
api_key_metadata = await get_api_key_metadata(
prisma_client, api_keys, _spend_logs_window(frozenset(record.date for record in records))
)
return await asyncio.to_thread(
_aggregate_spend_records_sync,
@ -1094,7 +1123,9 @@ async def _aggregate_grouping_sets_records(
api_key_metadata: dict[str, _KeyMetadataDict] = {}
if api_keys:
api_key_metadata = await get_api_key_metadata(prisma_client, api_keys)
api_key_metadata = await get_api_key_metadata(
prisma_client, api_keys, _spend_logs_window(frozenset(r.date for r in records))
)
return await asyncio.to_thread(
_aggregate_grouping_sets_records_sync,
@ -1357,7 +1388,9 @@ async def get_daily_activity_aggregated(
r.api_key for r in entity_records if r.api_key and r.api_key != PTU_SENTINEL_API_KEY
)
entity_key_metadata: Final = (
await get_api_key_metadata(prisma_client, entity_api_keys)
await get_api_key_metadata(
prisma_client, entity_api_keys, _spend_logs_window(frozenset(r.date for r in entity_records))
)
if entity_api_keys
else {} # mutable-ok: matches the helper's dict return
)

View file

@ -1,5 +1,7 @@
import asyncio
from collections.abc import Awaitable, Callable, Mapping, Sequence
from collections.abc import Set as AbstractSet
from datetime import datetime, timedelta
from types import MappingProxyType
from typing import Final, TypeVar
@ -7,6 +9,13 @@ from pydantic import BaseModel, TypeAdapter
from typing_extensions import ReadOnly, TypedDict
from litellm._logging import verbose_proxy_logger
from litellm.caching.in_memory_cache import InMemoryCache
from litellm.constants import (
SPEND_LOG_KEY_METADATA_CACHE_MAX_ITEMS,
SPEND_LOG_KEY_METADATA_CACHE_TTL,
SPEND_LOG_KEY_METADATA_MISS_CACHE_TTL,
SPEND_LOG_KEY_METADATA_QUERY_TIMEOUT_MS,
)
from litellm.litellm_core_utils.litellm_logging import is_valid_sha256_hash
from litellm.proxy.utils import PrismaClient
from litellm.repositories.user_repository import UserRepository
@ -27,6 +36,33 @@ WHERE encode(sha256(convert_to(token, 'UTF8')), 'hex') = ANY($1::text[])
ORDER BY token, deleted_at DESC
"""
_SPEND_LOG_ALIAS_SQL: Final = """
SELECT api_key AS digest,
MIN(key_alias) AS first_alias,
MAX(key_alias) AS last_alias,
MIN(team_id) AS first_team,
MAX(team_id) AS last_team,
MIN(user_id) AS first_owner,
MAX(user_id) AS last_owner
FROM (
SELECT api_key,
NULLIF(metadata->>'user_api_key_alias', '') AS key_alias,
COALESCE(NULLIF(team_id, ''), NULLIF(metadata->>'user_api_key_team_id', '')) AS team_id,
COALESCE(NULLIF("user", ''), NULLIF(metadata->>'user_api_key_user_id', '')) AS user_id
FROM "LiteLLM_SpendLogs"
WHERE api_key = ANY($1::text[])
AND "startTime" >= $2::timestamp
AND "startTime" < $3::timestamp
) named
WHERE COALESCE(key_alias, user_id, team_id) IS NOT NULL
GROUP BY api_key
"""
_SPEND_LOG_STATEMENT_TIMEOUT_SQL: Final = f"SET LOCAL statement_timeout = {SPEND_LOG_KEY_METADATA_QUERY_TIMEOUT_MS}"
_SPEND_LOG_TRANSACTION_TIMEOUT: Final = timedelta(milliseconds=2 * SPEND_LOG_KEY_METADATA_QUERY_TIMEOUT_MS)
_HASHED_JWT_PREFIX: Final = "hashed-jwt-"
class KeyMetadataDict(TypedDict, total=False):
key_alias: ReadOnly[str | None]
@ -42,7 +78,35 @@ class _TokenDigestRow(BaseModel):
user_id: str | None = None
def _unanimous(first: str | None, last: str | None) -> str | None:
return first if first == last else None
class _SpendLogDigestRow(BaseModel):
digest: str
first_alias: str | None = None
last_alias: str | None = None
first_team: str | None = None
last_team: str | None = None
first_owner: str | None = None
last_owner: str | None = None
def metadata(self) -> KeyMetadataDict:
return KeyMetadataDict(
key_alias=_unanimous(self.first_alias, self.last_alias),
team_id=_unanimous(self.first_team, self.last_team),
user_id=_unanimous(self.first_owner, self.last_owner),
)
_TOKEN_DIGEST_ROWS: Final = TypeAdapter(tuple[_TokenDigestRow, ...])
_SPEND_LOG_DIGEST_ROWS: Final = TypeAdapter(tuple[_SpendLogDigestRow, ...])
_CACHED_KEY_METADATA: Final = TypeAdapter(KeyMetadataDict)
_SPEND_LOG_METADATA_CACHE: Final = InMemoryCache(
max_size_in_memory=SPEND_LOG_KEY_METADATA_CACHE_MAX_ITEMS,
default_ttl=SPEND_LOG_KEY_METADATA_CACHE_TTL,
)
_SPEND_LOG_QUERY_LOCK: Final = asyncio.Lock()
_EMPTY_KEY_METADATA: Final[Mapping[str, KeyMetadataDict]] = MappingProxyType({})
_EMPTY_EMAILS: Final[Mapping[str, str]] = MappingProxyType({})
@ -138,14 +202,6 @@ async def recover_double_hashed_key_metadata(
prisma_client: PrismaClient,
missing_keys: AbstractSet[str],
) -> Mapping[str, KeyMetadataDict]:
"""
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. 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
@ -168,6 +224,117 @@ async def recover_double_hashed_key_metadata(
return MappingProxyType({**from_active, **from_deleted})
def _is_spend_log_digest(key: str) -> bool:
return is_valid_sha256_hash(key.removeprefix(_HASHED_JWT_PREFIX))
def _spend_log_cache_key(digest: str, window: tuple[datetime, datetime]) -> str:
start, end = window
return f"spend_log_key_metadata:{digest}:{start.isoformat()}:{end.isoformat()}"
def _cached_spend_log_metadata(
cache: InMemoryCache,
digests: AbstractSet[str],
window: tuple[datetime, datetime],
) -> Mapping[str, KeyMetadataDict]:
return MappingProxyType(
{
digest: _CACHED_KEY_METADATA.validate_python(cached)
for digest in digests
for cached in (cache.get_cache(_spend_log_cache_key(digest, window)),)
if cached is not None
}
)
async def _spend_log_rows_within_the_statement_timeout(
prisma_client: PrismaClient,
digests: AbstractSet[str],
window: tuple[datetime, datetime],
) -> Sequence[Mapping[str, object]]:
start, end = window
async with prisma_client.db.tx(timeout=_SPEND_LOG_TRANSACTION_TIMEOUT) as transaction:
await transaction.execute_raw(_SPEND_LOG_STATEMENT_TIMEOUT_SQL)
return await transaction.query_raw(_SPEND_LOG_ALIAS_SQL, sorted(digests), start, end)
async def _query_spend_log_metadata(
prisma_client: PrismaClient,
digests: AbstractSet[str],
window: tuple[datetime, datetime],
) -> Mapping[str, KeyMetadataDict] | None:
rows: Final = await _db_or_empty(
lambda: _spend_log_rows_within_the_statement_timeout(prisma_client, digests, window),
"Failed spend-log alias recovery for %d missing keys: %s",
len(digests),
)
if rows is None:
return None
return MappingProxyType(
{
row.digest: meta
for row in _SPEND_LOG_DIGEST_ROWS.validate_python(rows)
for meta in (row.metadata(),)
if row.digest in digests and any(meta.values())
}
)
def _remember_spend_log_metadata(
cache: InMemoryCache, digest: str, window: tuple[datetime, datetime], meta: KeyMetadataDict | None
) -> None:
key: Final = _spend_log_cache_key(digest, window)
if meta is not None:
cache.set_cache(key, meta)
return
missed_before: Final = f"{key}:missed-before"
if cache.get_cache(missed_before) is not None:
cache.set_cache(key, KeyMetadataDict())
return
cache.set_cache(key, KeyMetadataDict(), ttl=SPEND_LOG_KEY_METADATA_MISS_CACHE_TTL)
cache.set_cache(missed_before, True)
async def _spend_log_metadata_one_query_at_a_time(
prisma_client: PrismaClient,
cache: InMemoryCache,
lock: asyncio.Lock,
digests: AbstractSet[str],
window: tuple[datetime, datetime],
) -> Mapping[str, KeyMetadataDict]:
async with lock:
settled: Final = _cached_spend_log_metadata(cache, digests, window)
pending: Final = digests - frozenset(settled)
fresh: Final = (
await _query_spend_log_metadata(prisma_client, pending, window) if pending else _EMPTY_KEY_METADATA
)
found: Final = fresh if fresh is not None else _EMPTY_KEY_METADATA
for digest in pending:
_remember_spend_log_metadata(cache, digest, window, found.get(digest))
return MappingProxyType({**settled, **found})
async def recover_key_metadata_from_spend_logs(
prisma_client: PrismaClient,
missing_keys: AbstractSet[str],
window: tuple[datetime, datetime],
cache: InMemoryCache = _SPEND_LOG_METADATA_CACHE,
lock: asyncio.Lock = _SPEND_LOG_QUERY_LOCK,
) -> Mapping[str, KeyMetadataDict]:
digests: Final = frozenset(key for key in missing_keys if _is_spend_log_digest(key))
if not digests:
return _EMPTY_KEY_METADATA
cached: Final = _cached_spend_log_metadata(cache, digests, window)
uncached: Final = digests - frozenset(cached)
settled: Final = (
await _spend_log_metadata_one_query_at_a_time(prisma_client, cache, lock, uncached, window)
if uncached
else _EMPTY_KEY_METADATA
)
return MappingProxyType({digest: meta for digest, meta in (*cached.items(), *settled.items()) if meta})
def _row_with_recovered_fields(
row: Mapping[str, object],
recovered: Mapping[str, KeyMetadataDict],

View file

@ -492,7 +492,7 @@ async def test_get_api_key_metadata_recovers_double_hashed_key_via_reverse_hash(
@pytest.mark.asyncio
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."""
"""Without a spend-log window a dirty key no table can explain costs two digest lookups and never a token page walk."""
from litellm.proxy.utils import hash_token
double_hashed = hash_token("b" * 64)
@ -518,6 +518,93 @@ async def test_get_api_key_metadata_permanent_miss_never_pages_tokens_or_reads_s
assert all("take" not in call.kwargs and "skip" not in call.kwargs for call in token_lookups)
def _spend_log_transaction(mock_prisma: MagicMock, rows: list[dict[str, str | None]]) -> AsyncMock:
transaction = MagicMock()
transaction.execute_raw = AsyncMock(return_value=0)
transaction.query_raw = AsyncMock(return_value=rows)
mock_prisma.db.tx.return_value.__aenter__.return_value = transaction
return transaction.query_raw
def _spend_log_row(digest: str, key_alias: str, user_id: str) -> dict[str, str | None]:
return {
"digest": digest,
"first_alias": key_alias,
"last_alias": key_alias,
"first_team": None,
"last_team": None,
"first_owner": user_id,
"last_owner": user_id,
}
@pytest.mark.asyncio
async def test_get_api_key_metadata_permanent_miss_with_a_window_reads_spend_logs_once_within_it():
from litellm.proxy.utils import hash_token
double_hashed = hash_token("permanent-miss-with-window-6852")
window = (datetime(2024, 1, 1), datetime(2024, 1, 4))
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.litellm_usertable.find_many = AsyncMock(return_value=[])
mock_prisma.db.query_raw = AsyncMock(return_value=[])
spend_log_query_raw = _spend_log_transaction(mock_prisma, [])
result = await get_api_key_metadata(prisma_client=mock_prisma, api_keys={double_hashed}, spend_logs_window=window)
assert double_hashed not in result
assert mock_prisma.db.query_raw.await_count == 2
((_, digests, start, end),) = [call.args for call in spend_log_query_raw.call_args_list]
assert digests == [double_hashed]
assert (start, end) == window
@pytest.mark.asyncio
async def test_get_daily_activity_recovers_a_session_key_alias_from_spend_logs_around_the_page_dates():
from litellm.proxy.utils import hash_token
session_digest = hash_token("cli-session-daily-activity-6852")
records = [_daily_user_spend_record(user_id="session-user", api_key=session_digest, spend=1.5)]
mock_prisma = MagicMock()
mock_prisma.db = MagicMock()
mock_table = MagicMock()
mock_table.count = AsyncMock(return_value=len(records))
mock_table.find_many = AsyncMock(return_value=records)
mock_prisma.db.litellm_dailyuserspend = mock_table
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="session-user", user_email="session@example.com")]
)
mock_prisma.db.query_raw = AsyncMock(return_value=[])
spend_log_query_raw = _spend_log_transaction(
mock_prisma, [_spend_log_row(session_digest, "cli-session-alias", "session-user")]
)
result = await get_daily_activity(
prisma_client=mock_prisma,
table_name="litellm_dailyuserspend",
entity_id_field="user_id",
entity_id=None,
entity_metadata_field=None,
start_date="2024-01-01",
end_date="2024-01-01",
model=None,
api_key=None,
page=1,
page_size=1000,
)
key_metadata = result.results[0].breakdown.api_keys[session_digest].metadata
assert key_metadata.key_alias == "cli-session-alias"
assert key_metadata.user_email == "session@example.com"
((_, digests, start, end),) = [call.args for call in spend_log_query_raw.call_args_list]
assert digests == [session_digest]
assert (start, end) == (datetime(2023, 12, 31), datetime(2024, 1, 3))
def test_key_metadata_includes_recovered_user_email():
from litellm.proxy.management_endpoints.common_daily_activity import _key_metadata
@ -2105,3 +2192,48 @@ async def test_get_daily_activity_aggregated_with_entity_breakdown():
# Rollups with the entity bit set must still land in their usual buckets
assert daily.breakdown.models["gpt-4o"].metrics.spend == 18.0
assert daily.breakdown.api_keys["key-1"].metrics.spend == 12.0
@pytest.mark.asyncio
async def test_get_api_key_metadata_resolves_session_key_via_spend_log_window():
from litellm.proxy.utils import hash_token
session_digest = hash_token("cli-session-user-42")
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.litellm_usertable.find_many = AsyncMock(
return_value=[SimpleNamespace(user_id="user-42", user_email="user42@example.com")]
)
mock_prisma.db.query_raw = AsyncMock(return_value=[])
spend_log_query_raw = _spend_log_transaction(
mock_prisma, [_spend_log_row(session_digest, "cli-session-user-42", "user-42")]
)
result = await get_api_key_metadata(
prisma_client=mock_prisma,
api_keys={session_digest},
spend_logs_window=(datetime(2026, 9, 7), datetime(2026, 9, 10)),
)
assert result[session_digest]["key_alias"] == "cli-session-user-42"
assert result[session_digest]["user_id"] == "user-42"
assert result[session_digest]["user_email"] == "user42@example.com"
((_, digests, start, end),) = [call.args for call in spend_log_query_raw.call_args_list]
assert digests == [session_digest]
assert (start, end) == (datetime(2026, 9, 7), datetime(2026, 9, 10))
def test_spend_logs_window_pads_min_minus_one_day_and_max_plus_two_days():
from litellm.proxy.management_endpoints.common_daily_activity import _spend_logs_window
window = _spend_logs_window({"2026-09-08", "2026-09-05", "not-a-date"})
assert window == (datetime(2026, 9, 4), datetime(2026, 9, 10))
def test_spend_logs_window_is_none_when_no_date_parses():
from litellm.proxy.management_endpoints.common_daily_activity import _spend_logs_window
assert _spend_logs_window({"garbage", ""}) is None

View file

@ -1,21 +1,60 @@
import asyncio
import time
from collections.abc import Sequence
from datetime import datetime, timedelta
from types import SimpleNamespace
from unittest.mock import AsyncMock, MagicMock
import pytest
from prisma.errors import PrismaError
from litellm.caching.in_memory_cache import InMemoryCache
from litellm.constants import (
SPEND_LOG_KEY_METADATA_CACHE_TTL,
SPEND_LOG_KEY_METADATA_MISS_CACHE_TTL,
SPEND_LOG_KEY_METADATA_QUERY_TIMEOUT_MS,
)
from litellm.proxy.spend_tracking.key_metadata_recovery import (
fill_missing_api_key_aliases,
recover_double_hashed_key_metadata,
recover_key_metadata_from_spend_logs,
)
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]:
def _digest_row(digest: str, key_alias: str | None, 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_spend_logs(rows: Sequence[dict[str, str | None]]) -> AsyncMock:
async def query_raw(sql: str, *params: object) -> list[dict[str, str | None]]:
if '"LiteLLM_SpendLogs"' in sql:
return list(rows)
raise AssertionError(f"unexpected query: {sql}")
return AsyncMock(side_effect=query_raw)
def _spend_log_row(digest: str, key_alias: str | None, team_id: str | None, user_id: str | None) -> dict[str, str | None]:
return {
"digest": digest,
"first_alias": key_alias,
"last_alias": key_alias,
"first_team": team_id,
"last_team": team_id,
"first_owner": user_id,
"last_owner": user_id,
}
def _spend_log_transaction(mock_prisma: MagicMock, query_raw: AsyncMock) -> AsyncMock:
transaction = MagicMock()
transaction.execute_raw = AsyncMock(return_value=0)
transaction.query_raw = query_raw
mock_prisma.db.tx.return_value.__aenter__.return_value = transaction
return query_raw
def _query_raw_by_table(
active_rows: Sequence[dict[str, str | None]],
deleted_rows: Sequence[dict[str, str | None]],
@ -218,3 +257,334 @@ async def test_fill_missing_api_key_aliases_skips_named_keys_that_have_no_email(
assert filled == rows
mock_prisma.db.query_raw.assert_not_called()
@pytest.mark.asyncio
async def test_recover_key_metadata_from_spend_logs_resolves_session_token_from_metadata():
session_digest = hash_token("cli-session-repro-user-6852")
window = (datetime(2026, 9, 7), datetime(2026, 9, 10))
mock_prisma = MagicMock()
query_raw = _spend_log_transaction(
mock_prisma,
_query_raw_spend_logs(
[_spend_log_row(session_digest, "cli-session-repro-user-6852", None, "repro-user-6852")]
),
)
result = await recover_key_metadata_from_spend_logs(mock_prisma, {session_digest}, window, cache=InMemoryCache())
assert result[session_digest]["key_alias"] == "cli-session-repro-user-6852"
assert result[session_digest]["user_id"] == "repro-user-6852"
((_, digests, start, end),) = [call.args for call in query_raw.call_args_list]
assert digests == [session_digest]
assert (start, end) == window
@pytest.mark.asyncio
async def test_recover_key_metadata_from_spend_logs_skips_query_when_no_missing_keys():
window = (datetime(2026, 9, 7), datetime(2026, 9, 10))
mock_prisma = MagicMock()
query_raw = _spend_log_transaction(mock_prisma, AsyncMock(return_value=[]))
result = await recover_key_metadata_from_spend_logs(mock_prisma, set(), window, cache=InMemoryCache())
assert result == {}
query_raw.assert_not_called()
@pytest.mark.asyncio
async def test_recover_key_metadata_from_spend_logs_returns_empty_on_prisma_error():
window = (datetime(2026, 9, 7), datetime(2026, 9, 10))
mock_prisma = MagicMock()
query_raw = _spend_log_transaction(mock_prisma, AsyncMock(side_effect=PrismaError("db down")))
result = await recover_key_metadata_from_spend_logs(
mock_prisma, {hash_token("cli-session-x")}, window, cache=InMemoryCache()
)
assert result == {}
@pytest.mark.asyncio
async def test_recover_key_metadata_from_spend_logs_ignores_foreign_and_all_null_rows():
wanted = hash_token("cli-session-wanted")
all_null = hash_token("cli-session-null")
foreign = hash_token("cli-session-foreign")
window = (datetime(2026, 9, 7), datetime(2026, 9, 10))
mock_prisma = MagicMock()
query_raw = _spend_log_transaction(
mock_prisma,
_query_raw_spend_logs(
[
_spend_log_row(wanted, "kept-alias", None, "owner-1"),
_spend_log_row(all_null, None, None, None),
_spend_log_row(foreign, "foreign-alias", None, "owner-2"),
]
),
)
result = await recover_key_metadata_from_spend_logs(mock_prisma, {wanted, all_null}, window, cache=InMemoryCache())
assert set(result) == {wanted}
assert result[wanted]["key_alias"] == "kept-alias"
assert result[wanted]["user_id"] == "owner-1"
@pytest.mark.asyncio
async def test_recover_key_metadata_from_spend_logs_skips_non_sha256_keys():
window = (datetime(2026, 9, 7), datetime(2026, 9, 10))
mock_prisma = MagicMock()
query_raw = _spend_log_transaction(mock_prisma, AsyncMock(return_value=[]))
result = await recover_key_metadata_from_spend_logs(
mock_prisma, {"cli-session-raw-1798", "key-hash-short"}, window, cache=InMemoryCache()
)
assert result == {}
query_raw.assert_not_called()
@pytest.mark.asyncio
async def test_recover_key_metadata_from_spend_logs_accepts_hashed_jwt_digests():
jwt_digest = f"hashed-jwt-{hash_token('jwt-subject-1')}"
window = (datetime(2026, 9, 7), datetime(2026, 9, 10))
mock_prisma = MagicMock()
query_raw = _spend_log_transaction(mock_prisma, _query_raw_spend_logs([_spend_log_row(jwt_digest, None, "team-jwt", "jwt-user")]))
result = await recover_key_metadata_from_spend_logs(mock_prisma, {jwt_digest}, window, cache=InMemoryCache())
assert result[jwt_digest]["team_id"] == "team-jwt"
assert result[jwt_digest]["user_id"] == "jwt-user"
((_, digests, _, _),) = [call.args for call in query_raw.call_args_list]
assert digests == [jwt_digest]
@pytest.mark.asyncio
async def test_recover_key_metadata_from_spend_logs_serves_repeat_lookups_from_the_cache():
found = hash_token("cli-session-found")
unknown = hash_token("cli-session-unknown")
window = (datetime(2026, 9, 7), datetime(2026, 9, 10))
cache = InMemoryCache()
mock_prisma = MagicMock()
query_raw = _spend_log_transaction(mock_prisma, _query_raw_spend_logs([_spend_log_row(found, "found-alias", None, "owner-1")]))
first = await recover_key_metadata_from_spend_logs(mock_prisma, {found, unknown}, window, cache=cache)
second = await recover_key_metadata_from_spend_logs(mock_prisma, {found, unknown}, window, cache=cache)
assert first == second
assert set(first) == {found}
assert first[found]["key_alias"] == "found-alias"
assert query_raw.await_count == 1
@pytest.mark.asyncio
async def test_recover_key_metadata_from_spend_logs_only_queries_digests_the_cache_has_not_seen():
cached_digest = hash_token("cli-session-cached")
new_digest = hash_token("cli-session-new")
window = (datetime(2026, 9, 7), datetime(2026, 9, 10))
cache = InMemoryCache()
mock_prisma = MagicMock()
query_raw = _spend_log_transaction(mock_prisma, _query_raw_spend_logs([_spend_log_row(cached_digest, "cached-alias", None, None)]))
await recover_key_metadata_from_spend_logs(mock_prisma, {cached_digest}, window, cache=cache)
query_raw = _spend_log_transaction(mock_prisma, _query_raw_spend_logs([_spend_log_row(new_digest, "new-alias", None, None)]))
result = await recover_key_metadata_from_spend_logs(mock_prisma, {cached_digest, new_digest}, window, cache=cache)
assert result[cached_digest]["key_alias"] == "cached-alias"
assert result[new_digest]["key_alias"] == "new-alias"
((_, digests, _, _),) = [call.args for call in query_raw.call_args_list]
assert digests == [new_digest]
@pytest.mark.asyncio
async def test_recover_key_metadata_from_spend_logs_rescans_when_the_window_changes():
digest = hash_token("cli-session-windowed")
cache = InMemoryCache()
mock_prisma = MagicMock()
query_raw = _spend_log_transaction(mock_prisma, _query_raw_spend_logs([]))
await recover_key_metadata_from_spend_logs(
mock_prisma, {digest}, (datetime(2026, 9, 1), datetime(2026, 9, 4)), cache=cache
)
query_raw = _spend_log_transaction(mock_prisma, _query_raw_spend_logs([_spend_log_row(digest, "later-alias", None, None)]))
result = await recover_key_metadata_from_spend_logs(
mock_prisma, {digest}, (datetime(2026, 9, 7), datetime(2026, 9, 10)), cache=cache
)
assert result[digest]["key_alias"] == "later-alias"
assert query_raw.await_count == 1
@pytest.mark.asyncio
async def test_recover_key_metadata_from_spend_logs_retries_a_failed_query_only_after_the_miss_ttl():
digest = hash_token("cli-session-retry")
window = (datetime(2026, 9, 7), datetime(2026, 9, 10))
cache = InMemoryCache(default_ttl=SPEND_LOG_KEY_METADATA_CACHE_TTL)
mock_prisma = MagicMock()
query_raw = _spend_log_transaction(mock_prisma, AsyncMock(side_effect=PrismaError("statement timeout")))
started = time.time()
assert await recover_key_metadata_from_spend_logs(mock_prisma, {digest}, window, cache=cache) == {}
query_raw = _spend_log_transaction(mock_prisma, _query_raw_spend_logs([_spend_log_row(digest, "back-online", None, None)]))
assert await recover_key_metadata_from_spend_logs(mock_prisma, {digest}, window, cache=cache) == {}
query_raw.assert_not_awaited()
miss_key = next(key for key in cache.ttl_dict if digest in key and not key.endswith(":missed-before"))
assert cache.ttl_dict[miss_key] - started <= SPEND_LOG_KEY_METADATA_MISS_CACHE_TTL + 1
cache.ttl_dict[miss_key] = time.time() - 1
result = await recover_key_metadata_from_spend_logs(mock_prisma, {digest}, window, cache=cache)
assert result[digest]["key_alias"] == "back-online"
@pytest.mark.asyncio
async def test_recover_key_metadata_from_spend_logs_drops_the_owner_of_a_digest_shared_by_several_users():
shared_ui_digest = hash_token("ui-token")
window = (datetime(2026, 9, 7), datetime(2026, 9, 10))
mock_prisma = MagicMock()
query_raw = _spend_log_transaction(
mock_prisma,
_query_raw_spend_logs(
[{**_spend_log_row(shared_ui_digest, "ui-token", "litellm-dashboard", None), "first_owner": "alice", "last_owner": "bob"}]
),
)
result = await recover_key_metadata_from_spend_logs(
mock_prisma, {shared_ui_digest}, window, cache=InMemoryCache()
)
assert result[shared_ui_digest] == {"key_alias": "ui-token", "team_id": "litellm-dashboard", "user_id": None}
@pytest.mark.asyncio
async def test_recover_key_metadata_from_spend_logs_keeps_the_owner_when_every_named_row_agrees():
digest = hash_token("cli-session-one-owner")
window = (datetime(2026, 9, 7), datetime(2026, 9, 10))
mock_prisma = MagicMock()
query_raw = _spend_log_transaction(
mock_prisma,
_query_raw_spend_logs(
[_spend_log_row(digest, None, None, "carol")]
),
)
result = await recover_key_metadata_from_spend_logs(mock_prisma, {digest}, window, cache=InMemoryCache())
assert result[digest]["user_id"] == "carol"
@pytest.mark.asyncio
async def test_recover_key_metadata_from_spend_logs_forgets_a_miss_long_before_a_hit():
found = hash_token("cli-session-found")
unknown = hash_token("cli-session-unknown")
window = (datetime(2026, 9, 7), datetime(2026, 9, 10))
cache = InMemoryCache(default_ttl=SPEND_LOG_KEY_METADATA_CACHE_TTL)
mock_prisma = MagicMock()
query_raw = _spend_log_transaction(mock_prisma, _query_raw_spend_logs([_spend_log_row(found, "found-alias", None, None)]))
started = time.time()
await recover_key_metadata_from_spend_logs(mock_prisma, {found, unknown}, window, cache=cache)
hit_expires = next(deadline for key, deadline in cache.ttl_dict.items() if found in key)
miss_expires = next(deadline for key, deadline in cache.ttl_dict.items() if unknown in key)
assert miss_expires - started <= SPEND_LOG_KEY_METADATA_MISS_CACHE_TTL + 1
assert hit_expires - started >= SPEND_LOG_KEY_METADATA_CACHE_TTL - 1
@pytest.mark.asyncio
async def test_recover_key_metadata_from_spend_logs_runs_one_query_for_concurrent_lookups():
digest = hash_token("cli-session-shared")
window = (datetime(2026, 9, 7), datetime(2026, 9, 10))
cache = InMemoryCache()
lock = asyncio.Lock()
mock_prisma = MagicMock()
async def slow_query_raw(sql: str, *params: object) -> list[dict[str, str | None]]:
await asyncio.sleep(0.01)
return [_spend_log_row(digest, "shared-alias", None, None)]
query_raw = _spend_log_transaction(mock_prisma, AsyncMock(side_effect=slow_query_raw))
results = await asyncio.gather(
*(
recover_key_metadata_from_spend_logs(mock_prisma, {digest}, window, cache=cache, lock=lock)
for _ in range(9)
)
)
assert all(result[digest]["key_alias"] == "shared-alias" for result in results)
assert query_raw.await_count == 1
@pytest.mark.asyncio
async def test_recover_key_metadata_from_spend_logs_keeps_a_repeated_miss_as_long_as_a_hit():
unknown = hash_token("cli-session-never-named")
window = (datetime(2026, 9, 7), datetime(2026, 9, 10))
cache = InMemoryCache(default_ttl=SPEND_LOG_KEY_METADATA_CACHE_TTL)
mock_prisma = MagicMock()
query_raw = _spend_log_transaction(mock_prisma, _query_raw_spend_logs([]))
await recover_key_metadata_from_spend_logs(mock_prisma, {unknown}, window, cache=cache)
first_miss_key = next(key for key in cache.ttl_dict if unknown in key and not key.endswith(":missed-before"))
cache.ttl_dict[first_miss_key] = time.time() - 1
started = time.time()
await recover_key_metadata_from_spend_logs(mock_prisma, {unknown}, window, cache=cache)
assert query_raw.await_count == 2
assert cache.ttl_dict[first_miss_key] - started >= SPEND_LOG_KEY_METADATA_CACHE_TTL - 1
@pytest.mark.asyncio
async def test_recover_key_metadata_from_spend_logs_keeps_the_owner_older_rows_agree_on_when_the_newest_is_nameless():
digest = hash_token("cli-session-owner-from-older-rows")
window = (datetime(2026, 9, 7), datetime(2026, 9, 10))
mock_prisma = MagicMock()
_spend_log_transaction(mock_prisma, _query_raw_spend_logs([_spend_log_row(digest, None, "team-x", "alice")]))
result = await recover_key_metadata_from_spend_logs(mock_prisma, {digest}, window, cache=InMemoryCache())
assert result[digest] == {"key_alias": None, "team_id": "team-x", "user_id": "alice"}
@pytest.mark.asyncio
async def test_recover_key_metadata_from_spend_logs_names_nothing_for_a_field_whose_rows_disagree():
digest = hash_token("cli-session-disagreeing-rows")
window = (datetime(2026, 9, 7), datetime(2026, 9, 10))
mock_prisma = MagicMock()
_spend_log_transaction(
mock_prisma,
_query_raw_spend_logs(
[
{
**_spend_log_row(digest, None, None, "carol"),
"first_alias": "old-alias",
"last_alias": "renamed-alias",
"first_team": "team-a",
"last_team": "team-b",
}
]
),
)
result = await recover_key_metadata_from_spend_logs(mock_prisma, {digest}, window, cache=InMemoryCache())
assert result[digest] == {"key_alias": None, "team_id": None, "user_id": "carol"}
@pytest.mark.asyncio
async def test_recover_key_metadata_from_spend_logs_bounds_the_scan_with_a_statement_timeout():
digest = hash_token("cli-session-bounded-scan")
window = (datetime(2026, 9, 7), datetime(2026, 9, 10))
mock_prisma = MagicMock()
calls: list[str] = []
transaction = MagicMock()
transaction.execute_raw = AsyncMock(side_effect=lambda sql: calls.append(sql) or 0)
transaction.query_raw = AsyncMock(side_effect=lambda sql, *args: calls.append("scan") or [])
mock_prisma.db.tx.return_value.__aenter__.return_value = transaction
await recover_key_metadata_from_spend_logs(mock_prisma, {digest}, window, cache=InMemoryCache())
assert calls == [f"SET LOCAL statement_timeout = {SPEND_LOG_KEY_METADATA_QUERY_TIMEOUT_MS}", "scan"]
assert mock_prisma.db.tx.call_args.kwargs["timeout"] == timedelta(
milliseconds=2 * SPEND_LOG_KEY_METADATA_QUERY_TIMEOUT_MS
)