From 3864e1241560ff8e2fc98bcc3c84bd03f6f63d8a Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Wed, 12 Aug 2026 20:10:50 -0700 Subject: [PATCH] fix(spend): stop losing spend log rows when a flush is cancelled (#34826) --- litellm/proxy/proxy_server.py | 22 +- litellm/proxy/utils.py | 58 +++++- .../proxy/proxy_server/test_lifecycle.py | 44 ++++ .../proxy/utils/prisma_and_spend/conftest.py | 1 + .../prisma_and_spend/test_spend_functions.py | 194 ++++++++++++++++++ 5 files changed, 310 insertions(+), 9 deletions(-) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 3ac002ce8cc..64a91b880cd 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -845,6 +845,22 @@ def cleanup_router_config_variables(): prisma_client = None +async def _flush_spend_logs_queue_on_shutdown() -> None: + if prisma_client is None: + return + + try: + from litellm.proxy.utils import drain_spend_logs_queue + + await drain_spend_logs_queue( + prisma_client=prisma_client, + db_writer_client=db_writer_client, + proxy_logging_obj=proxy_logging_obj, + ) + except Exception as e: # noqa: BLE001 # shutdown must continue even if the drain fails + verbose_proxy_logger.exception("Error flushing spend logs queue on shutdown: %s", e) + + async def proxy_shutdown_event(): global prisma_client, master_key, user_custom_auth, user_custom_key_generate, user_custom_key_update verbose_proxy_logger.info("Shutting down LiteLLM Proxy Server") @@ -1255,6 +1271,8 @@ async def proxy_startup_event(app: FastAPI): except Exception as e: verbose_proxy_logger.error("Error stopping DB health watchdog task: %s", e) + await _flush_spend_logs_queue_on_shutdown() + await proxy_config.stop_config_sync_subscriber() await proxy_config.stop_auth_cache_invalidation_subscriber() @@ -8731,14 +8749,14 @@ class ProxyStartupEvent: if general_settings.get("disable_spend_logs", False) is False: from litellm.proxy.utils import _monitor_spend_logs_queue - # Start background task to monitor spend logs queue size - asyncio.create_task( + monitor_task: Final = asyncio.create_task( _monitor_spend_logs_queue( prisma_client=prisma_client, db_writer_client=db_writer_client, proxy_logging_obj=proxy_logging_obj, ) ) + prisma_client.spend_logs_queue_monitor_task = monitor_task # rebind-ok: the client owns its monitor handle ### ADD NEW MODELS ### store_model_in_db = get_secret_bool("STORE_MODEL_IN_DB", store_model_in_db) or store_model_in_db diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index cce8379ab25..8a1fae42789 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -1,4 +1,5 @@ import asyncio +import contextlib import copy import hashlib import inspect @@ -3006,6 +3007,7 @@ async def prefetch_config_params(prisma_client: "PrismaClient | None", param_nam class PrismaClient: spend_log_transactions: list = [] _spend_log_transactions_lock = asyncio.Lock() + spend_logs_queue_monitor_task: "asyncio.Task[None] | None" = None tool_usage_transactions: list["ToolUsageTransaction"] = [] _tool_usage_transactions_lock = asyncio.Lock() autorouter_turn_transactions: ClassVar[ @@ -5722,13 +5724,22 @@ async def update_spend_logs_job( logs_to_process: Final = prisma_client.spend_log_transactions[:MAX_LOGS_PER_INTERVAL] prisma_client.spend_log_transactions = prisma_client.spend_log_transactions[len(logs_to_process) :] - await ProxyUpdateSpend.update_spend_logs( - n_retry_times=n_retry_times, - prisma_client=prisma_client, - proxy_logging_obj=proxy_logging_obj, - db_writer_client=db_writer_client, - logs_to_process=logs_to_process, - ) + try: + await ProxyUpdateSpend.update_spend_logs( + n_retry_times=n_retry_times, + prisma_client=prisma_client, + proxy_logging_obj=proxy_logging_obj, + db_writer_client=db_writer_client, + logs_to_process=logs_to_process, + ) + except asyncio.CancelledError: + async with prisma_client._spend_log_transactions_lock: + prisma_client.spend_log_transactions[:0] = logs_to_process + verbose_proxy_logger.warning( + "Spend tracking - spend log write cancelled, requeued %d rows for the next flush", + len(logs_to_process), + ) + raise # Guardrail/policy usage tracking (same batch, outside spend-logs update) try: @@ -5787,6 +5798,39 @@ async def update_spend_logs_job( ) +MAX_SPEND_LOG_DRAIN_ITERATIONS: Final = 20 + + +async def drain_spend_logs_queue( + prisma_client: PrismaClient, + db_writer_client: "AsyncHTTPHandler | None", + proxy_logging_obj: ProxyLogging, +) -> None: + monitor_task: Final = prisma_client.spend_logs_queue_monitor_task + if monitor_task is not None: + monitor_task.cancel() + with contextlib.suppress(asyncio.CancelledError): + await monitor_task + prisma_client.spend_logs_queue_monitor_task = None # rebind-ok: the client owns its monitor handle + + for _ in range(MAX_SPEND_LOG_DRAIN_ITERATIONS): + if await _total_queued_spend_transactions(prisma_client) == 0: + return + await update_spend_logs_job( + prisma_client=prisma_client, + db_writer_client=db_writer_client, + proxy_logging_obj=proxy_logging_obj, + ) + + remaining: Final = await _total_queued_spend_transactions(prisma_client) + if remaining > 0: + spend_log_error( + "Spend tracking - %d spend log rows still queued after %d drain passes", + remaining, + MAX_SPEND_LOG_DRAIN_ITERATIONS, + ) + + async def _monitor_spend_logs_queue( prisma_client: PrismaClient, db_writer_client: AsyncHTTPHandler | None, diff --git a/tests/test_litellm/proxy/proxy_server/test_lifecycle.py b/tests/test_litellm/proxy/proxy_server/test_lifecycle.py index 40ca7e3a64e..ad41592db1e 100644 --- a/tests/test_litellm/proxy/proxy_server/test_lifecycle.py +++ b/tests/test_litellm/proxy/proxy_server/test_lifecycle.py @@ -205,6 +205,50 @@ async def test_proxy_shutdown_event_prisma_disconnect_raises_error(monkeypatch): await proxy_shutdown_event() +# --------------------------------------------------------------------------- +# _flush_spend_logs_queue_on_shutdown +# --------------------------------------------------------------------------- + + +@pytest.mark.asyncio +async def test_flush_spend_logs_queue_on_shutdown_drains_before_disconnect(monkeypatch): + fake_prisma = MagicMock() + monkeypatch.setattr(ps, "prisma_client", fake_prisma, raising=False) + monkeypatch.setattr(ps, "db_writer_client", None, raising=False) + + drain = AsyncMock() + import litellm.proxy.utils as utils_mod + + monkeypatch.setattr(utils_mod, "drain_spend_logs_queue", drain) + + await ps._flush_spend_logs_queue_on_shutdown() + + observed = { + "drain_calls": drain.await_count, + "drain_prisma": drain.await_args.kwargs["prisma_client"] is fake_prisma, + } + assert observed == { + "drain_calls": 1, + "drain_prisma": True, + } + + +@pytest.mark.asyncio +async def test_flush_spend_logs_queue_on_shutdown_swallows_drain_errors(monkeypatch): + monkeypatch.setattr(ps, "prisma_client", MagicMock(), raising=False) + monkeypatch.setattr(ps, "db_writer_client", None, raising=False) + + import litellm.proxy.utils as utils_mod + + monkeypatch.setattr( + utils_mod, + "drain_spend_logs_queue", + AsyncMock(side_effect=RuntimeError("db gone")), + ) + + await ps._flush_spend_logs_queue_on_shutdown() + + # --------------------------------------------------------------------------- # _initialize_shared_aiohttp_session # --------------------------------------------------------------------------- diff --git a/tests/test_litellm/proxy/utils/prisma_and_spend/conftest.py b/tests/test_litellm/proxy/utils/prisma_and_spend/conftest.py index 74c9abd9978..19abcb5d66d 100644 --- a/tests/test_litellm/proxy/utils/prisma_and_spend/conftest.py +++ b/tests/test_litellm/proxy/utils/prisma_and_spend/conftest.py @@ -128,6 +128,7 @@ def mock_prisma_client() -> MagicMock: client.proxy_logging_obj.failure_handler = AsyncMock() client.spend_log_transactions = [] client._spend_log_transactions_lock = asyncio.Lock() + client.spend_logs_queue_monitor_task = None client.tool_usage_transactions = [] client._tool_usage_transactions_lock = asyncio.Lock() client.jsonify_object = lambda data: dict(data) diff --git a/tests/test_litellm/proxy/utils/prisma_and_spend/test_spend_functions.py b/tests/test_litellm/proxy/utils/prisma_and_spend/test_spend_functions.py index d9eeb168611..54d59e690f9 100644 --- a/tests/test_litellm/proxy/utils/prisma_and_spend/test_spend_functions.py +++ b/tests/test_litellm/proxy/utils/prisma_and_spend/test_spend_functions.py @@ -17,8 +17,10 @@ from unittest.mock import AsyncMock, MagicMock import pytest from litellm.proxy.utils import ( + MAX_SPEND_LOG_DRAIN_ITERATIONS, _monitor_spend_logs_queue, _raise_failed_update_spend_exception, + drain_spend_logs_queue, update_daily_tag_spend, update_spend, update_spend_logs_job, @@ -263,6 +265,198 @@ async def test_update_spend_logs_job_processes_and_clears_queue( } +@pytest.mark.asyncio +async def test_update_spend_logs_job_requeues_popped_rows_when_write_cancelled( + mock_prisma_client: Any, make_spend_log_row: Any +) -> None: + proxy_logging = MagicMock() + proxy_logging.failure_handler = AsyncMock() + mock_prisma_client.spend_log_transactions = [ + make_spend_log_row(request_id="r1"), + make_spend_log_row(request_id="r2"), + ] + + row_arriving_mid_flush = make_spend_log_row(request_id="r3") + + async def _cancel_mid_write(*args: Any, **kwargs: Any) -> None: + mock_prisma_client.spend_log_transactions.append(row_arriving_mid_flush) + raise asyncio.CancelledError() + + mock_prisma_client.db.litellm_spendlogs.create_many = AsyncMock( + side_effect=_cancel_mid_write + ) + + with pytest.raises(asyncio.CancelledError): + await update_spend_logs_job( + prisma_client=mock_prisma_client, + db_writer_client=None, + proxy_logging_obj=proxy_logging, + ) + + assert [ + row["request_id"] for row in mock_prisma_client.spend_log_transactions + ] == ["r1", "r2", "r3"] + + +@pytest.mark.asyncio +async def test_update_spend_logs_job_does_not_requeue_when_cancelled_after_write( + mock_prisma_client: Any, make_spend_log_row: Any, monkeypatch: pytest.MonkeyPatch +) -> None: + """Rows are already committed once guardrail tracking runs, so replaying + them would double-count the non-idempotent daily guardrail increments. + """ + import litellm.proxy.guardrails.usage_tracking as guard_mod + + proxy_logging = MagicMock() + proxy_logging.failure_handler = AsyncMock() + mock_prisma_client.spend_log_transactions = [make_spend_log_row(request_id="r1")] + mock_prisma_client.db.litellm_spendlogs.create_many = AsyncMock() + + monkeypatch.setattr( + guard_mod, + "process_spend_logs_guardrail_usage", + AsyncMock(side_effect=asyncio.CancelledError()), + raising=False, + ) + + with pytest.raises(asyncio.CancelledError): + await update_spend_logs_job( + prisma_client=mock_prisma_client, + db_writer_client=None, + proxy_logging_obj=proxy_logging, + ) + + assert mock_prisma_client.spend_log_transactions == [] + + +@pytest.mark.asyncio +async def test_drain_spend_logs_queue_flushes_rows_queued_while_draining( + mock_prisma_client: Any, make_spend_log_row: Any, monkeypatch: pytest.MonkeyPatch +) -> None: + import litellm.proxy.db.spend_log_tool_index as tool_mod + import litellm.proxy.guardrails.usage_tracking as guard_mod + + monkeypatch.setattr( + guard_mod, "process_spend_logs_guardrail_usage", AsyncMock(), raising=False + ) + monkeypatch.setattr( + tool_mod, "process_spend_logs_tool_usage", AsyncMock(), raising=False + ) + + proxy_logging = MagicMock() + proxy_logging.failure_handler = AsyncMock() + mock_prisma_client.spend_log_transactions = [make_spend_log_row(request_id="r1")] + + written: list[str] = [] + + async def _write(*args: Any, **kwargs: Any) -> None: + written.extend(row["request_id"] for row in kwargs["data"]) + if len(written) == 1: + mock_prisma_client.spend_log_transactions.append( + make_spend_log_row(request_id="r2") + ) + + mock_prisma_client.db.litellm_spendlogs.create_many = AsyncMock(side_effect=_write) + + await drain_spend_logs_queue( + prisma_client=mock_prisma_client, + db_writer_client=None, + proxy_logging_obj=proxy_logging, + ) + + assert written == ["r1", "r2"] + assert mock_prisma_client.spend_log_transactions == [] + + +@pytest.mark.asyncio +async def test_drain_spend_logs_queue_stops_monitor_and_keeps_its_popped_rows( + mock_prisma_client: Any, make_spend_log_row: Any, monkeypatch: pytest.MonkeyPatch +) -> None: + import litellm.proxy.db.spend_log_tool_index as tool_mod + import litellm.proxy.guardrails.usage_tracking as guard_mod + + monkeypatch.setattr( + guard_mod, "process_spend_logs_guardrail_usage", AsyncMock(), raising=False + ) + monkeypatch.setattr( + tool_mod, "process_spend_logs_tool_usage", AsyncMock(), raising=False + ) + + proxy_logging = MagicMock() + proxy_logging.failure_handler = AsyncMock() + mock_prisma_client.spend_log_transactions = [make_spend_log_row(request_id="r1")] + + write_started = asyncio.Event() + written: list[str] = [] + write_calls = {"n": 0} + + async def _write(*args: Any, **kwargs: Any) -> None: + write_calls["n"] += 1 + if write_calls["n"] == 1: + write_started.set() + await asyncio.Event().wait() + written.extend(row["request_id"] for row in kwargs["data"]) + + mock_prisma_client.db.litellm_spendlogs.create_many = AsyncMock(side_effect=_write) + + async def _monitor() -> None: + await update_spend_logs_job( + prisma_client=mock_prisma_client, + db_writer_client=None, + proxy_logging_obj=proxy_logging, + ) + + mock_prisma_client.spend_logs_queue_monitor_task = asyncio.create_task(_monitor()) + await write_started.wait() + + await drain_spend_logs_queue( + prisma_client=mock_prisma_client, + db_writer_client=None, + proxy_logging_obj=proxy_logging, + ) + + assert written == ["r1"] + assert mock_prisma_client.spend_log_transactions == [] + assert mock_prisma_client.spend_logs_queue_monitor_task is None + + +@pytest.mark.asyncio +async def test_drain_spend_logs_queue_gives_up_after_max_passes( + mock_prisma_client: Any, make_spend_log_row: Any, monkeypatch: pytest.MonkeyPatch +) -> None: + import litellm.proxy.db.spend_log_tool_index as tool_mod + import litellm.proxy.guardrails.usage_tracking as guard_mod + + monkeypatch.setattr( + guard_mod, "process_spend_logs_guardrail_usage", AsyncMock(), raising=False + ) + monkeypatch.setattr( + tool_mod, "process_spend_logs_tool_usage", AsyncMock(), raising=False + ) + + proxy_logging = MagicMock() + proxy_logging.failure_handler = AsyncMock() + mock_prisma_client.spend_log_transactions = [make_spend_log_row(request_id="r1")] + + async def _write_and_refill(*args: Any, **kwargs: Any) -> None: + mock_prisma_client.spend_log_transactions.append(make_spend_log_row()) + + mock_prisma_client.db.litellm_spendlogs.create_many = AsyncMock( + side_effect=_write_and_refill + ) + + await drain_spend_logs_queue( + prisma_client=mock_prisma_client, + db_writer_client=None, + proxy_logging_obj=proxy_logging, + ) + + assert ( + mock_prisma_client.db.litellm_spendlogs.create_many.await_count + == MAX_SPEND_LOG_DRAIN_ITERATIONS + ) + + @pytest.mark.asyncio async def test_monitor_spend_logs_queue_invokes_job_when_queue_nonempty( mock_prisma_client: Any,