diff --git a/litellm/a2a_protocol/streaming_iterator.py b/litellm/a2a_protocol/streaming_iterator.py index 2a20979501a..3e69869e0d9 100644 --- a/litellm/a2a_protocol/streaming_iterator.py +++ b/litellm/a2a_protocol/streaming_iterator.py @@ -12,6 +12,7 @@ from litellm._logging import verbose_logger from litellm.a2a_protocol.cost_calculator import A2ACostCalculator from litellm.a2a_protocol.utils import A2ARequestUtils from litellm.litellm_core_utils.asyncify import asyncify +from litellm.litellm_core_utils.core_helpers import bind_budget_reservation_to_callbacks from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj if TYPE_CHECKING: @@ -130,6 +131,8 @@ class A2AStreamingIterator(Generic[_StreamChunk]): # Build result for logging result: Final = self._build_logging_result(usage) + bind_budget_reservation_to_callbacks(self.logging_obj.litellm_params) + # Call success handlers - they will build standard_logging_object asyncio.create_task( self.logging_obj.dispatch_success_handlers( diff --git a/tests/unit/a2a_protocol/test_a2a_streaming_iterator.py b/tests/unit/a2a_protocol/test_a2a_streaming_iterator.py index 2e883e91fda..9bf43537737 100644 --- a/tests/unit/a2a_protocol/test_a2a_streaming_iterator.py +++ b/tests/unit/a2a_protocol/test_a2a_streaming_iterator.py @@ -135,3 +135,40 @@ async def test_stream_completion_counts_tokens_off_the_event_loop(monkeypatch): assert usage.prompt_tokens > 100_000 assert usage.completion_tokens > 100_000 assert_loop_stayed_free(took, lags) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("outcome", ["success", "failure", "cancelled"]) +async def test_stream_reservation_survives_cleanup_only_when_billing_is_scheduled(outcome): + from unittest.mock import AsyncMock + + from litellm.proxy.spend_tracking.budget_reservation import release_unbound_budget_reservation + + reservation = {"reserved_cost": 0.01, "entries": [], "finalized": False} + logging_obj = SimpleNamespace( + litellm_params={"metadata": {"user_api_key_budget_reservation": reservation}}, + model_call_details={}, + dispatch_success_handlers=AsyncMock(), + ) + + async def stream(): + yield {"result": {"kind": "message", "parts": [{"kind": "text", "text": "hello"}]}} + if outcome == "failure": + raise RuntimeError("upstream failed") + if outcome == "cancelled": + raise asyncio.CancelledError() + + iterator = A2AStreamingIterator( + stream=stream(), + request=SimpleNamespace(params=SimpleNamespace(message={"parts": [{"kind": "text", "text": "hi"}]})), + logging_obj=logging_obj, + ) + if outcome == "success": + assert len([chunk async for chunk in iterator]) == 1 + else: + with pytest.raises(RuntimeError if outcome == "failure" else asyncio.CancelledError): + _ = [chunk async for chunk in iterator] + await release_unbound_budget_reservation(reservation) + assert reservation["finalized"] is (outcome != "success") + await asyncio.sleep(0) + assert logging_obj.dispatch_success_handlers.await_count == (1 if outcome == "success" else 0)