fix(anthropic): let /v1/messages mid-stream failures reach the proxy failure boundary (#44800)

* fix(anthropic): let /v1/messages mid-stream failures reach the proxy failure boundary

AnthropicStreamWrapper.async_anthropic_sse_wrapper swallowed every upstream
exception and emitted its own Anthropic error frame, so a provider drop on a
streamed /v1/messages call never reached the proxy's streaming boundary — no
failure spend row, no failure callbacks.

Re-raise when the stream is proxy-managed (detached failure hook armed) so
async_streaming_data_generator runs post_call_failure_hook once and serializes
the error frame; when consumed standalone (SDK litellm.messages path), run the
logging object's async failure handler and keep the client-facing error frame.

Fixes #44742

* test(anthropic): satisfy PT012 single-statement rule in mid-stream error test

* fix(anthropic): dispatch both sync and async failure callbacks

* fix(anthropic): dispatch both sync and async failure callbacks

* test(anthropic): move the regression cases into the mapped test file

* refactor(anthropic): read the public detached-failure hook and drop the explanatory comments

* test(anthropic): drive the mid-stream failure tests through a real logging object

* fix(anthropic): re-raise the provider error from the chat wrapper envelope so the router keeps its pre-content retry

* test(integration): cover mid-stream failure bookkeeping of bridged /v1/messages streams

Nine integration cells for the chat bridge: raw httpx and the Anthropic SDK (sync and async) streams that fail after the first content delta now land exactly one failure SpendLogs row carrying the provider's error class, the non-streaming 500, the response-cache twin and a client that leaves mid-stream keep their bookkeeping, a 24-stream burst lands every call id once, and the standalone SDK stream reports the provider error to failure callbacks once.

---------

Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com>
This commit is contained in:
D41910 2026-10-08 08:42:35 +08:00 • committed by GitHub
parent 79f62db620
commit 87e961fad0
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 649 additions and 1 deletions

View file

@ -69,6 +69,12 @@ def _error_status_and_message(exc: Exception) -> tuple[int, str]:
return 500, str(exc) or "Upstream stream ended before completion"
def _provider_error(exc: Exception) -> Exception:
if isinstance(exc, MidStreamFallbackError) and exc.original_exception is not None:
return exc.original_exception
return exc
def _mid_stream_error_sse_event(exc: Exception) -> bytes:
from litellm.anthropic_interface.exceptions.exception_mapping_utils import (
anthropic_error_sse_frame,
@ -336,6 +342,7 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper):
self._message_id: str = f"msg_{uuid.uuid4()}"
if litellm_logging_obj is not None:
litellm_logging_obj.record_streamed_anthropic_message_id(self._message_id)
self.litellm_logging_obj = litellm_logging_obj
# Mapping of truncated tool names to original names (for OpenAI's 64-char limit)
self.tool_name_mapping = tool_name_mapping or {}
# Polyfill applied_edits on final message_delta.
@ -1041,7 +1048,25 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper):
else:
yield chunk
except Exception as e: # noqa: BLE001 # boundary before the socket: any upstream failure becomes an Anthropic error event
verbose_logger.exception("Anthropic Adapter - mid-stream error, emitting Anthropic error event: %s", e)
verbose_logger.exception("Anthropic Adapter - mid-stream error: %s", e)
logging_obj: Final = self.litellm_logging_obj
provider_error: Final = _provider_error(e)
if logging_obj is not None and logging_obj.on_detached_stream_failure is not None:
if provider_error is e:
raise
raise provider_error from e
if logging_obj is not None:
try:
await logging_obj.dispatch_failure_handlers(
exception=provider_error,
traceback_exception=traceback.format_exc(),
prefer_async_handlers=True,
)
except Exception as failure_handler_error: # noqa: BLE001 # a failing failure handler must not also drop the error frame
verbose_logger.exception(
"Anthropic Adapter - failure handler raised while reporting a mid-stream error: %s",
failure_handler_error,
)
yield _mid_stream_error_sse_event(e)
def _increment_content_block_index(self):

View file

