fix(bedrock): honor the per-request timeout on Converse and Invoke streaming (internal copy of #38210) (#44134)

* fix(bedrock): propagate timeout to streaming requests

* test(bedrock): prove streaming fails at the request timeout against a slow upstream

* test(bedrock): simulate the slow upstream in process instead of over a local socket

* test(bedrock): audit the Converse and Invoke stream timeout on the proxy

Two integration files drive the per-request timeout on Bedrock streams
through the real proxy against an owned wire peer: the wire file covers
every surface (chat, messages, responses, invoke, pass-through), the
sad, edge and precedence rows, and the chaos file covers bursts, a
dropping upstream, a killed worker and a proxy stopped mid-burst.

The wire peer gains Reply.drop_connection so a cell can close the
socket before any response, and the harness's graceful stop grace is
now INTEGRATION_PROXY_STOP_SECONDS (default unchanged at 30), since a
two-worker supervisor's interpreter finalization takes longer than that
on a loaded box.

* test(bedrock): pin the fallback audit cell to one proxy worker

The fallback cell created both deployments through /model/new on one
worker and sent the chat request to the other, whose registry
read-through loads only the requested model, so the fallback target
was unknown there until the periodic DB poll. The cell now warms the
fallback model and sends the request over one keep-alive client, so
one TCP connection stays with one uvicorn worker, and it expects the
fallback upstream to see both requests.

---------

Co-authored-by: Sainyam Kapoor <hello@sainyam.me>
Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-10-03 19:29:15 +00:00 • committed by GitHub
parent 9e9c29f404
commit 1b98748528
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
17 changed files with 1189 additions and 5 deletions

View file

@ -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

View file

@ -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.

View file

@ -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,

View file

@ -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:

View file

@ -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,

View file

@ -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={})

View file

@ -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(

View file

@ -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.

View file

@ -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={})

View file

@ -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:

View file

@ -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 (

View file

@ -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():

View file

@ -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

View file

@ -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

View file

@ -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())

View file

@ -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())

View file

@ -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)))