fix(anthropic_messages): bill partial spend when a queued pump error is never consumed

When the upstream errors while the client is still connected, the pump
forwards the exception through the relay queue so the proxy's failure
handling re-raises it. If the client disconnects before consuming that
queued exception, neither the failure hook nor billing ran and the spend
row was lost. The pump now waits for client detach and, if the exception
was never consumed, salvages partial spend like the post-disconnect
error path.

Also rewrites the bedrock disconnect logging test to the detached-pump
contract: billing fires after the upstream drain completes, not
synchronously at aclose().
This commit is contained in:
mateo-berri 2026-08-31 12:47:26 -07:00
parent 219b4ebe46
commit 46e090d2f3
3 changed files with 70 additions and 10 deletions

View file

@ -157,6 +157,17 @@ def _try_claim_detached_drain_slot() -> bool:
return True
def _exception_left_unconsumed(queue: "asyncio.Queue[bytes | None | BaseException]", exc: BaseException) -> bool:
"""After client detach the relay never reads the queue again, so drain it here.
The forwarded exception still sitting in the queue means the relay tore
down before re-raising it, so the proxy's failure handling never ran and
the caller must salvage spend itself.
"""
remaining: Final = tuple(queue.get_nowait() for _ in range(queue.qsize()))
return any(item is exc for item in remaining)
def _sse_event(event_type: str, payload: Mapping[str, object]) -> bytes:
return f"event: {event_type}\ndata: {json.dumps(payload)}\n\n".encode()
@ -637,13 +648,16 @@ class BaseAnthropicMessagesStreamingIterator:
Handing the original exception to the client-facing generator lets it
re-raise so the proxy's failure handling keeps the provider status and
owns logging (no success-bill). If the client already went away, no
failure hook runs, so bill the partial instead of dropping the request.
owns logging (no success-bill). If the client already went away, or
disconnects before ever consuming the queued exception, no failure hook
runs, so bill the partial instead of dropping the request.
"""
from litellm._logging import verbose_proxy_logger
if not client_detached.is_set() and await self._enqueue_for_client(queue, client_detached, exc):
return
await client_detached.wait()
if not _exception_left_unconsumed(queue, exc):
return
verbose_proxy_logger.warning(
"async_sse_wrapper upstream pump failed after client disconnect (%d chunks): %s(%s)",
len(collected_chunks),

View file

@ -601,6 +601,43 @@ async def test_async_sse_wrapper_salvages_partial_spend_on_upstream_error_after_
assert iterator.logged_chunks == received
@pytest.mark.asyncio
async def test_async_sse_wrapper_salvages_spend_when_queued_error_is_never_consumed():
"""
When the upstream errors while the client is still connected, the pump
forwards the exception through the queue expecting the relay to re-raise it
into the proxy's failure handling. If the client disconnects before
consuming that queued exception, the handoff never happens and no failure
hook runs, so the pump must notice the unconsumed exception at teardown and
salvage partial spend instead of dropping the row entirely.
"""
upstream_errored = asyncio.Event()
async def _failing_stream():
yield {"type": "message_start", "message": {"id": "msg_1", "usage": {"input_tokens": 52, "output_tokens": 1}}}
yield {"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": "partial"}}
upstream_errored.set()
raise _ProviderStreamError("mid-stream failure", status_code=500)
iterator = _RecordingLoggingIterator(
litellm_logging_obj=_make_logging_obj("test_salvage_on_unconsumed_queued_error"),
request_body={},
)
gen = iterator.async_sse_wrapper(_failing_stream())
received = [await gen.__anext__(), await gen.__anext__()]
await upstream_errored.wait() # exception is now queued behind the consumed chunks
await gen.aclose() # client disconnects without ever consuming the queued exception
for _ in range(100):
if iterator.logged_chunks:
break
await asyncio.sleep(0.01)
assert iterator.logging_call_count == 1
assert iterator.logged_chunks == received
@pytest.mark.asyncio
async def test_async_sse_wrapper_applies_backpressure_to_slow_client(monkeypatch):
"""

View file

@ -3066,17 +3066,21 @@ def test_bedrock_invoke_messages_allows_converted_websearch_function_tool():
async def test_bedrock_sse_wrapper_dispatches_logging_on_client_disconnect():
"""
Regression test for LIT-5839: closing the outer bedrock_sse_wrapper
mid-stream (what the proxy does on a client disconnect) must close the
inner async_sse_wrapper deterministically so the partial-stream logging
fires. `completion_start_time` is only stamped on the logging object by
that dispatch, so it observing a value proves the whole chain ran.
mid-stream (what the proxy does on a client disconnect) must not lose the
stream's spend logging. Since the detached-pump relay, the upstream read
survives the disconnect and billing fires once the provider stream ends,
so the dispatch is awaited after releasing the upstream instead of being
observed synchronously at aclose(). `completion_start_time` is only
stamped on the logging object by that dispatch, so it observing a value
proves the whole chain ran.
"""
cfg = AmazonAnthropicClaudeMessagesConfig()
release_upstream = asyncio.Event()
async def _hanging_stream():
async def _gated_stream():
yield {"type": "message_start", "message": {"id": "msg_1", "usage": {"input_tokens": 25, "output_tokens": 1}}}
yield {"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": "partial"}}
await asyncio.Event().wait()
await release_upstream.wait()
logging_obj = LiteLLMLoggingObj(
model="bedrock/invoke/anthropic.claude-3-sonnet-20240229-v1:0",
@ -3087,11 +3091,16 @@ async def test_bedrock_sse_wrapper_dispatches_logging_on_client_disconnect():
litellm_call_id="test_bedrock_sse_wrapper_disconnect_logging",
function_id="test_bedrock_sse_wrapper_disconnect_logging",
)
wrapped = cfg.bedrock_sse_wrapper(_hanging_stream(), litellm_logging_obj=logging_obj, request_body={})
wrapped = cfg.bedrock_sse_wrapper(_gated_stream(), litellm_logging_obj=logging_obj, request_body={})
await wrapped.__anext__()
await wrapped.__anext__()
assert logging_obj.completion_start_time is None
await wrapped.aclose()
release_upstream.set()
for _ in range(500):
if logging_obj.completion_start_time is not None:
break
await asyncio.sleep(0.01)
assert logging_obj.completion_start_time is not None