diff --git a/litellm/llms/base_llm/chat/transformation.py b/litellm/llms/base_llm/chat/transformation.py index f1b41a2302d..3885e4ad38a 100644 --- a/litellm/llms/base_llm/chat/transformation.py +++ b/litellm/llms/base_llm/chat/transformation.py @@ -395,6 +395,7 @@ class BaseConfig(ABC): signed_json_body: bytes | None = None, *, litellm_params: Mapping[str, object], + timeout: float | httpx.Timeout | None = None, ) -> "CustomStreamWrapper": raise NotImplementedError @@ -412,6 +413,7 @@ class BaseConfig(ABC): signed_json_body: bytes | None = None, *, litellm_params: Mapping[str, object], + timeout: float | httpx.Timeout | None = None, ) -> "CustomStreamWrapper": raise NotImplementedError diff --git a/litellm/llms/bedrock/chat/agentcore/transformation.py b/litellm/llms/bedrock/chat/agentcore/transformation.py index 2ad23e84e9f..27102f4289f 100644 --- a/litellm/llms/bedrock/chat/agentcore/transformation.py +++ b/litellm/llms/bedrock/chat/agentcore/transformation.py @@ -645,6 +645,7 @@ class AmazonAgentCoreConfig(BaseConfig, BaseAWSLLM): signed_json_body: bytes | None = None, *, litellm_params: Mapping[str, object], + timeout: float | httpx.Timeout | None = None, ) -> "CustomStreamWrapper": """ Simplified sync streaming - returns a generator that yields ModelResponse chunks. @@ -866,6 +867,7 @@ class AmazonAgentCoreConfig(BaseConfig, BaseAWSLLM): signed_json_body: bytes | None = None, *, litellm_params: Mapping[str, object], + timeout: float | httpx.Timeout | None = None, ) -> "CustomStreamWrapper": """ Simplified async streaming - returns an async generator that yields ModelResponse chunks. diff --git a/litellm/llms/bedrock/chat/converse_handler.py b/litellm/llms/bedrock/chat/converse_handler.py index 28df1af4bc0..d54fb1b804f 100644 --- a/litellm/llms/bedrock/chat/converse_handler.py +++ b/litellm/llms/bedrock/chat/converse_handler.py @@ -34,6 +34,7 @@ def make_sync_call( json_mode: bool | None = False, fake_stream: bool = False, stream_chunk_size: int | None = None, + timeout: float | httpx.Timeout | None = None, ) -> tuple[Any, httpx.Headers]: if client is None: client = _get_httpx_client() # Create a new client if none provided @@ -44,6 +45,7 @@ def make_sync_call( data=data, stream=not fake_stream, logging_obj=logging_obj, + timeout=timeout, ) if response.status_code != 200: @@ -152,6 +154,7 @@ class BedrockConverseLLM(BaseAWSLLM): fake_stream=fake_stream, json_mode=json_mode, stream_chunk_size=stream_chunk_size, + timeout=timeout, ) streaming_response: Final = CustomStreamWrapper( completion_stream=completion_stream, @@ -456,6 +459,7 @@ class BedrockConverseLLM(BaseAWSLLM): json_mode=json_mode, fake_stream=fake_stream, stream_chunk_size=stream_chunk_size, + timeout=timeout, ) streaming_response: Final = CustomStreamWrapper( completion_stream=completion_stream, diff --git a/litellm/llms/bedrock/chat/invoke_handler.py b/litellm/llms/bedrock/chat/invoke_handler.py index 9f3d27cbfb9..8a2291d33be 100644 --- a/litellm/llms/bedrock/chat/invoke_handler.py +++ b/litellm/llms/bedrock/chat/invoke_handler.py @@ -201,6 +201,7 @@ async def make_call( json_mode: bool | None = False, bedrock_invoke_provider: litellm.BEDROCK_INVOKE_PROVIDERS_LITERAL | None = None, stream_chunk_size: int | None = None, + timeout: float | httpx.Timeout | None = None, ) -> "tuple[MockResponseIterator | AsyncIterator[GChunk | ModelResponseStream | dict], httpx.Headers]": try: if client is None: @@ -219,6 +220,7 @@ async def make_call( data=data, stream=not fake_stream, logging_obj=logging_obj, + timeout=timeout, ) if response.status_code != 200: @@ -291,6 +293,7 @@ def make_sync_call( json_mode: bool | None = False, bedrock_invoke_provider: litellm.BEDROCK_INVOKE_PROVIDERS_LITERAL | None = None, stream_chunk_size: int | None = None, + timeout: float | httpx.Timeout | None = None, ) -> "tuple[MockResponseIterator | Iterator[GChunk | ModelResponseStream | dict], httpx.Headers]": try: if client is None: @@ -308,6 +311,7 @@ def make_sync_call( data=signed_json_body if signed_json_body is not None else data, stream=not fake_stream, logging_obj=logging_obj, + timeout=timeout, ) if response.status_code != 200: diff --git a/litellm/llms/bedrock/chat/invoke_transformations/base_invoke_transformation.py b/litellm/llms/bedrock/chat/invoke_transformations/base_invoke_transformation.py index 9baf8110b4e..a1fa7963ab5 100644 --- a/litellm/llms/bedrock/chat/invoke_transformations/base_invoke_transformation.py +++ b/litellm/llms/bedrock/chat/invoke_transformations/base_invoke_transformation.py @@ -455,6 +455,7 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM): signed_json_body: bytes | None = None, *, litellm_params: Mapping[str, object], + timeout: float | httpx.Timeout | None = None, ) -> CustomStreamWrapper: chunk_size: Final = stored_control_options(litellm_params).stream_chunk_size completion_stream, response_headers = await make_call( @@ -469,6 +470,7 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM): bedrock_invoke_provider=self.get_bedrock_invoke_provider(model), json_mode=json_mode, stream_chunk_size=chunk_size, + timeout=timeout, ) streaming_response: Final = CustomStreamWrapper( completion_stream=completion_stream, @@ -494,6 +496,7 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM): signed_json_body: bytes | None = None, *, litellm_params: Mapping[str, object], + timeout: float | httpx.Timeout | None = None, ) -> CustomStreamWrapper: sync_client: Final = ( _get_httpx_client(params={}) if client is None or isinstance(client, AsyncHTTPHandler) else client @@ -512,6 +515,7 @@ class AmazonInvokeConfig(BaseConfig, BaseAWSLLM): bedrock_invoke_provider=self.get_bedrock_invoke_provider(model), json_mode=json_mode, stream_chunk_size=chunk_size, + timeout=timeout, ) streaming_response: Final = CustomStreamWrapper( completion_stream=completion_stream, diff --git a/litellm/llms/bytez/chat/transformation.py b/litellm/llms/bytez/chat/transformation.py index 5846ba560a8..12338d83017 100644 --- a/litellm/llms/bytez/chat/transformation.py +++ b/litellm/llms/bytez/chat/transformation.py @@ -261,6 +261,7 @@ class BytezChatConfig(BaseConfig): signed_json_body: bytes | None = None, *, litellm_params: Mapping[str, object], + timeout: float | httpx.Timeout | None = None, ) -> "BytezCustomStreamWrapper": if client is None or isinstance(client, AsyncHTTPHandler): client = _get_httpx_client(params={}) @@ -305,6 +306,7 @@ class BytezChatConfig(BaseConfig): signed_json_body: bytes | None = None, *, litellm_params: Mapping[str, object], + timeout: float | httpx.Timeout | None = None, ) -> "BytezCustomStreamWrapper": if client is None or isinstance(client, HTTPHandler): client = get_async_httpx_client(llm_provider=LlmProviders.BYTEZ, params={}) diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 1fc9ffaacb2..72bbb4f9556 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -814,6 +814,7 @@ class BaseLLMHTTPHandler: client=client, json_mode=json_mode, litellm_params=litellm_params, + timeout=timeout, ) completion_stream, headers = self.make_sync_call( provider_config=provider_config, @@ -978,6 +979,7 @@ class BaseLLMHTTPHandler: json_mode=json_mode, signed_json_body=signed_json_body, litellm_params=litellm_params, + timeout=timeout, ) completion_stream, _response_headers = await self.make_async_call_stream_helper( diff --git a/litellm/llms/langgraph/chat/transformation.py b/litellm/llms/langgraph/chat/transformation.py index 293672f1ca9..07c90a73a9d 100644 --- a/litellm/llms/langgraph/chat/transformation.py +++ b/litellm/llms/langgraph/chat/transformation.py @@ -288,6 +288,7 @@ class LangGraphConfig(BaseConfig): signed_json_body: bytes | None = None, *, litellm_params: Mapping[str, object], + timeout: float | httpx.Timeout | None = None, ) -> CustomStreamWrapper: """ Get a CustomStreamWrapper for synchronous streaming. @@ -349,6 +350,7 @@ class LangGraphConfig(BaseConfig): signed_json_body: bytes | None = None, *, litellm_params: Mapping[str, object], + timeout: float | httpx.Timeout | None = None, ) -> CustomStreamWrapper: """ Get a CustomStreamWrapper for asynchronous streaming. diff --git a/litellm/llms/oci/chat/transformation.py b/litellm/llms/oci/chat/transformation.py index 24f3ddd5162..01b39ba3c40 100644 --- a/litellm/llms/oci/chat/transformation.py +++ b/litellm/llms/oci/chat/transformation.py @@ -644,6 +644,7 @@ class OCIChatConfig(BaseConfig): signed_json_body: bytes | None = None, *, litellm_params: Mapping[str, object], + timeout: float | httpx.Timeout | None = None, ) -> "OCIStreamWrapper": if client is None or isinstance(client, AsyncHTTPHandler): client = _get_httpx_client(params={}) @@ -685,6 +686,7 @@ class OCIChatConfig(BaseConfig): signed_json_body: bytes | None = None, *, litellm_params: Mapping[str, object], + timeout: float | httpx.Timeout | None = None, ) -> "OCIStreamWrapper": if client is None or isinstance(client, HTTPHandler): client = get_async_httpx_client(llm_provider=LlmProviders.OCI, params={}) diff --git a/litellm/llms/sagemaker/chat/transformation.py b/litellm/llms/sagemaker/chat/transformation.py index f99a3f9e1bc..8f320343de1 100644 --- a/litellm/llms/sagemaker/chat/transformation.py +++ b/litellm/llms/sagemaker/chat/transformation.py @@ -152,6 +152,7 @@ class SagemakerChatConfig(OpenAIGPTConfig, BaseAWSLLM): signed_json_body: bytes | None = None, *, litellm_params: Mapping[str, object], + timeout: float | httpx.Timeout | None = None, ) -> CustomStreamWrapper: if client is None or isinstance(client, AsyncHTTPHandler): client = _get_httpx_client(params={}) @@ -196,6 +197,7 @@ class SagemakerChatConfig(OpenAIGPTConfig, BaseAWSLLM): signed_json_body: bytes | None = None, *, litellm_params: Mapping[str, object], + timeout: float | httpx.Timeout | None = None, ) -> CustomStreamWrapper: if client is None or isinstance(client, HTTPHandler): try: diff --git a/litellm/llms/vertex_ai/agent_engine/transformation.py b/litellm/llms/vertex_ai/agent_engine/transformation.py index cf889c481a3..fd1ff24fd52 100644 --- a/litellm/llms/vertex_ai/agent_engine/transformation.py +++ b/litellm/llms/vertex_ai/agent_engine/transformation.py @@ -368,6 +368,7 @@ class VertexAgentEngineConfig(BaseConfig, VertexBase): signed_json_body: bytes | None = None, *, litellm_params: Mapping[str, object], + timeout: float | httpx.Timeout | None = None, ) -> "CustomStreamWrapper": """Get a CustomStreamWrapper for synchronous streaming.""" from litellm.llms.custom_httpx.http_handler import ( @@ -428,6 +429,7 @@ class VertexAgentEngineConfig(BaseConfig, VertexBase): signed_json_body: bytes | None = None, *, litellm_params: Mapping[str, object], + timeout: float | httpx.Timeout | None = None, ) -> "CustomStreamWrapper": """Get a CustomStreamWrapper for asynchronous streaming.""" from litellm.llms.custom_httpx.http_handler import ( diff --git a/tests/integration/_support/wire.py b/tests/integration/_support/wire.py index 1201a156c00..c807852ef6b 100644 --- a/tests/integration/_support/wire.py +++ b/tests/integration/_support/wire.py @@ -30,6 +30,7 @@ class Reply: gate_after_first: threading.Event | None = None pause_between_chunks: float = 0 headers: Mapping[str, str] = MappingProxyType({}) + drop_connection: bool = False @dataclass(frozen=True, slots=True) @@ -81,6 +82,10 @@ def wire_server( except Exception as error: errors.put(error) reply = Reply(status=500) + if reply.drop_connection: + self.close_connection = True + disconnected.put(request.target) + return self.send_response(reply.status) self.send_header("content-type", reply.content_type) for name, value in reply.headers.items(): diff --git a/tests/integration/providers/test_bedrock_stream_timeout_chaos.py b/tests/integration/providers/test_bedrock_stream_timeout_chaos.py new file mode 100644 index 00000000000..2632b81b4dc --- /dev/null +++ b/tests/integration/providers/test_bedrock_stream_timeout_chaos.py @@ -0,0 +1,388 @@ +import asyncio +import json +import re +import signal +import threading +import uuid +from collections.abc import Iterator, Mapping +from contextlib import contextmanager +from dataclasses import dataclass +from pathlib import Path +from typing import Final, Literal +from urllib.parse import unquote, urlsplit + +import httpx +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 +from integration._support.wire import Reply, Request, Wire, wire_server +from pydantic import JsonValue + +_CONVERSE_MODEL: Final = "bedrock/converse/global.moonshotai.kimi-k3" +_EVENT_STREAM: Final = "application/vnd.amazon.eventstream" +_ANSWER: Final = "bedrock stream timeout chaos control" +_PROMPT: Final = "How long does the gateway wait for this stream?" +_USER_TURN: Final[dict[str, JsonValue]] = {"role": "user", "content": _PROMPT} +_TIMEOUT_SECONDS: Final = 1 +_KILL_TIMEOUT_SECONDS: Final = 10 +_LONG_TIMEOUT_SECONDS: Final = 30 +_HOLD_SECONDS: Final = 120.0 +_CLIENT_WINDOW: Final = 10.0 +_SHORT_WINDOW: Final = 3.0 +_BURST_WINDOW: Final = 40.0 +_BARE_MODEL: Final = "bedrock-stream-timeout-bare" +_SLOW_MODEL: Final = "bedrock-stream-timeout-thirty" +_KILL_MODEL: Final = "bedrock-stream-timeout-ten" +_DEPLOYMENTS: Final = ((_BARE_MODEL, None), (_SLOW_MODEL, _LONG_TIMEOUT_SECONDS), (_KILL_MODEL, _KILL_TIMEOUT_SECONDS)) +_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}} +_TIMEOUT_PASSED: Final = re.compile(r"Timeout passed=(?:Timeout\(timeout=)?(-?\d+\.\d+)") +_STARTED_WORKER: Final = re.compile(r"Started server process \[(\d+)\]") +_USAGE: Final[dict[str, JsonValue]] = {"usage": {"inputTokens": 11, "outputTokens": 4, "totalTokens": 15}} +_CONVERSE_RESPONSE: Final = json.dumps( + { + "output": {"message": {"role": "assistant", "content": [{"text": _ANSWER}]}}, + "stopReason": "end_turn", + **_USAGE, + "metrics": {"latencyMs": 1}, + } +).encode() + +Endpoint = Literal["chat", "messages", "responses"] +_ENDPOINTS: Final[tuple[Endpoint, ...]] = ("chat", "messages", "responses") + + +def _frame(event_type: str, payload: Mapping[str, JsonValue]) -> bytes: + return _aws_event_frame(event_type, payload, "sc", "u") + + +_CONVERSE_FRAMES: 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), + ) +) + + +@dataclass(frozen=True, slots=True) +class _Peer: + wire: Wire + hold: threading.Event + drop: threading.Event + released: threading.Event + + def respond(self, request: Request) -> Reply: + if self.hold.is_set(): + self.released.wait(timeout=_HOLD_SECONDS) + if self.drop.is_set(): + return Reply(drop_connection=True) + if unquote(request.target).endswith("converse-stream"): + return Reply(body=_CONVERSE_FRAMES, content_type=_EVENT_STREAM) + return Reply(body=_CONVERSE_RESPONSE) + + +@contextmanager +def _peer() -> Iterator[_Peer]: + hold: Final = threading.Event() + drop: Final = threading.Event() + released: Final = threading.Event() + + def respond(request: Request) -> Reply: + return _Peer(wire, hold, drop, released).respond(request) + + with wire_server(respond) as wire: + try: + yield _Peer(wire, hold, drop, released) + finally: + released.set() + + +@dataclass(frozen=True, slots=True) +class _Call: + endpoint: Endpoint + stream: bool = True + + +@dataclass(frozen=True, slots=True) +class _Served: + call: _Call + status: int + text: str + + +def _calls(count: int) -> tuple[_Call, ...]: + return tuple(_Call(_ENDPOINTS[index % len(_ENDPOINTS)]) for index in range(count)) + + +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(model: str, call: _Call) -> dict[str, JsonValue]: + match call.endpoint: + case "chat": + return {"model": model, "messages": [_USER_TURN], "stream": call.stream, **_EXTRA} + case "messages": + return {"model": model, "messages": [_USER_TURN], "max_tokens": 16, "stream": call.stream, **_EXTRA} + case "responses": + return {"model": model, "input": _PROMPT, "stream": call.stream, **_EXTRA} + + +def _proxy_url(gateway: Gateway) -> str: + return str(gateway.client.base_url).rstrip("/") + + +async def _send(client: httpx.AsyncClient, key: str, model: str, call: _Call) -> _Served: + response: Final = await client.post( + _path(call.endpoint), json=_body(model, call), headers={"Authorization": f"Bearer {key}"} + ) + return _Served(call, response.status_code, response.text) + + +async def _burst( + gateway: Gateway, model: str, calls: tuple[_Call, ...], *, window: float, tolerate_transport_errors: bool = False +) -> tuple[_Served, ...]: + async with httpx.AsyncClient(base_url=_proxy_url(gateway), timeout=window, trust_env=False) as client: + results: Final = await asyncio.gather( + *(_send(client, gateway.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)) + + +async def _one(gateway: Gateway, model: str, call: _Call, *, window: float = _CLIENT_WINDOW) -> _Served: + (served,) = await _burst(gateway, model, (call,), window=window) + return served + + +def _timeout_passed(text: str) -> str: + found: Final = _TIMEOUT_PASSED.search(text) + assert found is not None, text + return found.group(1) + + +def _assert_timed_out(served: _Served, seconds: float = _TIMEOUT_SECONDS) -> None: + assert served.status == 408, served.text + assert _timeout_passed(served.text) == f"{seconds:.1f}", served.text + assert _ANSWER not in served.text, served.text + + +def _assert_answered(served: _Served) -> None: + assert served.status == 200, served.text + assert _ANSWER in served.text, served.text + assert "Timeout" not in served.text, served.text + + +def _assert_cut_off(served: _Served) -> None: + assert served.status >= 500, served.text + assert "error" in served.text, served.text + assert _ANSWER not in served.text, served.text + + +async def _wait_for_received(peer: _Peer, expected: int) -> None: + await asyncio.to_thread(eventually, peer.wire.received.qsize, lambda size: size == expected, 60) + + +def _failure_rows(model: str, expected: int) -> None: + rows: Final = eventually( + lambda: read_rows('SELECT request_id, status FROM "LiteLLM_SpendLogs" WHERE model_group=%s', (model,)), + lambda found: len(found) >= expected, + seconds=60, + ) + assert [row["status"] for row in rows] == ["failure"] * expected, rows + assert len({row["request_id"] for row in rows}) == expected, rows + + +def _converse(scenario: Scenario, wire: Wire, **extra: JsonValue) -> str: + return scenario.model(model=_CONVERSE_MODEL, api_base=wire.url, api_key=None, **_AWS, **extra) + + +def _deployment(name: str, wire: Wire, timeout: int | None) -> dict[str, JsonValue]: + return { + "model_name": name, + "litellm_params": { + "model": _CONVERSE_MODEL, + "api_base": wire.url, + **_AWS, + **({} if timeout is None else {"timeout": timeout}), + }, + } + + +def _config(wire: Wire, directory: Path, *, router_timeout: int | None = None) -> Path: + base: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + router_settings: Final = { + **base["router_settings"], + **({} if router_timeout is None else {"timeout": router_timeout}), + } + config: Final = { + **base, + "model_list": [_deployment(name, wire, timeout) for name, timeout in _DEPLOYMENTS], + "router_settings": router_settings, + } + path: Final = directory / f"bedrock-stream-timeout-{uuid.uuid4().hex[:8]}.yaml" + path.write_text(yaml.safe_dump(config)) + return path + + +def _worker_pids(log: Path) -> tuple[int, ...]: + return eventually( + lambda: tuple(int(pid) for pid in _STARTED_WORKER.findall(log.read_text())), + lambda pids: len(pids) == 2, + seconds=30, + ) + + +def _open_upstream_connections(pid: int, upstream: str) -> int: + port: Final = urlsplit(upstream).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 + ) + + +@pytest.mark.timeout(300) +async def test_p7_p8_the_global_request_timeout_bounds_a_stream_only_when_the_deployment_sets_none( + gateway: Gateway, tmp_path: Path +) -> None: + with _peer() as peer: + peer.hold.set() + path: Final = _config(peer.wire, tmp_path) + with owned_proxy_process(gateway, tmp_path, {"REQUEST_TIMEOUT": "1"}, config=path, workers=2) as owned: + _assert_timed_out(await _one(owned.gateway, _BARE_MODEL, _Call("chat"))) + with pytest.raises(httpx.ReadTimeout): + await _one(owned.gateway, _SLOW_MODEL, _Call("messages"), window=_SHORT_WINDOW) + await _wait_for_received(peer, 2) + peer.released.set() + + +@pytest.mark.timeout(300) +async def test_p9_under_router_settings_timeout_the_global_request_timeout_precedes_the_deployment_for_streams( + gateway: Gateway, tmp_path: Path +) -> None: + with _peer() as peer: + peer.hold.set() + path: Final = _config(peer.wire, tmp_path, router_timeout=_LONG_TIMEOUT_SECONDS) + with owned_proxy_process(gateway, tmp_path, {"REQUEST_TIMEOUT": "1"}, config=path, workers=2) as owned: + _assert_timed_out(await _one(owned.gateway, _SLOW_MODEL, _Call("chat"))) + with pytest.raises(httpx.ReadTimeout): + await _one(owned.gateway, _SLOW_MODEL, _Call("chat", stream=False), window=_SHORT_WINDOW) + await _wait_for_received(peer, 2) + peer.released.set() + + +@pytest.mark.timeout(120) +async def test_x1_thirty_stalled_streams_across_every_endpoint_each_time_out_while_the_proxy_stays_live( + gateway: Gateway, +) -> None: + calls: Final = _calls(30) + with _peer() as peer, gateway.scenario() as scenario: + peer.hold.set() + model: Final = _converse(scenario, peer.wire, timeout=_TIMEOUT_SECONDS) + burst: Final = asyncio.create_task(_burst(gateway, model, calls, window=_CLIENT_WINDOW)) + await _wait_for_received(peer, 30) + async with httpx.AsyncClient(base_url=_proxy_url(gateway), timeout=2, trust_env=False) as client: + liveliness: Final = await client.get("/health/liveliness") + assert liveliness.status_code == 200, liveliness.text + served: Final = await burst + assert len(served) == 30 + for item in served: + _assert_timed_out(item) + assert len(peer.wire.drain()) == 30 + _failure_rows(model, 30) + + +@pytest.mark.timeout(120) +async def test_x2_an_upstream_dropping_every_held_connection_fails_each_stream_once_and_then_recovers( + gateway: Gateway, +) -> None: + calls: Final = _calls(20) + with _peer() as peer, gateway.scenario() as scenario: + peer.hold.set() + peer.drop.set() + model: Final = _converse(scenario, peer.wire, timeout=_LONG_TIMEOUT_SECONDS) + burst: Final = asyncio.create_task( + _burst(gateway, model, calls, window=_CLIENT_WINDOW, tolerate_transport_errors=True) + ) + await _wait_for_received(peer, 20) + peer.released.set() + await asyncio.to_thread(eventually, peer.wire.disconnected.qsize, lambda size: size == 20, 30) + served: Final = await burst + assert len(served) == 20 + for item in served: + _assert_cut_off(item) + peer.drop.clear() + peer.hold.clear() + _assert_answered(await _one(gateway, model, _Call("chat"))) + assert len(peer.wire.drain()) == 21 + + +@pytest.mark.timeout(360) +async def test_x3_killing_one_worker_mid_burst_leaves_the_sibling_timing_its_streams_out( + gateway: Gateway, tmp_path: Path +) -> None: + calls: Final = _calls(20) + with _peer() as peer: + peer.hold.set() + path: Final = _config(peer.wire, tmp_path) + with owned_proxy_process(gateway, tmp_path, {}, config=path, workers=2) as owned: + workers: Final = _worker_pids(owned.log) + burst: Final = asyncio.create_task( + _burst(owned.gateway, _KILL_MODEL, calls, window=_BURST_WINDOW, tolerate_transport_errors=True) + ) + await _wait_for_received(peer, 20) + held_by: Final = {pid: _open_upstream_connections(pid, peer.wire.url) for pid in workers} + assert sum(held_by.values()) == 20, held_by + victim_pid, survivor_pid = sorted(workers, key=held_by.__getitem__) + victim: Final = psutil.Process(victim_pid) + victim.suspend() + victim.send_signal(signal.SIGKILL) + 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_timed_out(item, _KILL_TIMEOUT_SECONDS) + peer.hold.clear() + _assert_answered(await _one(owned.gateway, _KILL_MODEL, _Call("chat"))) + assert len(peer.wire.drain()) == 21 + + +@pytest.mark.timeout(480) +async def test_x4_a_proxy_stopped_mid_burst_still_times_its_held_streams_out_and_replays_none( + gateway: Gateway, tmp_path: Path +) -> None: + calls: Final = _calls(20) + with _peer() as peer: + peer.hold.set() + path: Final = _config(peer.wire, tmp_path) + with owned_proxy_process(gateway, tmp_path, {}, config=path, workers=2) as first: + burst: Final = asyncio.create_task(_burst(first.gateway, _KILL_MODEL, calls, window=_BURST_WINDOW)) + await _wait_for_received(peer, 20) + served: Final = await burst + assert len(served) == 20 + for item in served: + _assert_timed_out(item, _KILL_TIMEOUT_SECONDS) + assert len(peer.wire.drain()) == 20 + peer.hold.clear() + with owned_proxy_process(gateway, tmp_path, {}, config=path, workers=2) as second: + _assert_answered(await _one(second.gateway, _KILL_MODEL, _Call("chat"))) + assert len(peer.wire.drain()) == 1 diff --git a/tests/integration/providers/test_bedrock_stream_timeout_wire.py b/tests/integration/providers/test_bedrock_stream_timeout_wire.py new file mode 100644 index 00000000000..e7a76052558 --- /dev/null +++ b/tests/integration/providers/test_bedrock_stream_timeout_wire.py @@ -0,0 +1,673 @@ +import base64 +import json +import re +import threading +from collections.abc import Iterator, Mapping +from contextlib import contextmanager +from dataclasses import dataclass +from typing import Final, Literal +from urllib.parse import unquote + +import anthropic +import httpx +import openai +import pytest +from integration._support.client import Gateway, Scenario, eventually, object_value, string_value +from integration._support.database import read_rows +from integration._support.upstream import _aws_event_frame +from integration._support.wire import Reply, Request, Wire, wire_server +from pydantic import JsonValue, TypeAdapter + +_CONVERSE_MODEL_ID: Final = "global.moonshotai.kimi-k3" +_CONVERSE_MODEL: Final = f"bedrock/converse/{_CONVERSE_MODEL_ID}" +_INVOKE_MODEL_ID: Final = "anthropic.claude-3-haiku-20240307-v1:0" +_INVOKE_MODEL: Final = f"bedrock/invoke/{_INVOKE_MODEL_ID}" +_EVENT_STREAM: Final = "application/vnd.amazon.eventstream" +_ANSWER_HEAD: Final = "scripted answer part one " +_ANSWER_TAIL: Final = "scripted answer part two" +_ANSWER: Final = _ANSWER_HEAD + _ANSWER_TAIL +_PROMPT: Final = "How long does the gateway wait for this stream?" +_USER_TURN: Final[dict[str, JsonValue]] = {"role": "user", "content": _PROMPT} +_TIMEOUT_SECONDS: Final = 1 +_PAUSE_SECONDS: Final = 3.0 +_HOLD_SECONDS: Final = 40.0 +_CLIENT_WINDOW: Final = 10.0 +_RETRY_WINDOW: Final = 25.0 +_SHORT_WINDOW: Final = 3.0 +_AWS: Final[dict[str, JsonValue]] = { + "api_key": None, + "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}} +_JSON: Final = TypeAdapter(dict[str, JsonValue]) +_MODELS: Final = TypeAdapter(list[dict[str, JsonValue]]) +_TIMEOUT_PASSED: Final = re.compile(r"Timeout passed=(?:Timeout\(timeout=)?(-?\d+\.\d+)") +_USAGE: Final[dict[str, JsonValue]] = {"usage": {"inputTokens": 11, "outputTokens": 4, "totalTokens": 15}} +_CONVERSE_RESPONSE: Final = json.dumps( + { + "output": {"message": {"role": "assistant", "content": [{"text": _ANSWER}]}}, + "stopReason": "end_turn", + **_USAGE, + "metrics": {"latencyMs": 1}, + } +).encode() +_INVOKE_RESPONSE: Final = json.dumps( + { + "id": "msg_invoke", + "type": "message", + "role": "assistant", + "content": [{"type": "text", "text": _ANSWER}], + "model": _INVOKE_MODEL_ID, + "stop_reason": "end_turn", + "usage": {"input_tokens": 11, "output_tokens": 4}, + } +).encode() + +Endpoint = Literal["chat", "messages", "responses"] +Mode = Literal["fast", "stall", "pause", "cut"] + + +def _frame(event_type: str, payload: Mapping[str, JsonValue]) -> bytes: + return _aws_event_frame(event_type, payload, "sc", "u") + + +def _invoke_chunk(event: Mapping[str, JsonValue]) -> bytes: + return _frame("chunk", {"bytes": base64.b64encode(json.dumps(event).encode()).decode()}) + + +_CONVERSE_FRAMES: Final = ( + _frame("messageStart", {"role": "assistant"}), + _frame("contentBlockDelta", {"delta": {"text": _ANSWER_HEAD}, "contentBlockIndex": 0}), + _frame("contentBlockDelta", {"delta": {"text": _ANSWER_TAIL}, "contentBlockIndex": 0}), + _frame("contentBlockStop", {"contentBlockIndex": 0}), + _frame("messageStop", {"stopReason": "end_turn"}), + _frame("metadata", _USAGE), +) +_INVOKE_FRAMES: Final = ( + _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_HEAD}}), + _invoke_chunk({"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": _ANSWER_TAIL}}), + _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"}), +) + + +def _is_stream(target: str) -> bool: + return target.endswith("converse-stream") or target.endswith("invoke-with-response-stream") + + +def _frames_for(target: str) -> tuple[bytes, ...]: + return _INVOKE_FRAMES if "/invoke" in target else _CONVERSE_FRAMES + + +def _frames_before_the_pause(target: str, mode: Mode) -> int: + if mode == "pause": + return 1 + return 3 if "/invoke" in target else 2 + + +def _json_for(target: str) -> bytes: + return _INVOKE_RESPONSE if "/invoke" in target else _CONVERSE_RESPONSE + + +@dataclass(frozen=True, slots=True) +class _Peer: + mode: Mode + release: threading.Event + + def __call__(self, request: Request) -> Reply: + target: Final = unquote(request.target) + if self.mode == "stall": + self.release.wait(timeout=_HOLD_SECONDS) + if not _is_stream(target): + return Reply(body=_json_for(target)) + frames: Final = _frames_for(target) + if self.mode in ("pause", "cut"): + sent: Final = _frames_before_the_pause(target, self.mode) + return Reply( + chunks=(b"".join(frames[:sent]), b"".join(frames[sent:])), + content_type=_EVENT_STREAM, + pause_between_chunks=_PAUSE_SECONDS, + ) + return Reply(body=b"".join(frames), content_type=_EVENT_STREAM) + + +@contextmanager +def _peer(mode: Mode) -> Iterator[Wire]: + release: Final = threading.Event() + with wire_server(_Peer(mode, release)) as wire: + try: + yield wire + finally: + release.set() + + +def _converse(scenario: Scenario, wire: Wire, **extra: JsonValue) -> str: + return scenario.model(model=_CONVERSE_MODEL, api_base=wire.url, **_AWS, **extra) + + +def _invoke(scenario: Scenario, wire: Wire, **extra: JsonValue) -> str: + return scenario.model( + model=_INVOKE_MODEL, api_base=wire.url, aws_bedrock_runtime_endpoint=wire.url, **_AWS, **extra + ) + + +def _proxy_url(gateway: Gateway) -> str: + return str(gateway.client.base_url).rstrip("/") + + +def _auth(gateway: Gateway) -> dict[str, str]: + return {"Authorization": f"Bearer {gateway.key}"} + + +def _openai(gateway: Gateway, window: float = _CLIENT_WINDOW) -> openai.OpenAI: + return openai.OpenAI(base_url=f"{_proxy_url(gateway)}/v1", api_key=gateway.key, max_retries=0, timeout=window) + + +def _async_openai(gateway: Gateway) -> openai.AsyncOpenAI: + return openai.AsyncOpenAI( + base_url=f"{_proxy_url(gateway)}/v1", api_key=gateway.key, max_retries=0, timeout=_CLIENT_WINDOW + ) + + +def _anthropic(gateway: Gateway) -> anthropic.Anthropic: + return anthropic.Anthropic(base_url=_proxy_url(gateway), api_key=gateway.key, max_retries=0, timeout=_CLIENT_WINDOW) + + +def _async_anthropic(gateway: Gateway) -> anthropic.AsyncAnthropic: + return anthropic.AsyncAnthropic( + base_url=_proxy_url(gateway), api_key=gateway.key, max_retries=0, timeout=_CLIENT_WINDOW + ) + + +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, *, stream: bool = True, **extra: JsonValue) -> dict[str, JsonValue]: + match endpoint: + case "chat": + return {"model": model, "messages": [_USER_TURN], "stream": stream, **_EXTRA, **extra} + case "messages": + return {"model": model, "messages": [_USER_TURN], "max_tokens": 16, "stream": stream, **_EXTRA, **extra} + case "responses": + return {"model": model, "input": _PROMPT, "stream": stream, **_EXTRA, **extra} + + +@dataclass(frozen=True, slots=True) +class _Streamed: + status: int + headers: Mapping[str, str] + text: str + + +def _streamed( + gateway: Gateway, + endpoint: Endpoint, + body: Mapping[str, JsonValue] | None = None, + *, + content: bytes | None = None, + headers: Mapping[str, str] | None = None, + key: str | None = None, + window: float = _CLIENT_WINDOW, +) -> _Streamed: + with httpx.Client(base_url=_proxy_url(gateway), timeout=window, trust_env=False) as client: + return _streamed_on(client, gateway, endpoint, body, content=content, headers=headers, key=key) + + +def _streamed_on( + client: httpx.Client, + gateway: Gateway, + endpoint: Endpoint, + body: Mapping[str, JsonValue] | None = None, + *, + content: bytes | None = None, + headers: Mapping[str, str] | None = None, + key: str | None = None, +) -> _Streamed: + request_headers: Final = { + "Authorization": f"Bearer {key or gateway.key}", + **(headers or {}), + **({"Content-Type": "application/json"} if content is not None else {}), + } + with client.stream("POST", _path(endpoint), json=body, content=content, headers=request_headers) as response: + lines: Final = tuple(line for line in response.iter_lines() if line) + return _Streamed(response.status_code, dict(response.headers), "\n".join(lines)) + + +def _timeout_passed(text: str) -> str: + found: Final = _TIMEOUT_PASSED.search(text) + assert found is not None, text + return found.group(1) + + +_STREAMS_FROM_THE_FIRST_FRAME: Final = frozenset({"messages", "responses"}) + + +def _assert_no_answer(text: str) -> None: + assert _ANSWER_HEAD not in text, text + assert _ANSWER_TAIL not in text, text + + +def _assert_timed_out(status: int, text: str, seconds: float = _TIMEOUT_SECONDS) -> None: + assert status == 408, text + assert _timeout_passed(text) == f"{seconds:.1f}", text + _assert_no_answer(text) + + +def _assert_mid_stream_timeout(served: _Streamed, endpoint: Endpoint) -> None: + if endpoint in _STREAMS_FROM_THE_FIRST_FRAME: + assert served.status == 200, served.text + assert "error" in served.text, served.text + else: + assert served.status >= 400, served.text + assert "Timeout" in served.text, served.text + _assert_no_answer(served.text) + + +def _assert_cut_off_mid_answer(served: _Streamed) -> None: + assert served.status == 200, served.text + assert _ANSWER_HEAD in served.text, served.text + assert _ANSWER_TAIL not in served.text, served.text + assert "Timeout" in served.text, served.text + + +def _assert_answered(served: _Streamed) -> None: + assert served.status == 200, served.text + assert _ANSWER_HEAD in served.text, served.text + assert _ANSWER_TAIL in served.text, served.text + assert "Timeout" not in served.text, served.text + + +def _failure_rows(model: str, expected: int) -> list[dict[str, JsonValue]]: + rows: Final = eventually( + lambda: read_rows('SELECT request_id, status FROM "LiteLLM_SpendLogs" WHERE model_group=%s', (model,)), + lambda found: len(found) >= expected, + seconds=60, + ) + assert [row["status"] for row in rows] == ["failure"] * expected, rows + assert len({row["request_id"] for row in rows}) == expected, rows + return rows + + +def _received(wire: Wire, expected: int) -> None: + eventually(wire.received.qsize, lambda size: size == expected, 10) + assert len(wire.drain()) == expected + + +def _model_identity(gateway: Gateway, model: str) -> str: + entries: Final = _MODELS.validate_python(gateway.get("/model/info")["data"]) + (entry,) = tuple(item for item in entries if item["model_name"] == model) + return string_value(object_value(entry["model_info"])["id"]) + + +def _model_timeout(gateway: Gateway, model: str) -> JsonValue: + entries: Final = _MODELS.validate_python(gateway.get("/model/info")["data"]) + (entry,) = tuple(item for item in entries if item["model_name"] == model) + return object_value(entry["litellm_params"]).get("timeout") + + +def test_h1_chat_stream_through_the_openai_sdk_fails_at_the_deployment_timeout(gateway: Gateway) -> None: + with _peer("stall") as wire, gateway.scenario() as scenario: + model: Final = _converse(scenario, wire, timeout=_TIMEOUT_SECONDS) + with pytest.raises(openai.APIStatusError) as caught: + _openai(gateway).chat.completions.create(model=model, messages=[_USER_TURN], stream=True, extra_body=_EXTRA) + _assert_timed_out(caught.value.status_code, caught.value.response.text) + _received(wire, 1) + _failure_rows(model, 1) + + +async def test_h2_chat_stream_through_the_async_openai_sdk_fails_at_the_deployment_timeout(gateway: Gateway) -> None: + with _peer("stall") as wire, gateway.scenario() as scenario: + model: Final = _converse(scenario, wire, timeout=_TIMEOUT_SECONDS) + with pytest.raises(openai.APIStatusError) as caught: + await _async_openai(gateway).chat.completions.create( + model=model, messages=[_USER_TURN], stream=True, extra_body=_EXTRA + ) + _assert_timed_out(caught.value.status_code, caught.value.response.text) + _received(wire, 1) + + +def test_h3_messages_stream_through_the_anthropic_sdk_fails_at_the_deployment_timeout(gateway: Gateway) -> None: + with _peer("stall") as wire, gateway.scenario() as scenario: + model: Final = _converse(scenario, wire, timeout=_TIMEOUT_SECONDS) + with pytest.raises(anthropic.APIStatusError) as caught: + _anthropic(gateway).messages.create( + model=model, max_tokens=16, messages=[_USER_TURN], stream=True, extra_body=_EXTRA + ) + _assert_timed_out(caught.value.status_code, caught.value.response.text) + _received(wire, 1) + _failure_rows(model, 1) + + +async def test_h4_messages_stream_through_the_async_anthropic_sdk_fails_at_the_deployment_timeout( + gateway: Gateway, +) -> None: + with _peer("stall") as wire, gateway.scenario() as scenario: + model: Final = _converse(scenario, wire, timeout=_TIMEOUT_SECONDS) + with pytest.raises(anthropic.APIStatusError) as caught: + await _async_anthropic(gateway).messages.create( + model=model, max_tokens=16, messages=[_USER_TURN], stream=True, extra_body=_EXTRA + ) + _assert_timed_out(caught.value.status_code, caught.value.response.text) + _received(wire, 1) + + +def test_h5_responses_stream_through_the_openai_sdk_fails_at_the_deployment_timeout(gateway: Gateway) -> None: + with _peer("stall") as wire, gateway.scenario() as scenario: + model: Final = _converse(scenario, wire, timeout=_TIMEOUT_SECONDS) + with pytest.raises(openai.APIStatusError) as caught: + _openai(gateway).responses.create(model=model, input=_PROMPT, stream=True, extra_body=_EXTRA) + _assert_timed_out(caught.value.status_code, caught.value.response.text) + _received(wire, 1) + _failure_rows(model, 1) + + +async def test_h6_responses_stream_through_raw_httpx_fails_at_the_deployment_timeout(gateway: Gateway) -> None: + with _peer("stall") as wire, gateway.scenario() as scenario: + model: Final = _converse(scenario, wire, timeout=_TIMEOUT_SECONDS) + async with httpx.AsyncClient(base_url=_proxy_url(gateway), timeout=_CLIENT_WINDOW, trust_env=False) as client: + response: Final = await client.post("/v1/responses", json=_body("responses", model), headers=_auth(gateway)) + _assert_timed_out(response.status_code, response.text) + _received(wire, 1) + + +async def test_h7_invoke_stream_through_the_async_openai_sdk_fails_at_the_deployment_timeout(gateway: Gateway) -> None: + with _peer("stall") as wire, gateway.scenario() as scenario: + model: Final = _invoke(scenario, wire, timeout=_TIMEOUT_SECONDS) + with pytest.raises(openai.APIStatusError) as caught: + await _async_openai(gateway).chat.completions.create( + model=model, messages=[_USER_TURN], stream=True, extra_body=_EXTRA + ) + _assert_timed_out(caught.value.status_code, caught.value.response.text) + _received(wire, 1) + _failure_rows(model, 1) + + +@pytest.mark.parametrize("endpoint", ("chat", "messages", "responses")) +def test_h8_to_h10_a_converse_stream_that_pauses_past_the_timeout_ends_with_a_timeout_error( + gateway: Gateway, endpoint: Endpoint +) -> None: + with _peer("pause") as wire, gateway.scenario() as scenario: + model: Final = _converse(scenario, wire, timeout=_TIMEOUT_SECONDS) + _assert_mid_stream_timeout(_streamed(gateway, endpoint, _body(endpoint, model)), endpoint) + _received(wire, 1) + + +def test_h11_an_invoke_stream_that_pauses_past_the_timeout_ends_with_a_timeout_error(gateway: Gateway) -> None: + with _peer("pause") as wire, gateway.scenario() as scenario: + model: Final = _invoke(scenario, wire, timeout=_TIMEOUT_SECONDS) + _assert_mid_stream_timeout(_streamed(gateway, "chat", _body("chat", model)), "chat") + _received(wire, 1) + + +@pytest.mark.parametrize("endpoint", ("chat", "messages", "responses")) +def test_h12_a_converse_stream_cut_off_mid_answer_reports_the_timeout_in_band( + gateway: Gateway, endpoint: Endpoint +) -> None: + with _peer("cut") as wire, gateway.scenario() as scenario: + model: Final = _converse(scenario, wire, timeout=_TIMEOUT_SECONDS) + _assert_cut_off_mid_answer(_streamed(gateway, endpoint, _body(endpoint, model))) + _received(wire, 1) + + +def test_h13_an_invoke_stream_cut_off_mid_answer_reports_the_timeout_in_band(gateway: Gateway) -> None: + with _peer("cut") as wire, gateway.scenario() as scenario: + model: Final = _invoke(scenario, wire, timeout=_TIMEOUT_SECONDS) + _assert_cut_off_mid_answer(_streamed(gateway, "chat", _body("chat", model))) + _received(wire, 1) + + +def test_c1_chat_without_streaming_already_fails_at_the_deployment_timeout(gateway: Gateway) -> None: + with _peer("stall") as wire, gateway.scenario() as scenario: + model: Final = _converse(scenario, wire, timeout=_TIMEOUT_SECONDS) + with pytest.raises(openai.APIStatusError) as caught: + _openai(gateway).chat.completions.create(model=model, messages=[_USER_TURN], extra_body=_EXTRA) + _assert_timed_out(caught.value.status_code, caught.value.response.text) + _received(wire, 1) + _failure_rows(model, 1) + + +def test_c2_invoke_without_streaming_already_fails_at_the_deployment_timeout(gateway: Gateway) -> None: + with _peer("stall") as wire, gateway.scenario() as scenario: + model: Final = _invoke(scenario, wire, timeout=_TIMEOUT_SECONDS) + with pytest.raises(openai.APIStatusError) as caught: + _openai(gateway).chat.completions.create(model=model, messages=[_USER_TURN], extra_body=_EXTRA) + _assert_timed_out(caught.value.status_code, caught.value.response.text) + _received(wire, 1) + + +@pytest.mark.parametrize("endpoint", ("chat", "messages", "responses")) +def test_c3_to_c5_a_prompt_converse_stream_under_the_timeout_is_answered(gateway: Gateway, endpoint: Endpoint) -> None: + with _peer("fast") as wire, gateway.scenario() as scenario: + model: Final = _converse(scenario, wire, timeout=_TIMEOUT_SECONDS) + _assert_answered(_streamed(gateway, endpoint, _body(endpoint, model))) + _received(wire, 1) + + +def test_c6_a_prompt_invoke_stream_under_the_timeout_is_answered(gateway: Gateway) -> None: + with _peer("fast") as wire, gateway.scenario() as scenario: + model: Final = _invoke(scenario, wire, timeout=_TIMEOUT_SECONDS) + _assert_answered(_streamed(gateway, "chat", _body("chat", model))) + _received(wire, 1) + + +_PASSTHROUGH_BODY: Final[dict[str, JsonValue]] = {"messages": [{"role": "user", "content": [{"text": _PROMPT}]}]} + + +def test_c7_the_bedrock_passthrough_stream_is_relayed_verbatim(gateway: Gateway) -> None: + with _peer("fast") as wire, gateway.scenario() as scenario: + model: Final = _converse(scenario, wire, timeout=_TIMEOUT_SECONDS) + response: Final = gateway.request("POST", f"/bedrock/model/{model}/converse-stream", _PASSTHROUGH_BODY) + assert response.status_code == 200, response.text + assert response.headers.get("content-type") == _EVENT_STREAM, dict(response.headers) + assert response.content == b"".join(_CONVERSE_FRAMES), response.text + _received(wire, 1) + + +def test_c8_the_bedrock_passthrough_stream_already_retries_at_the_deployment_timeout(gateway: Gateway) -> None: + with _peer("stall") as wire, gateway.scenario() as scenario: + model: Final = _converse(scenario, wire, timeout=_TIMEOUT_SECONDS) + with httpx.Client(base_url=_proxy_url(gateway), timeout=_RETRY_WINDOW, trust_env=False) as client: + response: Final = client.post( + f"/bedrock/model/{model}/converse-stream", json=_PASSTHROUGH_BODY, headers=_auth(gateway) + ) + assert response.status_code >= 400, response.text + assert "Timeout" in response.text, response.text + assert response.headers["x-litellm-timeout"] == f"{_TIMEOUT_SECONDS:.1f}", response.headers + _assert_no_answer(response.text) + _received(wire, 3) + + +def test_c9_model_info_reports_the_deployment_timeout(gateway: Gateway) -> None: + with _peer("fast") as wire, gateway.scenario() as scenario: + model: Final = _converse(scenario, wire, timeout=_TIMEOUT_SECONDS) + assert _model_timeout(gateway, model) == _TIMEOUT_SECONDS + + +def test_c10_a_deployment_without_any_timeout_keeps_waiting_on_a_stalled_stream(gateway: Gateway) -> None: + with _peer("stall") as wire, gateway.scenario() as scenario: + model: Final = _converse(scenario, wire) + with pytest.raises(httpx.ReadTimeout): + _streamed(gateway, "chat", _body("chat", model), window=_SHORT_WINDOW) + _received(wire, 1) + + +_SOURCES: Final = ( + pytest.param({"timeout": _TIMEOUT_SECONDS}, {}, id="p1-body-timeout"), + pytest.param({"request_timeout": _TIMEOUT_SECONDS}, {}, id="p2-body-request-timeout"), + pytest.param({}, {"x-litellm-timeout": str(_TIMEOUT_SECONDS)}, id="p3-header-timeout"), + pytest.param({"stream_timeout": _TIMEOUT_SECONDS}, {}, id="p4-body-stream-timeout"), + pytest.param({}, {"x-litellm-stream-timeout": str(_TIMEOUT_SECONDS)}, id="p5-header-stream-timeout"), +) + + +@pytest.mark.parametrize(("extra", "headers"), _SOURCES) +def test_p1_to_p5_every_request_level_timeout_source_bounds_the_stream( + gateway: Gateway, extra: Mapping[str, JsonValue], headers: Mapping[str, str] +) -> None: + with _peer("stall") as wire, gateway.scenario() as scenario: + model: Final = _converse(scenario, wire) + served: Final = _streamed(gateway, "chat", _body("chat", model, **extra), headers=headers) + _assert_timed_out(served.status, served.text) + _received(wire, 1) + + +def test_p6_the_request_timeout_beats_the_deployment_timeout(gateway: Gateway) -> None: + with _peer("stall") as wire, gateway.scenario() as scenario: + model: Final = _converse(scenario, wire, timeout=30) + served: Final = _streamed(gateway, "chat", _body("chat", model, timeout=_TIMEOUT_SECONDS)) + _assert_timed_out(served.status, served.text) + _received(wire, 1) + + +def test_r1_retries_each_wait_the_timeout_and_the_last_one_answers(gateway: Gateway) -> None: + with _peer("stall") as wire, gateway.scenario() as scenario: + model: Final = _converse(scenario, wire, timeout=_TIMEOUT_SECONDS) + served: Final = _streamed(gateway, "chat", _body("chat", model, num_retries=2), window=_RETRY_WINDOW) + _assert_timed_out(served.status, served.text) + _received(wire, 3) + + +def test_r2_a_stream_timeout_falls_back_to_the_next_deployment(gateway: Gateway) -> None: + with _peer("stall") as stalled, _peer("fast") as prompt, gateway.scenario() as scenario: + slow: Final = _converse(scenario, stalled, timeout=_TIMEOUT_SECONDS) + fast: Final = _converse(scenario, prompt) + with httpx.Client(base_url=_proxy_url(gateway), timeout=_RETRY_WINDOW, trust_env=False) as same_worker: + _assert_answered(_streamed_on(same_worker, gateway, "chat", _body("chat", fast))) + served: Final = _streamed_on(same_worker, gateway, "chat", _body("chat", slow, fallbacks=[fast])) + _assert_answered(served) + assert served.headers.get("x-litellm-attempted-fallbacks") == "1", served.headers + _received(stalled, 1) + _received(prompt, 2) + + +_BAD_VALUES: Final = ( + pytest.param("abc", id="s1-string"), + pytest.param([1], id="s2-list"), + pytest.param("x" * 5120, id="s4-five-kilobytes"), +) + + +@pytest.mark.parametrize("value", _BAD_VALUES) +def test_s1_s2_s4_an_unusable_timeout_value_is_an_error_the_caller_sees(gateway: Gateway, value: JsonValue) -> None: + with _peer("fast") as wire, gateway.scenario() as scenario: + model: Final = _converse(scenario, wire) + served: Final = _streamed(gateway, "chat", _body("chat", model, timeout=value)) + assert served.status >= 400, served.text + assert "error" in served.text, served.text + assert wire.received.qsize() == 0, wire.drain() + _assert_answered(_streamed(gateway, "chat", _body("chat", model))) + _received(wire, 1) + + +def test_s3_an_empty_timeout_string_is_ignored(gateway: Gateway) -> None: + with _peer("fast") as wire, gateway.scenario() as scenario: + model: Final = _converse(scenario, wire) + _assert_answered(_streamed(gateway, "chat", _body("chat", model, timeout=""))) + _received(wire, 1) + + +def test_s5_the_last_duplicated_timeout_key_wins(gateway: Gateway) -> None: + with _peer("stall") as wire, gateway.scenario() as scenario: + model: Final = _converse(scenario, wire) + duplicated: Final = ( + json.dumps({**_body("chat", model), "timeout": 30})[:-1] + f', "timeout": {_TIMEOUT_SECONDS}}}' + ) + served: Final = _streamed(gateway, "chat", content=duplicated.encode()) + _assert_timed_out(served.status, served.text) + _received(wire, 1) + + +def test_s6_a_zero_timeout_leaves_the_deployment_timeout_in_force(gateway: Gateway) -> None: + with _peer("stall") as wire, gateway.scenario() as scenario: + model: Final = _converse(scenario, wire, timeout=_TIMEOUT_SECONDS) + served: Final = _streamed(gateway, "chat", _body("chat", model, timeout=0)) + _assert_timed_out(served.status, served.text) + _received(wire, 1) + + +def test_s7_a_negative_timeout_has_already_expired_for_streams_and_non_streams_alike(gateway: Gateway) -> None: + with _peer("fast") as wire, gateway.scenario() as scenario: + model: Final = _converse(scenario, wire) + streamed: Final = _streamed(gateway, "chat", _body("chat", model, timeout=-1)) + _assert_timed_out(streamed.status, streamed.text, -1) + plain: Final = _streamed(gateway, "chat", _body("chat", model, stream=False, timeout=-1)) + _assert_timed_out(plain.status, plain.text, -1) + + +def test_s8_a_malformed_deployment_timeout_fails_only_its_own_deployment(gateway: Gateway) -> None: + with _peer("fast") as wire, gateway.scenario() as scenario: + broken: Final = _converse(scenario, wire, timeout="fast") + healthy: Final = _converse(scenario, wire, timeout=_TIMEOUT_SECONDS) + served: Final = _streamed(gateway, "chat", _body("chat", broken)) + assert served.status >= 400, served.text + assert "error" in served.text, served.text + assert wire.received.qsize() == 0, wire.drain() + _assert_answered(_streamed(gateway, "chat", _body("chat", healthy))) + _received(wire, 1) + + +def test_s9_an_unauthenticated_stream_never_reaches_the_deployment(gateway: Gateway) -> None: + with _peer("fast") as wire, gateway.scenario() as scenario: + model: Final = _converse(scenario, wire) + served: Final = _streamed(gateway, "chat", _body("chat", model, timeout=_TIMEOUT_SECONDS), key="sk-not-a-key") + assert served.status == 401, served.text + assert wire.received.qsize() == 0, wire.drain() + + +def test_e1_a_null_timeout_reads_as_missing(gateway: Gateway) -> None: + with _peer("stall") as wire, gateway.scenario() as scenario: + model: Final = _converse(scenario, wire, timeout=_TIMEOUT_SECONDS) + served: Final = _streamed(gateway, "chat", _body("chat", model, timeout=None)) + _assert_timed_out(served.status, served.text) + _received(wire, 1) + + +def test_e2_a_timeout_updated_while_traffic_flows_applies_to_the_next_stream(gateway: Gateway) -> None: + with _peer("stall") as wire, gateway.scenario() as scenario: + model: Final = _converse(scenario, wire, timeout=_TIMEOUT_SECONDS) + _assert_timed_out(*_status_and_text(_streamed(gateway, "chat", _body("chat", model)))) + gateway.post( + "/model/update", + { + "model_name": model, + "litellm_params": {"timeout": 3}, + "model_info": {"id": _model_identity(gateway, model)}, + }, + ) + eventually(lambda: _model_timeout(gateway, model), lambda value: value == 3, 30) + eventually( + lambda: _timeout_passed(_streamed(gateway, "chat", _body("chat", model)).text), + lambda passed: passed == "3.0", + 45, + ) + received: Final = wire.drain() + assert len(received) >= 2, received + + +def _status_and_text(served: _Streamed) -> tuple[int, str]: + return served.status, served.text diff --git a/tests/unit/llms/bedrock/chat/test_bedrock_converse_handler.py b/tests/unit/llms/bedrock/chat/test_bedrock_converse_handler.py index 08bcac33a35..c025a368e51 100644 --- a/tests/unit/llms/bedrock/chat/test_bedrock_converse_handler.py +++ b/tests/unit/llms/bedrock/chat/test_bedrock_converse_handler.py @@ -16,11 +16,17 @@ import pytest from botocore.credentials import Credentials from botocore.exceptions import ClientError +import litellm from litellm.llms.bedrock.chat.converse_handler import BedrockConverseLLM from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler from litellm.rust_bridge import configuration from litellm.types.utils import ModelResponse from tests.unit.llms.bedrock.event_loop_probe import EventLoopProbe +from tests.unit.llms.bedrock.slow_upstream import ( + STREAM_TIMEOUT_SECONDS, + slow_upstream_async_client, + slow_upstream_sync_client, +) RESOLVED_CREDENTIALS = Credentials( access_key="AKIARESOLVED", @@ -273,3 +279,26 @@ def test_session_tags_sign_the_request_and_stay_out_of_the_body(monkeypatch): sent = client.post.call_args.kwargs assert "Credential=ASIACONVERSETAGGED/" in sent["headers"]["Authorization"] assert "aws_session_tags" not in sent["data"] + + +def _converse_streaming_kwargs() -> dict[str, object]: + return { + "model": "bedrock/anthropic.claude-sonnet-4-5-v1:0", + "messages": [{"role": "user", "content": "hi"}], + "stream": True, + "timeout": STREAM_TIMEOUT_SECONDS, + "aws_access_key_id": "fake", + "aws_secret_access_key": "fake", + "aws_region_name": "us-east-1", + } + + +@pytest.mark.asyncio +async def test_async_converse_streaming_fails_at_the_request_timeout_not_the_upstreams_pace() -> None: + with pytest.raises(litellm.Timeout): + await litellm.acompletion(client=slow_upstream_async_client(), **_converse_streaming_kwargs()) + + +def test_sync_converse_streaming_fails_at_the_request_timeout_not_the_upstreams_pace() -> None: + with pytest.raises(litellm.Timeout): + litellm.completion(client=slow_upstream_sync_client(), **_converse_streaming_kwargs()) diff --git a/tests/unit/llms/bedrock/chat/test_invoke_handler.py b/tests/unit/llms/bedrock/chat/test_invoke_handler.py index 43b689e499d..dfe1c06edb5 100644 --- a/tests/unit/llms/bedrock/chat/test_invoke_handler.py +++ b/tests/unit/llms/bedrock/chat/test_invoke_handler.py @@ -1,7 +1,7 @@ import base64 import binascii -import itertools import datetime +import itertools import json import re import struct @@ -13,6 +13,7 @@ import httpx import pytest import litellm +from litellm.exceptions import MidStreamFallbackError from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper from litellm.llms.bedrock.chat.invoke_handler import ( @@ -21,10 +22,14 @@ from litellm.llms.bedrock.chat.invoke_handler import ( make_call, make_sync_call, ) -from litellm.exceptions import MidStreamFallbackError 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 +from tests.unit.llms.bedrock.slow_upstream import ( + STREAM_TIMEOUT_SECONDS, + slow_upstream_async_client, + slow_upstream_sync_client, +) def test_transform_thinking_blocks_with_redacted_content(): @@ -215,9 +220,7 @@ def test_bedrock_converse_streaming_consistent_id(): expected_id = f"chatcmpl-{native_conversation_id}" for response in parsed_responses: - assert ( - response.id == expected_id - ), "All chunk IDs must match the one captured from the messageStart event" + assert response.id == expected_id, "All chunk IDs must match the one captured from the messageStart event" def test_converse_streaming_usage_uses_provider_thinking_tokens(): @@ -1125,3 +1128,26 @@ def test_converse_stream_made_only_of_unknown_events_raises_instead_of_an_empty_ assert isinstance(exc_info.value.original_exception, litellm.BadGatewayError) assert "somethingBedrockAddedLater" in str(exc_info.value) assert _UPSTREAM_REJECTION in str(exc_info.value) + + +def _invoke_streaming_kwargs() -> dict[str, object]: + return { + "model": "bedrock/invoke/anthropic.claude-sonnet-4-6", + "messages": [{"role": "user", "content": "hi"}], + "stream": True, + "timeout": STREAM_TIMEOUT_SECONDS, + "aws_access_key_id": "fake", + "aws_secret_access_key": "fake", + "aws_region_name": "us-east-1", + } + + +@pytest.mark.asyncio +async def test_async_invoke_streaming_fails_at_the_request_timeout_not_the_upstreams_pace() -> None: + with pytest.raises(litellm.Timeout): + await litellm.acompletion(client=slow_upstream_async_client(), **_invoke_streaming_kwargs()) + + +def test_sync_invoke_streaming_fails_at_the_request_timeout_not_the_upstreams_pace() -> None: + with pytest.raises(litellm.Timeout): + litellm.completion(client=slow_upstream_sync_client(), **_invoke_streaming_kwargs()) diff --git a/tests/unit/llms/bedrock/slow_upstream.py b/tests/unit/llms/bedrock/slow_upstream.py new file mode 100644 index 00000000000..ed262d0f3ed --- /dev/null +++ b/tests/unit/llms/bedrock/slow_upstream.py @@ -0,0 +1,35 @@ +"""An upstream whose first byte takes longer than the request allows, the way a slow model does. + +httpx hands every transport the request's timeout in ``request.extensions["timeout"]``, so this one honours it +in process the way a socket would: a read timeout shorter than the first byte's latency times out, a longer one +gets the answer. +""" + +from __future__ import annotations + +from typing import Final + +import httpx +from pydantic import TypeAdapter + +from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler + +STREAM_TIMEOUT_SECONDS: Final = 0.5 +UPSTREAM_FIRST_BYTE_SECONDS: Final = 2.0 + +_TIMEOUT_EXTENSION: Final = TypeAdapter(dict[str, float | None]) + + +def _answer_once_the_first_byte_is_due(request: httpx.Request) -> httpx.Response: + read_timeout: Final = _TIMEOUT_EXTENSION.validate_python(request.extensions["timeout"])["read"] + if read_timeout is not None and read_timeout < UPSTREAM_FIRST_BYTE_SECONDS: + raise httpx.ReadTimeout(f"no byte arrived within {read_timeout}s", request=request) + return httpx.Response(200, content=b"", request=request) + + +def slow_upstream_async_client() -> AsyncHTTPHandler: + return AsyncHTTPHandler(transport=httpx.MockTransport(_answer_once_the_first_byte_is_due)) + + +def slow_upstream_sync_client() -> HTTPHandler: + return HTTPHandler(client=httpx.Client(transport=httpx.MockTransport(_answer_once_the_first_byte_is_due)))