From 6e7984e537ed04a882a540b8f5aa3bc2ecf6bbd3 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Fri, 14 Aug 2026 20:49:45 -0700 Subject: [PATCH] fix(proxy): requeue spend logs when the DB write fails with a transport error (#36716) * fix(proxy): requeue spend logs when the DB write fails with a transport error Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): hardcode the spend log queue cap and drop the stale re-export Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor(proxy): keep the spend log requeue within the type discipline budget Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): apply the spend log queue cap to producer appends too Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): lower the spend log queue cap to 1k and make it env configurable Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): bound the spend log queue by bytes instead of row count A row cap cannot bound memory: a row carries the whole prompt under store_prompts_in_spend_logs, so a cap that rides out an outage of counter-only rows is an OOM once prompts are stored. Every enqueue and dequeue now goes through one pair that tracks what the queue costs and drops the oldest rows past a 64 MB budget. Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): make the spend log queue byte budget env configurable Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): use a string default for the spend log queue byte budget env read Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): make the spend log queue byte total a public attribute The queue it accounts for is already public, and a private name only bought reportPrivateUsage errors at every call site. Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Co-authored-by: shivam --- litellm/constants.py | 1 + litellm/proxy/db/db_spend_update_writer.py | 5 +- litellm/proxy/db/spend_log_batching.py | 40 +++++ litellm/proxy/utils.py | 78 +++++++--- tests/proxy_unit_tests/test_update_spend.py | 2 +- .../proxy/db/test_spend_log_batching.py | 27 ++++ .../test_proxy_update_spend.py | 144 +++++++++++++++++- 7 files changed, 276 insertions(+), 21 deletions(-) diff --git a/litellm/constants.py b/litellm/constants.py index 6449834d6a4..14ef572888f 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -1499,6 +1499,7 @@ SPEND_LOG_PARTITION_INTERVAL: Final = os.getenv("SPEND_LOG_PARTITION_INTERVAL", SPEND_LOG_PARTITION_PRECREATE_AHEAD: Final = int(os.getenv("SPEND_LOG_PARTITION_PRECREATE_AHEAD", 7)) SPEND_LOG_WRITE_BATCH_MAX_BYTES: Final = max(1, int(os.getenv("SPEND_LOG_WRITE_BATCH_MAX_BYTES", 2_000_000))) SPEND_LOG_QUEUE_SIZE_THRESHOLD: Final = int(os.getenv("SPEND_LOG_QUEUE_SIZE_THRESHOLD", 100)) +SPEND_LOG_QUEUE_MAX_BYTES: Final = max(1, int(os.getenv("SPEND_LOG_QUEUE_MAX_BYTES", "64000000"))) SPEND_LOG_QUEUE_POLL_INTERVAL: Final = float(os.getenv("SPEND_LOG_QUEUE_POLL_INTERVAL", 2.0)) SPEND_COUNTER_RESEED_LOCKS_MAX_SIZE: Final = int(os.getenv("SPEND_COUNTER_RESEED_LOCKS_MAX_SIZE", 10000)) DEFAULT_CRON_JOB_LOCK_TTL_SECONDS: Final = int(os.getenv("DEFAULT_CRON_JOB_LOCK_TTL_SECONDS", 60)) # 1 minute diff --git a/litellm/proxy/db/db_spend_update_writer.py b/litellm/proxy/db/db_spend_update_writer.py index 471067b16f8..fe68e837a8e 100644 --- a/litellm/proxy/db/db_spend_update_writer.py +++ b/litellm/proxy/db/db_spend_update_writer.py @@ -787,8 +787,9 @@ class DBSpendUpdateWriter: ) ) if prisma_client is not None and spend_logs_url is not None or prisma_client is not None: - async with prisma_client._spend_log_transactions_lock: - prisma_client.spend_log_transactions.append(payload) + from litellm.proxy.utils import enqueue_spend_logs + + await enqueue_spend_logs(prisma_client, (payload,)) else: verbose_proxy_logger.debug("prisma_client is None. Skipping writing spend logs to db.") diff --git a/litellm/proxy/db/spend_log_batching.py b/litellm/proxy/db/spend_log_batching.py index 94deb6e6a1e..a8fced5485d 100644 --- a/litellm/proxy/db/spend_log_batching.py +++ b/litellm/proxy/db/spend_log_batching.py @@ -18,6 +18,7 @@ byte budget tracks what the engine actually allocates. import json from collections.abc import Iterator, Mapping, Sequence +from itertools import accumulate from typing import Final SpendLogRow = Mapping[str, object] @@ -56,6 +57,45 @@ def _row_payload_bytes(row: SpendLogRow) -> int: return 0 +def spend_log_row_bytes(row: SpendLogRow) -> int: + """Bytes this row costs, measured the same way the write budget measures it.""" + return _row_payload_bytes(row) + + +def spend_log_queue_within_budget( + rows: Sequence[SpendLogRow], + queued_bytes: int, + max_bytes: int, +) -> tuple[Sequence[SpendLogRow], int]: + """Drop the oldest rows until the queue costs at most ``max_bytes``. + + Returns the rows to keep and what they cost, so a caller tracking the total + across calls does not have to re-measure the rows it kept. ``queued_bytes`` + is that running total for ``rows``; only the rows actually dropped are + measured here, which is what keeps an append off an O(queue) path. + + A queue is bounded by bytes rather than by row count because a row's size + swings by orders of magnitude with ``store_prompts_in_spend_logs``, so any + row cap generous enough to ride out an outage of counter-only rows is an + OOM once prompts are stored. + + The newest row is kept whatever it costs, for the same reason a statement + over budget is still written: the budget is a memory guardrail, not an + admission filter, and losing spend data to protect RSS is the worse failure. + """ + if queued_bytes <= max_bytes or len(rows) <= 1: + return rows, queued_bytes + droppable: Final = rows[:-1] + remaining_by_drops: Final = ( + queued_bytes - freed for freed in accumulate(_row_payload_bytes(row) for row in droppable) + ) + fits: Final = next( + ((drops, remaining) for drops, remaining in enumerate(remaining_by_drops, start=1) if remaining <= max_bytes), + (len(droppable), _row_payload_bytes(rows[-1])), + ) + return rows[fits[0] :], fits[1] + + def spend_log_write_batches( rows: Sequence[SpendLogRow], max_bytes: int, diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index ec1aa262736..498b6d7ee3d 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -24,10 +24,10 @@ from litellm.constants import ( DEFAULT_MODEL_CREATED_AT_TIME, LITELLM_LOGGING_NO_UPSTREAM_LLM_CALL, MAX_TEAM_LIST_LIMIT, + SPEND_LOG_QUEUE_MAX_BYTES, SPEND_LOG_WRITE_BATCH_MAX_BYTES, ) from litellm.proxy._types import ( - DB_CONNECTION_ERROR_TYPES, DB_RETRY_SAFE_ERROR_TYPES, CommonProxyErrors, ProxyErrorTypes, @@ -121,7 +121,11 @@ from litellm.proxy.db.prisma_client import ( parse_iam_endpoint_from_url, ) from litellm.proxy.db.routing_prisma_wrapper import RoutingPrismaWrapper -from litellm.proxy.db.spend_log_batching import spend_log_write_batches +from litellm.proxy.db.spend_log_batching import ( + spend_log_queue_within_budget, + spend_log_row_bytes, + spend_log_write_batches, +) from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import ( UnifiedLLMGuardrails, ) @@ -3066,6 +3070,7 @@ class _StaleReadEngine: class PrismaClient: spend_log_transactions: list = [] _spend_log_transactions_lock = asyncio.Lock() + spend_log_queue_bytes: ClassVar[int] = 0 spend_logs_queue_monitor_task: "asyncio.Task[None] | None" = None tool_usage_transactions: list["ToolUsageTransaction"] = [] _tool_usage_transactions_lock = asyncio.Lock() @@ -5702,6 +5707,53 @@ def _hash_token_if_needed(token: str) -> str: return token +async def enqueue_spend_logs( + prisma_client: PrismaClient, + logs: Sequence[Mapping[str, object]], + *, + at_head: bool = False, + max_bytes: int = SPEND_LOG_QUEUE_MAX_BYTES, +) -> None: + """Queue spend logs for the next flush, held under ``SPEND_LOG_QUEUE_MAX_BYTES``. + + ``at_head`` replays a batch the DB refused, so it flushes before the logs + that piled up during the outage. Past the budget the oldest logs are + dropped, which keeps a long outage from growing the queue until the pod + dies. + """ + added: Final = sum(spend_log_row_bytes(row) for row in logs) + async with prisma_client._spend_log_transactions_lock: + queued: Final = ( + tuple(logs) + tuple(prisma_client.spend_log_transactions) + if at_head + else tuple(prisma_client.spend_log_transactions) + tuple(logs) + ) + kept, kept_bytes = spend_log_queue_within_budget(queued, PrismaClient.spend_log_queue_bytes + added, max_bytes) + prisma_client.spend_log_transactions[:] = kept + PrismaClient.spend_log_queue_bytes = kept_bytes + if len(kept) < len(queued): + verbose_proxy_logger.error( + "Spend tracking - spend log queue is at its %d byte budget; dropped the %d oldest spend logs", + max_bytes, + len(queued) - len(kept), + ) + + +async def dequeue_spend_logs(prisma_client: PrismaClient, limit: int) -> list[dict[str, object]]: + """Take up to ``limit`` of the oldest queued spend logs off the queue. + + Every enqueue and dequeue goes through this pair so the byte total the + queue is bounded by stays in step with what the queue actually holds. + """ + async with prisma_client._spend_log_transactions_lock: + popped: Final = prisma_client.spend_log_transactions[:limit] + prisma_client.spend_log_transactions[:] = prisma_client.spend_log_transactions[limit:] + PrismaClient.spend_log_queue_bytes = max( + 0, PrismaClient.spend_log_queue_bytes - sum(spend_log_row_bytes(row) for row in popped) + ) + return popped + + class ProxyUpdateSpend: @staticmethod async def update_end_user_spend( @@ -5754,11 +5806,7 @@ class ProxyUpdateSpend: MAX_LOGS_PER_INTERVAL: Final = 10000 # Maximum number of logs to flush in a single interval popped_batch = False if logs_to_process is None: - # Atomically read and remove logs to process (protected by lock) - async with prisma_client._spend_log_transactions_lock: - logs_to_process = prisma_client.spend_log_transactions[:MAX_LOGS_PER_INTERVAL] - # Remove the logs we're about to process - prisma_client.spend_log_transactions = prisma_client.spend_log_transactions[len(logs_to_process) :] + logs_to_process = await dequeue_spend_logs(prisma_client, MAX_LOGS_PER_INTERVAL) popped_batch = True if len(logs_to_process) > 0: verbose_proxy_logger.info( @@ -5808,9 +5856,9 @@ class ProxyUpdateSpend: "%s logs processed. Remaining in queue: %s", len(logs_to_process), remaining_count ) break - except DB_CONNECTION_ERROR_TYPES as e: - if i is None: - i = 0 + except Exception as e: + if not PrismaDBExceptionHandler.is_database_transport_error(e): + raise verbose_proxy_logger.warning( "Spend tracking - DB connection error writing spend logs, retry %d/%d. logs_count=%d, error=%s", i + 1, @@ -5819,11 +5867,10 @@ class ProxyUpdateSpend: str(e), ) if i >= n_retry_times: + await enqueue_spend_logs(prisma_client, logs_to_process, at_head=True) raise await asyncio.sleep(2**i) except Exception as e: - # Logs already removed from queue at start - don't put them back - # This matches the original behavior where logs are removed even on error _raise_failed_update_spend_exception(e=e, start_time=start_time, proxy_logging_obj=proxy_logging_obj) finally: # Clean up logs_to_process only if we popped it (caller-owned otherwise) @@ -5965,9 +6012,7 @@ async def update_spend_logs_job( if await _total_queued_spend_transactions(prisma_client) == 0: return - async with prisma_client._spend_log_transactions_lock: - 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) :] + logs_to_process: Final = await dequeue_spend_logs(prisma_client, MAX_LOGS_PER_INTERVAL) try: await ProxyUpdateSpend.update_spend_logs( @@ -5978,8 +6023,7 @@ async def update_spend_logs_job( 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 + await enqueue_spend_logs(prisma_client, logs_to_process, at_head=True) verbose_proxy_logger.warning( "Spend tracking - spend log write cancelled, requeued %d rows for the next flush", len(logs_to_process), diff --git a/tests/proxy_unit_tests/test_update_spend.py b/tests/proxy_unit_tests/test_update_spend.py index 96a57c427e7..0d1d6dcf3c6 100644 --- a/tests/proxy_unit_tests/test_update_spend.py +++ b/tests/proxy_unit_tests/test_update_spend.py @@ -15,7 +15,7 @@ from unittest.mock import MagicMock, patch, AsyncMock import httpx -from litellm.proxy.utils import update_spend, DB_CONNECTION_ERROR_TYPES +from litellm.proxy.utils import update_spend class MockPrismaClient: diff --git a/tests/test_litellm/proxy/db/test_spend_log_batching.py b/tests/test_litellm/proxy/db/test_spend_log_batching.py index a0fb4901a5c..2069490e7a0 100644 --- a/tests/test_litellm/proxy/db/test_spend_log_batching.py +++ b/tests/test_litellm/proxy/db/test_spend_log_batching.py @@ -6,6 +6,7 @@ exceeds the byte budget while every row is still written exactly once. Symbols pinned here: - ``spend_log_write_batches`` + - ``spend_log_queue_within_budget`` - ``_row_payload_bytes`` """ @@ -14,6 +15,7 @@ from typing import Any, Dict, List from litellm.proxy.db.spend_log_batching import ( _row_payload_bytes, + spend_log_queue_within_budget, spend_log_write_batches, ) @@ -140,6 +142,31 @@ def test_json_escaping_growth_is_counted() -> None: assert [len(batch) for batch in spend_log_write_batches([row, row], max_bytes=budget)] == [1, 1] +def test_queue_within_budget_drops_the_oldest_rows_and_reports_what_is_left() -> None: + """Trimming has to free enough bytes to get under the budget while keeping + the newest rows, and hand back the kept total so a queue tracking it across + appends never re-measures the rows it kept.""" + rows = [{"request_id": f"r{i}", "messages": "x" * 1000} for i in range(4)] + row_bytes = _row_payload_bytes(rows[0]) + + kept, kept_bytes = spend_log_queue_within_budget(rows, 4 * row_bytes, 2 * row_bytes) + + assert [row["request_id"] for row in kept] == ["r2", "r3"] + assert kept_bytes == 2 * row_bytes + + +def test_queue_within_budget_keeps_a_row_larger_than_the_whole_budget() -> None: + """A row over budget on its own is kept rather than dropped, the same call + the write batcher makes: the budget guards memory, and trading a spend + record for RSS is the worse failure.""" + row = {"request_id": "r", "messages": "x" * 10_000} + + kept, kept_bytes = spend_log_queue_within_budget([row], _row_payload_bytes(row), 100) + + assert list(kept) == [row] + assert kept_bytes == _row_payload_bytes(row) + + def test_unserialized_list_payloads_are_measured_not_ignored() -> None: """``jsonify_object`` only stringifies dicts, so a list-valued ``messages`` reaches the batcher raw; counting it as zero would let the largest rows diff --git a/tests/test_litellm/proxy/utils/prisma_and_spend/test_proxy_update_spend.py b/tests/test_litellm/proxy/utils/prisma_and_spend/test_proxy_update_spend.py index 9d6f53841ed..dd21bbc9e8a 100644 --- a/tests/test_litellm/proxy/utils/prisma_and_spend/test_proxy_update_spend.py +++ b/tests/test_litellm/proxy/utils/prisma_and_spend/test_proxy_update_spend.py @@ -10,13 +10,22 @@ from __future__ import annotations import asyncio import json +from collections.abc import Iterator from typing import Any, Dict, List from unittest.mock import AsyncMock, MagicMock import pytest import litellm.proxy.utils as utils_mod -from litellm.proxy.utils import ProxyUpdateSpend +from litellm.proxy.db.spend_log_batching import spend_log_row_bytes +from litellm.proxy.utils import PrismaClient, ProxyUpdateSpend, enqueue_spend_logs + + +@pytest.fixture(autouse=True) +def reset_spend_log_queue_bytes() -> Iterator[None]: + PrismaClient.spend_log_queue_bytes = 0 + yield + PrismaClient.spend_log_queue_bytes = 0 class _AsyncCM: @@ -358,6 +367,139 @@ async def test_update_spend_logs_reraises_connection_masquerade_dataerror( ) +@pytest.mark.asyncio +async def test_update_spend_logs_retries_and_requeues_batch_on_db_outage( + mock_prisma_client: Any, make_spend_log_row: Any, monkeypatch: pytest.MonkeyPatch +) -> None: + """A P1001 outage must be retried and, once retries exhaust, the batch goes + back to the head of the queue so the next flush persists it. Before the fix + prisma's ``DataError`` masquerade fell outside the retry clause, so the pod + dropped every queued spend log for the duration of the outage. + """ + + async def _fake_sleep(_: float) -> None: + return None + + monkeypatch.setattr(utils_mod.asyncio, "sleep", _fake_sleep) + mock_prisma_client.db.litellm_spendlogs.create_many = AsyncMock( + side_effect=_data_error("Can't reach database server at db-host:5432 (P1001)") + ) + proxy_logging = MagicMock() + proxy_logging.failure_handler = AsyncMock() + logs = [make_spend_log_row(request_id="a"), make_spend_log_row(request_id="b")] + queued_during_outage = make_spend_log_row(request_id="c") + mock_prisma_client.spend_log_transactions = [queued_during_outage] + + with pytest.raises(type(_data_error("x"))): + await ProxyUpdateSpend.update_spend_logs( + n_retry_times=2, + prisma_client=mock_prisma_client, + db_writer_client=None, + proxy_logging_obj=proxy_logging, + logs_to_process=logs, + ) + + assert mock_prisma_client.db.litellm_spendlogs.create_many.await_count == 3 + assert [row["request_id"] for row in mock_prisma_client.spend_log_transactions] == ["a", "b", "c"] + + +@pytest.mark.asyncio +async def test_requeue_after_outage_drops_oldest_logs_past_the_byte_budget( + mock_prisma_client: Any, make_spend_log_row: Any +) -> None: + """Requeueing must stay bounded by what the queue costs in memory, not by a + row count: a row carries the whole prompt under + ``store_prompts_in_spend_logs``, so a row cap that survives an outage of + counter-only rows is an OOM once prompts are stored. Past the budget the + oldest rows are the ones dropped. + """ + budget = 3 * spend_log_row_bytes(make_spend_log_row(request_id="new0")) + mock_prisma_client.spend_log_transactions = [] + await enqueue_spend_logs(mock_prisma_client, [make_spend_log_row(request_id="new0")], max_bytes=budget) + + await enqueue_spend_logs( + mock_prisma_client, + [make_spend_log_row(request_id=f"old{i}") for i in range(4)], + at_head=True, + max_bytes=budget, + ) + + assert [row["request_id"] for row in mock_prisma_client.spend_log_transactions] == ["old2", "old3", "new0"] + + +@pytest.mark.asyncio +async def test_enqueue_drops_oldest_logs_once_producers_fill_the_queue( + mock_prisma_client: Any, make_spend_log_row: Any +) -> None: + """The budget has to govern the producer side too. While a flush retries + against a dead DB, requests keep landing, so an append path that ignores the + budget leaves the outage OOM open no matter how well the requeue trims. + """ + budget = 2 * spend_log_row_bytes(make_spend_log_row(request_id="old0")) + mock_prisma_client.spend_log_transactions = [] + await enqueue_spend_logs( + mock_prisma_client, + [make_spend_log_row(request_id=f"old{i}") for i in range(2)], + max_bytes=budget, + ) + + await enqueue_spend_logs(mock_prisma_client, [make_spend_log_row(request_id="new0")], max_bytes=budget) + + assert [row["request_id"] for row in mock_prisma_client.spend_log_transactions] == ["old1", "new0"] + + +@pytest.mark.asyncio +async def test_flush_returns_the_bytes_it_took_off_the_queue(mock_prisma_client: Any, make_spend_log_row: Any) -> None: + """A flush has to give its bytes back to the budget. Accounting that only + ever grows would treat a healthy pod as permanently full and start dropping + fresh spend logs after the queue has already drained. + """ + budget = 2 * spend_log_row_bytes(make_spend_log_row(request_id="row0")) + proxy_logging = MagicMock() + proxy_logging.failure_handler = AsyncMock() + mock_prisma_client.spend_log_transactions = [] + await enqueue_spend_logs( + mock_prisma_client, + [make_spend_log_row(request_id=f"row{i}") for i in range(2)], + max_bytes=budget, + ) + + await ProxyUpdateSpend.update_spend_logs( + n_retry_times=0, + prisma_client=mock_prisma_client, + db_writer_client=None, + proxy_logging_obj=proxy_logging, + ) + await enqueue_spend_logs(mock_prisma_client, [make_spend_log_row(request_id="row9")], max_bytes=budget) + + assert [row["request_id"] for row in mock_prisma_client.spend_log_transactions] == ["row9"] + + +@pytest.mark.asyncio +async def test_update_spend_logs_does_not_requeue_non_transport_failures( + mock_prisma_client: Any, make_spend_log_row: Any +) -> None: + """Only transport failures are worth replaying. A rejection the DB will keep + rejecting must not be requeued, or it would wedge the queue forever. + """ + mock_prisma_client.db.litellm_spendlogs.create_many = AsyncMock(side_effect=ValueError("bad payload")) + proxy_logging = MagicMock() + proxy_logging.failure_handler = AsyncMock() + mock_prisma_client.spend_log_transactions = [] + + with pytest.raises(ValueError): + await ProxyUpdateSpend.update_spend_logs( + n_retry_times=1, + prisma_client=mock_prisma_client, + db_writer_client=None, + proxy_logging_obj=proxy_logging, + logs_to_process=[make_spend_log_row(request_id="a")], + ) + + assert mock_prisma_client.spend_log_transactions == [] + assert mock_prisma_client.db.litellm_spendlogs.create_many.await_count == 1 + + @pytest.mark.asyncio async def test_update_spend_logs_caps_isolation_attempts_under_poison_flood( mock_prisma_client: Any, make_spend_log_row: Any