perf(spend-tracking): recover key identity in one bounded spend-log pass with a per-worker cache

The spend-log lookup for permanently unresolvable digests is back to a single DISTINCT ON scan over the requested window, keeping only rows that carry an alias, user, or team so a newer nameless row cannot hide an older named one. Results and misses are cached per worker for ten minutes keyed by digest and window, failed queries are not cached, and JWT rows keyed hashed-jwt-<sha256> now pass the digest gate. Tests cover the JWT gate, cache reuse and partial misses, window changes, error handling, the with-window guard, and the daily activity wiring
This commit is contained in:
mateo-berri 2026-09-08 13:09:43 -07:00
parent cf0488316e
commit 550e733bf3
4 changed files with 248 additions and 32 deletions

View file

@ -1763,6 +1763,8 @@ 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 = int(os.getenv("SPEND_LOG_KEY_METADATA_CACHE_TTL", "600"))
SPEND_LOG_KEY_METADATA_CACHE_MAX_ITEMS: Final = int(os.getenv("SPEND_LOG_KEY_METADATA_CACHE_MAX_ITEMS", "10000"))
# 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

@ -8,6 +8,8 @@ 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
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
@ -29,23 +31,27 @@ ORDER BY token, deleted_at DESC
"""
_SPEND_LOG_ALIAS_SQL: Final = """
SELECT newest.digest, newest.key_alias, newest.team_id, newest.user_id
FROM unnest($1::text[]) AS missing(digest)
CROSS JOIN LATERAL (
SELECT
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 = missing.digest
AND "startTime" >= $2::timestamp
AND "startTime" < $3::timestamp
ORDER BY "startTime" DESC
LIMIT 1
) AS newest
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
AND COALESCE(
metadata->>'user_api_key_alias',
NULLIF("user", ''),
metadata->>'user_api_key_user_id',
NULLIF(team_id, ''),
metadata->>'user_api_key_team_id'
) IS NOT NULL
ORDER BY api_key, "startTime" DESC
"""
_HASHED_JWT_PREFIX: Final = "hashed-jwt-"
class KeyMetadataDict(TypedDict, total=False):
key_alias: ReadOnly[str | None]
@ -62,6 +68,11 @@ class _TokenDigestRow(BaseModel):
_TOKEN_DIGEST_ROWS: Final = TypeAdapter(tuple[_TokenDigestRow, ...])
_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,
)
_EMPTY_KEY_METADATA: Final[Mapping[str, KeyMetadataDict]] = MappingProxyType({})
_EMPTY_EMAILS: Final[Mapping[str, str]] = MappingProxyType({})
@ -179,31 +190,73 @@ 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],
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,
digest: str,
window: tuple[datetime, datetime],
) -> Mapping[str, KeyMetadataDict]:
sha_missing: Final = frozenset(key for key in missing_keys if is_valid_sha256_hash(key))
if not sha_missing:
return _EMPTY_KEY_METADATA
) -> KeyMetadataDict | None:
cached: Final[object] = cache.get_cache(_spend_log_cache_key(digest, window))
return None if cached is None else _CACHED_KEY_METADATA.validate_python(cached)
async def _query_spend_log_metadata(
prisma_client: PrismaClient,
digests: AbstractSet[str],
window: tuple[datetime, datetime],
) -> Mapping[str, KeyMetadataDict] | None:
start, end = window
rows: Final = await _db_or_empty(
lambda: prisma_client.db.query_raw(_SPEND_LOG_ALIAS_SQL, sorted(sha_missing), start, end),
lambda: prisma_client.db.query_raw(_SPEND_LOG_ALIAS_SQL, sorted(digests), start, end),
"Failed spend-log alias recovery for %d missing keys: %s",
len(sha_missing),
len(digests),
)
if rows is None:
return _EMPTY_KEY_METADATA
return None
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)
if row.digest in digests and (row.key_alias or row.user_id or row.team_id)
}
)
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,
) -> 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 = MappingProxyType(
{
digest: meta
for digest in digests
for meta in (_cached_spend_log_metadata(cache, digest, window),)
if meta is not None
}
)
uncached: Final = digests - frozenset(cached)
fresh: Final = await _query_spend_log_metadata(prisma_client, uncached, window) if uncached else _EMPTY_KEY_METADATA
if fresh is not None:
for digest in uncached:
cache.set_cache(_spend_log_cache_key(digest, window), fresh.get(digest, KeyMetadataDict()))
return MappingProxyType(
{digest: meta for digest, meta in (*cached.items(), *(fresh or _EMPTY_KEY_METADATA).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,78 @@ 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)
@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=[])
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 == 3
((_, digests, start, end),) = [
call.args for call in mock_prisma.db.query_raw.call_args_list if "LiteLLM_SpendLogs" in call.args[0]
]
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")]
)
async def query_raw(sql, *params):
if "LiteLLM_SpendLogs" in sql:
return [{"digest": session_digest, "key_alias": "cli-session-alias", "team_id": None, "user_id": "session-user"}]
return []
mock_prisma.db.query_raw = AsyncMock(side_effect=query_raw)
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 mock_prisma.db.query_raw.call_args_list if "LiteLLM_SpendLogs" in call.args[0]
]
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

View file

@ -6,6 +6,7 @@ from unittest.mock import AsyncMock, MagicMock
import pytest
from prisma.errors import PrismaError
from litellm.caching.in_memory_cache import InMemoryCache
from litellm.proxy.spend_tracking.key_metadata_recovery import (
fill_missing_api_key_aliases,
recover_double_hashed_key_metadata,
@ -240,7 +241,7 @@ async def test_recover_key_metadata_from_spend_logs_resolves_session_token_from_
[_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)
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"
@ -255,7 +256,7 @@ async def test_recover_key_metadata_from_spend_logs_skips_query_when_no_missing_
mock_prisma = MagicMock()
mock_prisma.db.query_raw = AsyncMock(return_value=[])
result = await recover_key_metadata_from_spend_logs(mock_prisma, set(), window)
result = await recover_key_metadata_from_spend_logs(mock_prisma, set(), window, cache=InMemoryCache())
assert result == {}
mock_prisma.db.query_raw.assert_not_called()
@ -267,7 +268,9 @@ async def test_recover_key_metadata_from_spend_logs_returns_empty_on_prisma_erro
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)
result = await recover_key_metadata_from_spend_logs(
mock_prisma, {hash_token("cli-session-x")}, window, cache=InMemoryCache()
)
assert result == {}
@ -287,7 +290,7 @@ async def test_recover_key_metadata_from_spend_logs_ignores_foreign_and_all_null
]
)
result = await recover_key_metadata_from_spend_logs(mock_prisma, {wanted, all_null}, window)
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"
@ -301,8 +304,94 @@ async def test_recover_key_metadata_from_spend_logs_skips_non_sha256_keys():
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
mock_prisma, {"cli-session-raw-1798", "key-hash-short"}, window, cache=InMemoryCache()
)
assert result == {}
mock_prisma.db.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()
mock_prisma.db.query_raw = _query_raw_spend_logs([_digest_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 mock_prisma.db.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()
mock_prisma.db.query_raw = _query_raw_spend_logs([_digest_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 mock_prisma.db.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()
mock_prisma.db.query_raw = _query_raw_spend_logs([_digest_row(cached_digest, "cached-alias", None, None)])
await recover_key_metadata_from_spend_logs(mock_prisma, {cached_digest}, window, cache=cache)
mock_prisma.db.query_raw = _query_raw_spend_logs([_digest_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 mock_prisma.db.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()
mock_prisma.db.query_raw = _query_raw_spend_logs([])
await recover_key_metadata_from_spend_logs(
mock_prisma, {digest}, (datetime(2026, 9, 1), datetime(2026, 9, 4)), cache=cache
)
mock_prisma.db.query_raw = _query_raw_spend_logs([_digest_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 mock_prisma.db.query_raw.await_count == 1
@pytest.mark.asyncio
async def test_recover_key_metadata_from_spend_logs_does_not_cache_a_failed_query():
digest = hash_token("cli-session-retry")
window = (datetime(2026, 9, 7), datetime(2026, 9, 10))
cache = InMemoryCache()
mock_prisma = MagicMock()
mock_prisma.db.query_raw = AsyncMock(side_effect=PrismaError("db down"))
assert await recover_key_metadata_from_spend_logs(mock_prisma, {digest}, window, cache=cache) == {}
mock_prisma.db.query_raw = _query_raw_spend_logs([_digest_row(digest, "back-online", None, None)])
result = await recover_key_metadata_from_spend_logs(mock_prisma, {digest}, window, cache=cache)
assert result[digest]["key_alias"] == "back-online"