fix(anthropic_messages): strip inline comments, add abort-upstream regression test

Strip net-new inline # blocks from streaming_iterator.py, the unit test file,
and the live-proxy regression test to comply with the no-new-comments rule.

Add test_async_sse_wrapper_aborts_upstream_when_detached_drain_cap_reached:
verifies that when the detached-drain cap is already full, the pump calls
aclose() on the upstream so the provider stops generating and billing
instead of continuing to stream while we record only the partial prefix.

Also fixes LIT001 (bare dict in AsyncIterator union) by replacing dict
with Mapping[str, object] across all three stream-type annotations, and
adds the required LIT003 reason strings to the three noqa: BLE001 directives.
This commit is contained in:
nuernber 2026-08-06 10:46:54 -07:00
parent 321779138e
commit a85a9e1186
3 changed files with 94 additions and 72 deletions

View file

@ -1,6 +1,6 @@
import asyncio
import json
from collections.abc import AsyncIterator
from collections.abc import AsyncIterator, Mapping
from datetime import datetime
from typing import Any, Final, Protocol, runtime_checkable
@ -22,18 +22,7 @@ from litellm.types.utils import GenericStreamingChunk, ModelResponseStream
GLOBAL_PASS_THROUGH_SUCCESS_HANDLER_OBJ: Final = PassThroughEndpointLogging()
# asyncio holds only a weak reference to a bare create_task() result, so a
# fire-and-forget task can be garbage-collected mid-run. The upstream pump
# below must outlive the client-facing generator (which is closed on client
# disconnect), so root every pump task in a module-level set per the stdlib
# guidance and drop it again from the done callback.
_UPSTREAM_PUMP_TASKS: Final[set[asyncio.Task[None]]] = set() # mutable-ok: stdlib strong-ref set for pump tasks
# Rooted set of pumps still draining upstream AFTER their client disconnected.
# Bounds how many detached drains run at once so a burst of slow/abandoned
# streams can't pin unbounded worker memory; a pump over the cap bills what it
# already collected instead of continuing to drain. Only ever touched from the
# event loop, so a plain set + len() check needs no lock.
_DETACHED_STREAM_DRAINS: Final[set[asyncio.Task[None]]] = set() # mutable-ok: bounded strong-ref set, detached drains
INCOMPLETE_STREAM_ERROR_MESSAGE: Final = (
@ -221,7 +210,7 @@ class BaseAnthropicMessagesStreamingIterator:
async def async_sse_wrapper(
self,
completion_stream: AsyncIterator[bytes | GenericStreamingChunk | ModelResponseStream | dict],
completion_stream: AsyncIterator[bytes | GenericStreamingChunk | ModelResponseStream | Mapping[str, object]],
) -> AsyncIterator[bytes]:
"""
Generic async SSE wrapper that converts streaming chunks to SSE format
@ -272,11 +261,6 @@ class BaseAnthropicMessagesStreamingIterator:
raise item
yield item
finally:
# Client-facing generator is being torn down (normal end, a
# re-raised upstream error, or a disconnect GeneratorExit). Signal
# the pump to stop enqueueing and unblock any backpressure-blocked
# put; the pump then either finishes billing or drains detached
# (subject to the cap) for accurate usage.
client_detached.set()
async def _bill_collected_chunks(
@ -295,6 +279,22 @@ class BaseAnthropicMessagesStreamingIterator:
exc,
)
@staticmethod
async def _abort_upstream(
completion_stream: AsyncIterator[bytes | GenericStreamingChunk | ModelResponseStream | Mapping[str, object]],
) -> None:
"""Close the upstream provider stream so it stops generating and billing."""
from litellm._logging import verbose_proxy_logger
try:
await aclose_if_supported(completion_stream)
except Exception as exc: # noqa: BLE001 # abort is best-effort; log and continue
verbose_proxy_logger.warning(
"async_sse_wrapper failed to abort upstream stream: %s(%s)",
type(exc).__name__,
exc,
)
@staticmethod
async def _enqueue_for_client(
queue: "asyncio.Queue[bytes | None | BaseException]",
@ -328,7 +328,7 @@ class BaseAnthropicMessagesStreamingIterator:
async def _pump_upstream_to_queue(
self,
completion_stream: AsyncIterator[bytes | GenericStreamingChunk | ModelResponseStream | dict],
completion_stream: AsyncIterator[bytes | GenericStreamingChunk | ModelResponseStream | Mapping[str, object]],
queue: "asyncio.Queue[bytes | None | BaseException]",
client_detached: "asyncio.Event",
) -> None:
@ -352,21 +352,19 @@ class BaseAnthropicMessagesStreamingIterator:
if not client_detached.is_set():
await self._enqueue_for_client(queue, client_detached, encoded_chunk)
continue
# Client has gone: keep draining only to reach the terminal usage
# event for billing, but claim a detached-drain slot first; over
# the cap, bill what we have rather than pinning more memory.
if not draining_detached:
if not _try_claim_detached_drain_slot():
verbose_proxy_logger.warning(
"async_sse_wrapper: detached-drain cap (%d) reached; billing %d partial "
"chunks without draining the rest of the upstream stream",
"chunks and aborting the upstream stream to stop provider billing",
ANTHROPIC_MESSAGES_MAX_DETACHED_STREAM_DRAINS,
len(collected_chunks),
)
await self._bill_collected_chunks(collected_chunks)
await self._abort_upstream(completion_stream)
return
draining_detached = True
except Exception as exc: # noqa: BLE001 # forward the provider error to a live client, else salvage spend
except Exception as exc: # noqa: BLE001 # upstream errors are handled/forwarded by _handle_pump_upstream_error
await self._handle_pump_upstream_error(queue, client_detached, collected_chunks, exc)
return
@ -383,17 +381,17 @@ class BaseAnthropicMessagesStreamingIterator:
collected_chunks: list[bytes], # mutable-ok: SSE buffer forwarded to list-typed _bill_collected_chunks
exc: BaseException,
) -> None:
"""Forward a provider error to a still-connected client, else salvage partial spend.
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.
"""
from litellm._logging import verbose_proxy_logger
if not client_detached.is_set() and await self._enqueue_for_client(queue, client_detached, exc):
# Preserve the provider-specific failure: the client-facing
# generator re-raises it and the proxy's failure handling (status
# code, post_call_failure_hook) runs. That path owns logging, so
# don't also success-bill.
return
# Client already gone (or disconnected before the error reached it): no
# failure hook will run, so salvage the partial spend instead of
# dropping the request entirely.
verbose_proxy_logger.warning(
"async_sse_wrapper upstream pump failed after client disconnect (%d chunks): %s(%s)",
len(collected_chunks),

View file

@ -96,14 +96,11 @@ async def test_v1_messages_streaming_disconnect_has_spend_log():
chunks_read = 0
# ---- send the streaming request and disconnect early ----
async with session.post(
f"{BASE_URL}/v1/messages", json=payload, headers=headers
) as resp:
assert resp.status == 200, f"/v1/messages failed: {await resp.text()}"
# Read a handful of SSE chunks, then break out (closes the
# connection, which is the "interruption").
async for raw_line in resp.content:
line = raw_line.decode("utf-8", errors="replace").strip()
if not line:
@ -111,7 +108,6 @@ async def test_v1_messages_streaming_disconnect_has_spend_log():
chunks_read += 1
print(f" chunk #{chunks_read}: {line[:120]}")
if chunks_read >= 5:
# We have received enough data — disconnect now.
break
assert chunks_read >= 3, (
@ -123,7 +119,6 @@ async def test_v1_messages_streaming_disconnect_has_spend_log():
f"Waiting for spend pipeline to flush …"
)
# ---- wait & poll for the spend-log entry ----
spend_data = None
max_retries = 4
for attempt in range(1, max_retries + 1):
@ -135,7 +130,6 @@ async def test_v1_messages_streaming_disconnect_has_spend_log():
break
print(" … not found yet")
# ---- assertions ----
assert spend_data is not None and len(spend_data) > 0, (
f"No spend-log entry found for spend_id={spend_id} "
f"after streaming disconnect. "
@ -148,12 +142,6 @@ async def test_v1_messages_streaming_disconnect_has_spend_log():
f"\nSpend-log entry:\n{json.dumps(log_entry, indent=2, default=str)}"
)
# A row alone is not enough: the earlier drop-in-finally attempt logged a
# row whose completion tokens reflected only the handful of chunks the
# client drained before disconnecting (~1-15), not the full response
# Bedrock generated and billed. The prompt is written to produce a long
# completion, so the recorded completion tokens must reflect the full
# upstream stream, well above what 5 SSE chunks could carry.
prompt_tokens = log_entry.get("prompt_tokens", 0)
completion_tokens = log_entry.get("completion_tokens", 0)
assert prompt_tokens > 0, (

View file

@ -247,12 +247,6 @@ def test_incomplete_stream_error_sse_event_is_valid_anthropic_error():
assert event.endswith("\n\n")
# The full stream a provider (Bedrock invoke) generates: a short prefix the
# client reads before disconnecting, then the tail (including the terminal
# ``message_delta`` carrying the real output_tokens) that arrives only after
# the client is gone. output_tokens=64 is the authoritative billed count; a
# naive "log whatever the client drained" implementation would instead see the
# ``message_start`` placeholder (output_tokens=1) and undercount ~64x.
_STREAM_PREFIX = (
{"type": "message_start", "message": {"id": "msg_1", "usage": {"input_tokens": 52, "output_tokens": 1}}},
{"type": "content_block_start", "index": 0, "content_block": {"type": "text", "text": ""}},
@ -299,7 +293,6 @@ async def test_async_sse_wrapper_bills_full_stream_after_client_disconnect():
async def _gated_stream():
for event in _STREAM_PREFIX:
yield event
# Block until the test releases the tail (after the client disconnects).
await tail_gated.wait()
for event in _STREAM_TAIL:
yield event
@ -312,31 +305,25 @@ async def test_async_sse_wrapper_bills_full_stream_after_client_disconnect():
gen = iterator.async_sse_wrapper(_gated_stream())
# Client reads the prefix, then disconnects (closes the generator).
client_chunks = []
async for chunk in gen:
client_chunks.append(chunk)
if len(client_chunks) == len(_STREAM_PREFIX):
break
await gen.aclose() # client disconnect tears down the client-facing generator
await gen.aclose()
# Now let the provider finish. The detached pump must still be alive.
tail_gated.set()
await asyncio.wait_for(upstream_fully_drained.wait(), timeout=5)
# Give the pump's finally (billing) a turn to run.
for _ in range(100):
if iterator.logged_chunks:
break
await asyncio.sleep(0.01)
# The client only ever saw the prefix.
assert len(client_chunks) == len(_STREAM_PREFIX)
# Billing saw the WHOLE stream, including the terminal usage event.
assert iterator.logged_chunks, "pump never billed after client disconnect"
assert _output_tokens_from_logged_chunks(iterator.logged_chunks) == 64
assert any(c.startswith(b"event: message_stop\n") for c in iterator.logged_chunks)
# No synthetic incomplete-stream error, because the real message_stop arrived.
assert not any(c.startswith(b"event: error\n") for c in iterator.logged_chunks)
@ -400,11 +387,9 @@ async def test_async_sse_wrapper_reraises_upstream_error_to_connected_client():
async for chunk in iterator.async_sse_wrapper(_failing_stream()):
received.append(chunk)
# Original exception + status preserved, not masked by a synthetic api_error.
assert excinfo.value.status_code == 529
assert received # the client still got the pre-error chunks
assert received
assert not any(c.startswith(b"event: error\n") for c in received)
# On the failure path we do NOT success-bill (failure handling owns logging).
assert iterator.logged_chunks == []
@ -439,7 +424,6 @@ async def test_async_sse_wrapper_salvages_partial_spend_on_upstream_error_after_
await asyncio.sleep(0.01)
assert len(received) == 2
# Partial spend was still recorded rather than the whole request being dropped.
assert iterator.logged_chunks == received
@ -466,12 +450,9 @@ async def test_async_sse_wrapper_applies_backpressure_to_slow_client(monkeypatch
iterator = _make_iterator("test_backpressure_slow_client")
gen = iterator.async_sse_wrapper(_fast_stream())
try:
await gen.__anext__() # read exactly one chunk, then stall
# Let the pump run as far as the bounded queue permits.
await gen.__anext__()
for _ in range(500):
await asyncio.sleep(0)
# Bounded by queue maxsize + the one in-flight put + the one delivered
# chunk; nowhere near the full 200-chunk stream.
assert produced <= 2 + 3, f"pump ran ahead unthrottled: produced {produced} of {total}"
assert produced < total
finally:
@ -493,7 +474,6 @@ async def test_async_sse_wrapper_bills_partial_when_detached_drain_cap_reached(m
monkeypatch.setattr(streaming_iterator_module, "ANTHROPIC_MESSAGES_MAX_DETACHED_STREAM_DRAINS", 1)
monkeypatch.setattr(streaming_iterator_module, "ANTHROPIC_MESSAGES_STREAM_RELAY_QUEUE_MAXSIZE", 4)
# Occupy the only detached-drain slot with a placeholder task.
async def _hold_slot():
await asyncio.sleep(3600)
@ -505,7 +485,6 @@ async def test_async_sse_wrapper_bills_partial_when_detached_drain_cap_reached(m
nonlocal tail_reached
yield {"type": "message_start", "message": {"id": "m", "usage": {"input_tokens": 5, "output_tokens": 1}}}
yield {"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": "x"}}
# These arrive only after the client has disconnected.
for i in range(100):
yield {"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": f"more{i}"}}
tail_reached = True
@ -524,10 +503,6 @@ async def test_async_sse_wrapper_bills_partial_when_detached_drain_cap_reached(m
break
await asyncio.sleep(0)
# Billed a small bounded partial (the prefix plus at most a queue's
# worth the pump ran ahead before disconnect) without draining the
# 100-chunk tail. The exact count depends on how far the bounded queue
# let the pump run ahead, so assert the bound, not an exact number.
assert iterator.logged_chunks, "capped pump never billed"
assert len(iterator.logged_chunks) <= 2 + streaming_iterator_module.ANTHROPIC_MESSAGES_STREAM_RELAY_QUEUE_MAXSIZE
assert len(iterator.logged_chunks) < 100
@ -538,6 +513,68 @@ async def test_async_sse_wrapper_bills_partial_when_detached_drain_cap_reached(m
streaming_iterator_module._DETACHED_STREAM_DRAINS.discard(holder)
@pytest.mark.asyncio
async def test_async_sse_wrapper_aborts_upstream_when_detached_drain_cap_reached(monkeypatch):
"""
Regression: when the cap is full and a disconnected pump bails, it must call
aclose on the upstream stream so the provider stops generating and billing,
not continue running the stream while we record only the partial prefix.
"""
monkeypatch.setattr(streaming_iterator_module, "ANTHROPIC_MESSAGES_MAX_DETACHED_STREAM_DRAINS", 1)
monkeypatch.setattr(streaming_iterator_module, "ANTHROPIC_MESSAGES_STREAM_RELAY_QUEUE_MAXSIZE", 4)
async def _hold_slot():
await asyncio.sleep(3600)
holder = asyncio.ensure_future(_hold_slot())
streaming_iterator_module._DETACHED_STREAM_DRAINS.add(holder)
class _AbortableStream:
def __init__(self):
self.aclose_called = False
self._remaining = iter(
(
{"type": "message_start", "message": {"id": "m", "usage": {"input_tokens": 5, "output_tokens": 1}}},
{"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": "x"}},
)
+ tuple(
{"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": f"t{i}"}}
for i in range(50)
)
)
def __aiter__(self):
return self
async def __anext__(self):
try:
return next(self._remaining)
except StopIteration:
raise StopAsyncIteration
async def aclose(self):
self.aclose_called = True
stream = _AbortableStream()
iterator = _RecordingLoggingIterator(litellm_logging_obj=_make_logging_obj("abort_upstream_at_cap"), request_body={})
try:
gen = iterator.async_sse_wrapper(stream)
await gen.__anext__()
await gen.__anext__()
await gen.aclose()
for _ in range(200):
if iterator.logged_chunks:
break
await asyncio.sleep(0)
assert iterator.logged_chunks, "capped pump never billed"
assert stream.aclose_called, "upstream aclose was not called when the detached-drain cap was reached"
finally:
holder.cancel()
streaming_iterator_module._DETACHED_STREAM_DRAINS.discard(holder)
@pytest.mark.asyncio
async def test_async_sse_wrapper_drains_detached_when_cap_available(monkeypatch):
"""Complement to the cap test: with a slot free, a disconnected pump drains
@ -565,5 +602,4 @@ async def test_async_sse_wrapper_drains_detached_when_cap_available(monkeypatch)
await asyncio.sleep(0.01)
assert any(c.startswith(b"event: message_stop\n") for c in iterator.logged_chunks)
# Slot released once the drain finished.
assert len(streaming_iterator_module._DETACHED_STREAM_DRAINS) == 0