diff --git a/litellm/proxy/spend_tracking/key_metadata_recovery.py b/litellm/proxy/spend_tracking/key_metadata_recovery.py index ce96dc62780..1686e055524 100644 --- a/litellm/proxy/spend_tracking/key_metadata_recovery.py +++ b/litellm/proxy/spend_tracking/key_metadata_recovery.py @@ -4,7 +4,7 @@ from collections.abc import Set as AbstractSet from dataclasses import dataclass from datetime import datetime, timedelta from types import MappingProxyType -from typing import Final, TypeVar +from typing import Final, Literal, TypeVar from pydantic import BaseModel, TypeAdapter from typing_extensions import ReadOnly, TypedDict @@ -39,26 +39,37 @@ 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 + +def _named_spend_log_edge_row_sql(direction: Literal["ASC", "DESC"]) -> str: + return f""" + SELECT key_alias, team_id, user_id + FROM ( + SELECT "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 = keys.digest + AND "startTime" >= $2::timestamp + AND "startTime" < $3::timestamp + ) named + WHERE COALESCE(key_alias, user_id, team_id) IS NOT NULL + ORDER BY "startTime" {direction} + LIMIT 1 + """ + + +_SPEND_LOG_ALIAS_SQL: Final = f""" +SELECT keys.digest, + first_row.key_alias AS first_alias, + last_row.key_alias AS last_alias, + first_row.team_id AS first_team, + last_row.team_id AS last_team, + first_row.user_id AS first_owner, + last_row.user_id AS last_owner +FROM unnest($1::text[]) AS keys(digest) +CROSS JOIN LATERAL ({_named_spend_log_edge_row_sql("ASC")}) first_row +CROSS JOIN LATERAL ({_named_spend_log_edge_row_sql("DESC")}) last_row """ _SPEND_LOG_STATEMENT_TIMEOUT_SQL: Final = f"SET LOCAL statement_timeout = {SPEND_LOG_KEY_METADATA_QUERY_TIMEOUT_MS}" @@ -84,7 +95,9 @@ class _TokenDigestRow(BaseModel): def _unanimous(first: str | None, last: str | None) -> str | None: - return first if first == last else None + if first is None: + return last + return first if last is None or first == last else None class _SpendLogDigestRow(BaseModel): 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 acd03964bf3..bf09649a89f 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 @@ -1,12 +1,20 @@ import asyncio +import re import time -from collections.abc import Sequence +from collections.abc import Mapping, Sequence from datetime import datetime, timedelta +from pathlib import Path from types import SimpleNamespace +from typing import Final from unittest.mock import AsyncMock, MagicMock +import litellm_proxy_extras +import psycopg import pytest from prisma.errors import PrismaError +from psycopg.rows import dict_row +from psycopg.types.json import Jsonb +from pytest_postgresql import factories from litellm.caching.in_memory_cache import InMemoryCache from litellm.constants import ( @@ -592,6 +600,147 @@ async def test_recover_key_metadata_from_spend_logs_bounds_the_scan_with_a_state ) +_spend_logs_postgresql_proc: Final = factories.postgresql_proc() +_spend_logs_postgresql: Final = factories.postgresql("_spend_logs_postgresql_proc") + +_SPEND_LOGS_DDL: Final = """ + CREATE TABLE "LiteLLM_SpendLogs" ( + request_id TEXT PRIMARY KEY, + api_key TEXT NOT NULL DEFAULT '', + "startTime" TIMESTAMP(3) NOT NULL, + "user" TEXT DEFAULT '', + team_id TEXT, + metadata JSONB DEFAULT '{}' + ) +""" + +_API_KEY_START_TIME_INDEX_MIGRATION: Final = ( + Path(litellm_proxy_extras.__file__).parent + / "migrations" + / "20260823000000_add_spend_logs_api_key_starttime_index" + / "migration.sql" +) + +_SPEND_LOG_ROWS_READ_IN_THIS_TRANSACTION_SQL: Final = """ + SELECT COALESCE(seq_tup_read, 0) + COALESCE(idx_tup_fetch, 0) AS rows_read + FROM pg_stat_xact_user_tables + WHERE relname = 'LiteLLM_SpendLogs' +""" + + +def _create_spend_logs_table(conn: psycopg.Connection) -> None: + conn.execute(_SPEND_LOGS_DDL) # pyright: ignore[reportArgumentType] # DDL literal + conn.execute(_API_KEY_START_TIME_INDEX_MIGRATION.read_text()) # pyright: ignore[reportArgumentType] # migration file + + +def _psycopg_prisma(conn: psycopg.Connection) -> MagicMock: + async def query_raw(sql: str, *params: object) -> list[dict[str, object]]: + with conn.cursor(row_factory=dict_row) as cur: + cur.execute( + re.sub(r"\$(\d+)", r"%(p\1)s", sql), # pyright: ignore[reportArgumentType] # proxy SQL is not a literal + {f"p{i}": v for i, v in enumerate(params, start=1)}, + ) + return cur.fetchall() + + async def execute_raw(sql: str) -> int: + conn.execute(sql) # pyright: ignore[reportArgumentType] # proxy SQL is not a literal + return 0 + + mock_prisma: Final = MagicMock() + transaction: Final = MagicMock() + transaction.query_raw = AsyncMock(side_effect=query_raw) + transaction.execute_raw = AsyncMock(side_effect=execute_raw) + mock_prisma.db.tx.return_value.__aenter__.return_value = transaction + return mock_prisma + + +def _named_spend_log( + digest: str, logged_at: datetime, alias: str | None, user: str | None, team: str | None = None +) -> tuple[str, str, datetime, str, str | None, Jsonb]: + return ( + f"{digest}-{logged_at.isoformat()}", + digest, + logged_at, + user or "", + team, + Jsonb({"user_api_key_alias": alias} if alias else {}), + ) + + +@pytest.mark.asyncio +async def test_recover_key_metadata_from_spend_logs_names_a_key_by_its_oldest_and_newest_named_rows_in_the_window( + _spend_logs_postgresql: psycopg.Connection, +): + conn: Final = _spend_logs_postgresql + _create_spend_logs_table(conn) + unnamed_edges, owner_logged_late, reowned, outside_window, never_named = ( + hash_token(f"cli-session-{name}") for name in ("edges", "late", "reowned", "window", "never") + ) + with conn.cursor() as cur: + cur.executemany( + 'INSERT INTO "LiteLLM_SpendLogs" (request_id, api_key, "startTime", "user", team_id, metadata)' + " VALUES (%s, %s, %s, %s, %s, %s)", + ( + _named_spend_log(unnamed_edges, datetime(2026, 9, 7, 1), None, None), + _named_spend_log(unnamed_edges, datetime(2026, 9, 8), "cli-a", "alice", "team-a"), + _named_spend_log(unnamed_edges, datetime(2026, 9, 9), "cli-a", "alice", "team-a"), + _named_spend_log(unnamed_edges, datetime(2026, 9, 9, 23), None, None), + _named_spend_log(owner_logged_late, datetime(2026, 9, 7, 1), "cli-b", None), + _named_spend_log(owner_logged_late, datetime(2026, 9, 9), "cli-b", "bob"), + _named_spend_log(reowned, datetime(2026, 9, 7, 1), "cli-c", "carol"), + _named_spend_log(reowned, datetime(2026, 9, 9), "cli-c", "dave"), + _named_spend_log(outside_window, datetime(2026, 9, 6), "stale-alias", "erin"), + _named_spend_log(outside_window, datetime(2026, 9, 8), "cli-d", "erin"), + _named_spend_log(outside_window, datetime(2026, 9, 10), "later-alias", "erin"), + _named_spend_log(never_named, datetime(2026, 9, 8), None, None), + ), + ) + conn.commit() + + result = await recover_key_metadata_from_spend_logs( + _psycopg_prisma(conn), + {unnamed_edges, owner_logged_late, reowned, outside_window, never_named}, + (datetime(2026, 9, 7), datetime(2026, 9, 10)), + cache=InMemoryCache(), + ) + + assert dict(result) == { + unnamed_edges: {"key_alias": "cli-a", "team_id": "team-a", "user_id": "alice"}, + owner_logged_late: {"key_alias": "cli-b", "team_id": None, "user_id": "bob"}, + reowned: {"key_alias": "cli-c", "team_id": None, "user_id": None}, + outside_window: {"key_alias": "cli-d", "team_id": None, "user_id": "erin"}, + } + + +@pytest.mark.asyncio +async def test_recover_key_metadata_from_spend_logs_reads_two_rows_per_key_however_many_the_key_logged( + _spend_logs_postgresql: psycopg.Connection, +): + conn: Final = _spend_logs_postgresql + _create_spend_logs_table(conn) + owners: Final[Mapping[str, str]] = {hash_token(f"cli-session-busy-{i}"): f"user-{i}" for i in range(5)} + for digest, owner in owners.items(): + conn.execute( + """ + INSERT INTO "LiteLLM_SpendLogs" (request_id, api_key, "startTime", "user", metadata) + SELECT %(digest)s || '-' || g, %(digest)s, %(start)s + g * interval '1 minute', %(owner)s, + jsonb_build_object('user_api_key_alias', 'cli-session-' || %(owner)s) + FROM generate_series(1, 2000) g + """, + {"digest": digest, "owner": owner, "start": datetime(2026, 9, 7)}, + ) + conn.execute('ANALYZE "LiteLLM_SpendLogs"') + conn.commit() + + result = await recover_key_metadata_from_spend_logs( + _psycopg_prisma(conn), frozenset(owners), (datetime(2026, 9, 7), datetime(2026, 9, 10)), cache=InMemoryCache() + ) + + assert {digest: meta.get("user_id") for digest, meta in result.items()} == owners + rows_read: Final = conn.execute(_SPEND_LOG_ROWS_READ_IN_THIS_TRANSACTION_SQL).fetchone() # pyright: ignore[reportArgumentType] # SQL literal + assert rows_read is not None and rows_read[0] <= 2 * len(owners) + + @pytest.mark.asyncio async def test_recover_cli_session_key_metadata_names_the_owner_only_when_the_suffix_is_a_real_user(): mock_prisma = MagicMock()