@ -0,0 +1,362 @@
import asyncio
import json
import uuid
from collections import deque
from collections.abc import Callable, Mapping
from dataclasses import dataclass, replace
from typing import Final, Literal
import anthropic
import httpx
import pytest
from anthropic.types import Message, MessageParam
from integration._support.anthropic_sse import (
ANTHROPIC_ERROR_TYPES,
SseEvent,
delta_text,
dropping_reply,
error_type,
event_types,
parse_sse,
stream_reply,
user_prompt,
)
from integration._support.client import Gateway, eventually, object_value, string_value
from integration._support.database import read_rows
from integration._support.openai_wire import answering_model_discovery, chat_stream, openai_error, posted_targets
from integration._support.wire import Reply, Request, wire_server
from pydantic import JsonValue, TypeAdapter
_BACKEND: Final = "gpt-4o-mini"
_PROVIDER_MODEL: Final = f"hosted_vllm/{_BACKEND}"
_PROVIDER_KEY: Final = "integration-provider-key"
_UPSTREAM_TARGET: Final = "/v1/chat/completions"
_TEXT: Final = "Hello"
_CALL_ID_HEADER: Final = "x-litellm-call-id"
_ROWS: Final = (
"SELECT metadata->>'litellm_call_id' AS call_id, request_id, status, cache_hit, metadata "
'FROM "LiteLLM_SpendLogs" WHERE model_group=%s'
)
_ROW_SECONDS: Final = 60
_SLOW_PAUSE: Final = 1.0
_BURST_PER_OUTCOME: Final = 8
Outcome = Literal[
"drop_after_content", "error_frame_after_content", "slow_drop_after_content", "succeeds", "rejected_before_any_body"
]
_OUTCOME: Final = TypeAdapter(Outcome)
_FAILING_STREAMS: Final[tuple[Outcome, ...]] = ("drop_after_content", "error_frame_after_content")
_BURST: Final[tuple[Outcome, ...]] = (
"slow_drop_after_content",
"error_frame_after_content",
"succeeds",
) * _BURST_PER_OUTCOME
@dataclass(frozen=True, slots=True)
class _Call:
outcome: Outcome
marker: str
prompt_tag: str
call_id: str
@property
def prompt(self) -> str:
return f"{self.outcome}:{self.marker}:{self.prompt_tag}"
@property
def streams(self) -> bool:
return self.outcome != "rejected_before_any_body"
@dataclass(frozen=True, slots=True)
class _SdkFailure:
text: str
error: anthropic.APIStatusError
def _call(outcome: Outcome, marker: str) -> _Call:
tag: Final = uuid.uuid4().hex
return _Call(outcome, marker, tag, tag)
def _marker() -> str:
return "bridge-mid-stream-" + uuid.uuid4().hex
def _error_frame(status: int) -> bytes:
error: Final = {"message": f"scripted mid-stream {status}", "type": "server_error", "code": status}
return b"data: " + json.dumps({"error": error}).encode() + b"\n\n"
def _reply(outcome: Outcome, chunks: tuple[bytes, bytes, bytes]) -> Reply:
match outcome:
case "drop_after_content":
return dropping_reply(chunks, abort_after=2)
case "error_frame_after_content":
return stream_reply((chunks[0], chunks[1], _error_frame(500) + b"data: [DONE]\n\n"))
case "slow_drop_after_content":
return stream_reply(chunks, abort_after=2, pause=_SLOW_PAUSE)
case "succeeds":
return stream_reply(chunks)
case "rejected_before_any_body":
return openai_error(500)
def _upstream(marker: str) -> Callable[[Request], Reply]:
def respond(request: Request) -> Reply:
assert (request.method, request.target) == ("POST", _UPSTREAM_TARGET), request
assert request.headers["authorization"] == f"Bearer {_PROVIDER_KEY}", request.headers
body: Final = object_value(json.loads(request.body))
assert body["model"] == _BACKEND, body
outcome, scripted_marker, prompt_tag = user_prompt(body).split(":")
assert scripted_marker == marker, body
scripted: Final = _OUTCOME.validate_python(outcome)
assert body.get("stream", False) is (scripted != "rejected_before_any_body"), body
return _reply(scripted, chat_stream(f"chunk-{prompt_tag}", _BACKEND, _TEXT))
return answering_model_discovery(respond)
def _body(model: str, call: _Call) -> dict[str, JsonValue]:
return {
"model": model,
"max_tokens": 16,
"stream": call.streams,
"messages": [{"role": "user", "content": call.prompt}],
}
def _headers(call: _Call) -> dict[str, str]:
return {_CALL_ID_HEADER: call.call_id}
def _stream(gateway: Gateway, model: str, call: _Call) -> tuple[SseEvent, ...]:
response: Final = gateway.request("POST", "/v1/messages", _body(model, call), headers=_headers(call))
assert response.status_code == 200, response.text
assert response.headers[_CALL_ID_HEADER] == call.call_id, response.headers
return parse_sse(response.text)
async def _stream_concurrently(client: httpx.AsyncClient, key: str, model: str, call: _Call) -> tuple[SseEvent, ...]:
response: Final = await client.post(
"/v1/messages", json=_body(model, call), headers={"Authorization": f"Bearer {key}", **_headers(call)}
)
assert response.status_code == 200, response.text
assert response.headers[_CALL_ID_HEADER] == call.call_id, response.headers
return parse_sse(response.text)
def _leave_after_the_first_content_delta(gateway: Gateway, model: str, call: _Call) -> str:
headers: Final = {"Authorization": f"Bearer {gateway.key}", **_headers(call)}
with gateway.client.stream("POST", "/v1/messages", json=_body(model, call), headers=headers) as response:
assert response.status_code == 200, response
assert response.headers[_CALL_ID_HEADER] == call.call_id, response.headers
return next(
line for line in response.iter_lines() if line.startswith("data:") and '"content_block_delta"' in line
)
def _call_id(row: dict[str, JsonValue]) -> str:
return string_value(row["call_id"])
def _rows(model: str, call_ids: frozenset[str]) -> dict[str, dict[str, JsonValue]]:
rows: Final = eventually(
lambda: read_rows(_ROWS, (model,)),
lambda rows: call_ids.issubset(_call_id(row) for row in rows),
seconds=_ROW_SECONDS,
)
by_call_id: Final = {_call_id(row): row for row in rows}
assert len(by_call_id) == len(rows), rows
return by_call_id
def _assert_failed_after_content(events: tuple[SseEvent, ...]) -> None:
types: Final = event_types(events)
assert types[0] == "message_start", events
assert delta_text(events) == _TEXT, events
assert types[-1] == "error" and "message_stop" not in types, events
def _assert_completed(events: tuple[SseEvent, ...]) -> None:
types: Final = event_types(events)
assert types[0] == "message_start" and types[-1] == "message_stop", events
assert "error" not in types, events
assert delta_text(events) == _TEXT, events
def _assert_outcome(outcome: Outcome, events: tuple[SseEvent, ...]) -> None:
if outcome == "succeeds":
_assert_completed(events)
return
_assert_failed_after_content(events)
def _expected_status(outcome: Outcome) -> str:
return "success" if outcome == "succeeds" else "failure"
def _assert_failure_row(row: Mapping[str, JsonValue], anthropic_error_type: str | None) -> None:
assert row["status"] == "failure", row
error: Final = object_value(object_value(row["metadata"])["error_information"])
assert error["error_class"] != "MidStreamFallbackError", error
assert ANTHROPIC_ERROR_TYPES[int(str(error["error_code"]))] == anthropic_error_type, (error, anthropic_error_type)
def _snapshot_text(snapshot: Message) -> str:
return "".join(block.text for block in snapshot.content if block.type == "text")
def _sdk_messages(call: _Call) -> list[MessageParam]:
return [{"role": "user", "content": call.prompt}]
def _sdk_error_type(error: anthropic.APIStatusError) -> str:
return string_value(object_value(object_value(error.body)["error"])["type"])
def _sdk_sync_failure(client: anthropic.Anthropic, model: str, call: _Call) -> _SdkFailure:
with client.messages.stream(
model=model, max_tokens=16, messages=_sdk_messages(call), extra_headers=_headers(call)
) as stream:
assert stream.response.headers[_CALL_ID_HEADER] == call.call_id, stream.response.headers
try:
deque(stream, maxlen=0)
except anthropic.APIStatusError as error:
return _SdkFailure(_snapshot_text(stream.current_message_snapshot), error)
pytest.fail(f"{call.call_id}: the bridged stream ended without the scripted mid-stream failure")
async def _sdk_async_failure(client: anthropic.AsyncAnthropic, model: str, call: _Call) -> _SdkFailure:
async with client.messages.stream(
model=model, max_tokens=16, messages=_sdk_messages(call), extra_headers=_headers(call)
) as stream:
assert stream.response.headers[_CALL_ID_HEADER] == call.call_id, stream.response.headers
try:
async for _ in stream:
pass
except anthropic.APIStatusError as error:
return _SdkFailure(_snapshot_text(stream.current_message_snapshot), error)
pytest.fail(f"{call.call_id}: the bridged stream ended without the scripted mid-stream failure")
def _assert_sdk_failure_bookkept(failure: _SdkFailure, model: str, call: _Call) -> None:
assert failure.text == _TEXT, failure
rows: Final = _rows(model, frozenset({call.call_id}))
assert rows.keys() == {call.call_id}, rows
_assert_failure_row(rows[call.call_id], _sdk_error_type(failure.error))
@pytest.mark.parametrize("outcome", _FAILING_STREAMS)
def test_bridged_stream_failing_after_content_writes_one_failure_row(gateway: Gateway, outcome: Outcome) -> None:
marker: Final = _marker()
with wire_server(_upstream(marker)) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(model=_PROVIDER_MODEL, api_base=wire.url + "/v1")
call: Final = _call(outcome, marker)
events: Final = _stream(gateway, model, call)
_assert_failed_after_content(events)
assert posted_targets(wire) == (_UPSTREAM_TARGET,)
rows: Final = _rows(model, frozenset({call.call_id}))
assert rows.keys() == {call.call_id}, rows
_assert_failure_row(rows[call.call_id], error_type(events))
def test_anthropic_sdk_sync_stream_failing_after_content_raises_and_writes_one_failure_row(gateway: Gateway) -> None:
marker: Final = _marker()
with wire_server(_upstream(marker)) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(model=_PROVIDER_MODEL, api_base=wire.url + "/v1")
call: Final = _call("drop_after_content", marker)
client: Final = anthropic.Anthropic(
base_url=str(gateway.client.base_url), api_key=gateway.key, max_retries=0, timeout=60
)
failure: Final = _sdk_sync_failure(client, model, call)
assert posted_targets(wire) == (_UPSTREAM_TARGET,)
_assert_sdk_failure_bookkept(failure, model, call)
async def test_anthropic_sdk_async_stream_failing_after_content_raises_and_writes_one_failure_row(
gateway: Gateway,
) -> None:
marker: Final = _marker()
with wire_server(_upstream(marker)) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(model=_PROVIDER_MODEL, api_base=wire.url + "/v1")
call: Final = _call("error_frame_after_content", marker)
client: Final = anthropic.AsyncAnthropic(
base_url=str(gateway.client.base_url), api_key=gateway.key, max_retries=0, timeout=60
)
failure: Final = await _sdk_async_failure(client, model, call)
assert posted_targets(wire) == (_UPSTREAM_TARGET,)
_assert_sdk_failure_bookkept(failure, model, call)
def test_bridged_request_rejected_before_any_body_keeps_its_error_and_failure_row(gateway: Gateway) -> None:
marker: Final = _marker()
with wire_server(_upstream(marker)) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(model=_PROVIDER_MODEL, api_base=wire.url + "/v1")
call: Final = _call("rejected_before_any_body", marker)
response: Final = gateway.request("POST", "/v1/messages", _body(model, call), headers=_headers(call))
assert response.status_code == 500, response.text
error: Final = object_value(object_value(json.loads(response.text))["error"])
assert error["type"] == ANTHROPIC_ERROR_TYPES[500], response.text
assert posted_targets(wire) == (_UPSTREAM_TARGET,)
rows: Final = _rows(model, frozenset({call.call_id}))
assert rows.keys() == {call.call_id}, rows
_assert_failure_row(rows[call.call_id], string_value(error["type"]))
def test_bridged_stream_served_from_the_response_cache_keeps_its_success_bookkeeping(gateway: Gateway) -> None:
marker: Final = _marker()
with wire_server(_upstream(marker)) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(model=_PROVIDER_MODEL, api_base=wire.url + "/v1")
first: Final = _call("succeeds", marker)
_assert_completed(_stream(gateway, model, first))
miss: Final = _rows(model, frozenset({first.call_id}))[first.call_id]
assert (miss["status"], miss["cache_hit"] != "True") == ("success", True), miss
twin: Final = replace(first, call_id=uuid.uuid4().hex)
_assert_completed(_stream(gateway, model, twin))
assert posted_targets(wire) == (_UPSTREAM_TARGET,)
rows: Final = _rows(model, frozenset({twin.call_id}))
hit: Final = rows[twin.call_id]
assert (hit["status"], hit["cache_hit"]) == ("success", "True"), hit
assert string_value(hit["request_id"]) != string_value(miss["request_id"]), rows
assert set(rows) == {first.call_id, twin.call_id}, rows
def test_client_leaving_a_bridged_stream_before_its_failure_leaves_the_proxy_serving(gateway: Gateway) -> None:
marker: Final = _marker()
with wire_server(_upstream(marker)) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(model=_PROVIDER_MODEL, api_base=wire.url + "/v1")
abandoned: Final = _call("slow_drop_after_content", marker)
first_delta: Final = _leave_after_the_first_content_delta(gateway, model, abandoned)
assert _TEXT in first_delta, first_delta
follow_up: Final = _call("succeeds", marker)
_assert_completed(_stream(gateway, model, follow_up))
assert posted_targets(wire) == (_UPSTREAM_TARGET,) * 2
rows: Final = _rows(model, frozenset({follow_up.call_id}))
assert rows[follow_up.call_id]["status"] == "success", rows
assert set(rows) <= {abandoned.call_id, follow_up.call_id}, rows
assert gateway.request("GET", "/health/liveliness").status_code == 200
async def test_concurrent_bridged_streams_failing_after_content_each_land_one_row(gateway: Gateway) -> None:
marker: Final = _marker()
with wire_server(_upstream(marker)) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(model=_PROVIDER_MODEL, api_base=wire.url + "/v1")
calls: Final = tuple(_call(outcome, marker) for outcome in _BURST)
async with httpx.AsyncClient(base_url=str(gateway.client.base_url), timeout=60, trust_env=False) as client:
streams: Final = await asyncio.gather(
*(_stream_concurrently(client, gateway.key, model, call) for call in calls)
)
for call, events in zip(calls, streams, strict=True):
_assert_outcome(call.outcome, events)
assert posted_targets(wire) == (_UPSTREAM_TARGET,) * len(calls)
rows: Final = _rows(model, frozenset(call.call_id for call in calls))
assert {call_id: row["status"] for call_id, row in rows.items()} == {
call.call_id: _expected_status(call.outcome) for call in calls
}, rows
for call, events in zip(calls, streams, strict=True):
if call.outcome != "succeeds":
_assert_failure_row(rows[call.call_id], error_type(events))
assert gateway.request("GET", "/health/liveliness").status_code == 200
_assert_completed(_stream(gateway, model, _call("succeeds", marker)))

