mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
test(python-bridge): cover native runtime boundaries
This commit is contained in:
parent
9d507383a9
commit
81202af458
8 changed files with 583 additions and 3 deletions
321
.github/scripts/test_native_routes.py
vendored
Normal file
321
.github/scripts/test_native_routes.py
vendored
Normal 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())
|
||||
5
.github/workflows/test-rust.yml
vendored
5
.github/workflows/test-rust.yml
vendored
|
|
@ -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
|
||||
|
|
|
|||
1
litellm-rust/Cargo.lock
generated
1
litellm-rust/Cargo.lock
generated
|
|
@ -1451,6 +1451,7 @@ dependencies = [
|
|||
"serde",
|
||||
"serde_json",
|
||||
"tokio",
|
||||
"tokio-tungstenite",
|
||||
]
|
||||
|
||||
[[package]]
|
||||
|
|
|
|||
|
|
@ -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(_)));
|
||||
}
|
||||
|
|
|
|||
|
|
@ -28,6 +28,7 @@ tokio.workspace = true
|
|||
|
||||
[dev-dependencies]
|
||||
criterion = "0.8.2"
|
||||
tokio-tungstenite.workspace = true
|
||||
|
||||
[[bench]]
|
||||
name = "serialization"
|
||||
|
|
|
|||
|
|
@ -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");
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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");
|
||||
});
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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"))
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue