mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
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:
parent
9e9c29f404
commit
1b98748528
17 changed files with 1189 additions and 5 deletions
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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={})
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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={})
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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 (
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
|
|
|
|||
388
tests/integration/providers/test_bedrock_stream_timeout_chaos.py
Normal file
388
tests/integration/providers/test_bedrock_stream_timeout_chaos.py
Normal 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
|
||||
673
tests/integration/providers/test_bedrock_stream_timeout_wire.py
Normal file
673
tests/integration/providers/test_bedrock_stream_timeout_wire.py
Normal 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
|
||||
|
|
@ -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())
|
||||
|
|
|
|||
|
|
@ -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())
|
||||
|
|
|
|||
35
tests/unit/llms/bedrock/slow_upstream.py
Normal file
35
tests/unit/llms/bedrock/slow_upstream.py
Normal 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)))
|
||||
Loading…
Add table
Reference in a new issue