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:
Mateo Wang 2026-08-31 16:54:41 -07:00 committed by GitHub
commit 99703a30f0
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
7 changed files with 990 additions and 79 deletions

View file

@ -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

View file

@ -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))

View file

@ -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)

View file

@ -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 = {

View file

@ -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()

View file

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

View file

@ -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