fix(bedrock): surface unrecognized converse-stream event frames instead of an empty turn (#44112)

* fix(bedrock): surface unrecognized converse-stream event frames instead of an empty turn

* fix(bedrock): track distinct unknown stream event types and cover the stream error builder

* test(bedrock): audit converse-stream event frames on the proxy and the sync decoder

---------

Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-10-02 16:06:26 -07:00 • committed by GitHub
parent aef0a53837
commit 8b4de39ad7
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
7 changed files with 1494 additions and 57 deletions

View file

@ -42,9 +42,15 @@ from litellm.types.utils import GenericStreamingChunk as GChunk
from ..common_utils import (
BedrockError,
BedrockEventStreamResponseDict,
bedrock_event_stream_header,
bedrock_event_stream_response,
bedrock_stream_event_error_status,
build_bedrock_stream_error,
build_bedrock_stream_event_error,
error_response_text,
get_bedrock_response_stream_shape,
get_bedrock_stream_event_statuses,
get_bedrock_tool_name,
)
@ -53,6 +59,7 @@ from litellm.llms.bedrock.chat.converse_transformation import AmazonConverseConf
if TYPE_CHECKING:
from botocore.eventstream import EventStreamMessage
from botocore.model import Shape
converse_config: Final = AmazonConverseConfig()
_STREAM_HEAD_BYTES: Final = 200
@ -365,10 +372,14 @@ def _response_header(response_headers: Mapping[str, str] | None, name: str) -> s
class _EventStreamTally:
def __init__(self) -> None:
def __init__(self, event_statuses: Mapping[str, int | None] | None) -> None:
self.event_statuses = event_statuses
self.bytes_received = 0
self.bytes_decoded = 0
self.events = 0
self.recognized_events = 0
self.unrecognized_event_types: frozenset[str] = frozenset()
self.unrecognized_head = b""
self.head = b""
def add_chunk(self, chunk: bytes) -> None:
@ -376,13 +387,24 @@ class _EventStreamTally:
if len(self.head) < _STREAM_HEAD_BYTES:
self.head = (self.head + chunk)[:_STREAM_HEAD_BYTES]
def add_event(self, event: "EventStreamMessage") -> None:
def add_event(self, event: "EventStreamMessage", headers: Mapping[str, object]) -> None:
self.events += 1
self.bytes_decoded += event.prelude.total_length
event_type: Final = bedrock_event_stream_header(headers, ":event-type")
if (
self.event_statuses is None
or bedrock_event_stream_header(headers, ":message-type") != "event"
or (event_type is not None and event_type in self.event_statuses)
):
self.recognized_events += 1
return
self.unrecognized_event_types = self.unrecognized_event_types | {event_type or "<missing>"}
if not self.unrecognized_head:
self.unrecognized_head = event.payload[:_STREAM_HEAD_BYTES]
def undecoded_stream_error(self, response_headers: Mapping[str, str] | None) -> BedrockError | None:
undecoded: Final = self.bytes_received - self.bytes_decoded
if self.events and not undecoded:
if self.recognized_events and not undecoded:
return None
detail: Final = (
f"content-type={_response_header(response_headers, 'content-type')!r}, "
@ -397,6 +419,15 @@ class _EventStreamTally:
f"({detail}, first bytes={self.head!r})"
),
)
if not self.recognized_events:
return BedrockError(
status_code=502,
message=(
f"Bedrock answered the stream with HTTP 200 but none of its {self.events} events carried a known "
f"event type (event types={sorted(self.unrecognized_event_types)}, {detail}, "
f"first payload={self.unrecognized_head!r})"
),
)
return BedrockError(
status_code=502,
message=f"Bedrock stream ended with {undecoded} undecoded bytes after {self.events} events ({detail})",
@ -749,13 +780,12 @@ class AWSEventStreamDecoder:
from botocore.eventstream import EventStreamBuffer
event_stream_buffer: Final = EventStreamBuffer()
tally: Final = _EventStreamTally()
tally: Final = _EventStreamTally(get_bedrock_stream_event_statuses())
for chunk in iterator:
event_stream_buffer.add_data(chunk)
tally.add_chunk(chunk)
for event in event_stream_buffer:
tally.add_event(event)
message = self._parse_message_from_event(event)
message = self._decode_event(event, tally)
if message:
# sse_event = ServerSentEvent(data=message, event="completion")
_data = json.loads(message)
@ -771,13 +801,12 @@ class AWSEventStreamDecoder:
from botocore.eventstream import EventStreamBuffer
event_stream_buffer: Final = EventStreamBuffer()
tally: Final = _EventStreamTally()
tally: Final = _EventStreamTally(get_bedrock_stream_event_statuses())
async for chunk in iterator:
event_stream_buffer.add_data(chunk)
tally.add_chunk(chunk)
for event in event_stream_buffer:
tally.add_event(event)
message = self._parse_message_from_event(event)
message = self._decode_event(event, tally)
if message:
_data = json.loads(message)
yield self._chunk_parser(chunk_data=_data)
@ -785,7 +814,7 @@ class AWSEventStreamDecoder:
if undecoded_stream_error is not None:
raise undecoded_stream_error
def _parse_message_from_event(self, event) -> str | None:
def _response_stream_shape(self) -> "Shape":
response_stream_shape: Final = get_bedrock_response_stream_shape()
if response_stream_shape is None:
raise BedrockError(
@ -795,11 +824,29 @@ class AWSEventStreamDecoder:
"Ensure botocore is correctly installed."
),
)
response_dict: Final = event.to_response_dict()
return response_stream_shape
def _decode_event(self, event: "EventStreamMessage", tally: _EventStreamTally) -> str | None:
response_stream_shape: Final = self._response_stream_shape()
response_dict: Final = bedrock_event_stream_response(event)
tally.add_event(event, response_dict["headers"])
return self._parse_message_from_response(response_dict, response_stream_shape)
def _parse_message_from_event(self, event: "EventStreamMessage") -> str | None:
response_stream_shape: Final = self._response_stream_shape()
return self._parse_message_from_response(bedrock_event_stream_response(event), response_stream_shape)
def _parse_message_from_response(
self, response_dict: BedrockEventStreamResponseDict, response_stream_shape: "Shape"
) -> str | None:
parsed_response: Final = self.parser.parse(response_dict, response_stream_shape)
if response_dict["status_code"] != 200:
raise build_bedrock_stream_error(response_dict, response_stream_shape)
event_type: Final = bedrock_event_stream_header(response_dict["headers"], ":event-type")
event_error_status: Final = bedrock_stream_event_error_status(event_type)
if event_type is not None and event_error_status is not None:
raise build_bedrock_stream_event_error(event_type, event_error_status, response_dict["body"])
if "chunk" in parsed_response:
chunk = parsed_response.get("chunk")
if not chunk:

View file

