litellm/tests/integration/_support/wire.py
devin-ai-integration[bot] 1b98748528
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>
2026-10-03 19:29:15 +00:00

158 lines
5.7 KiB
Python

from __future__ import annotations
import ssl
import threading
import time
from collections.abc import Callable, Generator, Mapping
from contextlib import contextmanager
from dataclasses import dataclass
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
from queue import SimpleQueue
from types import MappingProxyType
from typing import Final
@dataclass(frozen=True, slots=True)
class Request:
method: str
target: str
headers: Mapping[str, str]
body: bytes
@dataclass(frozen=True, slots=True)
class Reply:
status: int = 200
body: bytes = b"{}"
content_type: str = "application/json"
chunks: tuple[bytes, ...] | None = None
abort_after: int | None = None
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)
class Wire:
url: str
received: SimpleQueue[Request]
disconnected: SimpleQueue[str]
connected: SimpleQueue[str]
def drain(self) -> tuple[Request, ...]:
return tuple(self.received.get_nowait() for _ in range(self.received.qsize()))
def connections(self) -> int:
return self.connected.qsize()
@contextmanager
def wire_server(
respond: Callable[[Request], Reply],
tls: ssl.SSLContext | None = None,
port: int = 0,
keep_alive: bool = False,
) -> Generator[Wire, None, None]:
"""Owned TCP peer; requests traverse the real HTTP client and serialization. With `keep_alive` the
peer honours HTTP/1.1 persistent connections so `connections()` counts the client's TCP sessions."""
received: Final[SimpleQueue[Request]] = SimpleQueue()
errors: Final[SimpleQueue[Exception]] = SimpleQueue()
disconnected: Final[SimpleQueue[str]] = SimpleQueue()
connected: Final[SimpleQueue[str]] = SimpleQueue()
class Handler(BaseHTTPRequestHandler):
protocol_version = "HTTP/1.1"
timeout = 5
def setup(self) -> None:
super().setup()
connected.put(f"{self.client_address[0]}:{self.client_address[1]}")
def respond(self) -> None:
request: Final = Request(
self.command,
self.path,
{name.lower(): value for name, value in self.headers.items()},
self.rfile.read(int(self.headers.get("content-length", "0"))),
)
received.put(request)
try:
reply = respond(request)
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():
self.send_header(name, value)
if reply.chunks is None:
self.send_header("content-length", str(len(reply.body)))
else:
self.send_header("transfer-encoding", "chunked")
if not keep_alive:
self.send_header("connection", "close")
self.end_headers()
try:
if self.command == "HEAD":
self.wfile.flush()
elif reply.chunks is None:
self.wfile.write(reply.body)
else:
for index, chunk in enumerate(reply.chunks):
if reply.abort_after == index:
break
self.wfile.write(b"%x\r\n%s\r\n" % (len(chunk), chunk))
self.wfile.flush()
if index == 0 and reply.gate_after_first is not None:
assert reply.gate_after_first.wait(timeout=5), "Stream barrier was never released"
if reply.pause_between_chunks and index + 1 < len(reply.chunks):
time.sleep(reply.pause_between_chunks)
else:
self.wfile.write(b"0\r\n\r\n")
self.wfile.flush()
except (BrokenPipeError, ConnectionResetError):
disconnected.put(request.target)
except Exception as error:
errors.put(error)
self.close_connection = not keep_alive
do_POST = respond
do_PUT = respond
do_GET = respond
do_DELETE = respond
do_PATCH = respond
do_HEAD = respond
def log_message(self, format: str, *args: object) -> None:
pass
class OwnedHTTPServer(ThreadingHTTPServer):
daemon_threads = False
def server_bind(self) -> None:
super().server_bind()
if tls is not None:
self.socket = tls.wrap_socket(self.socket, server_side=True)
with OwnedHTTPServer(("127.0.0.1", port), Handler) as server:
thread: Final = threading.Thread(target=server.serve_forever, kwargs={"poll_interval": 0.05})
thread.start()
try:
yield Wire(
f"{'https' if tls is not None else 'http'}://127.0.0.1:{server.server_port}",
received,
disconnected,
connected,
)
finally:
server.shutdown()
thread.join(timeout=6)
assert not thread.is_alive(), "Owned HTTP server survived cleanup"
server.server_close()
failure: Final = None if errors.empty() else errors.get_nowait()
assert failure is None, f"Owned HTTP peer failed: {failure!r}"