Merge pull request #41213 from BerriAI/litellm_spend_log_cleanup_cancel_outcome

fix(proxy): record aborted outcome when spend-log cleanup is cancelled at shutdown
This commit is contained in:
yucheng-berri 2026-09-21 17:03:25 -07:00 • committed by GitHub
commit 99e284106d
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
6 changed files with 423 additions and 3 deletions

View file

@ -1749,6 +1749,12 @@ SPEND_LOG_CLEANUP_BATCH_FAILURE_BACKOFF_SECONDS: Final = float(
SPEND_LOG_CLEANUP_RUN_BUDGET_SECONDS: Final = float(os.getenv("SPEND_LOG_CLEANUP_RUN_BUDGET_SECONDS", "300"))
SPEND_LOG_CLEANUP_BATCH_TIMEOUT_SECONDS: Final = float(os.getenv("SPEND_LOG_CLEANUP_BATCH_TIMEOUT_SECONDS", "30"))
SPEND_LOG_CLEANUP_REMAINING_COUNT_CAP: Final = int(os.getenv("SPEND_LOG_CLEANUP_REMAINING_COUNT_CAP", "100000"))
SCHEDULED_JOB_SHUTDOWN_FINISH_TIMEOUT_SECONDS: Final = float(
os.getenv("SCHEDULED_JOB_SHUTDOWN_FINISH_TIMEOUT_SECONDS", "5")
)
SCHEDULED_JOB_SHUTDOWN_CANCEL_TIMEOUT_SECONDS: Final = float(
os.getenv("SCHEDULED_JOB_SHUTDOWN_CANCEL_TIMEOUT_SECONDS", "5")
)
TOOL_SPEND_TOP_TOOLS: Final = 100
SPEND_LOG_PARTITION_INTERVAL: Final = os.getenv("SPEND_LOG_PARTITION_INTERVAL", "day")
SPEND_LOG_PARTITION_PRECREATE_AHEAD: Final = int(os.getenv("SPEND_LOG_PARTITION_PRECREATE_AHEAD", 7))

View file

@ -1,5 +1,6 @@
import asyncio
import time
from contextvars import ContextVar
from dataclasses import dataclass
from datetime import datetime, timedelta, timezone
from typing import Final, Literal, TypeAlias
@ -40,6 +41,28 @@ class TableCleanupResult:
stop_reason: StopReason
class _RunProgress:
"""How far one cleanup run has got, reported if that run is cancelled"""
def __init__(self) -> None:
self.rows_deleted: int = 0
self.batches: int = 0
def record_batch(self, rows_deleted: int) -> None:
self.rows_deleted += rows_deleted
self.batches += 1
_run_progress: ContextVar[_RunProgress] = ContextVar("spend_log_cleanup_run_progress")
def _record_run_batch(rows_deleted: int) -> None:
"""Count a batch towards the run in progress, if a run is what issued it"""
progress: Final = _run_progress.get(None)
if progress is not None:
progress.record_batch(rows_deleted)
class _RemainingRow(BaseModel):
"""One row of the capped outstanding-rows probe, validated out of prisma's untyped result."""
@ -422,6 +445,7 @@ class SpendLogCleanup:
total_deleted += deleted_count
run_count += 1
_record_run_batch(deleted_count)
# Add a small sleep to prevent overwhelming the database
await asyncio.sleep(0.1)
@ -621,6 +645,9 @@ class SpendLogCleanup:
If no pod_lock_manager, runs cleanup without distributed locking.
"""
lock_acquired = False
run_started_at: Final = time.monotonic()
progress: Final = _RunProgress()
progress_token: Final = _run_progress.set(progress)
try:
verbose_proxy_logger.info("Cleanup job triggered at %s", datetime.now())
self._refresh_bounds()
@ -701,6 +728,15 @@ class SpendLogCleanup:
self._run_outcome(spend_log_results + session_results + health_check_results)
)
except asyncio.CancelledError:
verbose_proxy_logger.error(
"Spend log cleanup cancelled after %.2fs (rows_deleted=%d, batches=%d); the next run resumes from here",
time.monotonic() - run_started_at,
progress.rows_deleted,
progress.batches,
)
SpendLogCleanupMetrics.record_run("aborted")
raise
except Exception as e:
# .exception() captures the traceback; str(e) alone on a Prisma/DB
# timeout is often empty and gives operators no signal to diagnose.
@ -712,6 +748,7 @@ class SpendLogCleanup:
SpendLogCleanupMetrics.record_run("aborted")
return # Return after error handling
finally:
_run_progress.reset(progress_token)
# Only release the lock if it was actually acquired
if lock_acquired and self.pod_lock_manager and self.pod_lock_manager.redis_cache:
await self.pod_lock_manager.release_lock(cronjob_id=SPEND_LOG_CLEANUP_JOB_NAME)

View file

@ -722,6 +722,11 @@ from litellm.proxy.route_llm_request import route_request
from litellm.proxy.route_priority import hot_routes_first
from litellm.proxy.search_endpoints.endpoints import router as search_router
from litellm.proxy.shutdown.graceful_shutdown_manager import GracefulShutdownManager
from litellm.proxy.shutdown.scheduled_jobs import (
AwaitableAsyncIOExecutor,
pause_scheduled_jobs,
stop_in_flight_scheduler_jobs,
)
from litellm.proxy.spend_tracking.budget_reservation import get_budget_window_start
from litellm.proxy.spend_tracking.daily_global_spend_rollup import (
run_scheduled_daily_global_spend_reconcile,
@ -1495,6 +1500,10 @@ async def proxy_startup_event(app: FastAPI) -> AsyncGenerator[None, None]:
if model_info_scheduler is not scheduler:
model_info_scheduler.shutdown(wait=False)
# Shutdown event - stop starting scheduled jobs; the ones already running keep the drain window
if scheduler is not None:
pause_scheduled_jobs(scheduler)
# Shutdown event - drain in-flight requests before tearing down dependencies
# so SIGTERM (rolling update, scale-down, liveness kill) doesn't drop them.
GracefulShutdownManager.start_shutdown()
@ -1534,6 +1543,13 @@ async def proxy_startup_event(app: FastAPI) -> AsyncGenerator[None, None]:
await _drain_spend_event_producer_on_shutdown()
# Shutdown event - finish or cancel in-flight scheduled jobs before the shutdown flushes and the DB disconnect
if scheduler is not None and scheduler_executor is not None:
try:
await stop_in_flight_scheduler_jobs(scheduler, scheduler_executor)
except Exception as e:
verbose_proxy_logger.error("Error stopping in-flight scheduled jobs: %s", e)
await flush_spend_counters_on_shutdown()
await _flush_spend_logs_queue_on_shutdown()
@ -2536,6 +2552,7 @@ celery_app_conn: Final = None
celery_fn: Final = None # Redis Queue for handling requests
scheduler = None
scheduler_executor: AwaitableAsyncIOExecutor | None = None # rebind-ok: bound once the scheduler is built at startup
# Global variable for anthropic beta headers reload scheduling
last_anthropic_beta_headers_reload = None
@ -10094,7 +10111,7 @@ class ProxyStartupEvent:
proxy_logging_obj: ProxyLogging,
) -> ProxyWorkerHeartbeat:
"""Initializes scheduled background jobs"""
global heuristic_v1_tuning_baselines, store_model_in_db, scheduler # rebind-ok: startup publishes the one read-only baseline snapshot
global heuristic_v1_tuning_baselines, store_model_in_db, scheduler, scheduler_executor # rebind-ok: startup publishes the one read-only baseline snapshot
# MEMORY LEAK FIX: Configure scheduler with optimized settings
# Memray analysis showed APScheduler's normalize() and _apply_jitter() causing
@ -10103,9 +10120,9 @@ class ProxyStartupEvent:
# 1. Remove/minimize jitter to avoid normalize() memory explosion
# 2. Use larger misfire_grace_time to prevent backlog calculations
# 3. Set replace_existing=True to avoid duplicate jobs
from apscheduler.executors.asyncio import AsyncIOExecutor
from apscheduler.jobstores.memory import MemoryJobStore
scheduler_executor = AwaitableAsyncIOExecutor() # rebind-ok: shutdown awaits the jobs this executor runs
scheduler = AsyncIOScheduler(
job_defaults={
"coalesce": APSCHEDULER_COALESCE,
@ -10118,7 +10135,7 @@ class ProxyStartupEvent:
jobstores={"default": MemoryJobStore()}, # explicitly use memory job store
# Use simple executor to minimize overhead
executors={
"default": AsyncIOExecutor(),
"default": scheduler_executor,
},
# Disable timezone awareness to reduce computation
timezone=None,

View file

@ -0,0 +1,79 @@
# pyright: reportMissingTypeStubs=false # apscheduler ships no type information
import asyncio
from collections.abc import Collection
from typing import Final, Protocol
from apscheduler.executors.asyncio import AsyncIOExecutor
from litellm._logging import verbose_proxy_logger
from litellm.constants import (
SCHEDULED_JOB_SHUTDOWN_CANCEL_TIMEOUT_SECONDS,
SCHEDULED_JOB_SHUTDOWN_FINISH_TIMEOUT_SECONDS,
)
class StoppableScheduler(Protocol):
"""The slice of ``AsyncIOScheduler`` shutdown uses, which ships no type information"""
@property
def running(self) -> bool: ...
def pause(self) -> None: ...
def shutdown(self, wait: bool = ...) -> None: ...
class AwaitableAsyncIOExecutor(AsyncIOExecutor): # pyright: ignore[reportUntypedBaseClass] # apscheduler ships no type information and is absent from the type-check env
"""``AsyncIOExecutor`` whose in-flight job tasks can be awaited after ``shutdown`` cancels them"""
_pending_futures: Collection["asyncio.Future[object]"]
def in_flight_jobs(self) -> tuple["asyncio.Future[object]", ...]:
"""The job tasks that are running right now, as a snapshot"""
return tuple(future for future in self._pending_futures if not future.done())
def pause_scheduled_jobs(scheduler: StoppableScheduler) -> None:
"""Stop the scheduler from starting jobs that shutdown would only cancel; running jobs continue"""
if scheduler.running:
scheduler.pause()
async def stop_in_flight_scheduler_jobs(
scheduler: StoppableScheduler,
executor: AwaitableAsyncIOExecutor,
*,
finish_timeout_seconds: float = SCHEDULED_JOB_SHUTDOWN_FINISH_TIMEOUT_SECONDS,
cancel_timeout_seconds: float = SCHEDULED_JOB_SHUTDOWN_CANCEL_TIMEOUT_SECONDS,
) -> None:
"""
Let in-flight jobs finish for up to finish_timeout_seconds, then stop the scheduler and wait, bounded by
cancel_timeout_seconds, for the jobs it cancels.
Must run before the database is disconnected: a write job that finishes needs its connection,
and a job's cancellation handler is what records the run's outcome.
"""
if not scheduler.running:
return
in_flight: Final = executor.in_flight_jobs()
if in_flight:
verbose_proxy_logger.info(
"Waiting up to %ss for %d in-flight scheduled job(s) to finish",
finish_timeout_seconds,
len(in_flight),
)
still_running: Final = (
(await asyncio.wait(in_flight, timeout=finish_timeout_seconds))[1] if in_flight else frozenset()
)
scheduler.shutdown(wait=False)
if not still_running:
return
verbose_proxy_logger.info("Cancelling %d in-flight scheduled job(s) for shutdown", len(still_running))
_done, pending = await asyncio.wait(still_running, timeout=cancel_timeout_seconds)
if pending:
verbose_proxy_logger.warning(
"%d scheduled job(s) did not finish within %ss of cancellation; giving up on them",
len(pending),
cancel_timeout_seconds,
)

View file

@ -0,0 +1,158 @@
import asyncio
import logging
from collections.abc import AsyncIterator
from contextlib import asynccontextmanager
from datetime import datetime, timedelta
import pytest
from apscheduler.schedulers.asyncio import AsyncIOScheduler
from litellm.proxy.shutdown.scheduled_jobs import (
AwaitableAsyncIOExecutor,
pause_scheduled_jobs,
stop_in_flight_scheduler_jobs,
)
class _Job:
"""A scheduled job that blocks until cancelled, or for ``work_seconds``, and records what it observed"""
def __init__(self, swallow_cancellation: bool = False, work_seconds: float | None = None) -> None:
self.started = asyncio.Event()
self.events: list[str] = []
self.swallow_cancellation = swallow_cancellation
self.work_seconds = work_seconds
async def run(self) -> None:
self.started.set()
try:
if self.work_seconds is None:
await asyncio.Event().wait()
else:
await asyncio.sleep(self.work_seconds)
self.events.append("committed")
except asyncio.CancelledError:
self.events.append("cancelled")
if self.swallow_cancellation:
await asyncio.Event().wait()
raise
finally:
self.events.append("finished")
@asynccontextmanager
async def _running_scheduler(*jobs: _Job) -> AsyncIterator[tuple[AsyncIOScheduler, AwaitableAsyncIOExecutor]]:
"""A started scheduler with every job in flight, stopped on the way out whatever the test did"""
executor = AwaitableAsyncIOExecutor()
scheduler = AsyncIOScheduler(executors={"default": executor})
for index, job in enumerate(jobs):
scheduler.add_job(job.run, id=f"job-{index}", next_run_time=datetime.now())
scheduler.start()
try:
for job in jobs:
await asyncio.wait_for(job.started.wait(), timeout=5)
yield scheduler, executor
finally:
if scheduler.running:
scheduler.shutdown(wait=False)
stragglers = executor.in_flight_jobs()
for straggler in stragglers:
straggler.cancel()
await asyncio.gather(*stragglers, return_exceptions=True)
@pytest.mark.asyncio
async def test_in_flight_jobs_observe_cancellation_before_shutdown_returns():
"""The job's own CancelledError handler records how a run ended, so shutdown must wait for it"""
job = _Job()
async with _running_scheduler(job) as (scheduler, executor):
await stop_in_flight_scheduler_jobs(scheduler, executor)
assert job.events == ["cancelled", "finished"]
assert scheduler.running is False
assert executor.in_flight_jobs() == ()
@pytest.mark.asyncio
async def test_a_job_that_is_finishing_is_allowed_to_finish_rather_than_cancelled():
"""A spend write cancelled mid-commit drops the rows it popped, so short jobs get to finish first"""
write = _Job(work_seconds=0.2)
stuck = _Job()
async with _running_scheduler(write, stuck) as (scheduler, executor):
await stop_in_flight_scheduler_jobs(scheduler, executor, finish_timeout_seconds=2.0)
assert write.events == ["committed", "finished"]
assert stuck.events == ["cancelled", "finished"]
assert scheduler.running is False
@pytest.mark.asyncio
async def test_every_in_flight_job_is_cancelled_not_only_the_first():
first, second = _Job(), _Job()
async with _running_scheduler(first, second) as (scheduler, executor):
await stop_in_flight_scheduler_jobs(scheduler, executor)
assert first.events == ["cancelled", "finished"]
assert second.events == ["cancelled", "finished"]
@pytest.mark.asyncio
async def test_a_job_that_ignores_cancellation_is_abandoned_after_the_timeout(caplog):
"""A job that swallows CancelledError must not hold the pod past its termination grace period"""
job = _Job(swallow_cancellation=True)
async with _running_scheduler(job) as (scheduler, executor):
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
await stop_in_flight_scheduler_jobs(scheduler, executor, cancel_timeout_seconds=0.05)
assert job.events == ["cancelled"]
assert "1 scheduled job(s) did not finish within 0.05s of cancellation" in caplog.text
@pytest.mark.asyncio
async def test_shutdown_with_nothing_in_flight_still_stops_the_scheduler():
async with _running_scheduler() as (scheduler, executor):
await stop_in_flight_scheduler_jobs(scheduler, executor)
await asyncio.sleep(0)
assert scheduler.running is False
@pytest.mark.asyncio
async def test_a_scheduler_that_never_started_is_left_alone():
"""The proxy runs without a scheduler when it has no database"""
executor = AwaitableAsyncIOExecutor()
scheduler = AsyncIOScheduler(executors={"default": executor})
await stop_in_flight_scheduler_jobs(scheduler, executor)
assert scheduler.running is False
@pytest.mark.asyncio
async def test_pausing_stops_new_jobs_from_starting_but_leaves_running_ones_alone():
"""A job due during the shutdown drain would only be cancelled, so it must not start at all"""
running = _Job()
async with _running_scheduler(running) as (scheduler, executor):
late = _Job()
scheduler.add_job(late.run, id="late", next_run_time=datetime.now() + timedelta(seconds=0.1))
pause_scheduled_jobs(scheduler)
await asyncio.sleep(0.3)
assert late.started.is_set() is False
assert running.events == []
assert scheduler.running is True
await stop_in_flight_scheduler_jobs(scheduler, executor)
assert running.events == ["cancelled", "finished"]
assert late.started.is_set() is False
@pytest.mark.asyncio
async def test_pausing_a_scheduler_that_never_started_is_a_no_op():
scheduler = AsyncIOScheduler(executors={"default": AwaitableAsyncIOExecutor()})
pause_scheduled_jobs(scheduler)
assert scheduler.running is False

View file

@ -7,6 +7,7 @@ import math
import time
from contextlib import asynccontextmanager
from datetime import datetime, timedelta, timezone
from typing import Final
from unittest.mock import AsyncMock, MagicMock
import pytest
@ -1422,3 +1423,125 @@ def test_the_reported_run_outcome_is_the_most_significant_reason_in_any_order(st
"""
results = tuple(TableCleanupResult(rows_deleted=0, stop_reason=reason) for reason in stop_reasons)
assert SpendLogCleanup._run_outcome(results) == expected
_OTHER_OUTCOMES: Final = ("completed", "budget_exhausted", "batch_cap_reached", "skipped_locked", "skipped_disabled")
def _runs_recorded(outcome: str) -> float:
"""The real ``litellm_spend_log_cleanup_runs_total`` sample for one outcome, 0 when unset"""
from prometheus_client import REGISTRY
return REGISTRY.get_sample_value("litellm_spend_log_cleanup_runs_total", {"outcome": outcome}) or 0.0
@pytest.mark.asyncio
async def test_a_cancelled_run_records_aborted_and_logs_its_progress_before_re_raising(monkeypatch):
"""A run cut short by shutdown must leave its outcome and how far it got behind"""
import litellm.proxy.db.db_transaction_queue.spend_log_cleanup as cleanup_module
mock_logger = MagicMock()
monkeypatch.setattr(cleanup_module, "verbose_proxy_logger", mock_logger)
aborted_runs_before = _runs_recorded("aborted")
other_runs_before = {outcome: _runs_recorded(outcome) for outcome in _OTHER_OUTCOMES}
third_batch_reached = asyncio.Event()
async def _execute_raw(sql, *args):
if third_batch_reached.is_set():
raise AssertionError("no batch may be issued after the cancelled one")
if _execute_raw.calls < 2:
_execute_raw.calls += 1
return 150
third_batch_reached.set()
await asyncio.Event().wait()
_execute_raw.calls = 0
mock_prisma_client = MagicMock()
_wire_tx(mock_prisma_client.db)
mock_prisma_client.db.execute_raw = _execute_raw
cleaner = SpendLogCleanup(general_settings={"maximum_spend_logs_retention_period": "7d"})
cleaner.pod_lock_manager = MagicMock()
cleaner.pod_lock_manager.redis_cache = MagicMock()
cleaner.pod_lock_manager.acquire_lock = AsyncMock(return_value=True)
cleaner.pod_lock_manager.release_lock = AsyncMock()
run = asyncio.ensure_future(cleaner.cleanup_old_spend_logs(mock_prisma_client))
await asyncio.wait_for(third_batch_reached.wait(), timeout=5)
run.cancel()
with pytest.raises(asyncio.CancelledError):
await run
assert _runs_recorded("aborted") == aborted_runs_before + 1
assert {outcome: _runs_recorded(outcome) for outcome in _OTHER_OUTCOMES} == other_runs_before
cleaner.pod_lock_manager.release_lock.assert_awaited_once()
mock_logger.exception.assert_not_called()
(error_call,) = mock_logger.error.call_args_list
rendered = error_call[0][0] % error_call[0][1:]
assert rendered.startswith("Spend log cleanup cancelled after ")
assert "s (rows_deleted=300, batches=2)" in rendered
@pytest.mark.asyncio
async def test_progress_reported_for_a_cancelled_run_is_that_run_only(monkeypatch):
"""The scheduler holds one cleaner for the life of the process, so progress must not carry over"""
import litellm.proxy.db.db_transaction_queue.spend_log_cleanup as cleanup_module
mock_logger = MagicMock()
monkeypatch.setattr(cleanup_module, "verbose_proxy_logger", mock_logger)
mock_prisma_client = MagicMock()
_wire_tx(mock_prisma_client.db)
mock_prisma_client.db.execute_raw = AsyncMock(side_effect=[150, 0, 0])
cleaner = SpendLogCleanup(general_settings={"maximum_spend_logs_retention_period": "7d"})
cleaner.pod_lock_manager = None
await cleaner.cleanup_old_spend_logs(mock_prisma_client)
mock_prisma_client.db.execute_raw = AsyncMock(side_effect=[150, asyncio.CancelledError()])
with pytest.raises(asyncio.CancelledError):
await cleaner.cleanup_old_spend_logs(mock_prisma_client)
(error_call,) = mock_logger.error.call_args_list
rendered = error_call[0][0] % error_call[0][1:]
assert "(rows_deleted=150, batches=1)" in rendered
@pytest.mark.asyncio
async def test_progress_reported_by_an_overlapping_run_is_its_own(monkeypatch):
"""With APSCHEDULER_MAX_INSTANCES above one, two runs share the cleaner but not their progress"""
import litellm.proxy.db.db_transaction_queue.spend_log_cleanup as cleanup_module
mock_logger = MagicMock()
monkeypatch.setattr(cleanup_module, "verbose_proxy_logger", mock_logger)
first_batch_done = asyncio.Event()
second_run_done = asyncio.Event()
async def _slow_execute_raw(sql, *args):
first_batch_done.set()
await second_run_done.wait()
return 100
slow_client = MagicMock()
_wire_tx(slow_client.db)
slow_client.db.execute_raw = _slow_execute_raw
fast_client = MagicMock()
_wire_tx(fast_client.db)
fast_client.db.execute_raw = AsyncMock(side_effect=[150, 150, 0, 0])
cleaner = SpendLogCleanup(general_settings={"maximum_spend_logs_retention_period": "7d"})
cleaner.pod_lock_manager = None
slow_run = asyncio.ensure_future(cleaner.cleanup_old_spend_logs(slow_client))
await asyncio.wait_for(first_batch_done.wait(), timeout=5)
await cleaner.cleanup_old_spend_logs(fast_client)
second_run_done.set()
await asyncio.sleep(0)
slow_run.cancel()
with pytest.raises(asyncio.CancelledError):
await slow_run
(error_call,) = mock_logger.error.call_args_list
rendered = error_call[0][0] % error_call[0][1:]
assert "(rows_deleted=100, batches=1)" in rendered