mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(agents): retain streamed spend until callback settlement
This commit is contained in:
parent
92e13728f5
commit
bcfa93d95c
2 changed files with 40 additions and 0 deletions
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue