fix(spend-tracking): recover key alias for session tokens from spend logs

CLI session tokens are in-memory only and never get a LiteLLM_VerificationToken
row, so the usage APIs could not resolve key_alias, team_id, or user_email for
their spend rows: the exact join and the reverse-hash recovery both miss. The
owner is written to LiteLLM_SpendLogs.metadata at request time under the same
hashed api_key, so read it back from there for keys still unresolved after the
token-table passes.

The lookup is sha256-gated like the existing reverse-hash recovery and bounded
to the records' startTime window (min date minus one day, max date plus two) so
it stays on the startTime index. No migration.

Also guard the window parser against the date=None rollup rows GROUPING SETS
aggregation emits, which raised TypeError from strptime and turned the
aggregated usage endpoints into HTTP 500s.
This commit is contained in:
mateo-berri 2026-09-08 12:02:39 -07:00
parent 82e6b84f5a
commit 2285640eea
4 changed files with 238 additions and 9 deletions

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,15 +434,37 @@ 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.
This ensures that key_alias and team_id are preserved in historical activity logs
even after a key is deleted or regenerated. Also recovers aliases for api_key
values that were double-hashed by the v1.99 spend-log provenance gate.
values that were double-hashed by the v1.99 spend-log provenance gate, and, when
spend_logs_window is given, for keys never written to either token table (CLI
session tokens) from the spend-log rows those requests wrote in that window.
"""
key_records: Sequence[PrismaVerificationToken] = await VerificationTokenRepository(prisma_client).table.find_many(
where={"token": {"in": list(api_keys)}}
@ -481,11 +504,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 +927,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 +1125,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 +1390,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,6 @@
from collections.abc import Awaitable, Callable, Mapping, Sequence
from collections.abc import Set as AbstractSet
from datetime import datetime
from types import MappingProxyType
from typing import Final, TypeVar
@ -27,6 +28,19 @@ WHERE encode(sha256(convert_to(token, 'UTF8')), 'hex') = ANY($1::text[])
ORDER BY token, deleted_at DESC
"""
_SPEND_LOG_ALIAS_SQL: Final = """
SELECT DISTINCT ON (api_key)
api_key AS digest,
metadata->>'user_api_key_alias' AS key_alias,
COALESCE(NULLIF(team_id, ''), metadata->>'user_api_key_team_id') AS team_id,
COALESCE(NULLIF("user", ''), 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
ORDER BY api_key, (metadata->>'user_api_key_alias') IS NULL, "startTime" DESC
"""
class KeyMetadataDict(TypedDict, total=False):
key_alias: ReadOnly[str | None]
@ -168,6 +182,40 @@ async def recover_double_hashed_key_metadata(
return MappingProxyType({**from_active, **from_deleted})
async def recover_key_metadata_from_spend_logs(
prisma_client: PrismaClient,
missing_keys: AbstractSet[str],
window: tuple[datetime, datetime],
) -> Mapping[str, KeyMetadataDict]:
"""
Recover key_alias/team_id/user_id for hashed api_key values absent from both
verification-token tables, e.g. in-memory CLI session tokens that never get a
token row. Their owner is written to LiteLLM_SpendLogs metadata at request
time under the same hashed api_key, so it is the only surviving source. Only
sha256 digests are looked up, matching the reverse-hash recovery gate, since
every current api_key value in spend logs is a token hash. The [start, end)
bound keeps the lookup on the startTime index instead of scanning the table.
"""
sha_missing: Final = frozenset(key for key in missing_keys if is_valid_sha256_hash(key))
if not sha_missing:
return _EMPTY_KEY_METADATA
start, end = window
rows: Final = await _db_or_empty(
lambda: prisma_client.db.query_raw(_SPEND_LOG_ALIAS_SQL, sorted(sha_missing), start, end),
"Failed spend-log alias recovery for %d missing keys: %s",
len(sha_missing),
)
if rows is None:
return _EMPTY_KEY_METADATA
return MappingProxyType(
{
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 sha_missing and (row.key_alias or row.user_id or row.team_id)
}
)
def _row_with_recovered_fields(
row: Mapping[str, object],
recovered: Mapping[str, KeyMetadataDict],

View file

@ -2105,3 +2105,56 @@ 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():
"""
A CLI session token has no verification-token row, so the active/deleted lookups
and reverse-hash all miss. Given a spend-log window, its alias and owner are
recovered from the spend-log metadata and its email is filled from the user table.
"""
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")]
)
async def query_raw(sql, *params):
if "LiteLLM_SpendLogs" in sql:
return [{"digest": session_digest, "key_alias": "cli-session-user-42", "team_id": None, "user_id": "user-42"}]
return []
mock_prisma.db.query_raw = AsyncMock(side_effect=query_raw)
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"
spend_log_calls = [call.args for call in mock_prisma.db.query_raw.call_args_list if "LiteLLM_SpendLogs" in call.args[0]]
((_, digests, start, end),) = spend_log_calls
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,4 +1,5 @@
from collections.abc import Sequence
from datetime import datetime
from types import SimpleNamespace
from unittest.mock import AsyncMock, MagicMock
@ -8,14 +9,24 @@ from prisma.errors import PrismaError
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 _query_raw_by_table(
active_rows: Sequence[dict[str, str | None]],
deleted_rows: Sequence[dict[str, str | None]],
@ -218,3 +229,85 @@ 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():
"""
CLI session tokens never get a verification-token row, so both the exact join and
the reverse-hash lookup miss them. Their owner survives only in the spend-log
metadata written at request time, keyed by the same hashed api_key.
"""
session_digest = hash_token("cli-session-repro-user-6852")
window = (datetime(2026, 9, 7), datetime(2026, 9, 10))
mock_prisma = MagicMock()
mock_prisma.db.query_raw = _query_raw_spend_logs(
[_digest_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)
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 mock_prisma.db.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()
mock_prisma.db.query_raw = AsyncMock(return_value=[])
result = await recover_key_metadata_from_spend_logs(mock_prisma, set(), window)
assert result == {}
mock_prisma.db.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()
mock_prisma.db.query_raw = AsyncMock(side_effect=PrismaError("db down"))
result = await recover_key_metadata_from_spend_logs(mock_prisma, {hash_token("cli-session-x")}, window)
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()
mock_prisma.db.query_raw = _query_raw_spend_logs(
[
_digest_row(wanted, "kept-alias", None, "owner-1"),
_digest_row(all_null, None, None, None),
_digest_row(foreign, "foreign-alias", None, "owner-2"),
]
)
result = await recover_key_metadata_from_spend_logs(mock_prisma, {wanted, all_null}, window)
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()
mock_prisma.db.query_raw = AsyncMock(return_value=[])
result = await recover_key_metadata_from_spend_logs(
mock_prisma, {"cli-session-raw-1798", "key-hash-short"}, window
)
assert result == {}
mock_prisma.db.query_raw.assert_not_called()