@ -9,11 +9,15 @@ import functools
import json
import os
import re
from collections.abc import Mapping, Sequence
from typing import TYPE_CHECKING, Any, Final, Literal, TypedDict
from collections.abc import Iterator, Mapping, Sequence
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, Literal
from typing_extensions import ReadOnly, TypedDict
if TYPE_CHECKING:
from botocore.model import Shape
from botocore.eventstream import EventStreamMessage
from botocore.model import ServiceModel, Shape
from litellm.types.llms.bedrock import BedrockCreateBatchRequest
@ -1468,10 +1472,77 @@ def get_bedrock_response_stream_shape():
return _load_bedrock_response_stream_shape()
_BEDROCK_STREAM_OUTPUT_SHAPES: Final = ("ConverseStreamOutput", "ResponseStream")
def _modeled_error_status(member: Shape) -> int | None:
status: Final = (member.metadata or {}).get("error", {}).get("httpStatusCode")
return None if status is None else int(status)
def _structure_members(shape: Shape | None) -> Mapping[str, Shape]:
from botocore.model import StructureShape
return shape.members if isinstance(shape, StructureShape) else {}
def _bedrock_stream_output_members(service_model: ServiceModel) -> Iterator[tuple[str, Shape]]:
for shape_name in _BEDROCK_STREAM_OUTPUT_SHAPES:
yield from _structure_members(service_model.shape_for(shape_name)).items()
def _load_bedrock_stream_event_statuses() -> Mapping[str, int | None] | None:
try:
from botocore.loaders import Loader
from botocore.model import ServiceModel
service_description: Final = TypeAdapter(Mapping[str, object]).validate_python(
Loader().load_service_model("bedrock-runtime", "service-2")
)
service_model: Final = ServiceModel(service_description)
return MappingProxyType(
{name: _modeled_error_status(member) for name, member in _bedrock_stream_output_members(service_model)}
)
except Exception as e:
verbose_logger.warning(
"litellm: could not load the bedrock-runtime stream event types, "
"so unrecognized Bedrock stream events will pass through undetected. Error: %s",
e,
)
return None
@functools.lru_cache(maxsize=1)
def get_bedrock_stream_event_statuses() -> Mapping[str, int | None] | None:
"""Every modeled Bedrock stream event type mapped to its error status (None for a content event)."""
return _load_bedrock_stream_event_statuses()
def bedrock_stream_event_error_status(event_type: str | None) -> int | None:
statuses: Final = get_bedrock_stream_event_statuses()
return None if event_type is None or statuses is None else statuses.get(event_type)
class BedrockEventStreamResponseDict(TypedDict):
status_code: int
headers: Mapping[str, str]
body: bytes
status_code: ReadOnly[int]
headers: ReadOnly[Mapping[str, object]]
body: ReadOnly[bytes]
_BEDROCK_EVENT_STREAM_RESPONSE: Final = TypeAdapter(BedrockEventStreamResponseDict)
def bedrock_event_stream_response(event: EventStreamMessage) -> BedrockEventStreamResponseDict:
return _BEDROCK_EVENT_STREAM_RESPONSE.validate_python(event.to_response_dict())
def bedrock_event_stream_header(headers: Mapping[str, object], name: str) -> str | None:
value: Final = headers.get(name)
return value if isinstance(value, str) else None
def build_bedrock_stream_event_error(event_type: str, status_code: int, body: bytes) -> BedrockError:
return BedrockError(status_code=status_code, message=f"{event_type} {body.decode(errors='replace')}")
def build_bedrock_stream_error(
@ -1484,19 +1555,14 @@ def build_bedrock_stream_error(
ResponseStream member's httpStatusCode is the real status. Resolve it from the
shape and fall back to the raw status when the type is not modeled.
"""
exception_type: Final = response_dict["headers"].get(":exception-type")
decoded_body: Final = response_dict["body"].decode()
message: Final = f"{exception_type} {decoded_body}" if exception_type else decoded_body
exception_type: Final = bedrock_event_stream_header(response_dict["headers"], ":exception-type")
if exception_type is None:
return BedrockError(status_code=response_dict["status_code"], message=response_dict["body"].decode())
status_code = response_dict["status_code"]
if exception_type is not None and response_stream_shape is not None:
member: Final = response_stream_shape.members.get(exception_type)
if member is not None:
modeled_status: Final = (member.metadata or {}).get("error", {}).get("httpStatusCode")
if modeled_status is not None:
status_code = int(modeled_status)
return BedrockError(status_code=status_code, message=message)
member: Final = _structure_members(response_stream_shape).get(exception_type)
modeled_status: Final = None if member is None else _modeled_error_status(member)
status_code: Final = response_dict["status_code"] if modeled_status is None else modeled_status
return build_bedrock_stream_event_error(exception_type, status_code, response_dict["body"])
class BedrockEventStreamDecoderBase:

View file

@ -83,25 +83,39 @@ def _aws_str_header(name: str, value: str) -> bytes:
)
def _aws_int_header(name: str, value: int) -> bytes:
name_bytes: Final = name.encode()
return struct.pack("!B", len(name_bytes)) + name_bytes + struct.pack("!B", 4) + struct.pack("!i", value)
def aws_event_stream_frame(headers: Mapping[str, str | int], payload: bytes) -> bytes:
"""One AWS event-stream frame: a string header is wire type 7, an int header wire type 4 (int32)."""
headers_bytes: Final = b"".join(
_aws_str_header(name, value) if isinstance(value, str) else _aws_int_header(name, value)
for name, value in headers.items()
)
total_length: Final = 12 + len(headers_bytes) + len(payload) + 4
prelude: Final = struct.pack("!II", total_length, len(headers_bytes))
prelude_crc: Final = struct.pack("!I", zlib.crc32(prelude) & 0xFFFFFFFF)
message: Final = prelude + prelude_crc + headers_bytes + payload
return message + struct.pack("!I", zlib.crc32(message) & 0xFFFFFFFF)
def _aws_event_frame(
event_type: str,
payload: Mapping[str, JsonValue],
scenario_id: str,
unique_id: str,
) -> bytes:
payload_bytes: Final = json.dumps(payload, separators=(",", ":")).replace(
"$REQUEST_ID", scenario_id
).replace("$UNIQUE_ID", unique_id).encode()
headers_bytes: Final = (
_aws_str_header(":event-type", event_type)
+ _aws_str_header(":content-type", "application/json")
+ _aws_str_header(":message-type", "event")
payload_bytes: Final = (
json.dumps(payload, separators=(",", ":"))
.replace("$REQUEST_ID", scenario_id)
.replace("$UNIQUE_ID", unique_id)
.encode()
)
return aws_event_stream_frame(
{":event-type": event_type, ":content-type": "application/json", ":message-type": "event"}, payload_bytes
)
total_length: Final = 12 + len(headers_bytes) + len(payload_bytes) + 4
prelude: Final = struct.pack("!II", total_length, len(headers_bytes))
prelude_crc: Final = struct.pack("!I", zlib.crc32(prelude) & 0xFFFFFFFF)
message: Final = prelude + prelude_crc + headers_bytes + payload_bytes
return message + struct.pack("!I", zlib.crc32(message) & 0xFFFFFFFF)
class ScenarioStore:
@ -147,7 +161,13 @@ class Provider:
status: Final = script.popleft()
if status != 200:
return JSONResponse(
{"error": {"message": "Controlled provider failure", "type": error_type(status), "code": str(status)}},
{
"error": {
"message": "Controlled provider failure",
"type": error_type(status),
"code": str(status),
}
},
status_code=status,
)
return await chat_completions(request)
@ -244,9 +264,7 @@ class Provider:
if raw_body:
body: Final = JSON_OBJECT.validate_json(raw_body)
if isinstance(body, dict):
self.observations.put(
Observation(request.url.path, request.headers.get("authorization", ""), body)
)
self.observations.put(Observation(request.url.path, request.headers.get("authorization", ""), body))
if isinstance(response, RoutedResponse):
route_key: Final = f"{request.method} /{'/'.join(segments[1:])}"
route: Final = next(
@ -300,11 +318,10 @@ class Provider:
match response:
case JsonResponse():
return Response(
content=json.dumps(response.body, separators=(",", ":")).replace(
"$REQUEST_ID", scenario_id
).replace(
"$UNIQUE_ID", unique_id
).encode(),
content=json.dumps(response.body, separators=(",", ":"))
.replace("$REQUEST_ID", scenario_id)
.replace("$UNIQUE_ID", unique_id)
.encode(),
media_type=response.content_type,
status_code=response.status,
)
@ -321,6 +338,7 @@ class Provider:
)
case SseResponse():
if response.frame_delay_ms > 0:
async def stream() -> AsyncIterator[bytes]:
for frame in response.frames:
yield (
@ -329,9 +347,11 @@ class Provider:
await asyncio.sleep(response.frame_delay_ms / 1000)
return StreamingResponse(stream(), media_type=response.content_type)
stream_body: Final = ("\n\n".join(response.frames) + "\n\n").replace(
"$REQUEST_ID", scenario_id
).replace("$UNIQUE_ID", unique_id)
stream_body: Final = (
("\n\n".join(response.frames) + "\n\n")
.replace("$REQUEST_ID", scenario_id)
.replace("$UNIQUE_ID", unique_id)
)
return Response(content=stream_body.encode(), media_type=response.content_type)
case EventStreamResponse():
events: Final = (

View file

@ -0,0 +1,935 @@
import asyncio
import base64
import json
import os
import re
import signal
import threading
import uuid
from collections.abc import Callable, Iterable, Mapping
from dataclasses import dataclass
from pathlib import Path
from queue import SimpleQueue
from types import MappingProxyType
from typing import Final, Literal
from urllib.parse import unquote, urlsplit
import anthropic
import httpx
import openai
import psutil
import pytest
import yaml
from integration._support.client import Gateway, Scenario, eventually
from integration._support.database import read_rows
from integration._support.process import owned_proxy_process
from integration._support.upstream import _aws_event_frame, aws_event_stream_frame
from integration._support.wire import Reply, Request, Wire, wire_server
from pydantic import JsonValue, TypeAdapter
from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_if_encrypted_with
_MODEL_ID: Final = "global.moonshotai.kimi-k3"
_CONVERSE_MODEL: Final = f"bedrock/converse/{_MODEL_ID}"
_STREAM_TARGET: Final = f"/model/{_MODEL_ID}/converse-stream"
_CONVERSE_TARGET: Final = f"/model/{_MODEL_ID}/converse"
_INVOKE_MODEL_ID: Final = "anthropic.claude-3-haiku-20240307-v1:0"
_INVOKE_MODEL: Final = f"bedrock/invoke/{_INVOKE_MODEL_ID}"
_INVOKE_STREAM_TARGET: Final = f"/model/{_INVOKE_MODEL_ID}/invoke-with-response-stream"
_EVENT_STREAM: Final = "application/vnd.amazon.eventstream"
_ANSWER: Final = "bedrock event frame control"
_REJECTION: Final = "structured output schema uses unsupported regex negative look-ahead"
_THROTTLED: Final = "Too many requests, please wait before trying again"
_UNKNOWN_TYPE: Final = "somethingBedrockAddedLater"
_UNKNOWN_ONLY: Final = "none of its 1 events carried a known event type"
_SIGNING_KEY: Final = os.environ.get("LITELLM_SALT_KEY", "sk-integration-salt")
_JSON_HEADERS: Final = MappingProxyType({":content-type": "application/json", ":message-type": "event"})
_USAGE: Final[dict[str, JsonValue]] = {"usage": {"inputTokens": 11, "outputTokens": 4, "totalTokens": 15}}
_RESPONSE: Final = json.dumps(
{
"output": {"message": {"role": "assistant", "content": [{"text": _ANSWER}]}},
"stopReason": "end_turn",
**_USAGE,
"metrics": {"latencyMs": 1},
}
).encode()
_JSON: Final = TypeAdapter(dict[str, JsonValue])
_AWS: Final[dict[str, JsonValue]] = {
"aws_access_key_id": "AKIASCRIPTEDPROVIDER",
"aws_secret_access_key": "scripted-secret",
"aws_region_name": "us-east-1",
}
_EXTRA: Final[dict[str, JsonValue]] = {"num_retries": 0, "cache": {"no-cache": True}}
_PROMPT: Final = "What does the gateway do with this stream?"
_USER_TURN: Final[dict[str, JsonValue]] = {"role": "user", "content": _PROMPT}
_CONVERSE_USER_TURN: Final[dict[str, JsonValue]] = {"role": "user", "content": [{"text": _PROMPT}]}
_STARTED_WORKER: Final = re.compile(r"Started server process \[(\d+)\]")
_CALL_INDEX: Final = re.compile(r"call-[0-9a-f]{32}-(\d+)")
_PLAIN: Final = "bedrock-event-frames-plain"
Endpoint = Literal["chat", "messages", "responses"]
def _error_body(message: str) -> str:
return json.dumps({"message": message}, separators=(",", ":"))
def _frame(event_type: str, payload: Mapping[str, JsonValue]) -> bytes:
return _aws_event_frame(event_type, payload, "sc", "u")
def _typed_frame(headers: Mapping[str, str | int], payload: bytes) -> bytes:
return aws_event_stream_frame({**headers, **_JSON_HEADERS}, payload)
def _exception_message_frame(exception_type: str, message: str) -> bytes:
return aws_event_stream_frame(
{":exception-type": exception_type, ":content-type": "application/json", ":message-type": "exception"},
_error_body(message).encode(),
)
def _text_frames(text: str) -> tuple[bytes, ...]:
return (
_frame("messageStart", {"role": "assistant"}),
_frame("contentBlockDelta", {"delta": {"text": text}, "contentBlockIndex": 0}),
_frame("contentBlockStop", {"contentBlockIndex": 0}),
)
_NORMAL: Final = b"".join(
(*_text_frames(_ANSWER), _frame("messageStop", {"stopReason": "end_turn"}), _frame("metadata", _USAGE))
)
_VALIDATION_FRAME: Final = _frame("validationException", {"message": _REJECTION})
_THROTTLING_FRAME: Final = _frame("throttlingException", {"message": _THROTTLED})
_UNKNOWN_FRAME: Final = _frame(_UNKNOWN_TYPE, {"future": True})
_UNKNOWN_BESIDE_KNOWN: Final = b"".join(
(
*_text_frames(_ANSWER),
_UNKNOWN_FRAME,
_frame("messageStop", {"stopReason": "end_turn"}),
_frame("metadata", _USAGE),
)
)
_THROTTLED_AFTER_TEXT: Final = b"".join((*_text_frames(_ANSWER), _THROTTLING_FRAME))
_EMPTY_DELTA_BESIDE_KNOWN: Final = b"".join(
(
_frame("messageStart", {"role": "assistant"}),
_frame("contentBlockDelta", {}),
_frame("contentBlockDelta", {"delta": {"text": _ANSWER}, "contentBlockIndex": 0}),
_frame("messageStop", {"stopReason": "end_turn"}),
_frame("metadata", _USAGE),
)
)
_EXCEPTION_MID_STREAM: Final = b"".join((*_text_frames(_ANSWER), _VALIDATION_FRAME))
_EXCEPTION_MESSAGE_MID_STREAM: Final = b"".join(
(*_text_frames(_ANSWER), _exception_message_frame("throttlingException", _THROTTLED))
)
def _invoke_chunk(event: Mapping[str, JsonValue]) -> bytes:
return _frame("chunk", {"bytes": base64.b64encode(json.dumps(event).encode()).decode()})
_INVOKE_STREAM: Final = b"".join(
(
_invoke_chunk(
{
"type": "message_start",
"message": {
"id": "msg_invoke",
"type": "message",
"role": "assistant",
"content": [],
"model": _INVOKE_MODEL_ID,
"stop_reason": None,
"usage": {"input_tokens": 11, "output_tokens": 1},
},
}
),
_invoke_chunk({"type": "content_block_start", "index": 0, "content_block": {"type": "text", "text": ""}}),
_invoke_chunk({"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": _ANSWER}}),
_invoke_chunk({"type": "content_block_stop", "index": 0}),
_invoke_chunk({"type": "message_delta", "delta": {"stop_reason": "end_turn"}, "usage": {"output_tokens": 4}}),
_invoke_chunk({"type": "message_stop"}),
)
)
@dataclass(frozen=True, slots=True)
class _Streamed:
status: int
call_id: str
lines: tuple[str, ...]
@property
def text(self) -> str:
return "\n".join(self.lines)
@dataclass(frozen=True, slots=True)
class _Call:
endpoint: Endpoint
user: str
index: int
@dataclass(frozen=True, slots=True)
class _Served:
call: _Call
status: int
call_id: str
text: str
def _stream_peer(frames: bytes) -> Callable[[Request], Reply]:
def respond(request: Request) -> Reply:
if unquote(request.target) in (_STREAM_TARGET, _INVOKE_STREAM_TARGET):
return Reply(body=frames, content_type=_EVENT_STREAM)
return Reply(body=_RESPONSE)
return respond
def _converse_deployment(scenario: Scenario, wire: Wire, **extra: JsonValue) -> str:
return scenario.model(model=_CONVERSE_MODEL, api_base=wire.url, **_AWS, **extra)
def _auth(gateway: Gateway) -> dict[str, str]:
return {"Authorization": f"Bearer {gateway.key}"}
def _proxy_url(gateway: Gateway) -> str:
return str(gateway.client.base_url).rstrip("/")
def _path(endpoint: Endpoint) -> str:
match endpoint:
case "chat":
return "/v1/chat/completions"
case "messages":
return "/v1/messages"
case "responses":
return "/v1/responses"
def _body(endpoint: Endpoint, model: str, *, user: str | None = None) -> dict[str, JsonValue]:
marker: Final[dict[str, JsonValue]] = {} if user is None else {"user": user}
match endpoint:
case "chat":
return {"model": model, "messages": [_USER_TURN], "max_tokens": 16, "stream": True, **_EXTRA, **marker}
case "messages":
return {"model": model, "messages": [_USER_TURN], "max_tokens": 16, "stream": True, **_EXTRA}
case "responses":
return {"model": model, "input": _PROMPT, "stream": True, **_EXTRA, **marker}
def _stream(
gateway: Gateway, endpoint: Endpoint, body: Mapping[str, JsonValue], *, key: str | None = None
) -> _Streamed:
headers: Final = _auth(gateway) if key is None else {"Authorization": f"Bearer {key}"}
with gateway.client.stream("POST", _path(endpoint), json=body, headers=headers) as response:
lines: Final = tuple(line for line in response.iter_lines() if line)
return _Streamed(response.status_code, response.headers.get("x-litellm-call-id", ""), lines)
def _sse_payloads(lines: Iterable[str]) -> tuple[dict[str, JsonValue], ...]:
return tuple(json.loads(line[6:]) for line in lines if line.startswith("data: ") and line != "data: [DONE]")
def _sse_events(lines: Iterable[str]) -> tuple[str, ...]:
return tuple(line[7:] for line in lines if line.startswith("event: "))
def _first_choice(chunk: Mapping[str, JsonValue]) -> dict[str, JsonValue] | None:
choices: Final = chunk.get("choices")
return _JSON.validate_python(choices[0]) if isinstance(choices, list) and choices else None
def _chat_text(chunks: Iterable[dict[str, JsonValue]]) -> str:
choices: Final = tuple(choice for choice in map(_first_choice, chunks) if choice is not None)
return "".join(str(_JSON.validate_python(choice["delta"]).get("content") or "") for choice in choices)
def _finish_reasons(chunks: Iterable[dict[str, JsonValue]]) -> tuple[JsonValue, ...]:
return tuple(choice.get("finish_reason") for choice in map(_first_choice, chunks) if choice is not None)
def _chat_id(chunks: Iterable[dict[str, JsonValue]]) -> str:
(identity,) = {str(chunk["id"]) for chunk in chunks if "id" in chunk}
return identity
def _message_id(payloads: Iterable[dict[str, JsonValue]]) -> str:
(started,) = tuple(payload for payload in payloads if payload.get("type") == "message_start")
return str(_JSON.validate_python(started["message"])["id"])
def _messages_text(payloads: Iterable[dict[str, JsonValue]]) -> str:
deltas: Final = tuple(payload for payload in payloads if payload.get("type") == "content_block_delta")
return "".join(str(_JSON.validate_python(delta["delta"]).get("text") or "") for delta in deltas)
def _responses_text(events: Iterable[dict[str, JsonValue]]) -> str:
return "".join(
str(event.get("delta") or "") for event in events if event.get("type") == "response.output_text.delta"
)
def _inner_response_id(identity: str) -> str:
managed: Final = decrypt_if_encrypted_with(identity.removeprefix("resp_"), _SIGNING_KEY)
assert managed is not None, identity
issued: Final = managed.split(";", 1)[0].rsplit("response_id:", 1)[1]
decoded: Final = base64.b64decode(issued.removeprefix("resp_")).decode()
return decoded.rsplit("response_id:", 1)[1]
def _completed_response_id(events: Iterable[dict[str, JsonValue]]) -> str:
(completed,) = tuple(event for event in events if event.get("type") == "response.completed")
return _inner_response_id(str(_JSON.validate_python(completed["response"])["id"]))
def _only_received(wire: Wire) -> tuple[str, dict[str, JsonValue]]:
(request,) = wire.drain()
return unquote(request.target), _JSON.validate_python(json.loads(request.body))
def _assert_stream_request(wire: Wire, target: str = _STREAM_TARGET) -> dict[str, JsonValue]:
received_target, received = _only_received(wire)
assert received_target == target, received_target
return received
def _spend_row(request_id: str) -> dict[str, JsonValue]:
assert request_id, "No id to look the spend row up by"
(row,) = eventually(
lambda: read_rows(
'SELECT request_id, status, call_type, end_user FROM "LiteLLM_SpendLogs" WHERE request_id=%s',
(request_id,),
),
lambda found: len(found) >= 1,
seconds=70,
)
return row
def _failure_row(call_id: str) -> dict[str, JsonValue]:
row: Final = _spend_row(call_id)
assert row["status"] == "failure", row
return row
def _success_row(request_id: str) -> dict[str, JsonValue]:
row: Final = _spend_row(request_id)
assert row["status"] == "success", row
return row
def _rows_for(request_ids: frozenset[str], *, expected: int) -> tuple[dict[str, JsonValue], ...]:
rows: Final = eventually(
lambda: read_rows(
'SELECT request_id, status FROM "LiteLLM_SpendLogs" WHERE request_id = ANY(string_to_array(%s, %s))',
(",".join(sorted(request_ids)), ","),
),
lambda found: len(found) >= expected,
seconds=90,
)
return tuple(rows)
def _assert_rejected_chat(streamed: _Streamed, status: int, message: str) -> None:
assert streamed.status == status, (streamed.status, streamed.text)
error: Final = _JSON.validate_python(json.loads(streamed.text)["error"])
assert message in str(error["message"]), streamed.text
assert str(error["code"]) == str(status), streamed.text
def _unescaped(text: str) -> str:
return text.replace('\\"', '"')
def _assert_rejected_stream_body(streamed: _Streamed, message: str) -> None:
assert streamed.status == 200, (streamed.status, streamed.text)
assert message in _unescaped(streamed.text), streamed.text
assert _ANSWER not in streamed.text, streamed.text
def test_r01_chat_stream_with_a_validation_exception_frame_first_is_a_400_and_a_failure_row(gateway: Gateway) -> None:
with wire_server(_stream_peer(_VALIDATION_FRAME)) as wire, gateway.scenario() as scenario:
model: Final = _converse_deployment(scenario, wire)
streamed: Final = _stream(gateway, "chat", _body("chat", model))
_assert_rejected_chat(streamed, 400, f"validationException {_error_body(_REJECTION)}")
received: Final = _assert_stream_request(wire)
assert received["messages"] == [_CONVERSE_USER_TURN], received
_failure_row(streamed.call_id)
async def _consume_openai_chat_stream(client: openai.AsyncOpenAI, model: str) -> None:
stream = await client.chat.completions.create(
model=model, messages=[_USER_TURN], max_tokens=16, stream=True, extra_body=_EXTRA
)
_ = [chunk async for chunk in stream]
async def test_r02_openai_async_sdk_stream_with_a_validation_exception_frame_first_raises_bad_request(
gateway: Gateway,
) -> None:
with wire_server(_stream_peer(_VALIDATION_FRAME)) as wire, gateway.scenario() as scenario:
model: Final = _converse_deployment(scenario, wire)
async with openai.AsyncOpenAI(
base_url=f"{_proxy_url(gateway)}/v1", api_key=gateway.key, max_retries=0
) as client:
with pytest.raises(openai.BadRequestError, match=re.escape(_REJECTION)) as raised:
await _consume_openai_chat_stream(client, model)
_assert_stream_request(wire)
_failure_row(raised.value.response.headers.get("x-litellm-call-id", ""))
def test_r03_chat_stream_throttled_after_text_delivers_the_text_then_the_error_and_no_stop_chunk(
gateway: Gateway,
) -> None:
with wire_server(_stream_peer(_THROTTLED_AFTER_TEXT)) as wire, gateway.scenario() as scenario:
model: Final = _converse_deployment(scenario, wire)
streamed: Final = _stream(gateway, "chat", _body("chat", model))
assert streamed.status == 200, streamed.text
chunks: Final = _sse_payloads(streamed.lines)
assert _chat_text(chunks) == _ANSWER, streamed.text
assert "throttlingException" in streamed.text and _THROTTLED in streamed.text, streamed.text
assert "stop" not in _finish_reasons(chunks), streamed.text
_assert_stream_request(wire)
_failure_row(streamed.call_id)
def test_r04_messages_stream_with_a_validation_exception_frame_first_emits_an_error_event_and_no_message_stop(
gateway: Gateway,
) -> None:
with wire_server(_stream_peer(_VALIDATION_FRAME)) as wire, gateway.scenario() as scenario:
model: Final = _converse_deployment(scenario, wire)
streamed: Final = _stream(gateway, "messages", _body("messages", model))
_assert_rejected_stream_body(streamed, f"validationException {_error_body(_REJECTION)}")
events: Final = _sse_events(streamed.lines)
assert "error" in events and "message_stop" not in events, streamed.text
_assert_stream_request(wire)
async def test_r05_anthropic_async_sdk_stream_with_a_validation_exception_frame_first_raises(
gateway: Gateway,
) -> None:
with wire_server(_stream_peer(_VALIDATION_FRAME)) as wire, gateway.scenario() as scenario:
model: Final = _converse_deployment(scenario, wire)
async with anthropic.AsyncAnthropic(base_url=_proxy_url(gateway), api_key=gateway.key, max_retries=0) as client:
with pytest.raises(anthropic.APIError, match=re.escape(_REJECTION)):
async with client.messages.stream(
model=model, max_tokens=16, messages=[_USER_TURN], extra_body=_EXTRA
) as stream:
_ = [event async for event in stream]
_assert_stream_request(wire)
def test_r06_responses_stream_with_a_validation_exception_frame_first_fails_the_response(gateway: Gateway) -> None:
with wire_server(_stream_peer(_VALIDATION_FRAME)) as wire, gateway.scenario() as scenario:
model: Final = _converse_deployment(scenario, wire)
streamed: Final = _stream(gateway, "responses", _body("responses", model))
_assert_rejected_stream_body(streamed, f"validationException {_error_body(_REJECTION)}")
types: Final = tuple(str(event["type"]) for event in _sse_payloads(streamed.lines))
assert "response.failed" in types and "response.completed" not in types, streamed.text
_assert_stream_request(wire)
_failure_row(streamed.call_id)
def test_r07_chat_stream_whose_only_frame_has_an_unknown_event_type_is_a_502_naming_the_type(
gateway: Gateway,
) -> None:
with wire_server(_stream_peer(_UNKNOWN_FRAME)) as wire, gateway.scenario() as scenario:
model: Final = _converse_deployment(scenario, wire)
streamed: Final = _stream(gateway, "chat", _body("chat", model))
_assert_rejected_chat(streamed, 502, f"{_UNKNOWN_ONLY} (event types=['{_UNKNOWN_TYPE}']")
_assert_stream_request(wire)
_failure_row(streamed.call_id)
def _assert_text_stream(gateway: Gateway, endpoint: Endpoint, frames: bytes) -> None:
marker: Final = f"call-{uuid.uuid4().hex}"
with wire_server(_stream_peer(frames)) as wire, gateway.scenario() as scenario:
model: Final = _converse_deployment(scenario, wire)
streamed: Final = _stream(gateway, endpoint, _body(endpoint, model, user=marker))
assert streamed.status == 200, streamed.text
payloads: Final = _sse_payloads(streamed.lines)
match endpoint:
case "chat":
assert _chat_text(payloads) == _ANSWER, streamed.text
assert _finish_reasons(payloads)[-1] == "stop", streamed.text
_success_row(_chat_id(payloads))
case "messages":
assert _messages_text(payloads) == _ANSWER, streamed.text
assert _sse_events(streamed.lines)[-1] == "message_stop", streamed.text
_success_row(_message_id(payloads))
case "responses":
assert _responses_text(payloads) == _ANSWER, streamed.text
assert str(payloads[-1]["type"]) == "response.completed", streamed.text
assert _success_row(_completed_response_id(payloads))["end_user"] == marker
_assert_stream_request(wire)
@pytest.mark.parametrize(
"endpoint", ("chat", "messages", "responses"), ids=("r08-chat", "r08-messages", "r08-responses")
)
def test_r08_an_unknown_frame_between_known_frames_leaves_the_text_and_the_success_row_intact(
gateway: Gateway, endpoint: Endpoint
) -> None:
_assert_text_stream(gateway, endpoint, _UNKNOWN_BESIDE_KNOWN)
@pytest.mark.parametrize(
"endpoint", ("chat", "messages", "responses"), ids=("r09-chat", "r09-messages", "r09-responses")
)
def test_r09_a_normal_converse_stream_delivers_the_text_and_a_success_row(gateway: Gateway, endpoint: Endpoint) -> None:
_assert_text_stream(gateway, endpoint, _NORMAL)
def test_r10_a_non_streaming_converse_call_is_untouched(gateway: Gateway) -> None:
with wire_server(_stream_peer(_NORMAL)) as wire, gateway.scenario() as scenario:
model: Final = _converse_deployment(scenario, wire)
response: Final = gateway.request(
"POST", "/v1/chat/completions", {"model": model, "messages": [_USER_TURN], "max_tokens": 16, **_EXTRA}
)
assert response.status_code == 200, response.text
assert response.json()["choices"][0]["message"]["content"] == _ANSWER, response.text
_assert_stream_request(wire, _CONVERSE_TARGET)
_success_row(str(response.json()["id"]))
def test_r11_an_invoke_framed_anthropic_stream_is_untouched(gateway: Gateway) -> None:
with wire_server(_stream_peer(_INVOKE_STREAM)) as wire, gateway.scenario() as scenario:
model: Final = scenario.model(
model=_INVOKE_MODEL, api_key=None, aws_bedrock_runtime_endpoint=wire.url, api_base=wire.url, **_AWS
)
streamed: Final = _stream(gateway, "chat", _body("chat", model))
assert streamed.status == 200, streamed.text
chunks: Final = _sse_payloads(streamed.lines)
assert _chat_text(chunks) == _ANSWER, streamed.text
assert _finish_reasons(chunks)[-1] == "stop", streamed.text
_assert_stream_request(wire, _INVOKE_STREAM_TARGET)
_success_row(_chat_id(chunks))
@pytest.mark.parametrize(
("exception_type", "expected"),
(("throttlingException", 429), ("somethingNewException", 400)),
ids=("r12", "r13"),
)
def test_r12_r13_an_exception_message_frame_keeps_its_modeled_status(
gateway: Gateway, exception_type: str, expected: int
) -> None:
with (
wire_server(_stream_peer(_exception_message_frame(exception_type, _THROTTLED))) as wire,
gateway.scenario() as scenario,
):
model: Final = _converse_deployment(scenario, wire)
streamed: Final = _stream(gateway, "chat", _body("chat", model))
_assert_rejected_chat(streamed, expected, _THROTTLED)
_assert_stream_request(wire)
_failure_row(streamed.call_id)
@pytest.mark.parametrize(
"frames", (_EXCEPTION_MID_STREAM, _EXCEPTION_MESSAGE_MID_STREAM), ids=("r14-event-frame", "r15-exception-message")
)
def test_r14_r15_passthrough_converse_stream_relays_an_exception_frame_byte_for_byte(
gateway: Gateway, frames: bytes
) -> None:
with wire_server(_stream_peer(frames)) as wire, gateway.scenario() as scenario:
deployment: Final = scenario.model(
model=f"bedrock/{_MODEL_ID}", api_base=wire.url, aws_bedrock_runtime_endpoint=wire.url, **_AWS
)
response: Final = gateway.request(
"POST", f"/bedrock/model/{deployment}/converse-stream", {"messages": [_CONVERSE_USER_TURN]}
)
assert response.status_code == 200, response.text
assert response.headers.get("content-type") == _EVENT_STREAM, dict(response.headers)
assert response.content == frames, response.content
_assert_stream_request(wire)
_HEADERLESS: Final = _typed_frame({}, b'{"future": true}')
_INT_TYPED: Final = _typed_frame({":event-type": 7}, b'{"future": true}')
_EMPTY_TYPED: Final = _typed_frame({":event-type": ""}, b'{"future": true}')
_LONG_TYPE: Final = "x" * 5120
_LONG_TYPED: Final = _typed_frame({":event-type": _LONG_TYPE}, b'{"future": true}')
@pytest.mark.parametrize(
("frames", "named"),
(
(_HEADERLESS, f"{_UNKNOWN_ONLY} (event types=['<missing>']"),
(_INT_TYPED, f"{_UNKNOWN_ONLY} (event types=['<missing>']"),
(_EMPTY_TYPED, f"{_UNKNOWN_ONLY} (event types=['<missing>']"),
(_LONG_TYPED, f"{_UNKNOWN_ONLY} (event types=['{_LONG_TYPE}']"),
(
_UNKNOWN_FRAME + _UNKNOWN_FRAME,
f"none of its 2 events carried a known event type (event types=['{_UNKNOWN_TYPE}']",
),
),
ids=(
"s01-no-event-type",
"s02-int-event-type",
"s03-empty-event-type",
"s04-5kb-event-type",
"s05-same-type-twice",
),
)
def test_s01_to_s05_odd_event_type_headers_alone_are_a_502_that_names_what_arrived(
gateway: Gateway, frames: bytes, named: str
) -> None:
with wire_server(_stream_peer(frames)) as wire, gateway.scenario() as scenario:
model: Final = _converse_deployment(scenario, wire)
streamed: Final = _stream(gateway, "chat", _body("chat", model))
_assert_rejected_chat(streamed, 502, named)
_assert_stream_request(wire)
_failure_row(streamed.call_id)
def test_s06_a_validation_exception_frame_with_a_non_utf8_body_is_a_400_with_the_bytes_replaced(
gateway: Gateway,
) -> None:
frame: Final = _typed_frame({":event-type": "validationException"}, b'{"message": "bad \xff\xfe bytes"}')
with wire_server(_stream_peer(frame)) as wire, gateway.scenario() as scenario:
model: Final = _converse_deployment(scenario, wire)
streamed: Final = _stream(gateway, "chat", _body("chat", model))
_assert_rejected_chat(streamed, 400, 'validationException {"message": "bad <20><> bytes"}')
_assert_stream_request(wire)
_failure_row(streamed.call_id)
def test_s07_a_validation_exception_frame_with_an_empty_body_is_still_a_400(gateway: Gateway) -> None:
with wire_server(_stream_peer(_typed_frame({":event-type": "validationException"}, b""))) as wire:
with gateway.scenario() as scenario:
model: Final = _converse_deployment(scenario, wire)
streamed: Final = _stream(gateway, "chat", _body("chat", model))
_assert_rejected_chat(streamed, 400, "validationException")
_assert_stream_request(wire)
_failure_row(streamed.call_id)
def test_s08_a_known_frame_with_an_empty_body_beside_normal_frames_keeps_the_text(gateway: Gateway) -> None:
with wire_server(_stream_peer(_EMPTY_DELTA_BESIDE_KNOWN)) as wire, gateway.scenario() as scenario:
model: Final = _converse_deployment(scenario, wire)
streamed: Final = _stream(gateway, "chat", _body("chat", model))
assert streamed.status == 200, streamed.text
chunks: Final = _sse_payloads(streamed.lines)
assert _chat_text(chunks) == _ANSWER, streamed.text
assert _finish_reasons(chunks)[-1] == "stop", streamed.text
_assert_stream_request(wire)
_success_row(_chat_id(chunks))
def test_s09_an_unauthenticated_stream_never_reaches_the_peer(gateway: Gateway) -> None:
with wire_server(_stream_peer(_VALIDATION_FRAME)) as wire, gateway.scenario() as scenario:
model: Final = _converse_deployment(scenario, wire)
streamed: Final = _stream(gateway, "chat", _body("chat", model), key="sk-integration-bogus")
assert streamed.status == 401, streamed.text
assert wire.drain() == (), streamed.text
def test_s10_a_rejected_stream_leaves_a_healthy_deployment_serving(gateway: Gateway) -> None:
with (
wire_server(_stream_peer(_VALIDATION_FRAME)) as rejecting,
wire_server(_stream_peer(_NORMAL)) as healthy,
gateway.scenario() as scenario,
):
rejected_model: Final = _converse_deployment(scenario, rejecting)
healthy_model: Final = _converse_deployment(scenario, healthy)
_assert_rejected_chat(_stream(gateway, "chat", _body("chat", rejected_model)), 400, "validationException")
streamed: Final = _stream(gateway, "chat", _body("chat", healthy_model))
assert streamed.status == 200, streamed.text
assert _chat_text(_sse_payloads(streamed.lines)) == _ANSWER, streamed.text
assert len(rejecting.drain()) == 1 and len(healthy.drain()) == 1
def test_e01_a_rejected_stream_is_never_served_from_the_response_cache(gateway: Gateway) -> None:
with wire_server(_stream_peer(_VALIDATION_FRAME)) as wire, gateway.scenario() as scenario:
model: Final = _converse_deployment(scenario, wire)
cached: Final = {key: value for key, value in _body("chat", model).items() if key != "cache"}
first: Final = _stream(gateway, "chat", cached)
second: Final = _stream(gateway, "chat", cached)
_assert_rejected_chat(first, 400, "validationException")
_assert_rejected_chat(second, 400, "validationException")
assert len(wire.drain()) == 2, (first.text, second.text)
def _sibling_deployment(scenario: Scenario, name: str, wire: Wire, **extra: JsonValue) -> None:
created: Final = scenario.gateway.post(
"/model/new",
{
"model_name": name,
"litellm_params": {"model": _CONVERSE_MODEL, "api_base": wire.url, **_AWS, **extra},
"model_info": {},
},
)
identity: Final = _JSON.validate_python(created["model_info"])["id"]
assert isinstance(identity, str), created
scenario.cleanups.callback(scenario.delete_model, identity)
def test_e02_a_throttling_frame_first_is_a_429_after_one_attempt_even_with_retries_and_a_sibling_deployment(
gateway: Gateway,
) -> None:
with wire_server(_stream_peer(_THROTTLING_FRAME)) as wire, gateway.scenario() as scenario:
model: Final = _converse_deployment(scenario, wire, num_retries=2)
_sibling_deployment(scenario, model, wire, num_retries=2)
streamed: Final = _stream(gateway, "chat", {**_body("chat", model), "num_retries": 2})
_assert_rejected_chat(streamed, 429, f"throttlingException {_error_body(_THROTTLED)}")
attempts: Final = len(wire.drain())
assert attempts == 1, attempts
_failure_row(streamed.call_id)
def test_e03_three_rejected_streams_land_one_failure_row_each(gateway: Gateway) -> None:
with wire_server(_stream_peer(_VALIDATION_FRAME)) as wire, gateway.scenario() as scenario:
model: Final = _converse_deployment(scenario, wire)
streamed: Final = tuple(_stream(gateway, "chat", _body("chat", model)) for _ in range(3))
for item in streamed:
_assert_rejected_chat(item, 400, "validationException")
call_ids: Final = frozenset(item.call_id for item in streamed)
assert len(call_ids) == 3, streamed
rows: Final = _rows_for(call_ids, expected=3)
assert {str(row["request_id"]) for row in rows} == call_ids, rows
assert all(row["status"] == "failure" for row in rows), rows
assert len(wire.drain()) == 3
def _calls(marker: str, endpoint: Endpoint, indexes: range) -> tuple[_Call, ...]:
return tuple(_Call(endpoint, f"{marker}-{index}", index) for index in indexes)
def _burst_body(model: str, call: _Call) -> dict[str, JsonValue]:
prompt: Final = f"{_PROMPT} {call.user}"
match call.endpoint:
case "chat":
return {**_body("chat", model, user=call.user), "messages": [{"role": "user", "content": prompt}]}
case "messages":
return {**_body("messages", model), "messages": [{"role": "user", "content": prompt}]}
case "responses":
return {**_body("responses", model, user=call.user), "input": prompt}
def _call_index(request: Request) -> int:
found: Final = _CALL_INDEX.search(request.body.decode())
assert found is not None, request.body
return int(found.group(1))
async def _send(client: httpx.AsyncClient, key: str, model: str, call: _Call) -> _Served:
async with client.stream(
"POST", _path(call.endpoint), json=_burst_body(model, call), headers={"Authorization": f"Bearer {key}"}
) as response:
raw: Final = await response.aread()
return _Served(call, response.status_code, response.headers.get("x-litellm-call-id", ""), raw.decode())
async def _burst(
base_url: str, key: str, model: str, calls: tuple[_Call, ...], *, tolerate_transport_errors: bool = False
) -> tuple[_Served, ...]:
async with httpx.AsyncClient(base_url=base_url, timeout=60, trust_env=False) as client:
results: Final = await asyncio.gather(
*(_send(client, key, model, call) for call in calls), return_exceptions=tolerate_transport_errors
)
for result in results:
assert not isinstance(result, BaseException) or isinstance(result, httpx.TransportError), repr(result)
return tuple(result for result in results if isinstance(result, _Served))
def _carries_text(served: _Served) -> bool:
return _ANSWER in served.text
def _is_rejection(served: _Served, message: str) -> bool:
return (
message in _unescaped(served.text) and not _carries_text(served) and '"finish_reason":"stop"' not in served.text
)
def _served_success_id(served: _Served) -> str:
payloads: Final = _sse_payloads(served.text.splitlines())
match served.call.endpoint:
case "chat":
return _chat_id(payloads)
case "messages":
return _message_id(payloads)
case "responses":
return _completed_response_id(payloads)
async def test_c01_a_mixed_burst_of_normal_rejected_and_unknown_streams_sorts_every_call_and_row(
gateway: Gateway,
) -> None:
marker: Final = f"call-{uuid.uuid4().hex}"
calls: Final = (
*_calls(marker, "chat", range(0, 10)),
*_calls(marker, "messages", range(10, 20)),
*_calls(marker, "responses", range(20, 30)),
)
def respond(request: Request) -> Reply:
match _call_index(request) % 3:
case 1:
return Reply(body=_VALIDATION_FRAME, content_type=_EVENT_STREAM)
case 2:
return Reply(body=_UNKNOWN_FRAME, content_type=_EVENT_STREAM)
case _:
return Reply(body=_NORMAL, content_type=_EVENT_STREAM)
with wire_server(respond) as wire, gateway.scenario() as scenario:
model: Final = _converse_deployment(scenario, wire)
served: Final = await _burst(_proxy_url(gateway), gateway.key, model, calls)
assert len(served) == 30
normal: Final = tuple(item for item in served if item.call.index % 3 == 0)
rejected: Final = tuple(item for item in served if item.call.index % 3 == 1)
unknown: Final = tuple(item for item in served if item.call.index % 3 == 2)
for item in normal:
assert item.status == 200 and _carries_text(item), (item.call, item.status, item.text)
for item in rejected:
assert _is_rejection(item, _REJECTION), (item.call, item.status, item.text)
for item in unknown:
assert _is_rejection(item, _UNKNOWN_ONLY), (item.call, item.status, item.text)
failed_ids: Final = frozenset(
item.call_id for item in (*rejected, *unknown) if item.call.endpoint != "messages"
)
assert len(failed_ids) == 13, failed_ids
failure_rows: Final = _rows_for(failed_ids, expected=13)
assert {str(row["request_id"]) for row in failure_rows} == failed_ids, failure_rows
assert all(row["status"] == "failure" for row in failure_rows), failure_rows
success_ids: Final = frozenset(_served_success_id(item) for item in normal)
assert len(success_ids) == 10, success_ids
success_rows: Final = _rows_for(success_ids, expected=10)
assert {str(row["request_id"]) for row in success_rows} == success_ids, success_rows
assert all(row["status"] == "success" for row in success_rows), success_rows
assert len(wire.drain()) == 30
def _owned_config(wire: Wire, directory: Path) -> Path:
base: Final = _JSON.validate_python(yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()))
config: Final[dict[str, JsonValue]] = {
**base,
"model_list": [
{
"model_name": _PLAIN,
"litellm_params": {
"model": _CONVERSE_MODEL,
"api_base": wire.url,
"api_key": "integration-provider-key",
**_AWS,
},
}
],
"router_settings": {**_JSON.validate_python(base["router_settings"]), "num_retries": 0},
}
path: Final = directory / "bedrock-event-frames.yaml"
path.write_text(yaml.safe_dump(config))
return path
def _open_peer_connections(pid: int, peer_url: str) -> int:
port: Final = urlsplit(peer_url).port
return sum(
1
for connection in psutil.Process(pid).net_connections(kind="tcp")
if connection.status == psutil.CONN_ESTABLISHED and connection.raddr and connection.raddr.port == port
)
def _held_rejection(release: threading.Event, held_indexes: SimpleQueue[int]) -> Callable[[Request], Reply]:
first: Final = _frame("messageStart", {"role": "assistant"})
def held(request: Request) -> Reply:
held_indexes.put(_call_index(request))
return Reply(content_type=_EVENT_STREAM, chunks=(first, _VALIDATION_FRAME), gate_after_first=release)
return held
@pytest.mark.timeout(600)
async def test_c02_worker_sigkill_mid_burst_leaves_the_sibling_rejecting_the_held_streams(
gateway: Gateway, tmp_path: Path
) -> None:
marker: Final = f"call-{uuid.uuid4().hex}"
again: Final = f"call-{uuid.uuid4().hex}"
calls: Final = _calls(marker, "chat", range(20))
release: Final = threading.Event()
held_indexes: Final[SimpleQueue[int]] = SimpleQueue()
with wire_server(_held_rejection(release, held_indexes)) as wire:
config: Final = _owned_config(wire, tmp_path)
with owned_proxy_process(gateway, tmp_path, {}, config=config, workers=2) as owned:
candidate: Final = owned.gateway
workers: Final = eventually(
lambda: tuple(int(pid) for pid in _STARTED_WORKER.findall(owned.log.read_text())),
lambda pids: len(pids) == 2,
seconds=30,
)
burst: Final = asyncio.create_task(
_burst(_proxy_url(candidate), candidate.key, _PLAIN, calls, tolerate_transport_errors=True)
)
await asyncio.to_thread(eventually, held_indexes.qsize, lambda size: size == 20, 60)
held_by: Final = MappingProxyType({pid: _open_peer_connections(pid, wire.url) for pid in workers})
assert sum(held_by.values()) == 20, held_by
victim_pid, survivor_pid = sorted(workers, key=held_by.__getitem__)
psutil.Process(victim_pid).send_signal(signal.SIGKILL)
release.set()
served: Final = await burst
assert held_by[survivor_pid] >= 10, held_by
assert len(served) == held_by[survivor_pid], (held_by, len(served))
for item in served:
assert item.status in (200, 400) and _is_rejection(item, _REJECTION), (
item.call,
item.status,
item.text,
)
eventually(
lambda: len(_STARTED_WORKER.findall(owned.log.read_text())), lambda count: count == 3, seconds=60
)
follow_up: Final = await _burst(
_proxy_url(candidate), candidate.key, _PLAIN, _calls(again, "chat", range(6))
)
assert len(follow_up) == 6
for item in follow_up:
assert item.status in (200, 400) and _is_rejection(item, _REJECTION), (
item.call,
item.status,
item.text,
)
assert len(wire.drain()) == 26
served_ids: Final = frozenset(item.call_id for item in (*served, *follow_up))
assert len(served_ids) == len(served) + 6, served_ids
rows: Final = _rows_for(served_ids, expected=len(served_ids))
assert {str(row["request_id"]) for row in rows} == served_ids, rows
assert all(row["status"] == "failure" for row in rows), rows
@pytest.mark.timeout(600)
async def test_c03_proxy_terminated_mid_burst_lands_every_served_rejection_at_most_once(
gateway: Gateway, tmp_path: Path
) -> None:
marker: Final = f"call-{uuid.uuid4().hex}"
calls: Final = _calls(marker, "chat", range(12))
release: Final = threading.Event()
held_indexes: Final[SimpleQueue[int]] = SimpleQueue()
with wire_server(_held_rejection(release, held_indexes)) as wire:
config: Final = _owned_config(wire, tmp_path)
with owned_proxy_process(gateway, tmp_path, {}, config=config, workers=1) as owned:
candidate: Final = owned.gateway
burst: Final = asyncio.create_task(
_burst(_proxy_url(candidate), candidate.key, _PLAIN, calls, tolerate_transport_errors=True)
)
await asyncio.to_thread(eventually, held_indexes.qsize, lambda size: size == 12, 60)
owned.process.terminate()
release.set()
served: Final = await burst
eventually(owned.process.poll, lambda code: code is not None, seconds=60)
assert len(served) <= 12
for item in served:
assert _is_rejection(item, _REJECTION), (item.call, item.status, item.text)
served_ids: Final = frozenset(item.call_id for item in served if item.call_id)
landed: Final = tuple(str(row["request_id"]) for row in _rows_for(served_ids, expected=0))
assert len(landed) == len(set(landed)), landed
assert set(landed) <= served_ids, (landed, served_ids)
assert len(wire.drain()) == 12

View file

@ -0,0 +1,100 @@
import re
from collections.abc import Callable, Mapping
from dataclasses import dataclass
from typing import Final
from urllib.parse import unquote
import pytest
from integration._support.upstream import _aws_event_frame
from integration._support.wire import Reply, Request, Wire, wire_server
from pydantic import JsonValue
import litellm
from litellm.exceptions import BadGatewayError, BadRequestError, MidStreamFallbackError
from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper
_MODEL_ID: Final = "global.moonshotai.kimi-k3"
_CONVERSE_MODEL: Final = f"bedrock/converse/{_MODEL_ID}"
_STREAM_TARGET: Final = f"/model/{_MODEL_ID}/converse-stream"
_EVENT_STREAM: Final = "application/vnd.amazon.eventstream"
_ANSWER: Final = "bedrock sync decoder control"
_REJECTION: Final = "structured output schema uses unsupported regex negative look-ahead"
_UNKNOWN_TYPE: Final = "somethingBedrockAddedLater"
_USAGE: Final[dict[str, JsonValue]] = {"usage": {"inputTokens": 11, "outputTokens": 4, "totalTokens": 15}}
def _frame(event_type: str, payload: Mapping[str, JsonValue]) -> bytes:
return _aws_event_frame(event_type, payload, "sc", "u")
_NORMAL: Final = b"".join(
(
_frame("messageStart", {"role": "assistant"}),
_frame("contentBlockDelta", {"delta": {"text": _ANSWER}, "contentBlockIndex": 0}),
_frame("contentBlockStop", {"contentBlockIndex": 0}),
_frame("messageStop", {"stopReason": "end_turn"}),
_frame("metadata", _USAGE),
)
)
_VALIDATION_FRAME: Final = _frame("validationException", {"message": _REJECTION})
_UNKNOWN_FRAME: Final = _frame(_UNKNOWN_TYPE, {"future": True})
@dataclass(frozen=True, slots=True)
class _Consumed:
text: str
finish_reasons: tuple[str | None, ...]
def _peer(frames: bytes) -> Callable[[Request], Reply]:
def respond(request: Request) -> Reply:
assert unquote(request.target) == _STREAM_TARGET, request.target
return Reply(body=frames, content_type=_EVENT_STREAM)
return respond
def _consume_sync_stream(wire: Wire) -> _Consumed:
response: Final = litellm.completion(
model=_CONVERSE_MODEL,
messages=[{"role": "user", "content": "What does the sync decoder do with this stream?"}],
max_tokens=16,
stream=True,
api_base=wire.url,
aws_access_key_id="AKIASCRIPTEDPROVIDER",
aws_secret_access_key="scripted-secret",
aws_region_name="us-east-1",
num_retries=0,
)
assert isinstance(response, CustomStreamWrapper), type(response)
chunks: Final = tuple(response)
return _Consumed(
text="".join(str(chunk.choices[0].delta.content or "") for chunk in chunks),
finish_reasons=tuple(chunk.choices[0].finish_reason for chunk in chunks),
)
def test_k01_sync_stream_with_a_validation_exception_frame_first_raises_a_400() -> None:
with wire_server(_peer(_VALIDATION_FRAME)) as wire:
with pytest.raises(BadRequestError, match=re.escape(_REJECTION)) as raised:
_consume_sync_stream(wire)
assert raised.value.status_code == 400, raised.value
assert len(wire.drain()) == 1
def test_k02_sync_stream_whose_only_frame_has_an_unknown_event_type_raises_a_502_naming_it() -> None:
with wire_server(_peer(_UNKNOWN_FRAME)) as wire:
with pytest.raises(MidStreamFallbackError, match=re.escape(_UNKNOWN_TYPE)) as raised:
_consume_sync_stream(wire)
assert raised.value.status_code == 502, raised.value
assert isinstance(raised.value.original_exception, BadGatewayError), raised.value.original_exception
assert "none of its 1 events carried a known event type" in str(raised.value), raised.value
assert len(wire.drain()) == 1
def test_k03_sync_stream_with_normal_frames_delivers_the_text_and_a_stop() -> None:
with wire_server(_peer(_NORMAL)) as wire:
consumed: Final = _consume_sync_stream(wire)
assert consumed.text == _ANSWER, consumed
assert consumed.finish_reasons[-1] == "stop", consumed
assert len(wire.drain()) == 1

View file

@ -3,6 +3,7 @@ import binascii
import itertools
import datetime
import json
import re
import struct
from collections.abc import AsyncIterator, Mapping, Sequence
from typing import Final
@ -21,7 +22,7 @@ from litellm.llms.bedrock.chat.invoke_handler import (
make_sync_call,
)
from litellm.exceptions import MidStreamFallbackError
from litellm.llms.bedrock.common_utils import BedrockError
from litellm.llms.bedrock.common_utils import BedrockError, get_bedrock_stream_event_statuses
from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler
from litellm.types.utils import ModelResponseStream
@ -717,19 +718,28 @@ async def test_async_invoke_streaming_non_200_forwards_bedrock_response_headers(
assert exc_info.value.response.headers["x-amzn-requestid"] == "req-non200-async"
def _bedrock_event_stream_frame(chunk: Mapping[str, object]) -> bytes:
def _event_stream_frame(event_type: str, payload: bytes) -> bytes:
def header(name: str, value: str) -> bytes:
return bytes([len(name)]) + name.encode() + bytes([7]) + struct.pack(">H", len(value)) + value.encode()
headers: Final = header(":event-type", "chunk") + header(":content-type", "application/json") + header(
headers: Final = header(":event-type", event_type) + header(":content-type", "application/json") + header(
":message-type", "event"
)
payload: Final = json.dumps({"bytes": base64.b64encode(json.dumps(chunk).encode()).decode()}).encode()
prelude: Final = struct.pack(">II", 12 + len(headers) + len(payload) + 4, len(headers))
body: Final = prelude + struct.pack(">I", binascii.crc32(prelude)) + headers + payload
return body + struct.pack(">I", binascii.crc32(body))
def _bedrock_event_stream_frame(chunk: Mapping[str, object]) -> bytes:
return _event_stream_frame(
"chunk", json.dumps({"bytes": base64.b64encode(json.dumps(chunk).encode()).decode()}).encode()
)
def _converse_event_frame(event_type: str, body: Mapping[str, object]) -> bytes:
return _event_stream_frame(event_type, json.dumps(body).encode())
def _openai_stream_chunk(delta: Mapping[str, str], finish_reason: str | None = None) -> Mapping[str, object]:
return {
"id": "chatcmpl-1",
@ -925,3 +935,193 @@ async def test_async_converse_stream_with_an_empty_200_body_raises_instead_of_an
_ = [chunk async for chunk in stream]
_assert_empty_stream_surfaced_as_bad_gateway(exc_info.value)
_UPSTREAM_REJECTION: Final = "structured output schema uses unsupported regex negative look-ahead"
_CUSTOMER_REJECTION_EVENT_TYPE: Final = "validationException"
def _modeled_exception_event_types() -> tuple[str, ...]:
statuses: Final = get_bedrock_stream_event_statuses()
assert statuses is not None
return tuple(sorted(name for name, status in statuses.items() if status is not None))
def _modeled_status(event_type: str) -> int:
statuses: Final = get_bedrock_stream_event_statuses()
assert statuses is not None
status: Final = statuses[event_type]
assert status is not None
return status
_CONVERSE_CONTENT_FRAMES: Final = (
_converse_event_frame("messageStart", {"role": "assistant"}),
_converse_event_frame("contentBlockDelta", {"contentBlockIndex": 0, "delta": {"text": "hi"}}),
_converse_event_frame("contentBlockStop", {"contentBlockIndex": 0}),
_converse_event_frame("messageStop", {"stopReason": "end_turn"}),
)
def _unknown_event_frame() -> bytes:
return _converse_event_frame("somethingBedrockAddedLater", {"message": _UPSTREAM_REJECTION})
@pytest.mark.parametrize("event_type", _modeled_exception_event_types())
def test_iter_bytes_raises_the_modeled_error_for_an_exception_named_event_frame(event_type: str) -> None:
decoder: Final = AWSEventStreamDecoder(model="us.moonshotai.kimi-k3")
frame: Final = _converse_event_frame(event_type, {"message": _UPSTREAM_REJECTION})
with pytest.raises(BedrockError) as exc_info:
list(decoder.iter_bytes(iter([frame]), response_headers=_event_stream_headers()))
assert exc_info.value.status_code == _modeled_status(event_type)
assert exc_info.value.status_code != 200
assert exc_info.value.message.startswith(event_type)
assert _UPSTREAM_REJECTION in exc_info.value.message
@pytest.mark.asyncio
async def test_aiter_bytes_raises_the_modeled_error_for_an_exception_named_event_frame() -> None:
event_type: Final = _CUSTOMER_REJECTION_EVENT_TYPE
async def _chunks() -> AsyncIterator[bytes]:
yield _converse_event_frame(event_type, {"message": _UPSTREAM_REJECTION})
decoder: Final = AWSEventStreamDecoder(model="us.moonshotai.kimi-k3")
with pytest.raises(BedrockError) as exc_info:
_ = [chunk async for chunk in decoder.aiter_bytes(_chunks(), response_headers=_event_stream_headers())]
assert exc_info.value.status_code == _modeled_status(event_type)
assert _UPSTREAM_REJECTION in exc_info.value.message
def _assert_unknown_event_stream_error(error: BedrockError, body: bytes) -> None:
assert error.status_code == 502
assert "HTTP 200" in error.message
assert "none of its 1 events carried a known event type" in error.message
assert "somethingBedrockAddedLater" in error.message
assert _UPSTREAM_REJECTION in error.message
assert f"{len(body)} bytes received" in error.message
assert "req-empty-1" in error.message
def test_iter_bytes_raises_when_no_event_carries_a_known_event_type() -> None:
decoder: Final = AWSEventStreamDecoder(model="us.moonshotai.kimi-k3")
body: Final = _unknown_event_frame()
with pytest.raises(BedrockError) as exc_info:
list(decoder.iter_bytes(iter([body]), response_headers=_event_stream_headers()))
_assert_unknown_event_stream_error(exc_info.value, body)
@pytest.mark.asyncio
async def test_aiter_bytes_raises_when_no_event_carries_a_known_event_type() -> None:
body: Final = _unknown_event_frame()
async def _chunks() -> AsyncIterator[bytes]:
yield body
decoder: Final = AWSEventStreamDecoder(model="us.moonshotai.kimi-k3")
with pytest.raises(BedrockError) as exc_info:
_ = [chunk async for chunk in decoder.aiter_bytes(_chunks(), response_headers=_event_stream_headers())]
_assert_unknown_event_stream_error(exc_info.value, body)
def test_iter_bytes_keeps_a_stream_whose_unknown_event_sits_beside_known_frames() -> None:
decoder: Final = AWSEventStreamDecoder(model="us.moonshotai.kimi-k3")
frames: Final = (_CONVERSE_CONTENT_FRAMES[0], _unknown_event_frame(), *_CONVERSE_CONTENT_FRAMES[1:])
chunks: Final = list(decoder.iter_bytes(iter(frames), response_headers=_event_stream_headers()))
texts: Final = [chunk.choices[0].delta.content for chunk in chunks if isinstance(chunk, ModelResponseStream)]
assert "".join(text or "" for text in texts) == "hi"
finish_reasons: Final = [
chunk.choices[0].finish_reason for chunk in chunks if isinstance(chunk, ModelResponseStream)
]
assert "stop" in finish_reasons
def _assert_exception_event_surfaced_with_its_modeled_status(error: BaseException, event_type: str) -> None:
assert not isinstance(error, litellm.BadGatewayError)
assert getattr(error, "status_code", None) == _modeled_status(event_type)
assert event_type in str(error)
assert _UPSTREAM_REJECTION in str(error)
def test_converse_stream_with_an_exception_event_frame_raises_instead_of_an_empty_turn(
_aws_test_credentials: None,
) -> None:
event_type: Final = _CUSTOMER_REJECTION_EVENT_TYPE
frame: Final = _converse_event_frame(event_type, {"message": _UPSTREAM_REJECTION})
response: Final = MagicMock(status_code=200, headers=_event_stream_headers())
response.iter_bytes = lambda chunk_size=None: iter([frame])
client: Final = HTTPHandler()
client.post = MagicMock(return_value=response)
with pytest.raises(Exception, match=re.escape(_UPSTREAM_REJECTION)) as exc_info:
list(
litellm.completion(
model="bedrock/us.moonshotai.kimi-k3",
messages=[{"role": "user", "content": "hi"}],
stream=True,
client=client,
)
)
_assert_exception_event_surfaced_with_its_modeled_status(exc_info.value, event_type)
@pytest.mark.asyncio
async def test_async_converse_stream_with_an_exception_event_frame_raises_instead_of_an_empty_turn(
_aws_test_credentials: None,
) -> None:
event_type: Final = _CUSTOMER_REJECTION_EVENT_TYPE
async def _aiter_bytes(chunk_size: int | None = None) -> AsyncIterator[bytes]:
yield _converse_event_frame(event_type, {"message": _UPSTREAM_REJECTION})
response: Final = MagicMock(status_code=200, headers=_event_stream_headers())
response.aiter_bytes = _aiter_bytes
client: Final = AsyncHTTPHandler()
client.post = AsyncMock(return_value=response)
stream: Final = await litellm.acompletion(
model="bedrock/us.moonshotai.kimi-k3",
messages=[{"role": "user", "content": "hi"}],
stream=True,
client=client,
)
with pytest.raises(Exception, match=re.escape(_UPSTREAM_REJECTION)) as exc_info:
_ = [chunk async for chunk in stream]
_assert_exception_event_surfaced_with_its_modeled_status(exc_info.value, event_type)
def test_converse_stream_made_only_of_unknown_events_raises_instead_of_an_empty_turn(
_aws_test_credentials: None,
) -> None:
response: Final = MagicMock(status_code=200, headers=_event_stream_headers())
response.iter_bytes = lambda chunk_size=None: iter([_unknown_event_frame()])
client: Final = HTTPHandler()
client.post = MagicMock(return_value=response)
with pytest.raises(MidStreamFallbackError) as exc_info:
list(
litellm.completion(
model="bedrock/us.moonshotai.kimi-k3",
messages=[{"role": "user", "content": "hi"}],
stream=True,
client=client,
)
)
assert exc_info.value.status_code == 502
assert exc_info.value.is_pre_first_chunk is True
assert isinstance(exc_info.value.original_exception, litellm.BadGatewayError)
assert "somethingBedrockAddedLater" in str(exc_info.value)
assert _UPSTREAM_REJECTION in str(exc_info.value)

View file

@ -981,3 +981,72 @@ def test_unmapped_openai_family_model_routes_to_converse():
assert BedrockModelInfo.get_bedrock_route(unmapped) == "converse"
imported: Final = "bedrock/openai/arn:aws:bedrock:us-east-1:123456789012:imported-model/abc123"
assert BedrockModelInfo.get_bedrock_route(imported) == "openai"
def test_bedrock_stream_event_statuses_cover_every_modeled_member_of_both_stream_shapes():
pytest.importorskip("botocore")
from botocore.loaders import Loader
from botocore.model import ServiceModel
import litellm.llms.bedrock.common_utils as mod
mod.get_bedrock_stream_event_statuses.cache_clear()
statuses = mod.get_bedrock_stream_event_statuses()
assert statuses is not None
service_model = ServiceModel(Loader().load_service_model("bedrock-runtime", "service-2"))
for shape_name in ("ConverseStreamOutput", "ResponseStream"):
for name, member in service_model.shape_for(shape_name).members.items():
modeled = (member.metadata or {}).get("error", {}).get("httpStatusCode")
assert statuses[name] == (None if modeled is None else int(modeled))
assert mod.bedrock_stream_event_error_status(name) == statuses[name]
assert any(status is not None for status in statuses.values())
assert any(status is None for status in statuses.values())
assert mod.bedrock_stream_event_error_status("notAModeledEvent") is None
assert mod.bedrock_stream_event_error_status(None) is None
def test_bedrock_stream_event_statuses_load_failure_returns_none():
from unittest.mock import patch
import litellm.llms.bedrock.common_utils as mod
pytest.importorskip("botocore")
mod.get_bedrock_stream_event_statuses.cache_clear()
with patch("botocore.loaders.Loader.load_service_model", side_effect=Exception("no data")):
assert mod._load_bedrock_stream_event_statuses() is None
assert mod.get_bedrock_stream_event_statuses() is None
assert mod.bedrock_stream_event_error_status("validationException") is None
mod.get_bedrock_stream_event_statuses.cache_clear()
@pytest.mark.parametrize(
("headers", "expected_status", "expected_message"),
[
({":message-type": "error"}, 400, '{"message":"upstream failed"}'),
(
{":message-type": "exception", ":exception-type": "somethingNotModeled"},
400,
'somethingNotModeled {"message":"upstream failed"}',
),
(
{":message-type": "exception", ":exception-type": "throttlingException"},
429,
'throttlingException {"message":"upstream failed"}',
),
],
)
def test_build_bedrock_stream_error_resolves_status_from_the_exception_type(
headers: dict[str, str], expected_status: int, expected_message: str
):
pytest.importorskip("botocore")
from litellm.llms.bedrock.common_utils import build_bedrock_stream_error, get_bedrock_response_stream_shape
error = build_bedrock_stream_error(
{"status_code": 400, "headers": headers, "body": b'{"message":"upstream failed"}'},
get_bedrock_response_stream_shape(),
)
assert error.status_code == expected_status
assert error.message == expected_message