mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(agents): keep cancelled stream refunds under one owner
This commit is contained in:
parent
a3ec702e2f
commit
1529494bed
2 changed files with 40 additions and 1 deletions
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue