mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-07 08:26:10 +00:00
Merge pull request #36008 from nuernber/litellm_bedrock_messages_disconnect_billing
fix(anthropic_messages): drain upstream in a detached pump so client …
This commit is contained in:
commit
99703a30f0
7 changed files with 990 additions and 79 deletions
|
|
@ -57,7 +57,7 @@
|
|||
"limit": 5611
|
||||
},
|
||||
"reportMissingTypeArgument": {
|
||||
"limit": 15350
|
||||
"limit": 15348
|
||||
},
|
||||
"reportMissingTypeStubs": {
|
||||
"limit": 40
|
||||
|
|
@ -99,19 +99,19 @@
|
|||
"limit": 0
|
||||
},
|
||||
"reportUnknownArgumentType": {
|
||||
"limit": 44368
|
||||
"limit": 44364
|
||||
},
|
||||
"reportUnknownLambdaType": {
|
||||
"limit": 109
|
||||
},
|
||||
"reportUnknownMemberType": {
|
||||
"limit": 38468
|
||||
"limit": 38465
|
||||
},
|
||||
"reportUnknownParameterType": {
|
||||
"limit": 19665
|
||||
"limit": 19663
|
||||
},
|
||||
"reportUnknownVariableType": {
|
||||
"limit": 30066
|
||||
"limit": 30064
|
||||
},
|
||||
"reportUnnecessaryCast": {
|
||||
"limit": 111
|
||||
|
|
|
|||
|
|
@ -484,6 +484,22 @@ FIREWORKS_AI_80_B: Final = int(os.getenv("FIREWORKS_AI_80_B", 80))
|
|||
#### Logging callback constants ####
|
||||
REDACTED_BY_LITELM_STRING: Final = "REDACTED_BY_LITELM"
|
||||
MAX_LANGFUSE_INITIALIZED_CLIENTS: Final = int(os.getenv("MAX_LANGFUSE_INITIALIZED_CLIENTS", 50))
|
||||
# Backpressure + lifetime bounds for the /v1/messages streaming relay (see
|
||||
# BaseAnthropicMessagesStreamingIterator.async_sse_wrapper). The relay queue is
|
||||
# bounded so a slow client throttles the upstream pump instead of letting it
|
||||
# buffer the whole response in memory; the detached-drain cap bounds how many
|
||||
# post-disconnect drains may run concurrently so client behavior can't create
|
||||
# unbounded worker state.
|
||||
ANTHROPIC_MESSAGES_STREAM_RELAY_QUEUE_MAXSIZE: Final = int(
|
||||
os.getenv("ANTHROPIC_MESSAGES_STREAM_RELAY_QUEUE_MAXSIZE", "1024")
|
||||
)
|
||||
# Setting this to 0 disables detached draining entirely: every post-disconnect
|
||||
# pump bills whatever partial output it has already collected and aborts the
|
||||
# upstream stream immediately, instead of continuing to drain for the real
|
||||
# terminal usage.
|
||||
ANTHROPIC_MESSAGES_MAX_DETACHED_STREAM_DRAINS: Final = int(
|
||||
os.getenv("ANTHROPIC_MESSAGES_MAX_DETACHED_STREAM_DRAINS", "100")
|
||||
)
|
||||
LOGGING_WORKER_CONCURRENCY: Final = int(os.getenv("LOGGING_WORKER_CONCURRENCY", 100)) # Must be above 0
|
||||
LOGGING_WORKER_MAX_QUEUE_SIZE: Final = int(os.getenv("LOGGING_WORKER_MAX_QUEUE_SIZE", 50_000))
|
||||
LOGGING_WORKER_MAX_TIME_PER_COROUTINE: Final = float(os.getenv("LOGGING_WORKER_MAX_TIME_PER_COROUTINE", 20.0))
|
||||
|
|
|
|||
|
|
@ -8,6 +8,10 @@ import httpx
|
|||
from pydantic import TypeAdapter
|
||||
from typing_extensions import TypedDict
|
||||
|
||||
from litellm.constants import (
|
||||
ANTHROPIC_MESSAGES_MAX_DETACHED_STREAM_DRAINS,
|
||||
ANTHROPIC_MESSAGES_STREAM_RELAY_QUEUE_MAXSIZE,
|
||||
)
|
||||
from litellm.litellm_core_utils.core_helpers import process_response_headers
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.litellm_core_utils.logging_worker import GLOBAL_LOGGING_WORKER
|
||||
|
|
@ -21,6 +25,9 @@ from litellm.types.utils import GenericStreamingChunk, ModelResponseStream
|
|||
|
||||
GLOBAL_PASS_THROUGH_SUCCESS_HANDLER_OBJ: Final = PassThroughEndpointLogging()
|
||||
|
||||
_UPSTREAM_PUMP_TASKS: Final[set[asyncio.Task[None]]] = set() # mutable-ok: stdlib strong-ref set for pump tasks
|
||||
_DETACHED_STREAM_DRAINS: Final[set[asyncio.Task[None]]] = set() # mutable-ok: bounded strong-ref set, detached drains
|
||||
|
||||
INCOMPLETE_STREAM_ERROR_MESSAGE: Final = (
|
||||
"Provider stream ended before emitting a message_stop event; "
|
||||
"the response is incomplete and any partial content (e.g. tool_use input JSON) may be truncated."
|
||||
|
|
@ -133,6 +140,34 @@ def _is_terminal_stream_chunk(chunk: object) -> bool:
|
|||
return _is_message_stop_chunk(chunk) or _is_provider_error_chunk(chunk)
|
||||
|
||||
|
||||
def _try_claim_detached_drain_slot() -> bool:
|
||||
"""Claim a detached-drain slot for the current task, bounding concurrency.
|
||||
|
||||
Returns True if a slot was claimed (the caller may keep draining upstream
|
||||
for billing) or False if the cap is already reached (the caller should stop
|
||||
and bill what it has). Only touched from the event loop, so the check +
|
||||
insert need no lock.
|
||||
"""
|
||||
if len(_DETACHED_STREAM_DRAINS) >= ANTHROPIC_MESSAGES_MAX_DETACHED_STREAM_DRAINS:
|
||||
return False
|
||||
current_task: Final = asyncio.current_task()
|
||||
if current_task is not None:
|
||||
_DETACHED_STREAM_DRAINS.add(current_task)
|
||||
current_task.add_done_callback(_DETACHED_STREAM_DRAINS.discard)
|
||||
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()
|
||||
|
||||
|
|
@ -414,17 +449,167 @@ 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
|
||||
and handles logging.
|
||||
|
||||
The upstream read runs in a detached background task (``_pump_upstream``)
|
||||
so that a client disconnect tears down only this client-facing generator,
|
||||
never the upstream drain + billing. The provider (e.g. Bedrock) keeps
|
||||
generating and billing the full response regardless of the client, so
|
||||
draining it to completion is what lets spend tracking see the real
|
||||
terminal ``message_delta`` / ``message_stop`` usage instead of a
|
||||
truncated placeholder count.
|
||||
|
||||
Chunks reach the client through a bounded queue. While the client is
|
||||
connected the pump blocks on a full queue (racing the disconnect
|
||||
signal), so a slow reader throttles the upstream read exactly as the old
|
||||
direct ``yield`` did instead of letting the whole response buffer in
|
||||
memory. Once the client goes away the pump stops enqueueing and only
|
||||
keeps a single ``collected_chunks`` copy for billing, and the number of
|
||||
such post-disconnect drains running at once is capped so client behavior
|
||||
can't create unbounded worker state; over the cap the pump bills what it
|
||||
has rather than draining further. Detached-drain lifetime is otherwise
|
||||
bounded by the upstream stream/read timeout.
|
||||
|
||||
An upstream failure (Bedrock read / decode / chunk-conversion error)
|
||||
that happens while the client is still connected is forwarded through
|
||||
the queue and re-raised here, so the original provider exception (and
|
||||
its status) reaches the proxy's failure handling unchanged rather than
|
||||
being masked by a generic incomplete-stream event.
|
||||
|
||||
This method provides the common logic for both Anthropic and Bedrock implementations.
|
||||
"""
|
||||
collected_chunks: Final = []
|
||||
saw_terminal_event = False
|
||||
queue: Final[asyncio.Queue[bytes | None | BaseException]] = asyncio.Queue(
|
||||
maxsize=ANTHROPIC_MESSAGES_STREAM_RELAY_QUEUE_MAXSIZE
|
||||
)
|
||||
client_detached: Final = asyncio.Event()
|
||||
|
||||
pump_task: Final = asyncio.create_task(self._pump_upstream_to_queue(completion_stream, queue, client_detached))
|
||||
_UPSTREAM_PUMP_TASKS.add(pump_task)
|
||||
pump_task.add_done_callback(_UPSTREAM_PUMP_TASKS.discard)
|
||||
|
||||
reached_end = False # rebind-ok: flipped once the relay consumes the end-of-stream sentinel
|
||||
try:
|
||||
while True:
|
||||
item = await queue.get()
|
||||
if item is None:
|
||||
reached_end = True
|
||||
break
|
||||
if isinstance(item, BaseException):
|
||||
raise item
|
||||
yield item
|
||||
finally:
|
||||
client_detached.set()
|
||||
if not reached_end:
|
||||
self._dispatch_pending_deferred_logging()
|
||||
|
||||
def _dispatch_pending_deferred_logging(self) -> None:
|
||||
"""Fire deferred billing that a torn-down response would otherwise drop.
|
||||
|
||||
When the pump finishes draining while the client is still connected it
|
||||
stores the logging coroutine for ProxyLogging._fire_deferred_stream_logging,
|
||||
which the proxy only fires on a normally completed response: a client
|
||||
disconnect (GeneratorExit / CancelledError) re-raises past it. Without
|
||||
this dispatch that window loses the spend row entirely.
|
||||
"""
|
||||
deferred_cb: Final = getattr(self.litellm_logging_obj, "_on_deferred_stream_complete", None)
|
||||
deferred_args: Final = getattr(self.litellm_logging_obj, "_deferred_stream_complete_args", None)
|
||||
if deferred_cb is None or deferred_args is None:
|
||||
return
|
||||
self.litellm_logging_obj._on_deferred_stream_complete = None
|
||||
self.litellm_logging_obj._deferred_stream_complete_args = None
|
||||
GLOBAL_LOGGING_WORKER.ensure_initialized_and_enqueue(async_coroutine=deferred_cb(*deferred_args))
|
||||
|
||||
async def _bill_collected_chunks(
|
||||
self,
|
||||
collected_chunks: list[bytes], # mutable-ok: SSE buffer forwarded to list-typed _handle_streaming_logging
|
||||
*,
|
||||
stream_teardown: bool,
|
||||
) -> None:
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
|
||||
try:
|
||||
await self._handle_streaming_logging(collected_chunks, stream_teardown=stream_teardown)
|
||||
except Exception as exc: # noqa: BLE001 # billing is best-effort; never crash the pump
|
||||
verbose_proxy_logger.warning(
|
||||
"async_sse_wrapper billing failed after %d chunks: %s(%s)",
|
||||
len(collected_chunks),
|
||||
type(exc).__name__,
|
||||
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]",
|
||||
client_detached: "asyncio.Event",
|
||||
item: bytes | None | BaseException,
|
||||
) -> bool:
|
||||
"""Deliver one item to the client, applying backpressure.
|
||||
|
||||
Returns True if the item was queued, False if the client disconnected
|
||||
before there was room (the item is then dropped, since a gone client
|
||||
can't receive it). Never blocks once the client has detached.
|
||||
"""
|
||||
if client_detached.is_set():
|
||||
return False
|
||||
try:
|
||||
queue.put_nowait(item)
|
||||
except asyncio.QueueFull:
|
||||
pass
|
||||
else:
|
||||
return True
|
||||
put_task: Final = asyncio.ensure_future(queue.put(item))
|
||||
detached_task: Final = asyncio.ensure_future(client_detached.wait())
|
||||
try:
|
||||
await asyncio.wait(frozenset((put_task, detached_task)), return_when=asyncio.FIRST_COMPLETED)
|
||||
finally:
|
||||
if not detached_task.done():
|
||||
detached_task.cancel()
|
||||
if put_task.done() and not put_task.cancelled():
|
||||
return True
|
||||
put_task.cancel()
|
||||
return False
|
||||
|
||||
async def _pump_upstream_to_queue(
|
||||
self,
|
||||
completion_stream: AsyncIterator[bytes | GenericStreamingChunk | ModelResponseStream | Mapping[str, object]],
|
||||
queue: "asyncio.Queue[bytes | None | BaseException]",
|
||||
client_detached: "asyncio.Event",
|
||||
) -> None:
|
||||
"""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. 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
|
||||
|
||||
collected_chunks: Final[list[bytes]] = [] # mutable-ok: SSE billing buffer appended to across the drain
|
||||
saw_terminal_event = False # rebind-ok: accumulates across the upstream loop
|
||||
draining_detached = False # rebind-ok: set once this pump claims a detached-drain slot
|
||||
try:
|
||||
async for chunk in completion_stream:
|
||||
if self.completion_start_time is None:
|
||||
|
|
@ -432,17 +617,62 @@ class BaseAnthropicMessagesStreamingIterator:
|
|||
saw_terminal_event = saw_terminal_event or _is_terminal_stream_chunk(chunk)
|
||||
encoded_chunk = self._convert_chunk_to_sse_format(chunk)
|
||||
collected_chunks.append(encoded_chunk)
|
||||
yield encoded_chunk
|
||||
except (GeneratorExit, asyncio.CancelledError):
|
||||
# A client disconnect tears the generator down at the yield, so the
|
||||
# post-loop logging below never runs and the tokens already streamed
|
||||
# (and billed by the provider) would never reach spend tracking. See LIT-5839.
|
||||
if collected_chunks:
|
||||
await self._handle_streaming_logging(collected_chunks, stream_teardown=True)
|
||||
raise
|
||||
if not client_detached.is_set():
|
||||
await self._enqueue_for_client(queue, client_detached, encoded_chunk)
|
||||
continue
|
||||
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 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, stream_teardown=True)
|
||||
await self._abort_upstream(completion_stream)
|
||||
return
|
||||
draining_detached = True
|
||||
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
|
||||
|
||||
if not saw_terminal_event:
|
||||
yield _incomplete_stream_error_sse_event()
|
||||
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()
|
||||
|
||||
# Handle logging after all chunks are processed
|
||||
await self._handle_streaming_logging(collected_chunks)
|
||||
async def _handle_pump_upstream_error(
|
||||
self,
|
||||
queue: "asyncio.Queue[bytes | None | BaseException]",
|
||||
client_detached: "asyncio.Event",
|
||||
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, 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):
|
||||
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),
|
||||
type(exc).__name__,
|
||||
exc,
|
||||
)
|
||||
await self._bill_collected_chunks(collected_chunks, stream_teardown=True)
|
||||
|
|
|
|||
|
|
@ -33,6 +33,13 @@ EXCLUDED_ROLLOUT_FLAGS = {
|
|||
"LITELLM_RUST",
|
||||
}
|
||||
|
||||
# Internal infrastructure tuning parameters for streaming/queue management
|
||||
# These are advanced settings with sensible defaults that most users should not modify
|
||||
EXCLUDED_INTERNAL_TUNING_VARS = {
|
||||
"ANTHROPIC_MESSAGES_MAX_DETACHED_STREAM_DRAINS",
|
||||
"ANTHROPIC_MESSAGES_STREAM_RELAY_QUEUE_MAXSIZE",
|
||||
}
|
||||
|
||||
EXCLUDED_TERMINAL_VARS = {
|
||||
"TERM",
|
||||
"TERM_PROGRAM",
|
||||
|
|
@ -50,7 +57,9 @@ EXCLUDED_TERMINAL_VARS = {
|
|||
"ALACRITTY_SOCKET",
|
||||
}
|
||||
|
||||
EXCLUDED_KEYS = frozenset(EXCLUDED_TERMINAL_VARS | EXCLUDED_GUARD_ONLY_VARS | EXCLUDED_ROLLOUT_FLAGS)
|
||||
EXCLUDED_KEYS = frozenset(
|
||||
EXCLUDED_TERMINAL_VARS | EXCLUDED_GUARD_ONLY_VARS | EXCLUDED_ROLLOUT_FLAGS | EXCLUDED_INTERNAL_TUNING_VARS
|
||||
)
|
||||
|
||||
# Directories to skip (dependencies, venvs, caches) - only scan litellm source
|
||||
SKIP_DIRS = {
|
||||
|
|
|
|||
|
|
@ -6,9 +6,7 @@ import pytest
|
|||
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.llms.anthropic.experimental_pass_through.messages import (
|
||||
streaming_iterator as streaming_iterator_module,
|
||||
)
|
||||
from litellm.llms.anthropic.experimental_pass_through.messages import streaming_iterator as streaming_iterator_module
|
||||
from litellm.llms.anthropic.experimental_pass_through.messages.streaming_iterator import (
|
||||
INCOMPLETE_STREAM_ERROR_MESSAGE,
|
||||
AnthropicMessagesStreamHiddenParams,
|
||||
|
|
@ -338,47 +336,6 @@ async def _events_then_hang(events):
|
|||
await asyncio.Event().wait()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_sse_wrapper_logs_partial_chunks_on_client_disconnect():
|
||||
"""
|
||||
Regression test for LIT-5839: a client disconnect tears the generator
|
||||
down with GeneratorExit at the yield, which used to skip the post-loop
|
||||
logging dispatch entirely, so the partial output tokens the provider
|
||||
already generated (and billed) never reached spend tracking.
|
||||
"""
|
||||
iterator = _RecordingLoggingIterator(
|
||||
litellm_logging_obj=_make_logging_obj("test_disconnect_logs_partial_chunks"),
|
||||
request_body={},
|
||||
)
|
||||
wrapped = iterator.async_sse_wrapper(_events_then_hang(TRUNCATED_TOOL_USE_EVENTS))
|
||||
streamed = [await wrapped.__anext__() for _ in range(len(TRUNCATED_TOOL_USE_EVENTS))]
|
||||
assert iterator.logging_call_count == 0
|
||||
|
||||
await wrapped.aclose()
|
||||
|
||||
assert iterator.logging_call_count == 1
|
||||
assert iterator.logged_chunks == streamed
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_sse_wrapper_logs_partial_chunks_on_cancellation():
|
||||
iterator = _RecordingLoggingIterator(
|
||||
litellm_logging_obj=_make_logging_obj("test_cancellation_logs_partial_chunks"),
|
||||
request_body={},
|
||||
)
|
||||
wrapped = iterator.async_sse_wrapper(_events_then_hang(TRUNCATED_TOOL_USE_EVENTS))
|
||||
streamed = [await wrapped.__anext__() for _ in range(len(TRUNCATED_TOOL_USE_EVENTS))]
|
||||
|
||||
consume_task = asyncio.ensure_future(wrapped.__anext__())
|
||||
await asyncio.sleep(0.01)
|
||||
consume_task.cancel()
|
||||
with pytest.raises(asyncio.CancelledError):
|
||||
await consume_task
|
||||
|
||||
assert iterator.logging_call_count == 1
|
||||
assert iterator.logged_chunks == streamed
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_sse_wrapper_skips_logging_on_disconnect_before_first_chunk():
|
||||
iterator = _RecordingLoggingIterator(
|
||||
|
|
@ -408,6 +365,561 @@ def test_incomplete_stream_error_sse_event_is_valid_anthropic_error():
|
|||
assert event.endswith("\n\n")
|
||||
|
||||
|
||||
_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": ""}},
|
||||
{"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": "The Roman"}},
|
||||
)
|
||||
_STREAM_TAIL = (
|
||||
{"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": " Empire ..."}},
|
||||
{"type": "content_block_stop", "index": 0},
|
||||
{"type": "message_delta", "delta": {"stop_reason": "end_turn"}, "usage": {"output_tokens": 64}},
|
||||
{"type": "message_stop"},
|
||||
)
|
||||
|
||||
|
||||
def _output_tokens_from_logged_chunks(chunks: list[bytes]) -> int | None:
|
||||
"""Read the last output_tokens the billing path would see from the SSE bytes."""
|
||||
latest: int | None = None
|
||||
for raw in chunks:
|
||||
for line in raw.decode().splitlines():
|
||||
if not line.startswith("data:"):
|
||||
continue
|
||||
data = json.loads(line[len("data:"):].strip())
|
||||
usage = data.get("usage") if isinstance(data, dict) else None
|
||||
if isinstance(usage, dict) and usage.get("output_tokens") is not None:
|
||||
latest = usage["output_tokens"]
|
||||
return latest
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_sse_wrapper_bills_full_stream_after_client_disconnect():
|
||||
"""
|
||||
Regression: on a client disconnect mid-stream the upstream provider keeps
|
||||
generating (and billing) the full response. The wrapper must keep draining
|
||||
that upstream to its terminal ``message_delta`` and bill the real
|
||||
output_tokens (64), not the partial count the client drained before leaving
|
||||
(the message_start placeholder, 1).
|
||||
|
||||
A ``tail_gated`` event holds back the stream tail until the client has
|
||||
disconnected, so the tail can only be captured by a drain that survives the
|
||||
client teardown - exactly the path the previous implementation dropped.
|
||||
"""
|
||||
tail_gated = asyncio.Event()
|
||||
upstream_fully_drained = asyncio.Event()
|
||||
|
||||
async def _gated_stream():
|
||||
for event in _STREAM_PREFIX:
|
||||
yield event
|
||||
await tail_gated.wait()
|
||||
for event in _STREAM_TAIL:
|
||||
yield event
|
||||
upstream_fully_drained.set()
|
||||
|
||||
iterator = _RecordingLoggingIterator(
|
||||
litellm_logging_obj=_make_logging_obj("test_bills_full_stream_after_disconnect"),
|
||||
request_body={},
|
||||
)
|
||||
|
||||
gen = iterator.async_sse_wrapper(_gated_stream())
|
||||
|
||||
client_chunks = []
|
||||
async for chunk in gen:
|
||||
client_chunks.append(chunk)
|
||||
if len(client_chunks) == len(_STREAM_PREFIX):
|
||||
break
|
||||
await gen.aclose()
|
||||
|
||||
tail_gated.set()
|
||||
await asyncio.wait_for(upstream_fully_drained.wait(), timeout=5)
|
||||
for _ in range(100):
|
||||
if iterator.logged_chunks:
|
||||
break
|
||||
await asyncio.sleep(0.01)
|
||||
|
||||
assert len(client_chunks) == len(_STREAM_PREFIX)
|
||||
|
||||
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)
|
||||
assert not any(c.startswith(b"event: error\n") for c in iterator.logged_chunks)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_sse_wrapper_bills_full_stream_when_client_reads_all():
|
||||
"""Happy path: when the client drains the whole stream, billing still sees
|
||||
the terminal output_tokens (64) and the client gets every chunk."""
|
||||
tail_gated = asyncio.Event()
|
||||
tail_gated.set() # no gating; full stream flows immediately
|
||||
|
||||
async def _full_stream():
|
||||
for event in (*_STREAM_PREFIX, *_STREAM_TAIL):
|
||||
yield event
|
||||
|
||||
iterator = _RecordingLoggingIterator(
|
||||
litellm_logging_obj=_make_logging_obj("test_bills_full_stream_happy_path"),
|
||||
request_body={},
|
||||
)
|
||||
client_chunks = [chunk async for chunk in iterator.async_sse_wrapper(_full_stream())]
|
||||
|
||||
for _ in range(100):
|
||||
if iterator.logged_chunks:
|
||||
break
|
||||
await asyncio.sleep(0.01)
|
||||
|
||||
assert len(client_chunks) == len(_STREAM_PREFIX) + len(_STREAM_TAIL)
|
||||
assert _output_tokens_from_logged_chunks(iterator.logged_chunks) == 64
|
||||
assert not any(c.startswith(b"event: error\n") for c in iterator.logged_chunks)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_sse_wrapper_dispatches_deferred_logging_when_client_disconnects_mid_tail():
|
||||
"""
|
||||
Regression: when the pump finishes draining while the client is still
|
||||
connected, ``_handle_streaming_logging`` defers billing for the proxy's
|
||||
post-response hook (``ProxyLogging._fire_deferred_stream_logging``), which
|
||||
only fires on a normally completed response. If the client then disconnects
|
||||
before consuming the queued tail, the response generator tears down via
|
||||
GeneratorExit and that hook never runs. The relay teardown must dispatch
|
||||
the stored deferred billing itself, or the request logs no spend at all.
|
||||
"""
|
||||
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()
|
||||
|
||||
logging_obj = _make_logging_obj("test_deferred_dispatch_on_disconnect_mid_tail")
|
||||
logging_obj._on_deferred_stream_complete = _deferred_stream_complete
|
||||
iterator = BaseAnthropicMessagesStreamingIterator(litellm_logging_obj=logging_obj, request_body={})
|
||||
|
||||
async def _full_stream():
|
||||
for event in (*_STREAM_PREFIX, *_STREAM_TAIL):
|
||||
yield event
|
||||
|
||||
gen = iterator.async_sse_wrapper(_full_stream())
|
||||
client_chunks = []
|
||||
async for chunk in gen:
|
||||
client_chunks.append(chunk)
|
||||
if len(client_chunks) == len(_STREAM_PREFIX):
|
||||
break
|
||||
|
||||
for _ in range(100):
|
||||
if getattr(logging_obj, "_deferred_stream_complete_args", None) is not None:
|
||||
break
|
||||
await asyncio.sleep(0.01)
|
||||
assert getattr(logging_obj, "_deferred_stream_complete_args", None) is not None, "pump never deferred billing"
|
||||
|
||||
await gen.aclose()
|
||||
|
||||
assert len(dispatched) == 1, "relay teardown did not dispatch the deferred billing"
|
||||
assert logging_obj._on_deferred_stream_complete is None
|
||||
assert logging_obj._deferred_stream_complete_args is None
|
||||
await asyncio.wait_for(deferred_fired.wait(), timeout=5)
|
||||
|
||||
|
||||
class _ProviderStreamError(Exception):
|
||||
"""Stand-in for a provider-specific streaming failure carrying a status code."""
|
||||
|
||||
def __init__(self, message: str, status_code: int):
|
||||
super().__init__(message)
|
||||
self.status_code = status_code
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_sse_wrapper_reraises_upstream_error_to_connected_client():
|
||||
"""
|
||||
Regression: an upstream failure (Bedrock read / decode / chunk-conversion)
|
||||
before message_stop must propagate the ORIGINAL provider exception to a
|
||||
still-connected client, so the proxy's failure handling keeps the
|
||||
provider-specific status. The pump must not swallow it into a generic
|
||||
api_error event + normal termination.
|
||||
"""
|
||||
|
||||
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"}}
|
||||
raise _ProviderStreamError("bedrock stream blew up", status_code=529)
|
||||
|
||||
iterator = _RecordingLoggingIterator(
|
||||
litellm_logging_obj=_make_logging_obj("test_reraises_upstream_error"),
|
||||
request_body={},
|
||||
)
|
||||
|
||||
received = []
|
||||
|
||||
async def _drain():
|
||||
async for chunk in iterator.async_sse_wrapper(_failing_stream()):
|
||||
received.append(chunk)
|
||||
|
||||
with pytest.raises(_ProviderStreamError) as excinfo:
|
||||
await _drain()
|
||||
|
||||
assert excinfo.value.status_code == 529
|
||||
assert received
|
||||
assert not any(c.startswith(b"event: error\n") for c in received)
|
||||
assert iterator.logged_chunks == []
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_sse_wrapper_salvages_partial_spend_on_upstream_error_after_disconnect():
|
||||
"""
|
||||
When the upstream errors AFTER the client has already disconnected there is
|
||||
no live client to re-raise to and no failure hook will run, so the pump
|
||||
salvages partial spend from what it collected instead of dropping the row.
|
||||
"""
|
||||
tail_gated = asyncio.Event()
|
||||
|
||||
async def _gated_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"}}
|
||||
await tail_gated.wait()
|
||||
raise _ProviderStreamError("late failure", status_code=500)
|
||||
|
||||
iterator = _RecordingLoggingIterator(
|
||||
litellm_logging_obj=_make_logging_obj("test_salvage_partial_on_late_error"),
|
||||
request_body={},
|
||||
)
|
||||
|
||||
gen = iterator.async_sse_wrapper(_gated_failing_stream())
|
||||
received = [await gen.__anext__(), await gen.__anext__()]
|
||||
await gen.aclose() # client disconnects before the upstream error
|
||||
|
||||
tail_gated.set() # let the upstream raise now, after disconnect
|
||||
for _ in range(100):
|
||||
if iterator.logged_chunks:
|
||||
break
|
||||
await asyncio.sleep(0.01)
|
||||
|
||||
assert len(received) == 2
|
||||
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):
|
||||
"""
|
||||
Regression: the relay queue is bounded, so a slow client throttles the
|
||||
upstream read instead of letting the pump buffer the whole response in
|
||||
memory. With a tiny queue and a client that reads a single chunk, the pump
|
||||
must stall after producing only a queue's worth of chunks ahead, not race
|
||||
to the end of a large stream.
|
||||
"""
|
||||
monkeypatch.setattr(streaming_iterator_module, "ANTHROPIC_MESSAGES_STREAM_RELAY_QUEUE_MAXSIZE", 2)
|
||||
|
||||
total = 200
|
||||
produced = 0
|
||||
|
||||
async def _fast_stream():
|
||||
nonlocal produced
|
||||
for i in range(total):
|
||||
produced += 1
|
||||
yield {"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": f"t{i}"}}
|
||||
|
||||
iterator = _make_iterator("test_backpressure_slow_client")
|
||||
gen = iterator.async_sse_wrapper(_fast_stream())
|
||||
try:
|
||||
await gen.__anext__()
|
||||
for _ in range(500):
|
||||
await asyncio.sleep(0)
|
||||
assert produced <= 2 + 3, f"pump ran ahead unthrottled: produced {produced} of {total}"
|
||||
assert produced < total
|
||||
finally:
|
||||
await gen.aclose()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_sse_wrapper_bills_partial_when_detached_drain_cap_reached(monkeypatch):
|
||||
"""
|
||||
Regression: when the concurrent detached-drain cap is already reached, a
|
||||
pump whose client has disconnected must bill what it collected instead of
|
||||
continuing to drain (and accumulating) the rest of a large upstream stream,
|
||||
so slow/abandoned clients can't pin unbounded worker state.
|
||||
|
||||
The cap slot set is pre-occupied so the single slot is unavailable when this
|
||||
pump reaches its first post-disconnect chunk; that isolates the cap decision
|
||||
from multi-pump scheduling races.
|
||||
"""
|
||||
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)
|
||||
tail_reached = False
|
||||
|
||||
async def _long_stream():
|
||||
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"}}
|
||||
for i in range(100):
|
||||
yield {"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": f"more{i}"}}
|
||||
tail_reached = True
|
||||
yield {"type": "message_delta", "delta": {"stop_reason": "end_turn"}, "usage": {"output_tokens": 42}}
|
||||
yield {"type": "message_stop"}
|
||||
|
||||
iterator = _RecordingLoggingIterator(litellm_logging_obj=_make_logging_obj("drain_cap_full"), request_body={})
|
||||
try:
|
||||
gen = iterator.async_sse_wrapper(_long_stream())
|
||||
await gen.__anext__() # message_start
|
||||
await gen.__anext__() # first delta
|
||||
await gen.aclose() # client disconnects; 100+ chunks remain upstream
|
||||
|
||||
for _ in range(200):
|
||||
if iterator.logged_chunks:
|
||||
break
|
||||
await asyncio.sleep(0)
|
||||
|
||||
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
|
||||
assert not any(c.startswith(b"event: message_stop\n") for c in iterator.logged_chunks)
|
||||
assert tail_reached is False, "pump kept draining past the cap instead of stopping"
|
||||
finally:
|
||||
holder.cancel()
|
||||
streaming_iterator_module._DETACHED_STREAM_DRAINS.discard(holder)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_async_sse_wrapper_bills_partial_when_detached_drains_disabled(monkeypatch):
|
||||
"""
|
||||
Regression: ANTHROPIC_MESSAGES_MAX_DETACHED_STREAM_DRAINS=0 must disable
|
||||
detached draining entirely, not just shrink the cap. With no slots ever
|
||||
available, the very first post-disconnect chunk must fall back to partial
|
||||
spend logging instead of hanging on a cap that's unreachable.
|
||||
"""
|
||||
monkeypatch.setattr(streaming_iterator_module, "ANTHROPIC_MESSAGES_MAX_DETACHED_STREAM_DRAINS", 0)
|
||||
monkeypatch.setattr(streaming_iterator_module, "ANTHROPIC_MESSAGES_STREAM_RELAY_QUEUE_MAXSIZE", 4)
|
||||
|
||||
tail_reached = False
|
||||
|
||||
async def _long_stream():
|
||||
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"}}
|
||||
for i in range(100):
|
||||
yield {"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": f"more{i}"}}
|
||||
tail_reached = True
|
||||
yield {"type": "message_delta", "delta": {"stop_reason": "end_turn"}, "usage": {"output_tokens": 42}}
|
||||
yield {"type": "message_stop"}
|
||||
|
||||
iterator = _RecordingLoggingIterator(litellm_logging_obj=_make_logging_obj("drains_disabled"), request_body={})
|
||||
gen = iterator.async_sse_wrapper(_long_stream())
|
||||
await gen.__anext__() # message_start
|
||||
await gen.__anext__() # first delta
|
||||
await gen.aclose() # client disconnects; 100+ chunks remain upstream
|
||||
|
||||
for _ in range(200):
|
||||
if iterator.logged_chunks:
|
||||
break
|
||||
await asyncio.sleep(0)
|
||||
|
||||
assert iterator.logged_chunks, "pump never billed with detached drains disabled"
|
||||
assert len(iterator.logged_chunks) <= 2 + streaming_iterator_module.ANTHROPIC_MESSAGES_STREAM_RELAY_QUEUE_MAXSIZE
|
||||
assert len(iterator.logged_chunks) < 100
|
||||
assert not any(c.startswith(b"event: message_stop\n") for c in iterator.logged_chunks)
|
||||
assert tail_reached is False, "pump kept draining despite detached drains being disabled"
|
||||
assert len(streaming_iterator_module._DETACHED_STREAM_DRAINS) == 0
|
||||
|
||||
|
||||
@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_abort_upstream_logs_warning_when_aclose_raises(caplog):
|
||||
"""_abort_upstream must swallow and log any exception from aclose()."""
|
||||
import logging
|
||||
|
||||
class _ExplodingStream:
|
||||
def __aiter__(self):
|
||||
return self
|
||||
|
||||
async def __anext__(self):
|
||||
raise StopAsyncIteration
|
||||
|
||||
async def aclose(self):
|
||||
raise RuntimeError("aclose exploded")
|
||||
|
||||
with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"):
|
||||
await BaseAnthropicMessagesStreamingIterator._abort_upstream(_ExplodingStream())
|
||||
|
||||
assert any("abort" in r.message and "RuntimeError" in r.message for r in caplog.records)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_enqueue_for_client_returns_false_when_already_detached():
|
||||
"""_enqueue_for_client must return False immediately (without touching the queue)
|
||||
when client_detached is already set before the call."""
|
||||
from litellm.llms.anthropic.experimental_pass_through.messages.streaming_iterator import (
|
||||
BaseAnthropicMessagesStreamingIterator,
|
||||
)
|
||||
|
||||
queue: asyncio.Queue[bytes | None | BaseException] = asyncio.Queue(maxsize=1)
|
||||
client_detached = asyncio.Event()
|
||||
client_detached.set()
|
||||
|
||||
result = await BaseAnthropicMessagesStreamingIterator._enqueue_for_client(queue, client_detached, b"chunk")
|
||||
assert result is False
|
||||
assert queue.empty()
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_enqueue_for_client_returns_false_when_client_detaches_while_queue_full():
|
||||
"""_enqueue_for_client must return False (and cancel the put) when the queue
|
||||
is full and client_detached fires before space becomes available."""
|
||||
from litellm.llms.anthropic.experimental_pass_through.messages.streaming_iterator import (
|
||||
BaseAnthropicMessagesStreamingIterator,
|
||||
)
|
||||
|
||||
queue: asyncio.Queue[bytes | None | BaseException] = asyncio.Queue(maxsize=1)
|
||||
queue.put_nowait(b"already-full")
|
||||
|
||||
client_detached = asyncio.Event()
|
||||
|
||||
async def _set_detached_soon():
|
||||
await asyncio.sleep(0.01)
|
||||
client_detached.set()
|
||||
|
||||
asyncio.create_task(_set_detached_soon())
|
||||
result = await BaseAnthropicMessagesStreamingIterator._enqueue_for_client(queue, client_detached, b"new-chunk")
|
||||
assert result is False
|
||||
assert queue.qsize() == 1
|
||||
assert queue.get_nowait() == b"already-full"
|
||||
|
||||
|
||||
@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
|
||||
the full upstream and bills the terminal usage, and releases its slot after."""
|
||||
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 _stream():
|
||||
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"}}
|
||||
for i in range(20):
|
||||
yield {"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": f"m{i}"}}
|
||||
yield {"type": "message_delta", "delta": {"stop_reason": "end_turn"}, "usage": {"output_tokens": 42}}
|
||||
yield {"type": "message_stop"}
|
||||
|
||||
iterator = _RecordingLoggingIterator(litellm_logging_obj=_make_logging_obj("drain_cap_free"), request_body={})
|
||||
gen = iterator.async_sse_wrapper(_stream())
|
||||
await gen.__anext__()
|
||||
await gen.__anext__()
|
||||
await gen.aclose()
|
||||
|
||||
for _ in range(300):
|
||||
if iterator.logged_chunks:
|
||||
break
|
||||
await asyncio.sleep(0.01)
|
||||
|
||||
assert any(c.startswith(b"event: message_stop\n") for c in iterator.logged_chunks)
|
||||
assert len(streaming_iterator_module._DETACHED_STREAM_DRAINS) == 0
|
||||
|
||||
|
||||
def _decode_sse_events(events: tuple[bytes, ...]) -> list[tuple[str, dict]]:
|
||||
decoded = []
|
||||
for event in events:
|
||||
|
|
@ -599,20 +1111,35 @@ async def test_normal_end_with_deferred_dispatch_armed_parks_logging_coroutine(m
|
|||
@pytest.mark.asyncio
|
||||
async def test_client_disconnect_enqueues_immediately_even_when_deferred_dispatch_armed(monkeypatch):
|
||||
"""
|
||||
On client disconnect the guardrail end-of-stream scan never runs, so
|
||||
deferral would strand the spend log; the teardown path must keep
|
||||
enqueueing immediately (LIT-5839) even when the deferred callback is armed.
|
||||
Regression: on client disconnect the guardrail end-of-stream scan never
|
||||
runs, so deferral would strand the spend log. The detached pump's
|
||||
post-disconnect bill must bypass the deferred-dispatch park and enqueue
|
||||
immediately (LIT-5839) even when the deferred callback is armed (LIT-6409).
|
||||
"""
|
||||
worker = _RecordingLoggingWorker()
|
||||
monkeypatch.setattr(streaming_iterator_module, "GLOBAL_LOGGING_WORKER", worker)
|
||||
iterator = _make_iterator("test_disconnect_enqueues_when_armed")
|
||||
iterator.litellm_logging_obj._on_deferred_stream_complete = _noop_deferred_dispatch
|
||||
|
||||
wrapped = iterator.async_sse_wrapper(_events_then_hang(TRUNCATED_TOOL_USE_EVENTS))
|
||||
tail_gated = asyncio.Event()
|
||||
|
||||
async def _gated_stream():
|
||||
for event in TRUNCATED_TOOL_USE_EVENTS:
|
||||
yield event
|
||||
await tail_gated.wait()
|
||||
yield {"type": "message_stop"}
|
||||
|
||||
wrapped = iterator.async_sse_wrapper(_gated_stream())
|
||||
for _ in range(len(TRUNCATED_TOOL_USE_EVENTS)):
|
||||
await wrapped.__anext__()
|
||||
await wrapped.aclose()
|
||||
|
||||
tail_gated.set()
|
||||
for _ in range(100):
|
||||
if worker.enqueued:
|
||||
break
|
||||
await asyncio.sleep(0.01)
|
||||
|
||||
assert len(worker.enqueued) == 1
|
||||
assert getattr(iterator.litellm_logging_obj, "_deferred_stream_complete_args", None) is None
|
||||
worker.close_enqueued()
|
||||
|
|
@ -629,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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -1,9 +1,9 @@
|
|||
{
|
||||
"LIT001": {
|
||||
"limit": 22521
|
||||
"limit": 22519
|
||||
},
|
||||
"LIT002": {
|
||||
"limit": 26820
|
||||
"limit": 26818
|
||||
},
|
||||
"LIT003": {
|
||||
"limit": 269
|
||||
|
|
@ -27,7 +27,7 @@
|
|||
"limit": 0
|
||||
},
|
||||
"LIT010": {
|
||||
"limit": 16546
|
||||
"limit": 16544
|
||||
},
|
||||
"LIT011": {
|
||||
"limit": 5575
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue