fix(proxy): bound the newest-row probe at where the oldest probe stopped

The newest-row probe now starts at the row where the oldest-row probe gave up, so a key with under 200 rows in the window is read once instead of twice, and the lookup transaction turns bitmap scans off so the planner walks the (api_key, startTime) index instead of every row of a busy key when statistics or the visibility map are stale.
This commit is contained in:
mateo-berri 2026-09-29 14:44:01 -07:00
parent bea6f49c9c
commit f2d7062482
2 changed files with 152 additions and 16 deletions

View file

@ -41,9 +41,11 @@ ORDER BY token, deleted_at DESC
"""
def _named_spend_log_edge_row_sql(direction: Literal["ASC", "DESC"]) -> str:
def _named_spend_log_edge_row_sql(
direction: Literal["ASC", "DESC"], since: Literal["$2::timestamp", "oldest_probe.stopped_at"]
) -> str:
return f"""
SELECT key_alias, team_id, user_id
SELECT "startTime", key_alias, team_id, user_id
FROM (
SELECT "startTime",
NULLIF(metadata->>'user_api_key_alias', '') AS key_alias,
@ -53,7 +55,7 @@ def _named_spend_log_edge_row_sql(direction: Literal["ASC", "DESC"]) -> str:
SELECT "startTime", metadata, team_id, "user"
FROM "LiteLLM_SpendLogs"
WHERE api_key = keys.digest
AND "startTime" >= $2::timestamp
AND "startTime" >= {since}
AND "startTime" < $3::timestamp
ORDER BY "startTime" {direction}
LIMIT {SPEND_LOG_KEY_METADATA_ROWS_PER_PROBE}
@ -65,6 +67,19 @@ def _named_spend_log_edge_row_sql(direction: Literal["ASC", "DESC"]) -> str:
"""
_OLDEST_PROBE_STOPPED_AT_SQL: Final = f"""
SELECT COALESCE(first_row."startTime", (
SELECT "startTime"
FROM "LiteLLM_SpendLogs"
WHERE api_key = keys.digest
AND "startTime" >= $2::timestamp
AND "startTime" < $3::timestamp
ORDER BY "startTime" ASC
OFFSET {SPEND_LOG_KEY_METADATA_ROWS_PER_PROBE - 1}
LIMIT 1
)) AS stopped_at
"""
_SPEND_LOG_ALIAS_SQL: Final = f"""
SELECT keys.digest,
first_row.key_alias AS first_alias,
@ -74,11 +89,13 @@ SELECT keys.digest,
first_row.user_id AS first_owner,
last_row.user_id AS last_owner
FROM unnest($1::text[]) AS keys(digest)
LEFT JOIN LATERAL ({_named_spend_log_edge_row_sql("ASC")}) first_row ON true
LEFT JOIN LATERAL ({_named_spend_log_edge_row_sql("DESC")}) last_row ON true
LEFT JOIN LATERAL ({_named_spend_log_edge_row_sql("ASC", "$2::timestamp")}) first_row ON true
LEFT JOIN LATERAL ({_OLDEST_PROBE_STOPPED_AT_SQL}) oldest_probe ON true
LEFT JOIN LATERAL ({_named_spend_log_edge_row_sql("DESC", "oldest_probe.stopped_at")}) last_row ON true
"""
_SPEND_LOG_STATEMENT_TIMEOUT_SQL: Final = f"SET LOCAL statement_timeout = {SPEND_LOG_KEY_METADATA_QUERY_TIMEOUT_MS}"
_SPEND_LOG_NO_BITMAP_SCAN_SQL: Final = "SET LOCAL enable_bitmapscan = off"
_SPEND_LOG_TRANSACTION_TIMEOUT: Final = timedelta(milliseconds=2 * SPEND_LOG_KEY_METADATA_QUERY_TIMEOUT_MS)
_HASHED_JWT_PREFIX: Final = "hashed-jwt-"
@ -336,6 +353,7 @@ async def _spend_log_rows_within_the_statement_timeout(
start, end = window
async with prisma_client.db.tx(timeout=_SPEND_LOG_TRANSACTION_TIMEOUT) as transaction:
await transaction.execute_raw(_SPEND_LOG_STATEMENT_TIMEOUT_SQL)
await transaction.execute_raw(_SPEND_LOG_NO_BITMAP_SCAN_SQL)
return await transaction.query_raw(_SPEND_LOG_ALIAS_SQL, sorted(digests), start, end)

View file

@ -595,7 +595,11 @@ async def test_recover_key_metadata_from_spend_logs_bounds_the_scan_with_a_state
await recover_key_metadata_from_spend_logs(mock_prisma, {digest}, window, cache=InMemoryCache())
assert calls == [f"SET LOCAL statement_timeout = {SPEND_LOG_KEY_METADATA_QUERY_TIMEOUT_MS}", "scan"]
assert calls == [
f"SET LOCAL statement_timeout = {SPEND_LOG_KEY_METADATA_QUERY_TIMEOUT_MS}",
"SET LOCAL enable_bitmapscan = off",
"scan",
]
assert mock_prisma.db.tx.call_args.kwargs["timeout"] == timedelta(
milliseconds=2 * SPEND_LOG_KEY_METADATA_QUERY_TIMEOUT_MS
)
@ -655,6 +659,24 @@ def _psycopg_prisma(conn: psycopg.Connection) -> MagicMock:
return mock_prisma
def _commit_and_vacuum(conn: psycopg.Connection) -> None:
conn.commit()
conn.set_autocommit(True)
conn.execute('VACUUM (ANALYZE) "LiteLLM_SpendLogs"')
conn.set_autocommit(False)
def _insert_nameless_spend_logs(conn: psycopg.Connection, digest: str, rows: int) -> None:
conn.execute(
"""
INSERT INTO "LiteLLM_SpendLogs" (request_id, api_key, "startTime")
SELECT %(digest)s || '-' || g, %(digest)s, %(start)s + g * interval '1 minute'
FROM generate_series(1, %(rows)s) g
""",
{"digest": digest, "start": datetime(2026, 9, 7), "rows": rows},
)
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]:
@ -748,18 +770,10 @@ async def test_recover_key_metadata_from_spend_logs_walks_a_bounded_number_of_na
):
conn: Final = _spend_logs_postgresql
_create_spend_logs_table(conn)
nameless_rows: Final = 3 * SPEND_LOG_KEY_METADATA_ROWS_PER_PROBE
named_late: Final[Mapping[str, str]] = {hash_token(f"cli-session-late-{i}"): f"user-{i}" for i in range(3)}
never_named: Final = frozenset(hash_token(f"cli-session-never-{i}") for i in range(3))
for digest in (*named_late, *never_named):
conn.execute(
"""
INSERT INTO "LiteLLM_SpendLogs" (request_id, api_key, "startTime")
SELECT %(digest)s || '-' || g, %(digest)s, %(start)s + g * interval '1 minute'
FROM generate_series(1, %(rows)s) g
""",
{"digest": digest, "start": datetime(2026, 9, 7), "rows": nameless_rows},
)
_insert_nameless_spend_logs(conn, digest, 3 * SPEND_LOG_KEY_METADATA_ROWS_PER_PROBE)
for digest, owner in named_late.items():
conn.execute(
"""
@ -782,7 +796,111 @@ async def test_recover_key_metadata_from_spend_logs_walks_a_bounded_number_of_na
assert {digest: meta.get("user_id") for digest, meta in result.items()} == named_late
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
assert rows_read[0] <= 2 * SPEND_LOG_KEY_METADATA_ROWS_PER_PROBE * (len(named_late) + len(never_named))
assert rows_read[0] <= 3 * SPEND_LOG_KEY_METADATA_ROWS_PER_PROBE * (len(named_late) + len(never_named))
@pytest.mark.asyncio
async def test_recover_key_metadata_from_spend_logs_reads_a_short_nameless_key_once(
_spend_logs_postgresql: psycopg.Connection,
):
conn: Final = _spend_logs_postgresql
_create_spend_logs_table(conn)
rows_per_key: Final = SPEND_LOG_KEY_METADATA_ROWS_PER_PROBE // 2
never_named: Final = frozenset(hash_token(f"cli-session-short-{i}") for i in range(20))
for digest in never_named:
_insert_nameless_spend_logs(conn, digest, rows_per_key)
_commit_and_vacuum(conn)
result = await recover_key_metadata_from_spend_logs(
_psycopg_prisma(conn), never_named, (datetime(2026, 9, 7), datetime(2026, 9, 10)), cache=InMemoryCache()
)
assert dict(result) == {}
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
assert rows_read[0] <= rows_per_key * len(never_named)
@pytest.mark.asyncio
async def test_recover_key_metadata_from_spend_logs_bounds_a_busy_nameless_key_among_short_keys_before_any_vacuum(
_spend_logs_postgresql: psycopg.Connection,
):
conn: Final = _spend_logs_postgresql
_create_spend_logs_table(conn)
busy: Final = frozenset(hash_token(f"cli-session-busy-nameless-{i}") for i in range(3))
for short_key in range(200):
_insert_nameless_spend_logs(
conn, hash_token(f"cli-session-short-{short_key}"), SPEND_LOG_KEY_METADATA_ROWS_PER_PROBE // 5
)
for digest in busy:
_insert_nameless_spend_logs(conn, digest, 30 * SPEND_LOG_KEY_METADATA_ROWS_PER_PROBE)
conn.execute('ANALYZE "LiteLLM_SpendLogs"')
conn.commit()
result = await recover_key_metadata_from_spend_logs(
_psycopg_prisma(conn), busy, (datetime(2026, 9, 7), datetime(2026, 9, 10)), cache=InMemoryCache()
)
assert dict(result) == {}
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
assert rows_read[0] <= 3 * SPEND_LOG_KEY_METADATA_ROWS_PER_PROBE * len(busy)
@pytest.mark.asyncio
async def test_recover_key_metadata_from_spend_logs_finds_a_name_logged_where_the_oldest_probe_stopped(
_spend_logs_postgresql: psycopg.Connection,
):
conn: Final = _spend_logs_postgresql
_create_spend_logs_table(conn)
start: Final = datetime(2026, 9, 7)
past_the_stop, tied_with_the_stop = (hash_token(f"cli-session-{name}") for name in ("past", "tied"))
same_millisecond: Final = tuple(
start + timedelta(minutes=SPEND_LOG_KEY_METADATA_ROWS_PER_PROBE, microseconds=n) for n in (100, 200, 300)
)
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(past_the_stop, start + timedelta(minutes=minute), None, None)
for minute in range(1, SPEND_LOG_KEY_METADATA_ROWS_PER_PROBE + 20)
),
_named_spend_log(
past_the_stop,
start + timedelta(minutes=SPEND_LOG_KEY_METADATA_ROWS_PER_PROBE + 20),
"cli-p",
"pat",
),
*(
_named_spend_log(past_the_stop, start + timedelta(minutes=minute), None, None)
for minute in range(
SPEND_LOG_KEY_METADATA_ROWS_PER_PROBE + 21, SPEND_LOG_KEY_METADATA_ROWS_PER_PROBE + 51
)
),
*(
_named_spend_log(tied_with_the_stop, start + timedelta(minutes=minute), None, None)
for minute in range(1, SPEND_LOG_KEY_METADATA_ROWS_PER_PROBE)
),
_named_spend_log(tied_with_the_stop, same_millisecond[0], None, None),
_named_spend_log(tied_with_the_stop, same_millisecond[1], None, None),
_named_spend_log(tied_with_the_stop, same_millisecond[2], "cli-t", "tess"),
),
)
conn.commit()
result = await recover_key_metadata_from_spend_logs(
_psycopg_prisma(conn),
{past_the_stop, tied_with_the_stop},
(start, datetime(2026, 9, 10)),
cache=InMemoryCache(),
)
assert dict(result) == {
past_the_stop: {"key_alias": "cli-p", "team_id": None, "user_id": "pat"},
tied_with_the_stop: {"key_alias": "cli-t", "team_id": None, "user_id": "tess"},
}
@pytest.mark.asyncio