mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
Merge 64be9089fd into 431ecd8920
This commit is contained in:
commit
5717aa3b9b
4 changed files with 140 additions and 13 deletions
|
|
@ -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)
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
"""
|
||||
|
|
|
|||
103
tests/integration/database/test_cleanup_transactions.py
Normal file
103
tests/integration/database/test_cleanup_transactions.py
Normal 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)))
|
||||
|
|
@ -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):
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue