test(python-bridge): cover native runtime boundaries

This commit is contained in:
Yujong Lee 2026-08-31 17:19:00 -07:00
parent 9d507383a9
commit 81202af458
8 changed files with 583 additions and 3 deletions

321
.github/scripts/test_native_routes.py vendored Normal file
View file

@ -0,0 +1,321 @@
from __future__ import annotations
import asyncio
import importlib.util
import json
import os
import signal
import subprocess
import sys
import tempfile
import threading
import zipfile
from http.client import HTTPMessage
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
from pathlib import Path
from socket import socket as Socket
from typing import Final
REQUEST_STARTED: Final = threading.Event()
REQUEST_CANCELLED: Final = threading.Event()
ANTHROPIC_RESPONSE: Final = (
b'{"id":"msg_native","type":"message","role":"assistant",'
b'"model":"claude-sonnet-4-5","content":[{"type":"text","text":"native-message"}],'
b'"stop_reason":"end_turn","stop_sequence":null,'
b'"usage":{"input_tokens":2,"output_tokens":3}}'
)
class NativeRouteHandler(BaseHTTPRequestHandler):
protocol_version = "HTTP/1.1"
def do_POST(self) -> None:
content_length: Final = int(self.headers.get("content-length", "0"))
body: Final = json.loads(self.rfile.read(content_length))
route: Final = self.headers.get("x-test-route")
outcome: Final = self.headers.get("x-test-outcome")
assert_native_request(route, outcome, self.path, self.headers, body)
if outcome == "hang":
REQUEST_STARTED.set()
self.connection.settimeout(5)
if connection_was_cancelled(self.connection):
REQUEST_CANCELLED.set()
return
status: Final = 429 if outcome == "429" else 200
response_body: Final = native_response(status, route)
self.send_response(status)
self.send_header("content-type", "application/json")
self.send_header("content-length", str(len(response_body)))
self.send_header("connection", "close")
self.end_headers()
self.wfile.write(response_body)
def log_message(self, _message_format: str, *_args: object) -> None:
pass
def connection_was_cancelled(connection: Socket) -> bool:
try:
return connection.recv(1) == b""
except TimeoutError:
return False
except OSError:
return True
def assert_native_request(
route: str | None,
outcome: str | None,
path: str,
headers: HTTPMessage,
body: object,
) -> None:
if route not in {"ocr", "transcription", "messages", "chat_completions"}:
raise AssertionError(f"unexpected route marker: {route!r}")
if outcome not in {"success", "429", "hang"}:
raise AssertionError(f"unexpected outcome marker: {outcome!r}")
if not isinstance(body, dict):
raise TypeError(f"{route} sent {type(body).__name__}, expected a JSON object")
if route == "ocr":
assert path == "/v1/ocr"
assert headers.get("authorization") == "Bearer sk-native"
assert body["model"] == "mistral-ocr-latest"
assert body["document"]["document_url"] == "https://example.com/document.pdf"
assert body["include_image_base64"] is True
return
if route == "transcription":
assert path == "/model/mistral.voxtral-mini-3b-2507/converse"
assert headers.get("authorization", "").startswith("AWS4-HMAC-SHA256 ")
assert headers.get("x-amz-date")
assert body["messages"][0]["content"][0]["audio"]["source"]["bytes"] == "AQI="
assert "The audio language is en" in body["messages"][0]["content"][1]["text"]
return
assert path == "/v1/messages"
assert headers.get("x-api-key") == "sk-native"
assert body["model"] == "claude-sonnet-4-5"
if route == "messages":
assert body["max_tokens"] == 16
assert body["messages"][0]["content"] == "hello-from-messages"
return
assert body["max_tokens"] == 17
assert body["messages"][0]["content"] == [{"type": "text", "text": "hello-from-chat"}]
def native_response(status: int, route: str | None) -> bytes:
if status == 429:
return b'{"error":"native-rate-limit"}'
if route == "ocr":
return b'{"pages":[{"index":0,"markdown":"native-ocr"}]}'
if route == "transcription":
return b'{"output":{"message":{"content":[{"text":"native-transcription"}]}}}'
return ANTHROPIC_RESPONSE
def load_native(native_path: Path) -> object:
module_spec: Final = importlib.util.spec_from_file_location("litellm.rust_bridge._native", native_path)
if module_spec is None or module_spec.loader is None:
raise RuntimeError("cannot create native extension import specification")
native_module: Final = importlib.util.module_from_spec(module_spec)
module_spec.loader.exec_module(native_module)
return native_module
def route_kwargs(route: str, api_base: str, outcome: str) -> dict[str, object]:
common: Final = {
"api_base": api_base,
"extra_headers": {"x-test-outcome": outcome, "x-test-route": route},
"timeout_seconds": 3.0,
}
if route == "ocr":
return common | {
"model": "mistral-ocr-latest",
"document": {"type": "document_url", "document_url": "https://example.com/document.pdf"},
"api_key": "sk-native",
"custom_llm_provider": "mistral",
"optional_params": {"include_image_base64": True},
}
if route == "transcription":
return common | {
"model": "mistral.voxtral-mini-3b-2507",
"audio": {"data": "AQI=", "format": "wav", "filename": "audio.wav"},
"custom_llm_provider": "bedrock",
"optional_params": {
"aws_access_key_id": "native-access-key",
"aws_secret_access_key": "native-secret-key",
"aws_region_name": "us-east-1",
"language": "en",
},
}
if route == "messages":
return common | {
"model": "claude-sonnet-4-5",
"body": {
"model": "claude-sonnet-4-5",
"max_tokens": 16,
"messages": [{"role": "user", "content": "hello-from-messages"}],
},
"api_key": "sk-native",
"custom_llm_provider": "anthropic",
}
if route == "chat_completions":
return common | {
"model": "anthropic/claude-sonnet-4-5",
"messages": [{"role": "user", "content": "hello-from-chat"}],
"optional_params": {"max_tokens": 17},
"api_key": "sk-native",
}
raise AssertionError(f"unknown route: {route}")
def assert_success(route: str, response: object) -> None:
if not isinstance(response, dict):
raise TypeError(f"{route} returned {type(response).__name__}, expected dict")
actual: Final = success_value(route, response)
expected: Final = (
"native-ocr"
if route == "ocr"
else "native-transcription"
if route == "transcription"
else "native-message"
)
if actual != expected:
raise AssertionError(f"{route} returned {actual!r}, expected {expected!r}")
def success_value(route: str, response: dict[object, object]) -> object:
if route == "ocr":
return response["pages"][0]["markdown"]
if route == "transcription":
return response["text"]
if route == "messages":
return response["content"][0]["text"]
return response["choices"][0]["message"]["content"]
def assert_rate_limit(native: object, route: str, error: BaseException) -> None:
if route in {"messages", "chat_completions"}:
upstream_error: Final = native.RustUpstreamError
if not isinstance(error, upstream_error) or error.args[0] != 429:
raise AssertionError(f"{route} returned the wrong 429 error: {error!r}")
return
if not isinstance(error, RuntimeError) or "429" not in str(error):
raise AssertionError(f"{route} returned the wrong 429 error: {error!r}")
def exercise_sync(native: object, api_base: str) -> None:
for route in ("ocr", "transcription", "messages", "chat_completions"):
function: Final = getattr(native, route)
assert_success(route, function(**route_kwargs(route, api_base, "success")))
try:
function(**route_kwargs(route, api_base, "429"))
except (RuntimeError, native.RustUpstreamError) as error:
assert_rate_limit(native, route, error)
else:
raise AssertionError(f"{route} accepted a 429 response")
async def exercise_async(native: object, api_base: str) -> None:
for route in ("ocr", "transcription", "messages", "chat_completions"):
function: Final = getattr(native, f"a{route}")
assert_success(route, await function(**route_kwargs(route, api_base, "success")))
try:
await function(**route_kwargs(route, api_base, "429"))
except (RuntimeError, native.RustUpstreamError) as error:
assert_rate_limit(native, route, error)
else:
raise AssertionError(f"a{route} accepted a 429 response")
def exercise_routes(native_path: Path, api_base: str) -> object:
native: Final = load_native(native_path)
exercise_sync(native, api_base)
asyncio.run(exercise_async(native, api_base))
return native
def exercise_signal(native: object, api_base: str) -> int:
try:
native.messages(
**route_kwargs("messages", api_base, "hang"),
)
except KeyboardInterrupt:
sys.stdout.write("KeyboardInterrupt\n")
sys.stdout.flush()
sys.stdin.read(1)
return 0
raise AssertionError("sync native route ignored SIGINT")
def test_sigint(native_path: Path, api_base: str) -> None:
REQUEST_STARTED.clear()
REQUEST_CANCELLED.clear()
process: Final = subprocess.Popen(
(sys.executable, __file__, "child", str(native_path), api_base),
stdin=subprocess.PIPE,
stdout=subprocess.PIPE,
stderr=subprocess.PIPE,
text=True,
)
try:
if not REQUEST_STARTED.wait(30):
process.kill()
stdout, stderr = process.communicate(timeout=5)
raise AssertionError(
f"native route matrix did not reach the hanging upstream\nstdout:\n{stdout}\nstderr:\n{stderr}"
)
os.kill(process.pid, signal.SIGINT)
if not REQUEST_CANCELLED.wait(5):
raise AssertionError("interrupted native route did not cancel its upstream future")
if process.poll() is not None:
raise AssertionError("signal child exited before cancellation was observed")
stdout, stderr = process.communicate(input="\n", timeout=5)
if process.returncode != 0 or stdout != "KeyboardInterrupt\n":
raise AssertionError(
f"signal child exited with status {process.returncode}\nstdout:\n{stdout}\nstderr:\n{stderr}"
)
finally:
if process.poll() is None:
process.kill()
process.wait(timeout=5)
def test_wheel(wheel: Path) -> int:
with tempfile.TemporaryDirectory() as temporary_directory, zipfile.ZipFile(wheel) as archive:
native_members: Final = tuple(
member
for member in archive.infolist()
if member.filename.startswith("litellm/rust_bridge/_native.") and member.filename.endswith(".so")
)
if len(native_members) != 1:
raise AssertionError(f"expected one native extension, found {len(native_members)}")
native_path: Final = Path(temporary_directory) / Path(native_members[0].filename).name
native_path.write_bytes(archive.read(native_members[0]))
server: Final = ThreadingHTTPServer(("127.0.0.1", 0), NativeRouteHandler)
server_thread: Final = threading.Thread(target=server.serve_forever, daemon=True)
server_thread.start()
api_base: Final = f"http://127.0.0.1:{server.server_address[1]}"
try:
test_sigint(native_path, api_base)
finally:
server.shutdown()
server.server_close()
server_thread.join(timeout=5)
return 0
def main() -> int:
if len(sys.argv) == 2:
return test_wheel(Path(sys.argv[1]))
if len(sys.argv) == 4 and sys.argv[1] == "child":
native: Final = exercise_routes(Path(sys.argv[2]), sys.argv[3])
return exercise_signal(native, sys.argv[3])
sys.stderr.write(f"usage: {Path(sys.argv[0]).name} WHEEL\n")
return 2
if __name__ == "__main__":
sys.exit(main())

View file

@ -8,6 +8,7 @@ on:
- "pyproject.toml"
- "rust-toolchain.toml"
- ".github/scripts/smoke_test_native_wheel.py"
- ".github/scripts/test_native_routes.py"
- ".github/scripts/verify_linux_native_wheel.py"
- ".github/workflows/test-rust.yml"
pull_request:
@ -22,6 +23,7 @@ on:
- "pyproject.toml"
- "rust-toolchain.toml"
- ".github/scripts/smoke_test_native_wheel.py"
- ".github/scripts/test_native_routes.py"
- ".github/scripts/verify_linux_native_wheel.py"
- ".github/workflows/test-rust.yml"
@ -121,3 +123,6 @@ jobs:
env:
RELEASE_WHEEL_COMMIT_SHA: ${{ github.event.pull_request.head.sha || github.sha }}
run: python .github/scripts/verify_linux_native_wheel.py dist/*.whl
- name: Test native route wheel
run: python .github/scripts/test_native_routes.py dist/*.whl

View file

@ -1451,6 +1451,7 @@ dependencies = [
"serde",
"serde_json",
"tokio",
"tokio-tungstenite",
]
[[package]]

View file

@ -526,3 +526,43 @@ async fn messages_classifies_a_refused_connection_as_safe_to_fallback() {
assert!(matches!(error, CoreError::Connect(_)));
}
#[tokio::test]
async fn messages_classifies_an_established_request_timeout_as_network_error() {
let listener = TcpListener::bind("127.0.0.1:0").await.expect("binds");
let addr = listener.local_addr().expect("has an address");
let (request_received_tx, request_received_rx) = tokio::sync::oneshot::channel();
let (release_server_tx, release_server_rx) = tokio::sync::oneshot::channel();
let server = tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.expect("accepts request");
let request = read_http_request(&mut socket).await;
request_received_tx.send(request).expect("reports request");
release_server_rx.await.expect("server is released");
});
let error = tokio::time::timeout(
Duration::from_secs(2),
messages(MessagesRequest {
model: "claude-test",
body: json!({"model": "claude-test", "max_tokens": 8, "messages": []}),
api_key: Some("sk"),
api_base: Some(&format!("http://{addr}")),
custom_llm_provider: Some("anthropic"),
extra_headers: None,
timeout: Some(Duration::from_millis(100)),
}),
)
.await
.expect("client call completes")
.expect_err("established request times out");
let request = tokio::time::timeout(Duration::from_secs(2), request_received_rx)
.await
.expect("server observes request")
.expect("server reports request");
assert!(request.starts_with("POST /v1/messages "), "{request}");
release_server_tx.send(()).expect("releases server");
server.await.expect("server task completes");
assert!(matches!(error, CoreError::Network(_)));
}

View file

@ -28,6 +28,7 @@ tokio.workspace = true
[dev-dependencies]
criterion = "0.8.2"
tokio-tungstenite.workspace = true
[[bench]]
name = "serialization"

View file

@ -87,6 +87,14 @@ fn _native(module: &Bound<'_, PyModule>) -> PyResult<()> {
#[cfg(test)]
mod tests {
use std::ffi::CString;
use std::time::Duration;
use futures_util::{SinkExt, StreamExt};
use pyo3::types::PyDict;
use tokio::net::TcpListener;
use tokio_tungstenite::{accept_async, tungstenite::Message};
use super::*;
#[test]
@ -123,4 +131,69 @@ mod tests {
assert_eq!(public_names, expected);
});
}
#[test]
fn responses_websocket_connection_round_trips_through_python() {
Python::initialize();
let runtime = pyo3_async_runtimes::tokio::get_runtime();
let listener = runtime
.block_on(TcpListener::bind("127.0.0.1:0"))
.expect("listener should bind");
let address = listener
.local_addr()
.expect("listener should have an address");
let server = runtime.spawn(async move {
let (stream, _) = listener.accept().await.expect("server should accept");
let mut socket = accept_async(stream)
.await
.expect("handshake should succeed");
let message = socket
.next()
.await
.expect("client should send a frame")
.expect("client frame should be valid");
assert_eq!(message, Message::Text("from-python".into()));
socket
.send(Message::Text("from-server".into()))
.await
.expect("server should reply");
assert!(matches!(socket.next().await, Some(Ok(Message::Close(_)))));
});
Python::attach(|py| {
let module = PyModule::new(py, "_native").expect("module should be created");
_native(&module).expect("module should register");
let locals = PyDict::new(py);
locals
.set_item("native", &module)
.expect("module should enter Python locals");
locals
.set_item("url", format!("ws://{address}"))
.expect("URL should enter Python locals");
let code = CString::new(
r#"
import asyncio
async def exercise():
connection = await native.ResponsesWebSocketConnection.connect(url)
assert type(connection) is native.ResponsesWebSocketConnection
await connection.send_text("from-python")
assert await connection.recv_text() == "from-server"
await connection.close()
assert await connection.recv_text() is None
asyncio.run(asyncio.wait_for(exercise(), timeout=5))
"#,
)
.expect("Python source should not contain null bytes");
py.run(&code, Some(&locals), Some(&locals))
.expect("Python WebSocket methods should round trip");
});
runtime
.block_on(async { tokio::time::timeout(Duration::from_secs(5), server).await })
.expect("server should finish")
.expect("server task should not panic");
}
}

View file

@ -46,7 +46,7 @@ where
let result = release_gil(py, move || runtime.block_on(wait_for_sync_result(future)))?;
let result = map_core_result(result, map_error)?;
to_py(py, &result)
std::panic::catch_unwind(AssertUnwindSafe(|| to_py(py, &result))).map_err(panic_to_pyerr)?
}
pub(super) fn run_async<T, F>(
@ -108,7 +108,11 @@ where
mod tests {
use std::ffi::CString;
use std::future::poll_fn;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::sync::{Arc, mpsc};
use std::task::Poll;
use std::thread;
use std::time::Instant;
use pyo3::panic::PanicException;
use pyo3::types::{PyDict, PyModule};
@ -127,12 +131,14 @@ mod tests {
struct PanickingOutput;
static ASYNC_PROBE_COMPLETED: AtomicUsize = AtomicUsize::new(0);
impl Serialize for PanickingOutput {
fn serialize<S>(&self, _serializer: S) -> Result<S::Ok, S::Error>
where
S: Serializer,
{
panic!("async serializer panicked")
panic!("serializer panicked")
}
}
@ -141,6 +147,42 @@ mod tests {
run_async(py, async { Ok(PanickingOutput) }, runtime_error)
}
#[pyfunction]
fn async_runtime_probe(py: Python<'_>) -> PyResult<Bound<'_, PyAny>> {
run_async(
py,
async {
ASYNC_PROBE_COMPLETED.fetch_add(1, Ordering::SeqCst);
Ok(true)
},
runtime_error,
)
}
#[pyfunction]
fn runtime_worker_count() -> usize {
pyo3_async_runtimes::tokio::get_runtime()
.metrics()
.num_workers()
}
#[pyfunction]
fn runtime_is_responsive(_py: Python<'_>, expected_completions: usize) -> bool {
let completion_deadline = Instant::now() + Duration::from_secs(2);
while ASYNC_PROBE_COMPLETED.load(Ordering::SeqCst) < expected_completions {
if Instant::now() >= completion_deadline {
return false;
}
thread::sleep(Duration::from_millis(1));
}
let (heartbeat_tx, heartbeat_rx) = mpsc::sync_channel(1);
pyo3_async_runtimes::tokio::get_runtime().spawn(async move {
let _ = heartbeat_tx.send(());
});
heartbeat_rx.recv_timeout(Duration::from_secs(2)).is_ok()
}
fn extract_bool(py: Python<'_>, result: PyResult<Py<PyAny>>) -> bool {
result
.expect("route should complete")
@ -259,6 +301,51 @@ mod tests {
});
}
#[test]
fn sync_runner_surfaces_serializer_panics() {
Python::initialize();
Python::attach(|py| {
let error = run_sync(py, async { Ok(PanickingOutput) }, runtime_error)
.expect_err("serializer panic should become a Python exception");
assert!(error.is_instance_of::<PanicException>(py));
assert_eq!(error.to_string(), "PanicException: serializer panicked");
});
}
#[test]
fn sync_runner_supports_concurrent_callers_on_the_shared_runtime() {
Python::initialize();
let barrier = Arc::new(tokio::sync::Barrier::new(2));
let callers: Vec<_> = (0..2)
.map(|_| {
let barrier = Arc::clone(&barrier);
thread::spawn(move || {
Python::attach(|py| {
extract_bool(
py,
run_sync(
py,
async move {
Ok(tokio::time::timeout(Duration::from_secs(2), barrier.wait())
.await
.is_ok())
},
runtime_error,
),
)
})
})
})
.collect();
let results: Vec<_> = callers
.into_iter()
.map(|caller| caller.join().expect("caller should not panic"))
.collect();
assert_eq!(results, vec![true, true]);
}
#[test]
fn async_runner_surfaces_serializer_panics() {
Python::initialize();
@ -283,7 +370,7 @@ async def exercise():
await runtime.async_serialization_panic()
except BaseException as error:
assert type(error).__name__ == "PanicException"
assert str(error) == "async serializer panicked"
assert str(error) == "serializer panicked"
else:
raise AssertionError("serializer panic was not raised")
@ -295,4 +382,42 @@ asyncio.run(exercise())
.expect("serializer panic should reach the Python awaiter");
});
}
#[test]
fn async_result_delivery_does_not_stall_tokio_workers() {
Python::initialize();
ASYNC_PROBE_COMPLETED.store(0, Ordering::SeqCst);
Python::attach(|py| {
let module = PyModule::new(py, "runtime").expect("module should be created");
for function in [
wrap_pyfunction!(async_runtime_probe, &module).expect("function should wrap"),
wrap_pyfunction!(runtime_worker_count, &module).expect("function should wrap"),
wrap_pyfunction!(runtime_is_responsive, &module).expect("function should wrap"),
] {
module
.add_function(function)
.expect("function should register");
}
let locals = PyDict::new(py);
locals
.set_item("runtime", &module)
.expect("module should enter Python locals");
let code = CString::new(
r#"
import asyncio
async def exercise():
worker_count = runtime.runtime_worker_count()
awaitables = [runtime.async_runtime_probe() for _ in range(worker_count)]
assert runtime.runtime_is_responsive(worker_count)
assert await asyncio.gather(*awaitables) == [True] * worker_count
asyncio.run(exercise())
"#,
)
.expect("Python source should not contain null bytes");
py.run(&code, Some(&locals), Some(&locals))
.expect("result delivery should leave Tokio workers responsive");
});
}
}

View file

@ -294,6 +294,20 @@ async def test_gate_surfaces_an_upstream_failure_without_fallback(monkeypatch):
assert bridge.calls == 1
@pytest.mark.asyncio
async def test_gate_maps_statusless_upstream_failure_to_500_without_fallback(monkeypatch):
_install_fake_bridge_exceptions(monkeypatch)
bridge = RaisingAsyncMessages(FakeUpstreamError(0, "request timed out"))
litellm.use_litellm_rust(True, amessages=bridge)
with pytest.raises(APIError) as exc_info:
await _gate()
assert exc_info.value.status_code == 500
assert "request timed out" in str(exc_info.value)
assert bridge.calls == 1
@pytest.mark.asyncio
async def test_gate_reraises_an_unknown_bridge_failure():
bridge = RaisingAsyncMessages(RuntimeError("unknown bridge failure"))