fix(proxy): look up hashed key names with two spend log rows per key

The spend-log fallback for keys missing from the key table read every row per key to check that all named rows agreed, which passed the 5s statement timeout on busy keys even with the (api_key, startTime) index. Probe only the oldest and newest named row per key, so the lookup stays two index reads per key however much the key logged.
This commit is contained in:
mateo-berri 2026-09-28 19:01:12 -07:00
parent f4a217d005
commit dbeb4058e1
2 changed files with 185 additions and 23 deletions

View file

@ -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):

View file

@ -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()