mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
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:
parent
219b4ebe46
commit
46e090d2f3
3 changed files with 70 additions and 10 deletions
|
|
@ -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),
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
"""
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue