mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
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:
parent
79f62db620
commit
87e961fad0
4 changed files with 649 additions and 1 deletions
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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)))
|
||||
|
|
@ -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
|
||||
|
|
@ -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"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue