mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
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:
parent
f4a217d005
commit
dbeb4058e1
2 changed files with 185 additions and 23 deletions
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue