This commit is contained in:
devin-ai-integration[bot] 2026-09-30 21:50:09 -04:00 • committed by GitHub
commit 5717aa3b9b
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 140 additions and 13 deletions

View file

@ -27,6 +27,7 @@ from litellm.proxy.db.db_transaction_queue.spend_log_cleanup_metrics import (
from litellm.proxy.db.db_transaction_queue.spend_logs_partition_manager import (
RemainingTimeoutMs,
SpendLogsPartitionManager,
bounded_tx,
)
from litellm.proxy.utils import PrismaClient
@ -305,7 +306,7 @@ class SpendLogCleanup:
fault, so the caller stops instead of retrying.
"""
timeout_ms: Final = self._timeout_ms(deadline)
async with prisma_client.db.tx() as tx:
async with bounded_tx(prisma_client, timeout_ms) as tx:
await tx.execute_raw(f"SET LOCAL statement_timeout = {timeout_ms}")
await tx.execute_raw(f"SET LOCAL lock_timeout = {timeout_ms}")
deleted_result: Final = await tx.execute_raw(delete_sql, cutoff_date, self.batch_size)
@ -329,9 +330,10 @@ class SpendLogCleanup:
LIMIT $2
) capped
"""
timeout_ms: Final = self._timeout_ms(deadline)
try:
async with prisma_client.db.tx() as tx:
await tx.execute_raw(f"SET LOCAL statement_timeout = {self._timeout_ms(deadline)}")
async with bounded_tx(prisma_client, timeout_ms) as tx:
await tx.execute_raw(f"SET LOCAL statement_timeout = {timeout_ms}")
rows: Final = _REMAINING_ROWS.validate_python(
await tx.query_raw(count_sql, cutoff_date, SPEND_LOG_CLEANUP_REMAINING_COUNT_CAP)
)

View file

@ -126,7 +126,7 @@ def select_partitions_to_drop(partitions: list[tuple[str, datetime | None]], cut
_TX_COMMIT_SLACK: Final = timedelta(seconds=5)
def _bounded_tx(prisma_client: "PrismaClient", timeout_ms: int) -> "TransactionManager":
def bounded_tx(prisma_client: "PrismaClient", timeout_ms: int) -> "TransactionManager":
"""
Open an interactive transaction that outlives the statement bound it
carries. prisma's default 5s transaction timeout would close it mid
@ -159,7 +159,7 @@ class SpendLogsPartitionManager:
if budget_ms is None:
return False
try:
async with _bounded_tx(prisma_client, budget_ms) as tx:
async with bounded_tx(prisma_client, budget_ms) as tx:
await tx.execute_raw(f"SET LOCAL statement_timeout = {budget_ms}")
rows: Final = await tx.query_raw(
"""
@ -194,7 +194,7 @@ class SpendLogsPartitionManager:
wait for the lock and statement_timeout bounds the work itself, so a
partition this run cannot get is simply left for the next one.
"""
async with _bounded_tx(prisma_client, timeout_ms) as tx:
async with bounded_tx(prisma_client, timeout_ms) as tx:
await tx.execute_raw(f"SET LOCAL statement_timeout = {timeout_ms}")
await tx.execute_raw(f"SET LOCAL lock_timeout = {timeout_ms}")
await tx.execute_raw(statement)
@ -231,7 +231,7 @@ class SpendLogsPartitionManager:
async def _list_partitions(
self, prisma_client: "PrismaClient", timeout_ms: int
) -> list[tuple[str, datetime | None]]:
async with _bounded_tx(prisma_client, timeout_ms) as tx:
async with bounded_tx(prisma_client, timeout_ms) as tx:
await tx.execute_raw(f"SET LOCAL statement_timeout = {timeout_ms}")
rows: Final = await tx.query_raw(
"""

View file

@ -0,0 +1,103 @@
import asyncio
import logging
import os
import time
import uuid
from dataclasses import dataclass
from datetime import datetime, timedelta, timezone
from typing import Final
from urllib.parse import parse_qsl, urlencode, urlsplit, urlunsplit
import psycopg
import pytest
from integration._support.database import read_rows
from prisma import Prisma
from psycopg import sql
from litellm.proxy.db.db_transaction_queue.spend_log_cleanup import SpendLogCleanup
@dataclass(frozen=True)
class PartitionConnection:
db: Prisma
async def test_delete_batch_survives_witnessed_lock_past_transaction_default(caplog: pytest.LogCaptureFixture) -> None:
caplog.set_level(logging.ERROR)
schema: Final = f"integration_{uuid.uuid4().hex}"
url: Final = os.environ["DATABASE_URL"]
parsed: Final = urlsplit(url)
scoped_url: Final = urlunsplit(
parsed._replace(query=urlencode({**dict(parse_qsl(parsed.query)), "schema": schema}))
)
table: Final = sql.Identifier(schema, "LiteLLM_SpendLogs")
now: Final = datetime.now(timezone.utc).replace(tzinfo=None)
cutoff_date: Final = now - timedelta(days=30)
with psycopg.connect(url, autocommit=True) as setup:
setup.execute(sql.SQL("CREATE SCHEMA {}").format(sql.Identifier(schema)))
try:
setup.execute(
sql.SQL('CREATE TABLE {} (request_id text, "startTime" timestamp NOT NULL)').format(table)
)
setup.execute(
sql.SQL('INSERT INTO {} VALUES (%s, %s), (%s, %s), (%s, %s), (%s, %s)').format(table),
("expired-1", now - timedelta(days=60), "expired-2", now - timedelta(days=60),
"expired-3", now - timedelta(days=60), "fresh", now),
)
database: Final = Prisma(datasource={"url": scoped_url})
await database.connect()
try:
cleanup: Final = SpendLogCleanup(
general_settings={"maximum_spend_logs_retention_period": "30d"}
)
with psycopg.connect(url) as blocker:
blocker.execute(sql.SQL("LOCK TABLE {} IN SHARE ROW EXCLUSIVE MODE").format(table))
blocker_pid: Final = blocker.info.backend_pid
operation: Final = asyncio.create_task(
cleanup._delete_old_rows_batched(
PartitionConnection(database),
cutoff_date=cutoff_date,
table_name="LiteLLM_SpendLogs",
key_columns=("request_id",),
time_column="startTime",
deadline=time.monotonic() + 30,
)
)
wait_deadline: Final = time.monotonic() + 3
try:
while True:
witnesses: Final = read_rows(
"SELECT a.pid, extract(epoch FROM "
"clock_timestamp()-a.query_start)::double precision AS age "
"FROM pg_stat_activity a WHERE %s = ANY(pg_blocking_pids(a.pid)) "
"AND a.wait_event_type = 'Lock' AND a.query LIKE '%%DELETE FROM%%'",
(blocker_pid,),
)
if witnesses:
break
assert time.monotonic() < wait_deadline, "Cleanup DELETE never reached the held lock"
await asyncio.sleep(0.02)
assert len(witnesses) == 1
held_at: Final = time.monotonic()
age: Final = float(witnesses[0]["age"])
await asyncio.sleep(max(0, 5.6 - age))
held_seconds: Final = age + time.monotonic() - held_at
assert held_seconds >= 5.5, f"Lock released before the transaction boundary: {held_seconds}"
assert not operation.done(), "Cleanup DELETE completed while its required lock was held"
except BaseException:
operation.cancel()
await asyncio.gather(operation, return_exceptions=True)
raise
finally:
blocker.rollback()
result: Final = await asyncio.wait_for(operation, timeout=5)
assert result.rows_deleted == 3
assert not any("cleanup batch failed" in record.message for record in caplog.records)
assert read_rows(
f"SELECT request_id FROM {table.as_string(setup)} ORDER BY request_id",
(),
) == [{"request_id": "fresh"}]
finally:
await database.disconnect()
finally:
setup.execute(sql.SQL("DROP SCHEMA {} CASCADE").format(sql.Identifier(schema)))

View file

@ -33,7 +33,7 @@ def _far_deadline() -> float:
return time.monotonic() + 3600
def _wire_tx(db):
def _wire_tx(db, rejected_timeout=None):
"""
Model the prisma seam the cleanup job actually uses.
@ -47,7 +47,9 @@ def _wire_tx(db):
"""
@asynccontextmanager
async def _tx():
async def _tx(**kwargs):
if kwargs.get("timeout") == rejected_timeout:
raise TypeError("transaction timeout rejected by fake")
tx = MagicMock()
async def _execute_raw(sql, *args):
@ -100,6 +102,7 @@ def test_spend_log_cleanup_cron_scheduler_integration():
a real database connection.
"""
from unittest.mock import MagicMock
from apscheduler.triggers.cron import CronTrigger
# Mock scheduler
@ -206,6 +209,25 @@ async def test_should_delete_spend_logs():
assert cleaner._should_delete_spend_logs() is False
@pytest.mark.asyncio
async def test_delete_batch_transaction_timeout_covers_statement_timeout() -> None:
transaction: Final = MagicMock()
transaction.execute_raw = AsyncMock(side_effect=[0, 0, 3])
@asynccontextmanager
async def tx(timeout: timedelta):
yield transaction
prisma_client: Final = MagicMock()
prisma_client.db.tx = MagicMock(side_effect=tx)
cleaner: Final = SpendLogCleanup(general_settings={"maximum_spend_logs_retention_period": "30d"})
cleaner.batch_timeout_seconds = 12
await cleaner._execute_delete_batch(prisma_client, "DELETE FROM table", datetime.now(), _far_deadline())
assert prisma_client.db.tx.call_args.kwargs["timeout"] >= timedelta(seconds=12)
@pytest.mark.asyncio
async def test_cleanup_old_spend_logs_batch_deletion():
from unittest.mock import AsyncMock, MagicMock
@ -781,7 +803,7 @@ def _mock_prisma_for_retention(side_effect: list) -> "MagicMock":
from unittest.mock import AsyncMock, MagicMock
client = MagicMock()
_wire_tx(client.db)
_wire_tx(client.db, rejected_timeout=timedelta(seconds=10))
client.db.execute_raw = AsyncMock(side_effect=side_effect)
return client
@ -1004,7 +1026,7 @@ async def test_each_batch_carries_a_statement_and_lock_timeout():
mock_db = MagicMock()
@asynccontextmanager
async def _tx():
async def _tx(**kwargs):
tx = MagicMock()
async def _execute_raw(sql, *args):
@ -1213,7 +1235,7 @@ async def test_the_outstanding_rows_probe_carries_a_statement_timeout():
mock_db = MagicMock()
@asynccontextmanager
async def _tx():
async def _tx(**kwargs):
tx = MagicMock()
async def _execute_raw(sql, *args):
@ -1265,7 +1287,7 @@ async def test_a_statement_timeout_is_clamped_to_the_budget_that_is_left():
client = MagicMock()
@asynccontextmanager
async def _tx():
async def _tx(**kwargs):
tx = MagicMock()
async def _execute_raw(sql, *args):