From a170800f625219bb701837b330ecf262a97fa8f4 Mon Sep 17 00:00:00 2001 From: mubashir1osmani Date: Mon, 8 Jun 2026 19:10:24 -0700 Subject: [PATCH] add tests for spend writer deadlock for end users --- .../hooks/test_proxy_track_cost_callback.py | 14 +- .../test_proxy_update_spend.py | 169 ++++++++++++++++++ 2 files changed, 179 insertions(+), 4 deletions(-) diff --git a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py index 56d26ce9178..56e30b0bcca 100644 --- a/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py +++ b/tests/test_litellm/proxy/hooks/test_proxy_track_cost_callback.py @@ -17,6 +17,14 @@ from litellm.proxy.hooks.proxy_track_cost_callback import ( _update_database_and_spend_counters, ) +_UPDATE_DATABASE_RESPONSE_COST_ARG_INDEX = 9 + + +def _response_cost_from_update_database_call(call) -> float: + if "response_cost" in call.kwargs: + return call.kwargs["response_cost"] + return call.args[_UPDATE_DATABASE_RESPONSE_COST_ARG_INDEX] + @pytest.mark.asyncio async def test_async_post_call_failure_hook(): @@ -295,10 +303,8 @@ async def test_track_cost_callback_zeroes_response_cost_on_cache_hit(): ) mock_proxy_logging.db_spend_update_writer.update_database.assert_awaited_once() - recorded_cost = ( - mock_proxy_logging.db_spend_update_writer.update_database.await_args.kwargs[ - "response_cost" - ] + recorded_cost = _response_cost_from_update_database_call( + mock_proxy_logging.db_spend_update_writer.update_database.await_args ) assert recorded_cost == 0.0 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 6a4fd516c9b..ea00680826d 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 @@ -9,14 +9,27 @@ Symbols pinned here: from __future__ import annotations import asyncio +import types from typing import Any, Dict, List from unittest.mock import AsyncMock, MagicMock +import prisma.errors import pytest from litellm.proxy.utils import ProxyUpdateSpend +class _FailingTxCM: + def __init__(self, exc: Exception) -> None: + self._exc = exc + + async def __aenter__(self) -> Any: + raise self._exc + + async def __aexit__(self, *exc: Any) -> None: + return None + + class _AsyncCM: def __init__(self, target: Any) -> None: self.target = target @@ -115,6 +128,162 @@ async def test_update_end_user_spend_non_connection_error_raises_immediately( ) +@pytest.mark.parametrize( + "db_error", + [ + prisma.errors.PrismaError("deadlock detected"), + prisma.errors.PrismaError("pool timeout waiting for connection"), + ], +) +@pytest.mark.asyncio +async def test_update_end_user_spend_retries_db_contention_with_jittered_backoff( + mock_prisma_client: Any, + monkeypatch: pytest.MonkeyPatch, + db_error: Exception, +) -> None: + """Under Postgres contention, end-user spend must retry with jittered exponential + backoff, apply the increment exactly once, and never leave a partial upsert behind. + """ + import litellm.proxy.utils as utils_mod + + end_user_id = "customer-under-load" + response_cost = 0.75 + upsert_calls: list[dict[str, Any]] = [] + tx_attempts = {"count": 0} + + batcher = MagicMock() + + def _capture_upsert(**kwargs: Any) -> None: + upsert_calls.append(kwargs) + + batcher.litellm_endusertable.upsert = MagicMock(side_effect=_capture_upsert) + transaction = MagicMock() + transaction.batch_ = lambda: _AsyncCM(batcher) + + def _tx_factory(timeout: Any) -> Any: + tx_attempts["count"] += 1 + if tx_attempts["count"] == 1: + return _FailingTxCM(db_error) + return _AsyncCM(transaction) + + mock_prisma_client.db.tx = _tx_factory + + sleeps: list[float] = [] + + async def _fake_sleep(seconds: float) -> None: + sleeps.append(seconds) + + monkeypatch.setattr(utils_mod.asyncio, "sleep", _fake_sleep) + + jitter_calls: list[tuple[float, float]] = [] + + def _fake_uniform(lo: float, hi: float) -> float: + jitter_calls.append((lo, hi)) + return 1.25 + + monkeypatch.setattr( + utils_mod, + "random", + types.SimpleNamespace(uniform=_fake_uniform), + raising=False, + ) + + proxy_logging = MagicMock() + proxy_logging.failure_handler = AsyncMock() + + await ProxyUpdateSpend.update_end_user_spend( + n_retry_times=1, + prisma_client=mock_prisma_client, + proxy_logging_obj=proxy_logging, + end_user_list_transactions={end_user_id: response_cost}, + ) + + assert tx_attempts["count"] == 2 + assert jitter_calls == [(1.0, 2.0)] + assert sleeps == [1.25] + assert len(upsert_calls) == 1 + assert upsert_calls[0]["where"]["user_id"] == end_user_id + assert upsert_calls[0]["data"]["update"]["spend"]["increment"] == response_cost + assert upsert_calls[0]["data"]["create"]["spend"] == response_cost + + +@pytest.mark.asyncio +async def test_end_user_spend_queue_flush_retries_contention_without_leftover_items( + mock_prisma_client: Any, + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Queued end-user spend aggregated across requests must survive a contention + retry, write once, and leave the in-memory queue empty for that user. + """ + from litellm.proxy._types import Litellm_EntityType, SpendUpdateQueueItem + from litellm.proxy.db.db_spend_update_writer import DBSpendUpdateWriter + import litellm.proxy.utils as utils_mod + + end_user_id = "customer-queued" + writer = DBSpendUpdateWriter() + await writer.spend_update_queue.add_update( + SpendUpdateQueueItem( + entity_type=Litellm_EntityType.END_USER, + entity_id=end_user_id, + response_cost=0.10, + ) + ) + await writer.spend_update_queue.add_update( + SpendUpdateQueueItem( + entity_type=Litellm_EntityType.END_USER, + entity_id=end_user_id, + response_cost=0.15, + ) + ) + + upsert_calls: list[dict[str, Any]] = [] + tx_attempts = {"count": 0} + batcher = MagicMock() + + def _capture_upsert(**kwargs: Any) -> None: + upsert_calls.append(kwargs) + + batcher.litellm_endusertable.upsert = MagicMock(side_effect=_capture_upsert) + transaction = MagicMock() + transaction.batch_ = lambda: _AsyncCM(batcher) + + def _tx_factory(timeout: Any) -> Any: + tx_attempts["count"] += 1 + if tx_attempts["count"] == 1: + return _FailingTxCM(prisma.errors.PrismaError("deadlock detected")) + return _AsyncCM(transaction) + + mock_prisma_client.db.tx = _tx_factory + + async def _fake_sleep(_: float) -> None: + return None + + monkeypatch.setattr(utils_mod.asyncio, "sleep", _fake_sleep) + monkeypatch.setattr( + utils_mod, + "random", + types.SimpleNamespace(uniform=lambda lo, hi: lo), + raising=False, + ) + + proxy_logging = MagicMock() + proxy_logging.failure_handler = AsyncMock() + + await writer._commit_spend_updates_to_db_without_redis_buffer( + prisma_client=mock_prisma_client, + n_retry_times=1, + proxy_logging_obj=proxy_logging, + ) + + assert writer.spend_update_queue.update_queue.qsize() == 0 + assert tx_attempts["count"] == 2 + assert len(upsert_calls) == 1 + assert upsert_calls[0]["where"]["user_id"] == end_user_id + assert upsert_calls[0]["data"]["update"]["spend"]["increment"] == pytest.approx( + 0.25 + ) + + @pytest.mark.asyncio async def test_update_spend_logs_writes_batches_via_create_many( mock_prisma_client: Any, make_spend_log_row: Any