mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
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:
parent
321779138e
commit
a85a9e1186
3 changed files with 94 additions and 72 deletions
|
|
@ -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),
|
||||
|
|
|
|||
|
|
@ -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, (
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue