mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
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:
parent
aef0a53837
commit
8b4de39ad7
7 changed files with 1494 additions and 57 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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 = (
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
@ -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
|
||||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue