mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
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:
parent
82e6b84f5a
commit
2285640eea
4 changed files with 238 additions and 9 deletions
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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],
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue