fix(anthropic_messages): park deferred billing before the end-of-stream sentinel

At end of drain the pump enqueued the sentinel first and picked the
billing mode from client_detached afterward, so a client that consumed
the sentinel and tore the relay down before the pump resumed (possible
whenever the sentinel enqueue hit a full queue) had its fully delivered
response billed through the teardown path, skipping the proxy's
post-response hook. Bill or park before the sentinel goes out, and let
an unconsumed sentinel fall back to dispatching the parked billing.
This commit is contained in:
mateo-berri 2026-08-31 16:46:12 -07:00
parent 46e090d2f3
commit a99f62d1bc
2 changed files with 137 additions and 6 deletions

View file

@ -598,7 +598,12 @@ class BaseAnthropicMessagesStreamingIterator:
"""Drain the whole upstream into ``queue`` (backpressured) and bill once.
Runs detached so a client disconnect can't interrupt the upstream read;
see ``async_sse_wrapper`` for the full rationale. Returns after billing.
see ``async_sse_wrapper`` for the full rationale. On a completed drain
the success billing (or deferred park) happens before the end-of-stream
sentinel is enqueued: the relay can only tear down after consuming the
sentinel, so its teardown can never outrun the park and get mistaken
for a client disconnect, and a sentinel the client never consumes falls
back to dispatching the parked billing here.
"""
from litellm._logging import verbose_proxy_logger
@ -631,11 +636,17 @@ class BaseAnthropicMessagesStreamingIterator:
await self._handle_pump_upstream_error(queue, client_detached, collected_chunks, exc)
return
if not client_detached.is_set():
if not saw_terminal_event:
await self._enqueue_for_client(queue, client_detached, _incomplete_stream_error_sse_event())
await self._enqueue_for_client(queue, client_detached, None)
await self._bill_collected_chunks(collected_chunks, stream_teardown=client_detached.is_set())
if client_detached.is_set():
await self._bill_collected_chunks(collected_chunks, stream_teardown=True)
return
if not saw_terminal_event and not await self._enqueue_for_client(
queue, client_detached, _incomplete_stream_error_sse_event()
):
await self._bill_collected_chunks(collected_chunks, stream_teardown=True)
return
await self._bill_collected_chunks(collected_chunks, stream_teardown=False)
if not await self._enqueue_for_client(queue, client_detached, None):
self._dispatch_pending_deferred_logging()
async def _handle_pump_upstream_error(
self,

View file

@ -1156,3 +1156,123 @@ async def test_normal_end_without_deferred_dispatch_enqueues_immediately(monkeyp
assert len(worker.enqueued) == 1
assert getattr(iterator.litellm_logging_obj, "_deferred_stream_complete_args", None) is None
worker.close_enqueued()
def _backpressured_wrapper(iterator, upstream_exhausted: asyncio.Event):
async def _stream():
try:
for event in COMPLETE_STREAM_EVENTS:
yield event
finally:
upstream_exhausted.set()
return iterator.async_sse_wrapper(_stream())
async def _drain_with_pauses_until_upstream_exhausted(gen, upstream_exhausted: asyncio.Event) -> list:
received = []
while not upstream_exhausted.is_set():
received.append(await gen.__anext__())
for _ in range(25):
await asyncio.sleep(0)
assert len(received) <= len(COMPLETE_STREAM_EVENTS)
return received
@pytest.mark.asyncio
async def test_normal_end_parks_deferred_logging_even_when_sentinel_enqueue_backpressured(monkeypatch):
"""
Regression: with a full relay queue at end of stream, the pump suspends
while enqueueing the end-of-stream sentinel, and a client that then drains
the whole tail tears the relay down (setting ``client_detached``) before
the pump resumes. That teardown is a normally completed response, not a
disconnect: billing must still park for the proxy's post-response hook
(preserving post_call decoration such as guardrail_information) instead of
enqueueing immediately through the teardown path.
"""
monkeypatch.setattr(streaming_iterator_module, "ANTHROPIC_MESSAGES_STREAM_RELAY_QUEUE_MAXSIZE", 2)
worker = _RecordingLoggingWorker()
monkeypatch.setattr(streaming_iterator_module, "GLOBAL_LOGGING_WORKER", worker)
dispatched = []
async def _deferred_stream_complete(logging_coroutine):
dispatched.append(logging_coroutine)
logging_coroutine.close()
iterator = _make_iterator("test_sentinel_backpressure_normal_end")
iterator.litellm_logging_obj._on_deferred_stream_complete = _deferred_stream_complete
upstream_exhausted = asyncio.Event()
gen = _backpressured_wrapper(iterator, upstream_exhausted)
received = await _drain_with_pauses_until_upstream_exhausted(gen, upstream_exhausted)
while True:
try:
received.append(await gen.__anext__())
except StopAsyncIteration:
break
for _ in range(100):
if worker.enqueued or getattr(iterator.litellm_logging_obj, "_deferred_stream_complete_args", None):
break
await asyncio.sleep(0.01)
assert len(received) == len(COMPLETE_STREAM_EVENTS)
assert worker.enqueued == [], "fully delivered stream billed through the teardown path"
assert dispatched == []
parked = getattr(iterator.litellm_logging_obj, "_deferred_stream_complete_args", None)
assert parked is not None, "pump never parked deferred billing"
parked[0].close()
@pytest.mark.asyncio
async def test_relay_teardown_dispatches_deferred_billing_when_sentinel_never_consumed(monkeypatch):
"""
Regression: when the pump has parked deferred billing but its end-of-stream
sentinel never fits in the full relay queue (the client disconnects without
draining the tail), the proxy's post-response hook never fires. Exactly one
of the relay teardown or the pump's fallback must dispatch the parked
billing, or the request logs no spend at all.
"""
monkeypatch.setattr(streaming_iterator_module, "ANTHROPIC_MESSAGES_STREAM_RELAY_QUEUE_MAXSIZE", 2)
worker = _RecordingLoggingWorker()
monkeypatch.setattr(streaming_iterator_module, "GLOBAL_LOGGING_WORKER", worker)
dispatched = []
deferred_fired = asyncio.Event()
def _deferred_stream_complete(logging_coroutine):
dispatched.append(logging_coroutine)
async def _consume():
logging_coroutine.close()
deferred_fired.set()
return _consume()
iterator = _make_iterator("test_sentinel_never_consumed_dispatch")
iterator.litellm_logging_obj._on_deferred_stream_complete = _deferred_stream_complete
upstream_exhausted = asyncio.Event()
gen = _backpressured_wrapper(iterator, upstream_exhausted)
await _drain_with_pauses_until_upstream_exhausted(gen, upstream_exhausted)
for _ in range(100):
if getattr(iterator.litellm_logging_obj, "_deferred_stream_complete_args", None) is not None:
break
await asyncio.sleep(0.01)
await gen.aclose()
for _ in range(100):
if dispatched:
break
await asyncio.sleep(0.01)
assert len(dispatched) == 1, "parked billing was never dispatched"
assert getattr(iterator.litellm_logging_obj, "_deferred_stream_complete_args", None) is None
assert getattr(iterator.litellm_logging_obj, "_on_deferred_stream_complete", None) is None
assert len(worker.enqueued) == 1, "teardown billing enqueued alongside the deferred dispatch"
await worker.enqueued[0]
assert deferred_fired.is_set()