diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index e7a08711eb1..4005b5bebd8 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -3967,7 +3967,11 @@ class ProxyBaseLLMRequestProcessing: # Starlette closes on disconnect, so the nested iterator hook (which # only sees GeneratorExit on GC) cannot own the refund. client_disconnected = not stream_completed - if not delivered_chunk and not _withheld_provider_output(response): + if ( + not delivered_chunk + and not _withheld_provider_output(response) + and user_api_key_dict.agent_invocation_cost is None + ): from litellm.proxy.spend_tracking.budget_reservation import ( release_budget_reservation_on_cancel, ) diff --git a/tests/test_litellm/proxy/test_budget_reservation.py b/tests/test_litellm/proxy/test_budget_reservation.py index c8e4df1030f..6d2c10321e8 100644 --- a/tests/test_litellm/proxy/test_budget_reservation.py +++ b/tests/test_litellm/proxy/test_budget_reservation.py @@ -3720,3 +3720,38 @@ async def test_unreserved_model_access_group_is_charged_alongside_a_reserved_one assert counter_cache.in_memory_cache.get_cache( key=model_access_group_spend_counter_key("starter") ) == pytest.approx(4.2) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("invocation_cost", [0.0, 0.01]) +async def test_agent_stream_cancellation_leaves_refund_to_request_cleanup(spend_counter_state, invocation_cost): + from litellm.proxy.middleware.budget_reservation_release_middleware import BudgetReservationReleaseMiddleware + from litellm.proxy.spend_tracking.budget_reservation import release_unbound_budget_reservation + + counter_cache, _ = spend_counter_state + key = "spend:agent:cancelled-agent" + counter_cache.set_cache(key, invocation_cost) + reservation = {"reserved_cost": invocation_cost, "input_cost": 0.0, "finalized": False, "entries": [{"counter_key": key, "reserved_cost": invocation_cost}]} + auth = UserAPIKeyAuth() + auth.agent_invocation_cost = invocation_cost + auth.budget_reservation = reservation + + async def cancel_before_chunk(user_api_key_dict, response, request_data): + raise asyncio.CancelledError() + yield "unreachable" + + generator, logging = _drive_streaming_cancel(auth, cancel_before_chunk) + + async def app(scope, receive, send): + try: + await anext(generator) + finally: + assert reservation["finalized"] is False + assert await counter_cache.async_get_cache(key) == pytest.approx(invocation_cost) + + middleware = BudgetReservationReleaseMiddleware(app, release_unbound_budget_reservation) + with pytest.raises(asyncio.CancelledError): + await middleware({"type": "http", "state": {"budget_reservation": reservation}}, AsyncMock(), AsyncMock()) + assert reservation["finalized"] is True + assert await counter_cache.async_get_cache(key) == pytest.approx(0.0) + logging._arelease_max_parallel_requests_on_disconnect.assert_awaited_once()