View file

@ -0,0 +1,113 @@
import asyncio
import json
import threading
import uuid
from collections.abc import AsyncIterable, Callable
from typing import Final
from integration._support.anthropic_sse import (
SseEvent,
delta_text,
dropping_reply,
event_types,
parse_sse,
stream_reply,
user_prompt,
)
from integration._support.client import eventually, object_value
from integration._support.openai_wire import answering_model_discovery, chat_stream, posted_targets
from integration._support.wire import Reply, Request, Wire, wire_server
import litellm
from litellm.exceptions import MidStreamFallbackError
from litellm.integrations.custom_logger import CustomLogger
_BACKEND: Final = "gpt-4o-mini"
_PROVIDER_KEY: Final = "integration-provider-key"
_UPSTREAM_TARGET: Final = "/v1/chat/completions"
_TEXT: Final = "Hello"
_DROPS_AFTER_CONTENT: Final = "drops-after-content"
_SUCCEEDS: Final = "succeeds"
class _CallbackProbe(CustomLogger):
def __init__(self) -> None:
super().__init__()
self._lock: Final = threading.Lock()
self._failures: Final[list[object]] = []
self._successes: Final[list[object]] = []
async def async_log_failure_event(
self, kwargs: dict[str, object], response_obj: object, start_time: object, end_time: object
) -> None:
with self._lock:
self._failures.append(kwargs.get("exception"))
async def async_log_success_event(
self, kwargs: dict[str, object], response_obj: object, start_time: object, end_time: object
) -> None:
with self._lock:
self._successes.append(kwargs.get("litellm_call_id"))
def failures(self) -> tuple[object, ...]:
with self._lock:
return tuple(self._failures)
def successes(self) -> tuple[object, ...]:
with self._lock:
return tuple(self._successes)
def _upstream(marker: str) -> Callable[[Request], Reply]:
def respond(request: Request) -> Reply:
assert (request.method, request.target) == ("POST", _UPSTREAM_TARGET), request
assert request.headers["authorization"] == f"Bearer {_PROVIDER_KEY}", request.headers
body: Final = object_value(json.loads(request.body))
assert body["model"] == _BACKEND and body["stream"] is True, body
outcome, _, scripted_marker = user_prompt(body).partition(":")
assert scripted_marker == marker, body
chunks: Final = chat_stream(f"chunk-{outcome}-{marker}", _BACKEND, _TEXT)
if outcome == _DROPS_AFTER_CONTENT:
return dropping_reply(chunks, abort_after=2)
return stream_reply(chunks)
return answering_model_discovery(respond)
async def _stream(wire: Wire, prompt: str) -> tuple[SseEvent, ...]:
response: Final = await litellm.anthropic.messages.acreate(
model=f"hosted_vllm/{_BACKEND}",
api_base=wire.url + "/v1",
api_key=_PROVIDER_KEY,
max_tokens=16,
stream=True,
messages=[{"role": "user", "content": prompt}],
)
assert isinstance(response, AsyncIterable), response
frames: Final = [frame async for frame in response]
assert all(isinstance(frame, bytes) for frame in frames), frames
return parse_sse(b"".join(frame for frame in frames if isinstance(frame, bytes)).decode())
async def _settled(read: Callable[[], tuple[object, ...]]) -> tuple[object, ...]:
return await asyncio.to_thread(eventually, read, lambda seen: len(seen) >= 1)
async def test_sdk_bridged_stream_failing_after_content_reports_the_provider_error_once() -> None:
probe: Final = _CallbackProbe()
litellm.callbacks.append(probe)
marker: Final = uuid.uuid4().hex
with wire_server(_upstream(marker)) as wire:
failing: Final = await _stream(wire, f"{_DROPS_AFTER_CONTENT}:{marker}")
assert delta_text(failing) == _TEXT, failing
assert event_types(failing)[-1] == "error" and "message_stop" not in event_types(failing), failing
await _settled(probe.failures)
succeeding: Final = await _stream(wire, f"{_SUCCEEDS}:{marker}")
assert event_types(succeeding)[-1] == "message_stop", succeeding
assert delta_text(succeeding) == _TEXT, succeeding
await _settled(probe.successes)
assert posted_targets(wire) == (_UPSTREAM_TARGET,) * 2
failures: Final = probe.failures()
assert len(failures) == 1, failures
assert isinstance(failures[0], Exception), failures
assert not isinstance(failures[0], MidStreamFallbackError), failures

