diff --git a/litellm/proxy/spend_tracking/key_metadata_recovery.py b/litellm/proxy/spend_tracking/key_metadata_recovery.py index 319050b87cd..d555da99397 100644 --- a/litellm/proxy/spend_tracking/key_metadata_recovery.py +++ b/litellm/proxy/spend_tracking/key_metadata_recovery.py @@ -38,20 +38,23 @@ 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 - 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 + key_alias, + team_id, + user_id, + MIN(user_id) OVER (PARTITION BY api_key) AS first_owner, + MAX(user_id) OVER (PARTITION BY api_key) AS last_owner +FROM ( + SELECT api_key, + "startTime", + 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 ORDER BY api_key, "startTime" DESC """ @@ -72,7 +75,16 @@ class _TokenDigestRow(BaseModel): user_id: str | None = None +class _SpendLogDigestRow(_TokenDigestRow): + first_owner: str | None = None + last_owner: str | None = None + + def unanimous_owner(self) -> str | None: + return self.user_id if self.first_owner == self.last_owner else None + + _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, @@ -235,9 +247,10 @@ async def _query_spend_log_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 digests and (row.key_alias or row.user_id or row.team_id) + row.digest: KeyMetadataDict(key_alias=row.key_alias, team_id=row.team_id, user_id=owner) + for row in _SPEND_LOG_DIGEST_ROWS.validate_python(rows) + for owner in (row.unanimous_owner(),) + if row.digest in digests and (row.key_alias or owner or row.team_id) } ) @@ -270,11 +283,10 @@ async def _spend_log_metadata_one_query_at_a_time( fresh: Final = ( await _query_spend_log_metadata(prisma_client, pending, window) if pending else _EMPTY_KEY_METADATA ) - if fresh is None: - return settled + found: Final = fresh if fresh is not None else _EMPTY_KEY_METADATA for digest in pending: - _remember_spend_log_metadata(cache, digest, window, fresh.get(digest)) - return MappingProxyType({**settled, **fresh}) + _remember_spend_log_metadata(cache, digest, window, found.get(digest)) + return MappingProxyType({**settled, **found}) async def recover_key_metadata_from_spend_logs( diff --git a/tests/test_litellm/proxy/spend_tracking/test_key_metadata_recovery.py b/tests/test_litellm/proxy/spend_tracking/test_key_metadata_recovery.py index 5b45777502a..682b11f5bb7 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_key_metadata_recovery.py +++ b/tests/test_litellm/proxy/spend_tracking/test_key_metadata_recovery.py @@ -386,20 +386,63 @@ async def test_recover_key_metadata_from_spend_logs_rescans_when_the_window_chan @pytest.mark.asyncio -async def test_recover_key_metadata_from_spend_logs_does_not_cache_a_failed_query(): +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() + cache = InMemoryCache(default_ttl=SPEND_LOG_KEY_METADATA_CACHE_TTL) mock_prisma = MagicMock() - mock_prisma.db.query_raw = AsyncMock(side_effect=PrismaError("db down")) + mock_prisma.db.query_raw = AsyncMock(side_effect=PrismaError("statement timeout")) + started = time.time() 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)]) + assert await recover_key_metadata_from_spend_logs(mock_prisma, {digest}, window, cache=cache) == {} + mock_prisma.db.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() + mock_prisma.db.query_raw = _query_raw_spend_logs( + [ + { + **_digest_row(shared_ui_digest, "ui-token", "litellm-dashboard", "bob"), + "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() + mock_prisma.db.query_raw = _query_raw_spend_logs( + [{**_digest_row(digest, None, None, "carol"), "first_owner": "carol", "last_owner": "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")