diff --git a/litellm/proxy/db/db_transaction_queue/spend_log_cleanup.py b/litellm/proxy/db/db_transaction_queue/spend_log_cleanup.py index 06e4d06fca4..0f8f8f5e614 100644 --- a/litellm/proxy/db/db_transaction_queue/spend_log_cleanup.py +++ b/litellm/proxy/db/db_transaction_queue/spend_log_cleanup.py @@ -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) ) diff --git a/litellm/proxy/db/db_transaction_queue/spend_logs_partition_manager.py b/litellm/proxy/db/db_transaction_queue/spend_logs_partition_manager.py index b756bfeb6f6..c15784ba12d 100644 --- a/litellm/proxy/db/db_transaction_queue/spend_logs_partition_manager.py +++ b/litellm/proxy/db/db_transaction_queue/spend_logs_partition_manager.py @@ -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( """ diff --git a/tests/integration/database/test_cleanup_transactions.py b/tests/integration/database/test_cleanup_transactions.py new file mode 100644 index 00000000000..f33f819bc23 --- /dev/null +++ b/tests/integration/database/test_cleanup_transactions.py @@ -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))) diff --git a/tests/test_litellm/proxy/test_spend_log_cleanup.py b/tests/test_litellm/proxy/test_spend_log_cleanup.py index 46ac1234615..00d6b9e962a 100644 --- a/tests/test_litellm/proxy/test_spend_log_cleanup.py +++ b/tests/test_litellm/proxy/test_spend_log_cleanup.py @@ -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):