View file

@ -14,9 +14,13 @@ The async SSE wrapper must instead surface the failure as a well-formed
Anthropic ``error`` event so the stream stays valid and the client can retry.
"""
import asyncio
import json
import os
import sys
import threading
from datetime import datetime
from collections.abc import Callable
from typing import List, Optional
from unittest.mock import MagicMock
@ -25,6 +29,8 @@ import pytest
sys.path.insert(0, os.path.abspath("../../../../.."))
from litellm.exceptions import MidStreamFallbackError
from litellm.integrations.custom_logger import CustomLogger
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.llms.anthropic.pass_through.adapters.streaming_iterator import (
AnthropicStreamWrapper,
_mid_stream_error_sse_event,
@ -140,3 +146,145 @@ def test_error_event_preserves_midstream_fallback_error():
assert name == "error"
assert payload["error"]["type"] == "api_error"
assert "internalServerException" in payload["error"]["message"]
class _AsyncFailureRecorder(CustomLogger):
def __init__(self):
super().__init__()
self.exceptions: list[BaseException] = []
async def async_log_failure_event(self, kwargs, response_obj, start_time, end_time):
self.exceptions.append(kwargs["exception"])
class _SyncFailureRecorder:
def __init__(self):
self.exceptions: list[BaseException] = []
self.called = threading.Event()
def __call__(self, kwargs, completion_response, start_time, end_time):
self.exceptions.append(kwargs["exception"])
self.called.set()
def _make_logging_obj(
test_name: str,
async_recorder: _AsyncFailureRecorder,
sync_recorder: _SyncFailureRecorder,
) -> LiteLLMLoggingObj:
return LiteLLMLoggingObj(
model="bedrock-converse-sonnet-4-6",
messages=[{"role": "user", "content": "hi"}],
stream=True,
call_type="anthropic_messages",
start_time=datetime.now(),
litellm_call_id=test_name,
function_id=test_name,
dynamic_failure_callbacks=[sync_recorder],
dynamic_async_failure_callbacks=[async_recorder],
)
async def _proxy_boundary_hook(exc: Exception) -> None:
return None
async def _wait_for_sync_failure(sync_recorder: _SyncFailureRecorder) -> None:
for _ in range(500):
if sync_recorder.called.is_set():
return
await asyncio.sleep(0.01)
def _bedrock_drop() -> BedrockError:
return BedrockError(status_code=500, message="ConverseStream ended without messageStop")
def _chat_wrapper_envelope() -> MidStreamFallbackError:
provider_error = _bedrock_drop()
return MidStreamFallbackError(
message=str(provider_error),
model="bedrock-converse-sonnet-4-6",
llm_provider="bedrock",
original_exception=provider_error,
is_pre_first_chunk=False,
)
_RAISED_ERRORS = pytest.mark.parametrize(
"raised",
[_bedrock_drop, _chat_wrapper_envelope],
ids=["provider_error", "chat_wrapper_envelope"],
)
def _failing_wrapper(
logging_obj: LiteLLMLoggingObj | None,
raised: Callable[[], Exception] = _bedrock_drop,
) -> AnthropicStreamWrapper:
return AnthropicStreamWrapper(
completion_stream=_AsyncStreamThenRaise([_make_chunk(Delta(content="partial"))], raised()),
model="bedrock-converse-sonnet-4-6",
litellm_logging_obj=logging_obj,
)
@_RAISED_ERRORS
@pytest.mark.asyncio
async def test_mid_stream_error_reraises_the_provider_error_for_proxy_managed_stream(raised):
async_recorder = _AsyncFailureRecorder()
sync_recorder = _SyncFailureRecorder()
logging_obj = _make_logging_obj("proxy-managed", async_recorder, sync_recorder)
logging_obj.on_detached_stream_failure = _proxy_boundary_hook
with pytest.raises(BedrockError) as raised_info:
await _drain_sse(_failing_wrapper(logging_obj, raised))
assert str(raised_info.value) == "ConverseStream ended without messageStop"
await asyncio.sleep(0.1)
assert async_recorder.exceptions == []
assert sync_recorder.exceptions == []
@_RAISED_ERRORS
@pytest.mark.asyncio
async def test_mid_stream_error_dispatches_the_provider_error_to_failure_handlers_for_standalone_stream(raised):
async_recorder = _AsyncFailureRecorder()
sync_recorder = _SyncFailureRecorder()
wrapper = _failing_wrapper(_make_logging_obj("standalone", async_recorder, sync_recorder), raised)
events = await _drain_sse(wrapper)
await _wait_for_sync_failure(sync_recorder)
assert [str(exc) for exc in async_recorder.exceptions] == ["ConverseStream ended without messageStop"]
assert [str(exc) for exc in sync_recorder.exceptions] == ["ConverseStream ended without messageStop"]
assert _parse_sse(events[-1])[0] == "error"
class _BrokenFailureDispatchLogging(LiteLLMLoggingObj):
async def dispatch_failure_handlers(self, *args, **kwargs):
raise RuntimeError("failure sink is down")
@pytest.mark.asyncio
async def test_mid_stream_error_frame_survives_a_raising_failure_dispatch():
logging_obj = _BrokenFailureDispatchLogging(
model="bedrock-converse-sonnet-4-6",
messages=[{"role": "user", "content": "hi"}],
stream=True,
call_type="anthropic_messages",
start_time=datetime.now(),
litellm_call_id="broken-dispatch",
function_id="broken-dispatch",
)
events = await _drain_sse(_failing_wrapper(logging_obj))
assert _parse_sse(events[-1])[0] == "error"
@pytest.mark.asyncio
async def test_mid_stream_error_emits_error_event_without_logging_obj():
events = await _drain_sse(_failing_wrapper(None))
assert _parse_sse(events[-1])[0] == "error"