mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
fix(python-bridge): harden route runtime
This commit is contained in:
parent
cd46cb5478
commit
9d507383a9
21 changed files with 985 additions and 254 deletions
1
litellm-rust/Cargo.lock
generated
1
litellm-rust/Cargo.lock
generated
|
|
@ -1442,6 +1442,7 @@ name = "litellm-python-bridge"
|
|||
version = "0.1.0"
|
||||
dependencies = [
|
||||
"criterion",
|
||||
"futures-util",
|
||||
"litellm-ai-gateway",
|
||||
"litellm-core",
|
||||
"litellm-python-interop",
|
||||
|
|
|
|||
|
|
@ -1,7 +1,7 @@
|
|||
use serde_json::Value;
|
||||
|
||||
use crate::error::{CoreError, CoreResult};
|
||||
use crate::http_utils::truncate_error_body;
|
||||
use crate::error::{CoreError, CoreResult, as_response_error};
|
||||
use crate::http_utils::{classify_send_error, truncate_error_body};
|
||||
|
||||
use super::client::http_client;
|
||||
use super::transformation::ChatCompletionsAuth;
|
||||
|
|
@ -27,16 +27,7 @@ pub(super) async fn execute_chat_completions_provider_call(
|
|||
request_builder = request_builder.timeout(duration);
|
||||
}
|
||||
|
||||
let response = request_builder.send().await.map_err(|err| {
|
||||
// Failing to establish the connection means the request never went out,
|
||||
// so the host can still serve it. Everything else here, a timeout
|
||||
// above all, may have reached the provider and been answered.
|
||||
if err.is_connect() || err.is_builder() {
|
||||
CoreError::Connect(err.to_string())
|
||||
} else {
|
||||
CoreError::Network(err.to_string())
|
||||
}
|
||||
})?;
|
||||
let response = request_builder.send().await.map_err(classify_send_error)?;
|
||||
|
||||
let status = response.status();
|
||||
let text = response
|
||||
|
|
@ -60,22 +51,6 @@ pub(super) async fn execute_chat_completions_provider_call(
|
|||
.map_err(as_response_error)
|
||||
}
|
||||
|
||||
/// Re-tag an error raised while normalizing a response the provider already
|
||||
/// returned.
|
||||
///
|
||||
/// A config reports the same variants on either side of the call: a missing
|
||||
/// field or an unsupported block can mean "this request cannot be translated"
|
||||
/// during prepare and "this response cannot be normalized" here. Only the
|
||||
/// second kind has already been billed, and a host that keeps a reference
|
||||
/// implementation must not retry those, so collapse them to one variant that
|
||||
/// can only mean the provider was already called.
|
||||
pub(super) fn as_response_error(err: CoreError) -> CoreError {
|
||||
match err {
|
||||
already @ (CoreError::InvalidResponse(_) | CoreError::Http { .. }) => already,
|
||||
other => CoreError::InvalidResponse(other.to_string()),
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(feature = "bedrock-auth")]
|
||||
pub(super) async fn signed_headers(
|
||||
request: &ProviderChatCompletionsRequest,
|
||||
|
|
|
|||
|
|
@ -794,7 +794,7 @@ mod round_trip {
|
|||
|
||||
#[test]
|
||||
fn response_errors_collapse_to_one_variant_that_can_only_mean_already_sent() {
|
||||
use crate::chat_completions::handler::as_response_error;
|
||||
use crate::error::as_response_error;
|
||||
|
||||
for original in [
|
||||
CoreError::MissingField("usage"),
|
||||
|
|
|
|||
|
|
@ -38,6 +38,14 @@ pub enum CoreError {
|
|||
Unsupported(&'static str),
|
||||
}
|
||||
|
||||
/// Re-tag an error raised after the provider has already returned a response.
|
||||
pub(crate) fn as_response_error(err: CoreError) -> CoreError {
|
||||
match err {
|
||||
already @ (CoreError::InvalidResponse(_) | CoreError::Http { .. }) => already,
|
||||
other => CoreError::InvalidResponse(other.to_string()),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn json_type_name(value: &serde_json::Value) -> &'static str {
|
||||
match value {
|
||||
serde_json::Value::Null => "null",
|
||||
|
|
|
|||
|
|
@ -5,6 +5,14 @@ use serde_json::{Map, Value};
|
|||
use crate::constants::UPSTREAM_ERROR_BODY_MAX_CHARS;
|
||||
use crate::error::{CoreError, CoreResult, json_type_name};
|
||||
|
||||
pub(crate) fn classify_send_error(error: reqwest::Error) -> CoreError {
|
||||
if error.is_connect() || error.is_builder() {
|
||||
CoreError::Connect(error.to_string())
|
||||
} else {
|
||||
CoreError::Network(error.to_string())
|
||||
}
|
||||
}
|
||||
|
||||
/// Bound an upstream error body before it crosses a host boundary, so provider
|
||||
/// bodies stay data-minimized.
|
||||
pub fn truncate_error_body(body: &str) -> String {
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
use crate::constants::ANTHROPIC_MESSAGES_PROVIDER;
|
||||
use crate::error::{CoreError, CoreResult};
|
||||
use crate::error::{CoreError, CoreResult, as_response_error};
|
||||
use crate::http_utils::classify_send_error;
|
||||
|
||||
use super::client::http_client;
|
||||
use super::common_utils::truncate_error_body;
|
||||
|
|
@ -16,10 +17,7 @@ pub(super) async fn execute_messages_provider_call(
|
|||
request_builder = request_builder.timeout(duration);
|
||||
}
|
||||
|
||||
let response = request_builder
|
||||
.send()
|
||||
.await
|
||||
.map_err(|err| CoreError::Network(err.to_string()))?;
|
||||
let response = request_builder.send().await.map_err(classify_send_error)?;
|
||||
|
||||
let status = response.status();
|
||||
let text = response
|
||||
|
|
@ -37,7 +35,10 @@ pub(super) async fn execute_messages_provider_call(
|
|||
let response = serde_json::from_str(&text).map_err(|err| {
|
||||
CoreError::InvalidResponse(format!("invalid messages response JSON: {err}"))
|
||||
})?;
|
||||
request.config.transform_response(&request.model, response)
|
||||
request
|
||||
.config
|
||||
.transform_response(&request.model, response)
|
||||
.map_err(as_response_error)
|
||||
}
|
||||
|
||||
pub(super) async fn execute_messages_provider_stream(
|
||||
|
|
|
|||
|
|
@ -4,13 +4,46 @@ use serde_json::{Map, Value, json};
|
|||
use tokio::io::{AsyncReadExt, AsyncWriteExt};
|
||||
use tokio::net::{TcpListener, TcpStream};
|
||||
|
||||
use crate::error::CoreError;
|
||||
use crate::error::{CoreError, CoreResult};
|
||||
|
||||
use super::common_utils::{
|
||||
has_bearer_auth, has_header, messages_provider_config, string_headers, truncate_error_body,
|
||||
};
|
||||
use super::handler::execute_messages_provider_call;
|
||||
use super::messages;
|
||||
use super::types::MessagesRequest;
|
||||
use super::transformation::AnthropicMessagesProviderConfig;
|
||||
use super::types::{AnthropicMessagesResponse, MessagesRequest, ProviderMessagesRequest};
|
||||
|
||||
struct RejectingResponseConfig;
|
||||
|
||||
impl AnthropicMessagesProviderConfig for RejectingResponseConfig {
|
||||
fn complete_url(
|
||||
&self,
|
||||
_api_base: Option<&str>,
|
||||
_model: &str,
|
||||
_env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> CoreResult<String> {
|
||||
unreachable!()
|
||||
}
|
||||
|
||||
fn resolve_api_key(
|
||||
&self,
|
||||
_api_key: Option<&str>,
|
||||
_env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> CoreResult<String> {
|
||||
unreachable!()
|
||||
}
|
||||
|
||||
fn transform_response(
|
||||
&self,
|
||||
_model: &str,
|
||||
_response: AnthropicMessagesResponse,
|
||||
) -> CoreResult<AnthropicMessagesResponse> {
|
||||
Err(CoreError::MissingField("normalized_content"))
|
||||
}
|
||||
}
|
||||
|
||||
static REJECTING_RESPONSE_CONFIG: RejectingResponseConfig = RejectingResponseConfig;
|
||||
|
||||
async fn read_http_request(socket: &mut TcpStream) -> String {
|
||||
let mut request = Vec::new();
|
||||
|
|
@ -178,6 +211,39 @@ async fn messages_round_trip_builds_azure_request_and_passes_response_through()
|
|||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn post_response_transform_errors_are_non_retryable() {
|
||||
let listener = TcpListener::bind("127.0.0.1:0").await.expect("binds");
|
||||
let addr = listener.local_addr().expect("addr");
|
||||
|
||||
let server = tokio::spawn(async move {
|
||||
let (mut socket, _) = listener.accept().await.expect("accepts request");
|
||||
let _ = read_http_request(&mut socket).await;
|
||||
let response_body = r#"{"id":"msg_1","type":"message","role":"assistant","content":[],"model":"claude-test"}"#;
|
||||
socket
|
||||
.write_all(write_response(response_body).as_bytes())
|
||||
.await
|
||||
.expect("writes response");
|
||||
});
|
||||
|
||||
let error = execute_messages_provider_call(ProviderMessagesRequest {
|
||||
provider: "anthropic".to_string(),
|
||||
model: "claude-test".to_string(),
|
||||
config: &REJECTING_RESPONSE_CONFIG,
|
||||
url: format!("http://{addr}/v1/messages"),
|
||||
body: json!({}),
|
||||
upstream_headers: Vec::new(),
|
||||
timeout: Some(Duration::from_secs(5)),
|
||||
})
|
||||
.await
|
||||
.expect_err("response transform should fail");
|
||||
|
||||
server.await.expect("server task completes");
|
||||
assert!(
|
||||
matches!(error, CoreError::InvalidResponse(message) if message.contains("normalized_content"))
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn messages_round_trip_builds_native_anthropic_request() {
|
||||
let listener = TcpListener::bind("127.0.0.1:0").await.expect("binds");
|
||||
|
|
@ -439,3 +505,24 @@ async fn messages_rejects_unsupported_provider() {
|
|||
|
||||
assert!(matches!(err, CoreError::InvalidProvider(provider) if provider == "openai"));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn messages_classifies_a_refused_connection_as_safe_to_fallback() {
|
||||
let port = {
|
||||
let listener = TcpListener::bind("127.0.0.1:0").await.expect("binds");
|
||||
listener.local_addr().expect("has an address").port()
|
||||
};
|
||||
let error = messages(MessagesRequest {
|
||||
model: "claude-test",
|
||||
body: json!({"model": "claude-test", "max_tokens": 8, "messages": []}),
|
||||
api_key: Some("sk"),
|
||||
api_base: Some(&format!("http://127.0.0.1:{port}")),
|
||||
custom_llm_provider: Some("anthropic"),
|
||||
extra_headers: None,
|
||||
timeout: Some(Duration::from_secs(1)),
|
||||
})
|
||||
.await
|
||||
.expect_err("nothing is listening");
|
||||
|
||||
assert!(matches!(error, CoreError::Connect(_)));
|
||||
}
|
||||
|
|
|
|||
|
|
@ -16,6 +16,7 @@ extension-module = ["pyo3/extension-module"]
|
|||
panic-test = []
|
||||
|
||||
[dependencies]
|
||||
futures-util.workspace = true
|
||||
litellm-core = { workspace = true, features = ["bedrock-auth"] }
|
||||
litellm-ai-gateway = { workspace = true, default-features = false }
|
||||
litellm-python-interop.workspace = true
|
||||
|
|
|
|||
|
|
@ -33,7 +33,7 @@ pub(crate) fn core_error_to_pyerr(err: CoreError) -> PyErr {
|
|||
/// Everything raised before the request goes out is safe for the host to retry
|
||||
/// on its own path; anything after it is not, because the provider has already
|
||||
/// done the work and billed for it.
|
||||
pub(crate) fn chat_completions_error_to_pyerr(err: CoreError) -> PyErr {
|
||||
pub(crate) fn fallback_route_error_to_pyerr(err: CoreError) -> PyErr {
|
||||
match err {
|
||||
CoreError::Unsupported(_)
|
||||
| CoreError::Auth(_)
|
||||
|
|
@ -59,3 +59,60 @@ pub(crate) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> {
|
|||
module.add("RustBridgeDeclined", py.get_type::<RustBridgeDeclined>())?;
|
||||
module.add("RustUpstreamError", py.get_type::<RustUpstreamError>())
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn fallback_routes_distinguish_declines_from_upstream_failures() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let declines = [
|
||||
CoreError::Unsupported("unsupported"),
|
||||
CoreError::Auth("missing key".to_string()),
|
||||
CoreError::InvalidProvider("unsupported".to_string()),
|
||||
CoreError::InvalidRequest("invalid".to_string()),
|
||||
CoreError::InvalidType {
|
||||
expected: "string",
|
||||
actual: "number",
|
||||
},
|
||||
CoreError::MissingField("model"),
|
||||
CoreError::Routing("no route".to_string()),
|
||||
CoreError::Connect("connection refused".to_string()),
|
||||
];
|
||||
for error in declines {
|
||||
let mapped = fallback_route_error_to_pyerr(error);
|
||||
assert!(mapped.is_instance_of::<RustBridgeDeclined>(py));
|
||||
}
|
||||
|
||||
let upstream_failures = [
|
||||
(
|
||||
CoreError::Http {
|
||||
status: 429,
|
||||
body: "rate limited".to_string(),
|
||||
},
|
||||
(429, "429: rate limited"),
|
||||
),
|
||||
(
|
||||
CoreError::Network("request timed out".to_string()),
|
||||
(0, "request timed out"),
|
||||
),
|
||||
(
|
||||
CoreError::InvalidResponse("bad JSON".to_string()),
|
||||
(0, "bad JSON"),
|
||||
),
|
||||
];
|
||||
for (error, expected) in upstream_failures {
|
||||
let mapped = fallback_route_error_to_pyerr(error);
|
||||
assert!(mapped.is_instance_of::<RustUpstreamError>(py));
|
||||
let args: (u16, String) = mapped
|
||||
.value(py)
|
||||
.getattr("args")
|
||||
.and_then(|args| args.extract())
|
||||
.expect("upstream error should carry status and message");
|
||||
assert_eq!(args, (expected.0, expected.1.to_string()));
|
||||
}
|
||||
});
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -3,7 +3,10 @@ mod errors;
|
|||
mod marshal;
|
||||
mod routes;
|
||||
|
||||
use std::panic::{AssertUnwindSafe, catch_unwind};
|
||||
|
||||
use litellm_ai_gateway::io::responses_ws::ResponsesWebSocketConnection as RustResponsesWebSocketConnection;
|
||||
use litellm_python_interop::panic_to_pyerr;
|
||||
use pyo3::prelude::*;
|
||||
use pyo3::types::PyAny;
|
||||
|
||||
|
|
@ -15,6 +18,20 @@ struct ResponsesWebSocketConnection {
|
|||
inner: RustResponsesWebSocketConnection,
|
||||
}
|
||||
|
||||
struct NewResponsesWebSocketConnection(ResponsesWebSocketConnection);
|
||||
|
||||
impl<'py> IntoPyObject<'py> for NewResponsesWebSocketConnection {
|
||||
type Target = ResponsesWebSocketConnection;
|
||||
type Output = Bound<'py, ResponsesWebSocketConnection>;
|
||||
type Error = PyErr;
|
||||
|
||||
fn into_pyobject(self, py: Python<'py>) -> PyResult<Self::Output> {
|
||||
catch_unwind(AssertUnwindSafe(|| Py::new(py, self.0)))
|
||||
.map_err(panic_to_pyerr)?
|
||||
.map(|value| value.into_bound(py))
|
||||
}
|
||||
}
|
||||
|
||||
#[pymethods]
|
||||
impl ResponsesWebSocketConnection {
|
||||
#[classmethod]
|
||||
|
|
@ -32,7 +49,9 @@ impl ResponsesWebSocketConnection {
|
|||
let inner = RustResponsesWebSocketConnection::connect_url(&url, &headers, timeout)
|
||||
.await
|
||||
.map_err(core_error_to_pyerr)?;
|
||||
Python::attach(|py| Py::new(py, ResponsesWebSocketConnection { inner }))
|
||||
Ok(NewResponsesWebSocketConnection(
|
||||
ResponsesWebSocketConnection { inner },
|
||||
))
|
||||
})
|
||||
}
|
||||
|
||||
|
|
@ -93,13 +112,15 @@ mod tests {
|
|||
"gil_stats",
|
||||
];
|
||||
|
||||
for name in expected {
|
||||
assert!(
|
||||
module
|
||||
.hasattr(name)
|
||||
.expect("attribute lookup should succeed")
|
||||
);
|
||||
}
|
||||
let public_names: Vec<String> = module
|
||||
.dict()
|
||||
.keys()
|
||||
.extract::<Vec<String>>()
|
||||
.expect("module names should be strings")
|
||||
.into_iter()
|
||||
.filter(|name| !name.starts_with("__"))
|
||||
.collect();
|
||||
assert_eq!(public_names, expected);
|
||||
});
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -15,23 +15,24 @@ pub(crate) struct RouteOptions {
|
|||
pub(crate) timeout: Option<Duration>,
|
||||
}
|
||||
|
||||
pub(crate) struct RouteOptionsInputs {
|
||||
pub(crate) model: String,
|
||||
pub(crate) api_key: Option<String>,
|
||||
pub(crate) api_base: Option<String>,
|
||||
pub(crate) custom_llm_provider: Option<String>,
|
||||
pub(crate) extra_headers: Option<Py<PyAny>>,
|
||||
pub(crate) timeout_seconds: Option<f64>,
|
||||
}
|
||||
|
||||
impl RouteOptions {
|
||||
pub(crate) fn from_python(
|
||||
py: Python<'_>,
|
||||
model: String,
|
||||
api_key: Option<String>,
|
||||
api_base: Option<String>,
|
||||
custom_llm_provider: Option<String>,
|
||||
extra_headers: Option<Py<PyAny>>,
|
||||
timeout_seconds: Option<f64>,
|
||||
) -> PyResult<Self> {
|
||||
pub(crate) fn from_python(py: Python<'_>, inputs: RouteOptionsInputs) -> PyResult<Self> {
|
||||
Ok(Self {
|
||||
model,
|
||||
api_key,
|
||||
api_base,
|
||||
custom_llm_provider,
|
||||
extra_headers: optional_object(py, "extra_headers", extra_headers)?,
|
||||
timeout: optional_timeout(timeout_seconds),
|
||||
model: inputs.model,
|
||||
api_key: inputs.api_key,
|
||||
api_base: inputs.api_base,
|
||||
custom_llm_provider: inputs.custom_llm_provider,
|
||||
extra_headers: optional_object(py, "extra_headers", inputs.extra_headers)?,
|
||||
timeout: optional_timeout(inputs.timeout_seconds),
|
||||
})
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,41 +1,35 @@
|
|||
use std::future::Future;
|
||||
|
||||
use litellm_ai_gateway::io::audio_transcription::{
|
||||
AudioTranscriptionRequest, audio_transcription as run_audio_transcription,
|
||||
};
|
||||
use litellm_core::error::CoreResult;
|
||||
use litellm_python_interop::from_py;
|
||||
use pyo3::prelude::*;
|
||||
use serde_json::{Map, Value};
|
||||
use serde_json::Value;
|
||||
|
||||
use crate::errors::core_error_to_pyerr;
|
||||
use crate::marshal::{RouteOptions, object_or_empty};
|
||||
use crate::routes::BridgeRoute;
|
||||
use crate::marshal::{RouteOptions, RouteOptionsInputs, object_or_empty};
|
||||
|
||||
struct AudioTranscriptionCall {
|
||||
options: RouteOptions,
|
||||
audio: Value,
|
||||
optional_params: Map<String, Value>,
|
||||
}
|
||||
fn prepare_transcription(
|
||||
py: Python<'_>,
|
||||
inputs: AudioTranscriptionInputs,
|
||||
) -> PyResult<impl Future<Output = CoreResult<Value>> + Send + 'static> {
|
||||
let audio = from_py(inputs.audio.bind(py))?;
|
||||
let options = RouteOptions::from_python(
|
||||
py,
|
||||
RouteOptionsInputs {
|
||||
model: inputs.model,
|
||||
api_key: inputs.api_key,
|
||||
api_base: inputs.api_base,
|
||||
custom_llm_provider: inputs.custom_llm_provider,
|
||||
extra_headers: inputs.extra_headers,
|
||||
timeout_seconds: inputs.timeout_seconds,
|
||||
},
|
||||
)?;
|
||||
let optional_params = object_or_empty(py, "optional_params", inputs.optional_params)?;
|
||||
|
||||
impl BridgeRoute<AudioTranscriptionInputs> for AudioTranscriptionCall {
|
||||
type Output = Value;
|
||||
|
||||
fn from_python(py: Python<'_>, inputs: AudioTranscriptionInputs) -> PyResult<Self> {
|
||||
Ok(Self {
|
||||
options: RouteOptions::from_python(
|
||||
py,
|
||||
inputs.model,
|
||||
inputs.api_key,
|
||||
inputs.api_base,
|
||||
inputs.custom_llm_provider,
|
||||
inputs.extra_headers,
|
||||
inputs.timeout_seconds,
|
||||
)?,
|
||||
audio: from_py(inputs.audio.bind(py))?,
|
||||
optional_params: object_or_empty(py, "optional_params", inputs.optional_params)?,
|
||||
})
|
||||
}
|
||||
|
||||
async fn run(self) -> CoreResult<Value> {
|
||||
Ok(async move {
|
||||
let RouteOptions {
|
||||
model,
|
||||
api_key,
|
||||
|
|
@ -43,15 +37,15 @@ impl BridgeRoute<AudioTranscriptionInputs> for AudioTranscriptionCall {
|
|||
custom_llm_provider,
|
||||
extra_headers,
|
||||
timeout,
|
||||
} = self.options;
|
||||
} = options;
|
||||
run_audio_transcription(AudioTranscriptionRequest {
|
||||
model: &model,
|
||||
audio: self.audio,
|
||||
audio,
|
||||
api_key: api_key.as_deref(),
|
||||
api_base: api_base.as_deref(),
|
||||
custom_llm_provider: custom_llm_provider.as_deref(),
|
||||
extra_headers,
|
||||
optional_params: self.optional_params,
|
||||
optional_params,
|
||||
timeout,
|
||||
callbacks: Vec::new(),
|
||||
guardrails: Vec::new(),
|
||||
|
|
@ -59,7 +53,7 @@ impl BridgeRoute<AudioTranscriptionInputs> for AudioTranscriptionCall {
|
|||
litellm_call_id: None,
|
||||
})
|
||||
.await
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
bridge_route! {
|
||||
|
|
@ -78,6 +72,6 @@ bridge_route! {
|
|||
optional_params: Option<Py<PyAny>>,
|
||||
timeout_seconds: Option<f64>,
|
||||
},
|
||||
call = AudioTranscriptionCall,
|
||||
prepare = prepare_transcription,
|
||||
errors = core_error_to_pyerr,
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,3 +1,5 @@
|
|||
use std::future::Future;
|
||||
|
||||
use litellm_core::chat_completions::types::{ChatCompletionsRequest, ChatCompletionsResponse};
|
||||
use litellm_core::chat_completions::{
|
||||
chat_completions as run_chat_completions, chat_completions_decline_reason,
|
||||
|
|
@ -5,38 +7,30 @@ use litellm_core::chat_completions::{
|
|||
use litellm_core::error::CoreResult;
|
||||
use litellm_python_interop::from_py;
|
||||
use pyo3::prelude::*;
|
||||
use serde_json::{Map, Value};
|
||||
use serde_json::Value;
|
||||
|
||||
use crate::errors::chat_completions_error_to_pyerr;
|
||||
use crate::marshal::{RouteOptions, object_or_empty, required_value};
|
||||
use crate::routes::BridgeRoute;
|
||||
use crate::errors::fallback_route_error_to_pyerr;
|
||||
use crate::marshal::{RouteOptions, RouteOptionsInputs, object_or_empty, required_value};
|
||||
|
||||
struct ChatCompletionsCall {
|
||||
options: RouteOptions,
|
||||
messages: Value,
|
||||
optional_params: Map<String, Value>,
|
||||
}
|
||||
fn prepare_chat_completions(
|
||||
py: Python<'_>,
|
||||
inputs: ChatCompletionsInputs,
|
||||
) -> PyResult<impl Future<Output = CoreResult<ChatCompletionsResponse>> + Send + 'static> {
|
||||
let messages = required_value(py, "messages", inputs.messages, Value::is_array, "list")?;
|
||||
let optional_params = object_or_empty(py, "optional_params", inputs.optional_params)?;
|
||||
let options = RouteOptions::from_python(
|
||||
py,
|
||||
RouteOptionsInputs {
|
||||
model: inputs.model,
|
||||
api_key: inputs.api_key,
|
||||
api_base: inputs.api_base,
|
||||
custom_llm_provider: inputs.custom_llm_provider,
|
||||
extra_headers: inputs.extra_headers,
|
||||
timeout_seconds: inputs.timeout_seconds,
|
||||
},
|
||||
)?;
|
||||
|
||||
impl BridgeRoute<ChatCompletionsInputs> for ChatCompletionsCall {
|
||||
type Output = ChatCompletionsResponse;
|
||||
|
||||
fn from_python(py: Python<'_>, inputs: ChatCompletionsInputs) -> PyResult<Self> {
|
||||
Ok(Self {
|
||||
options: RouteOptions::from_python(
|
||||
py,
|
||||
inputs.model,
|
||||
inputs.api_key,
|
||||
inputs.api_base,
|
||||
inputs.custom_llm_provider,
|
||||
inputs.extra_headers,
|
||||
inputs.timeout_seconds,
|
||||
)?,
|
||||
messages: required_value(py, "messages", inputs.messages, Value::is_array, "list")?,
|
||||
optional_params: object_or_empty(py, "optional_params", inputs.optional_params)?,
|
||||
})
|
||||
}
|
||||
|
||||
async fn run(self) -> CoreResult<ChatCompletionsResponse> {
|
||||
Ok(async move {
|
||||
let RouteOptions {
|
||||
model,
|
||||
api_key,
|
||||
|
|
@ -44,11 +38,11 @@ impl BridgeRoute<ChatCompletionsInputs> for ChatCompletionsCall {
|
|||
custom_llm_provider,
|
||||
extra_headers,
|
||||
timeout,
|
||||
} = self.options;
|
||||
} = options;
|
||||
run_chat_completions(ChatCompletionsRequest {
|
||||
model: &model,
|
||||
messages: self.messages,
|
||||
optional_params: self.optional_params,
|
||||
messages,
|
||||
optional_params,
|
||||
api_key: api_key.as_deref(),
|
||||
api_base: api_base.as_deref(),
|
||||
custom_llm_provider: custom_llm_provider.as_deref(),
|
||||
|
|
@ -56,7 +50,7 @@ impl BridgeRoute<ChatCompletionsInputs> for ChatCompletionsCall {
|
|||
timeout,
|
||||
})
|
||||
.await
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
#[pyfunction]
|
||||
|
|
@ -95,7 +89,7 @@ bridge_route! {
|
|||
extra_headers: Option<Py<PyAny>>,
|
||||
timeout_seconds: Option<f64>,
|
||||
},
|
||||
call = ChatCompletionsCall,
|
||||
errors = chat_completions_error_to_pyerr,
|
||||
prepare = prepare_chat_completions,
|
||||
errors = fallback_route_error_to_pyerr,
|
||||
extra = [chat_completions_decline],
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,37 +1,32 @@
|
|||
use std::future::Future;
|
||||
|
||||
use litellm_core::error::CoreResult;
|
||||
use litellm_core::messages::messages as run_messages;
|
||||
use litellm_core::messages::types::{AnthropicMessagesResponse, MessagesRequest};
|
||||
use pyo3::prelude::*;
|
||||
use serde_json::Value;
|
||||
|
||||
use crate::errors::core_error_to_pyerr;
|
||||
use crate::marshal::{RouteOptions, required_value};
|
||||
use crate::routes::BridgeRoute;
|
||||
use crate::errors::fallback_route_error_to_pyerr;
|
||||
use crate::marshal::{RouteOptions, RouteOptionsInputs, required_value};
|
||||
|
||||
struct MessagesCall {
|
||||
options: RouteOptions,
|
||||
body: Value,
|
||||
}
|
||||
fn prepare_messages(
|
||||
py: Python<'_>,
|
||||
inputs: MessagesInputs,
|
||||
) -> PyResult<impl Future<Output = CoreResult<AnthropicMessagesResponse>> + Send + 'static> {
|
||||
let body = required_value(py, "body", inputs.body, Value::is_object, "dict")?;
|
||||
let options = RouteOptions::from_python(
|
||||
py,
|
||||
RouteOptionsInputs {
|
||||
model: inputs.model,
|
||||
api_key: inputs.api_key,
|
||||
api_base: inputs.api_base,
|
||||
custom_llm_provider: inputs.custom_llm_provider,
|
||||
extra_headers: inputs.extra_headers,
|
||||
timeout_seconds: inputs.timeout_seconds,
|
||||
},
|
||||
)?;
|
||||
|
||||
impl BridgeRoute<MessagesInputs> for MessagesCall {
|
||||
type Output = AnthropicMessagesResponse;
|
||||
|
||||
fn from_python(py: Python<'_>, inputs: MessagesInputs) -> PyResult<Self> {
|
||||
Ok(Self {
|
||||
options: RouteOptions::from_python(
|
||||
py,
|
||||
inputs.model,
|
||||
inputs.api_key,
|
||||
inputs.api_base,
|
||||
inputs.custom_llm_provider,
|
||||
inputs.extra_headers,
|
||||
inputs.timeout_seconds,
|
||||
)?,
|
||||
body: required_value(py, "body", inputs.body, Value::is_object, "dict")?,
|
||||
})
|
||||
}
|
||||
|
||||
async fn run(self) -> CoreResult<AnthropicMessagesResponse> {
|
||||
Ok(async move {
|
||||
let RouteOptions {
|
||||
model,
|
||||
api_key,
|
||||
|
|
@ -39,10 +34,10 @@ impl BridgeRoute<MessagesInputs> for MessagesCall {
|
|||
custom_llm_provider,
|
||||
extra_headers,
|
||||
timeout,
|
||||
} = self.options;
|
||||
} = options;
|
||||
run_messages(MessagesRequest {
|
||||
model: &model,
|
||||
body: self.body,
|
||||
body,
|
||||
api_key: api_key.as_deref(),
|
||||
api_base: api_base.as_deref(),
|
||||
custom_llm_provider: custom_llm_provider.as_deref(),
|
||||
|
|
@ -50,7 +45,7 @@ impl BridgeRoute<MessagesInputs> for MessagesCall {
|
|||
timeout,
|
||||
})
|
||||
.await
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
bridge_route! {
|
||||
|
|
@ -68,6 +63,6 @@ bridge_route! {
|
|||
extra_headers: Option<Py<PyAny>>,
|
||||
timeout_seconds: Option<f64>,
|
||||
},
|
||||
call = MessagesCall,
|
||||
errors = core_error_to_pyerr,
|
||||
prepare = prepare_messages,
|
||||
errors = fallback_route_error_to_pyerr,
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,29 +1,19 @@
|
|||
use std::future::Future;
|
||||
|
||||
use litellm_core::error::CoreResult;
|
||||
use pyo3::exceptions::PyRuntimeError;
|
||||
use pyo3::prelude::*;
|
||||
use serde::Serialize;
|
||||
use pyo3::types::PyCFunction;
|
||||
|
||||
mod runtime;
|
||||
|
||||
use runtime::{run_async, run_sync};
|
||||
|
||||
trait BridgeRoute<I>: Sized {
|
||||
type Output: Serialize + Send + 'static;
|
||||
|
||||
fn from_python(py: Python<'_>, inputs: I) -> PyResult<Self>;
|
||||
|
||||
fn run(self) -> impl Future<Output = CoreResult<Self::Output>> + Send + 'static;
|
||||
}
|
||||
|
||||
macro_rules! bridge_route {
|
||||
(
|
||||
sync = $sync_name:ident,
|
||||
asynchronous = $async_name:ident,
|
||||
inputs = $inputs:ident,
|
||||
required = { $($required_name:ident: $required_type:ty),* $(,)? },
|
||||
required = { $($required_name:ident: $required_type:ty),+ $(,)? },
|
||||
optional = { $($optional_name:ident: $optional_type:ty),* $(,)? },
|
||||
call = $call:ty,
|
||||
prepare = $prepare:path,
|
||||
errors = $map_error:path
|
||||
$(, extra = [$($extra:ident),* $(,)?])?
|
||||
$(,)?
|
||||
|
|
@ -41,15 +31,11 @@ macro_rules! bridge_route {
|
|||
$($required_name: $required_type,)*
|
||||
$($optional_name: $optional_type),*
|
||||
) -> pyo3::PyResult<pyo3::Py<pyo3::PyAny>> {
|
||||
let call = <$call as crate::routes::BridgeRoute<$inputs>>::from_python(py, $inputs {
|
||||
let future = $prepare(py, $inputs {
|
||||
$($required_name,)*
|
||||
$($optional_name),*
|
||||
})?;
|
||||
crate::routes::run_sync(
|
||||
py,
|
||||
<$call as crate::routes::BridgeRoute<$inputs>>::run(call),
|
||||
$map_error,
|
||||
)
|
||||
crate::routes::run_sync(py, future, $map_error)
|
||||
}
|
||||
|
||||
#[pyfunction]
|
||||
|
|
@ -60,47 +46,119 @@ macro_rules! bridge_route {
|
|||
$($required_name: $required_type,)*
|
||||
$($optional_name: $optional_type),*
|
||||
) -> pyo3::PyResult<pyo3::Bound<'_, pyo3::PyAny>> {
|
||||
let call = <$call as crate::routes::BridgeRoute<$inputs>>::from_python(py, $inputs {
|
||||
let future = $prepare(py, $inputs {
|
||||
$($required_name,)*
|
||||
$($optional_name),*
|
||||
})?;
|
||||
crate::routes::run_async(
|
||||
py,
|
||||
<$call as crate::routes::BridgeRoute<$inputs>>::run(call),
|
||||
$map_error,
|
||||
)
|
||||
crate::routes::run_async(py, future, $map_error)
|
||||
}
|
||||
|
||||
pub(super) fn register(
|
||||
module: &pyo3::Bound<'_, pyo3::types::PyModule>,
|
||||
) -> pyo3::PyResult<()> {
|
||||
module.add_function(pyo3::wrap_pyfunction!($sync_name, module)?)?;
|
||||
module.add_function(pyo3::wrap_pyfunction!($async_name, module)?)?;
|
||||
$($(module.add_function(pyo3::wrap_pyfunction!($extra, module)?)?;)*)?
|
||||
$($(crate::routes::add_function(module, pyo3::wrap_pyfunction!($extra, module)?)?;)*)?
|
||||
crate::routes::add_function(module, pyo3::wrap_pyfunction!($sync_name, module)?)?;
|
||||
crate::routes::add_function(module, pyo3::wrap_pyfunction!($async_name, module)?)?;
|
||||
Ok(())
|
||||
}
|
||||
};
|
||||
}
|
||||
|
||||
macro_rules! routes {
|
||||
($($route:ident),* $(,)?) => {
|
||||
$(mod $route;)*
|
||||
|
||||
pub(crate) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> {
|
||||
$($route::register(module)?;)*
|
||||
Ok(())
|
||||
}
|
||||
};
|
||||
fn add_function(module: &Bound<'_, PyModule>, function: Bound<'_, PyCFunction>) -> PyResult<()> {
|
||||
let name: String = function.getattr("__name__")?.extract()?;
|
||||
if module.hasattr(&name)? {
|
||||
return Err(PyRuntimeError::new_err(format!(
|
||||
"duplicate native route: {name}"
|
||||
)));
|
||||
}
|
||||
module.add_function(function)
|
||||
}
|
||||
|
||||
routes!(ocr, audio_transcription, messages, chat_completions);
|
||||
mod audio_transcription;
|
||||
mod chat_completions;
|
||||
mod messages;
|
||||
mod ocr;
|
||||
|
||||
pub(crate) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> {
|
||||
ocr::register(module)?;
|
||||
audio_transcription::register(module)?;
|
||||
messages::register(module)?;
|
||||
chat_completions::register(module)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::ffi::CString;
|
||||
use std::sync::atomic::{AtomicBool, Ordering};
|
||||
|
||||
use litellm_core::error::{CoreError, CoreResult};
|
||||
use pyo3::exceptions::PyLookupError;
|
||||
use pyo3::types::{PyDict, PyList};
|
||||
|
||||
use super::*;
|
||||
|
||||
mod synthetic {
|
||||
use std::future::{Future, pending};
|
||||
|
||||
use super::*;
|
||||
|
||||
static FUTURE_DROPPED: AtomicBool = AtomicBool::new(false);
|
||||
|
||||
struct DropGuard;
|
||||
|
||||
impl Drop for DropGuard {
|
||||
fn drop(&mut self) {
|
||||
FUTURE_DROPPED.store(true, Ordering::SeqCst);
|
||||
}
|
||||
}
|
||||
|
||||
#[pyfunction]
|
||||
fn future_dropped() -> bool {
|
||||
FUTURE_DROPPED.load(Ordering::SeqCst)
|
||||
}
|
||||
|
||||
bridge_route! {
|
||||
sync = echo,
|
||||
asynchronous = aecho,
|
||||
inputs = EchoInputs,
|
||||
required = { value: String },
|
||||
optional = {},
|
||||
prepare = prepare_echo,
|
||||
errors = map_error,
|
||||
extra = [future_dropped],
|
||||
}
|
||||
|
||||
fn prepare_echo(
|
||||
_py: Python<'_>,
|
||||
inputs: EchoInputs,
|
||||
) -> PyResult<impl Future<Output = CoreResult<String>> + Send + 'static> {
|
||||
FUTURE_DROPPED.store(false, Ordering::SeqCst);
|
||||
let drop_guard = (inputs.value == "pending").then_some(DropGuard);
|
||||
Ok(async move {
|
||||
let _drop_guard = drop_guard;
|
||||
tokio::task::yield_now().await;
|
||||
match inputs.value.as_str() {
|
||||
"error" => Err(CoreError::InvalidRequest("synthetic error".to_string())),
|
||||
"map_panic" => Err(CoreError::InvalidRequest("panic in mapper".to_string())),
|
||||
"panic" => panic!("synthetic panic"),
|
||||
"pending" => {
|
||||
pending::<()>().await;
|
||||
unreachable!()
|
||||
}
|
||||
_ => Ok(inputs.value),
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
fn map_error(error: CoreError) -> PyErr {
|
||||
if matches!(&error, CoreError::InvalidRequest(message) if message == "panic in mapper")
|
||||
{
|
||||
panic!("synthetic mapper panic")
|
||||
}
|
||||
PyLookupError::new_err(error.to_string())
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn sync_and_async_route_signatures_match_the_python_contract() {
|
||||
Python::initialize();
|
||||
|
|
@ -215,4 +273,205 @@ mod tests {
|
|||
}
|
||||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn route_input_validation_preserves_left_to_right_order() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let module = PyModule::new(py, "routes").expect("module should be created");
|
||||
register(&module).expect("routes should register");
|
||||
let invalid = PyList::empty(py);
|
||||
|
||||
let chat_kwargs = PyDict::new(py);
|
||||
chat_kwargs
|
||||
.set_item("optional_params", &invalid)
|
||||
.expect("kwargs should accept optional_params");
|
||||
chat_kwargs
|
||||
.set_item("extra_headers", &invalid)
|
||||
.expect("kwargs should accept extra_headers");
|
||||
let invalid_messages = PyDict::new(py);
|
||||
let error = module
|
||||
.getattr("chat_completions")
|
||||
.and_then(|function| {
|
||||
function.call(("model", &invalid_messages), Some(&chat_kwargs))
|
||||
})
|
||||
.expect_err("messages should be validated first");
|
||||
assert_eq!(error.to_string(), "ValueError: messages must be a list");
|
||||
|
||||
let valid_messages = PyList::empty(py);
|
||||
let error = module
|
||||
.getattr("chat_completions")
|
||||
.and_then(|function| function.call(("model", &valid_messages), Some(&chat_kwargs)))
|
||||
.expect_err("optional_params should be validated before headers");
|
||||
assert_eq!(
|
||||
error.to_string(),
|
||||
"ValueError: optional_params must be a dict"
|
||||
);
|
||||
|
||||
let headers_kwargs = PyDict::new(py);
|
||||
headers_kwargs
|
||||
.set_item("extra_headers", &invalid)
|
||||
.expect("kwargs should accept extra_headers");
|
||||
let invalid_body = PyList::empty(py);
|
||||
let error = module
|
||||
.getattr("messages")
|
||||
.and_then(|function| function.call(("model", &invalid_body), Some(&headers_kwargs)))
|
||||
.expect_err("body should be validated before headers");
|
||||
assert_eq!(error.to_string(), "ValueError: body must be a dict");
|
||||
|
||||
let invalid_payload =
|
||||
PyModule::new(py, "invalid_payload").expect("invalid payload should be created");
|
||||
for name in ["ocr", "transcription"] {
|
||||
let error = module
|
||||
.getattr(name)
|
||||
.and_then(|function| {
|
||||
function.call(("model", &invalid_payload), Some(&headers_kwargs))
|
||||
})
|
||||
.expect_err("payload should be validated before headers");
|
||||
assert!(!error.to_string().contains("extra_headers"));
|
||||
}
|
||||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn generated_routes_execute_sync_and_async_contracts() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let module = PyModule::new(py, "synthetic").expect("module should be created");
|
||||
synthetic::register(&module).expect("routes should register");
|
||||
|
||||
let sync_value: String = module
|
||||
.getattr("echo")
|
||||
.and_then(|function| function.call1(("sync",)))
|
||||
.and_then(|value| value.extract())
|
||||
.expect("sync route should return its value");
|
||||
assert_eq!(sync_value, "sync");
|
||||
|
||||
let sync_error = module
|
||||
.getattr("echo")
|
||||
.and_then(|function| function.call1(("error",)))
|
||||
.expect_err("sync route should map its error");
|
||||
assert!(sync_error.is_instance_of::<PyLookupError>(py));
|
||||
assert_eq!(
|
||||
sync_error.to_string(),
|
||||
"LookupError: invalid request: synthetic error"
|
||||
);
|
||||
|
||||
let locals = PyDict::new(py);
|
||||
locals
|
||||
.set_item("routes", &module)
|
||||
.expect("module should enter Python locals");
|
||||
let code = CString::new(
|
||||
r#"
|
||||
import asyncio
|
||||
|
||||
async def exercise():
|
||||
assert await routes.aecho("async") == "async"
|
||||
|
||||
try:
|
||||
await routes.aecho("error")
|
||||
except LookupError as error:
|
||||
assert str(error) == "invalid request: synthetic error"
|
||||
else:
|
||||
raise AssertionError("mapped error was not raised")
|
||||
|
||||
try:
|
||||
await routes.aecho("panic")
|
||||
except BaseException as error:
|
||||
assert type(error).__name__ == "PanicException"
|
||||
assert str(error) == "synthetic panic"
|
||||
else:
|
||||
raise AssertionError("panic was not raised")
|
||||
|
||||
try:
|
||||
await routes.aecho("map_panic")
|
||||
except BaseException as error:
|
||||
assert type(error).__name__ == "PanicException"
|
||||
assert str(error) == "synthetic mapper panic"
|
||||
else:
|
||||
raise AssertionError("mapper panic was not raised")
|
||||
|
||||
task = asyncio.ensure_future(routes.aecho("pending"))
|
||||
await asyncio.sleep(0)
|
||||
task.cancel()
|
||||
try:
|
||||
await task
|
||||
except asyncio.CancelledError:
|
||||
pass
|
||||
else:
|
||||
raise AssertionError("cancelled route completed")
|
||||
|
||||
for _ in range(100):
|
||||
if routes.future_dropped():
|
||||
break
|
||||
await asyncio.sleep(0.001)
|
||||
assert routes.future_dropped()
|
||||
|
||||
asyncio.run(exercise())
|
||||
"#,
|
||||
)
|
||||
.expect("Python source should not contain null bytes");
|
||||
py.run(&code, Some(&locals), Some(&locals))
|
||||
.expect("async route contract should hold");
|
||||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn messages_routes_map_declines_before_python_fallback() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let module = PyModule::new(py, "routes").expect("module should be created");
|
||||
register(&module).expect("routes should register");
|
||||
let body = PyDict::new(py);
|
||||
let kwargs = PyDict::new(py);
|
||||
kwargs
|
||||
.set_item("custom_llm_provider", "openai")
|
||||
.expect("kwargs should accept provider");
|
||||
|
||||
let sync_error = module
|
||||
.getattr("messages")
|
||||
.and_then(|function| function.call(("model", &body), Some(&kwargs)))
|
||||
.expect_err("unsupported provider should decline");
|
||||
assert!(sync_error.is_instance_of::<crate::errors::RustBridgeDeclined>(py));
|
||||
|
||||
let locals = PyDict::new(py);
|
||||
locals
|
||||
.set_item("routes", &module)
|
||||
.expect("module should enter Python locals");
|
||||
let code = CString::new(
|
||||
r#"
|
||||
import asyncio
|
||||
|
||||
async def exercise():
|
||||
try:
|
||||
await routes.amessages("model", {}, custom_llm_provider="openai")
|
||||
except Exception as error:
|
||||
assert type(error).__name__ == "RustBridgeDeclined"
|
||||
else:
|
||||
raise AssertionError("unsupported provider did not decline")
|
||||
|
||||
asyncio.run(exercise())
|
||||
"#,
|
||||
)
|
||||
.expect("Python source should not contain null bytes");
|
||||
py.run(&code, Some(&locals), Some(&locals))
|
||||
.expect("async route should preserve the decline contract");
|
||||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn route_registration_rejects_duplicate_python_names() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let module = PyModule::new(py, "synthetic").expect("module should be created");
|
||||
synthetic::register(&module).expect("first registration should succeed");
|
||||
let error = synthetic::register(&module)
|
||||
.expect_err("duplicate registration should be rejected");
|
||||
|
||||
assert_eq!(
|
||||
error.to_string(),
|
||||
"RuntimeError: duplicate native route: future_dropped"
|
||||
);
|
||||
});
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,39 +1,33 @@
|
|||
use std::future::Future;
|
||||
|
||||
use litellm_ai_gateway::io::ocr::{OcrRequest, ocr as run_ocr};
|
||||
use litellm_core::error::CoreResult;
|
||||
use litellm_python_interop::from_py;
|
||||
use pyo3::prelude::*;
|
||||
use serde_json::{Map, Value};
|
||||
use serde_json::Value;
|
||||
|
||||
use crate::errors::core_error_to_pyerr;
|
||||
use crate::marshal::{RouteOptions, object_or_empty};
|
||||
use crate::routes::BridgeRoute;
|
||||
use crate::marshal::{RouteOptions, RouteOptionsInputs, object_or_empty};
|
||||
|
||||
struct OcrCall {
|
||||
options: RouteOptions,
|
||||
document: Value,
|
||||
optional_params: Map<String, Value>,
|
||||
}
|
||||
fn prepare_ocr(
|
||||
py: Python<'_>,
|
||||
inputs: OcrInputs,
|
||||
) -> PyResult<impl Future<Output = CoreResult<Value>> + Send + 'static> {
|
||||
let document = from_py(inputs.document.bind(py))?;
|
||||
let options = RouteOptions::from_python(
|
||||
py,
|
||||
RouteOptionsInputs {
|
||||
model: inputs.model,
|
||||
api_key: inputs.api_key,
|
||||
api_base: inputs.api_base,
|
||||
custom_llm_provider: inputs.custom_llm_provider,
|
||||
extra_headers: inputs.extra_headers,
|
||||
timeout_seconds: inputs.timeout_seconds,
|
||||
},
|
||||
)?;
|
||||
let optional_params = object_or_empty(py, "optional_params", inputs.optional_params)?;
|
||||
|
||||
impl BridgeRoute<OcrInputs> for OcrCall {
|
||||
type Output = Value;
|
||||
|
||||
fn from_python(py: Python<'_>, inputs: OcrInputs) -> PyResult<Self> {
|
||||
Ok(Self {
|
||||
options: RouteOptions::from_python(
|
||||
py,
|
||||
inputs.model,
|
||||
inputs.api_key,
|
||||
inputs.api_base,
|
||||
inputs.custom_llm_provider,
|
||||
inputs.extra_headers,
|
||||
inputs.timeout_seconds,
|
||||
)?,
|
||||
document: from_py(inputs.document.bind(py))?,
|
||||
optional_params: object_or_empty(py, "optional_params", inputs.optional_params)?,
|
||||
})
|
||||
}
|
||||
|
||||
async fn run(self) -> CoreResult<Value> {
|
||||
Ok(async move {
|
||||
let RouteOptions {
|
||||
model,
|
||||
api_key,
|
||||
|
|
@ -41,15 +35,15 @@ impl BridgeRoute<OcrInputs> for OcrCall {
|
|||
custom_llm_provider,
|
||||
extra_headers,
|
||||
timeout,
|
||||
} = self.options;
|
||||
} = options;
|
||||
run_ocr(OcrRequest {
|
||||
model: &model,
|
||||
document: self.document,
|
||||
document,
|
||||
api_key: api_key.as_deref(),
|
||||
api_base: api_base.as_deref(),
|
||||
custom_llm_provider: custom_llm_provider.as_deref(),
|
||||
extra_headers,
|
||||
optional_params: self.optional_params,
|
||||
optional_params,
|
||||
timeout,
|
||||
callbacks: Vec::new(),
|
||||
guardrails: Vec::new(),
|
||||
|
|
@ -57,7 +51,7 @@ impl BridgeRoute<OcrInputs> for OcrCall {
|
|||
litellm_call_id: None,
|
||||
})
|
||||
.await
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
bridge_route! {
|
||||
|
|
@ -76,6 +70,6 @@ bridge_route! {
|
|||
optional_params: Option<Py<PyAny>>,
|
||||
timeout_seconds: Option<f64>,
|
||||
},
|
||||
call = OcrCall,
|
||||
prepare = prepare_ocr,
|
||||
errors = core_error_to_pyerr,
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,11 +1,15 @@
|
|||
use std::future::Future;
|
||||
use std::sync::mpsc::sync_channel;
|
||||
use std::panic::AssertUnwindSafe;
|
||||
use std::time::Duration;
|
||||
|
||||
use futures_util::FutureExt;
|
||||
use litellm_core::error::{CoreError, CoreResult};
|
||||
use litellm_python_interop::{release_gil, to_py};
|
||||
use litellm_python_interop::{Pythonized, panic_to_pyerr, release_gil, to_py};
|
||||
use pyo3::exceptions::PyRuntimeError;
|
||||
use pyo3::prelude::*;
|
||||
use serde::Serialize;
|
||||
use tokio::runtime::{Handle, Runtime};
|
||||
use tokio::time::{self, MissedTickBehavior};
|
||||
|
||||
pub(super) fn run_sync<T, F>(
|
||||
py: Python<'_>,
|
||||
|
|
@ -16,13 +20,32 @@ where
|
|||
T: Serialize + Send + 'static,
|
||||
F: Future<Output = CoreResult<T>> + Send + 'static,
|
||||
{
|
||||
let (sender, receiver) = sync_channel(1);
|
||||
pyo3_async_runtimes::tokio::get_runtime().spawn(async move {
|
||||
let _ = sender.send(future.await);
|
||||
});
|
||||
let result = release_gil(py, move || receiver.recv())
|
||||
.map_err(|_| PyRuntimeError::new_err("native route task terminated"))?
|
||||
.map_err(map_error)?;
|
||||
run_sync_on(
|
||||
py,
|
||||
pyo3_async_runtimes::tokio::get_runtime(),
|
||||
future,
|
||||
map_error,
|
||||
)
|
||||
}
|
||||
|
||||
fn run_sync_on<T, F>(
|
||||
py: Python<'_>,
|
||||
runtime: &Runtime,
|
||||
future: F,
|
||||
map_error: fn(CoreError) -> PyErr,
|
||||
) -> PyResult<Py<PyAny>>
|
||||
where
|
||||
T: Serialize + Send + 'static,
|
||||
F: Future<Output = CoreResult<T>> + Send + 'static,
|
||||
{
|
||||
if Handle::try_current().is_ok() {
|
||||
return Err(PyRuntimeError::new_err(
|
||||
"synchronous native routes cannot run from a Tokio context; use the async route",
|
||||
));
|
||||
}
|
||||
|
||||
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)
|
||||
}
|
||||
|
||||
|
|
@ -36,14 +59,61 @@ where
|
|||
F: Future<Output = CoreResult<T>> + Send + 'static,
|
||||
{
|
||||
pyo3_async_runtimes::tokio::future_into_py(py, async move {
|
||||
let result = future.await.map_err(map_error)?;
|
||||
Python::attach(|py| to_py(py, &result))
|
||||
let result = catch_route_panic(future).await?;
|
||||
let result = map_core_result(result, map_error)?;
|
||||
Ok(Pythonized(result))
|
||||
})
|
||||
}
|
||||
|
||||
fn map_core_result<T>(result: CoreResult<T>, map_error: fn(CoreError) -> PyErr) -> PyResult<T> {
|
||||
match result {
|
||||
Ok(value) => Ok(value),
|
||||
Err(error) => Err(
|
||||
std::panic::catch_unwind(AssertUnwindSafe(|| map_error(error)))
|
||||
.map_err(panic_to_pyerr)?,
|
||||
),
|
||||
}
|
||||
}
|
||||
|
||||
async fn catch_route_panic<T, F>(future: F) -> PyResult<CoreResult<T>>
|
||||
where
|
||||
F: Future<Output = CoreResult<T>>,
|
||||
{
|
||||
AssertUnwindSafe(future)
|
||||
.catch_unwind()
|
||||
.await
|
||||
.map_err(panic_to_pyerr)
|
||||
}
|
||||
|
||||
async fn wait_for_sync_result<T, F>(future: F) -> PyResult<CoreResult<T>>
|
||||
where
|
||||
F: Future<Output = CoreResult<T>>,
|
||||
{
|
||||
let future = catch_route_panic(future);
|
||||
tokio::pin!(future);
|
||||
|
||||
let signal_interval = Duration::from_millis(50);
|
||||
let mut signal_checks =
|
||||
time::interval_at(time::Instant::now() + signal_interval, signal_interval);
|
||||
signal_checks.set_missed_tick_behavior(MissedTickBehavior::Delay);
|
||||
loop {
|
||||
tokio::select! {
|
||||
result = &mut future => return result,
|
||||
_ = signal_checks.tick() => Python::attach(|py| py.check_signals())?,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use std::time::Duration;
|
||||
use std::ffi::CString;
|
||||
use std::future::poll_fn;
|
||||
use std::task::Poll;
|
||||
|
||||
use pyo3::panic::PanicException;
|
||||
use pyo3::types::{PyDict, PyModule};
|
||||
use serde::Serializer;
|
||||
use tokio::runtime::Builder;
|
||||
|
||||
use super::*;
|
||||
|
||||
|
|
@ -51,6 +121,26 @@ mod tests {
|
|||
PyRuntimeError::new_err(error.to_string())
|
||||
}
|
||||
|
||||
fn panicking_error_mapper(_error: CoreError) -> PyErr {
|
||||
panic!("error mapper panicked")
|
||||
}
|
||||
|
||||
struct PanickingOutput;
|
||||
|
||||
impl Serialize for PanickingOutput {
|
||||
fn serialize<S>(&self, _serializer: S) -> Result<S::Ok, S::Error>
|
||||
where
|
||||
S: Serializer,
|
||||
{
|
||||
panic!("async serializer panicked")
|
||||
}
|
||||
}
|
||||
|
||||
#[pyfunction]
|
||||
fn async_serialization_panic(py: Python<'_>) -> PyResult<Bound<'_, PyAny>> {
|
||||
run_async(py, async { Ok(PanickingOutput) }, runtime_error)
|
||||
}
|
||||
|
||||
fn extract_bool(py: Python<'_>, result: PyResult<Py<PyAny>>) -> bool {
|
||||
result
|
||||
.expect("route should complete")
|
||||
|
|
@ -60,13 +150,13 @@ mod tests {
|
|||
}
|
||||
|
||||
#[test]
|
||||
fn sync_runner_polls_future_on_tokio_worker() {
|
||||
fn sync_runner_polls_future_on_the_caller_thread() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let caller_thread = std::thread::current().id();
|
||||
let result = run_sync(
|
||||
py,
|
||||
async move { Ok(std::thread::current().id() != caller_thread) },
|
||||
async move { Ok(std::thread::current().id() == caller_thread) },
|
||||
runtime_error,
|
||||
);
|
||||
|
||||
|
|
@ -94,4 +184,115 @@ mod tests {
|
|||
assert!(extract_bool(py, result));
|
||||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn sync_runner_rejects_calls_from_a_tokio_context() {
|
||||
Python::initialize();
|
||||
let runtime = Builder::new_current_thread()
|
||||
.enable_all()
|
||||
.build()
|
||||
.expect("runtime should build");
|
||||
|
||||
let error = runtime.block_on(async {
|
||||
Python::attach(|py| {
|
||||
run_sync::<bool, _>(py, async { Ok(true) }, runtime_error)
|
||||
.expect_err("sync route should reject a nested Tokio runtime")
|
||||
})
|
||||
});
|
||||
|
||||
assert_eq!(
|
||||
error.to_string(),
|
||||
"RuntimeError: synchronous native routes cannot run from a Tokio context; use the async route"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn sync_runner_can_drive_a_current_thread_runtime() {
|
||||
Python::initialize();
|
||||
let runtime = Builder::new_current_thread()
|
||||
.enable_all()
|
||||
.build()
|
||||
.expect("runtime should build");
|
||||
Python::attach(|py| {
|
||||
let result = run_sync_on(
|
||||
py,
|
||||
&runtime,
|
||||
async {
|
||||
tokio::task::yield_now().await;
|
||||
Ok(true)
|
||||
},
|
||||
runtime_error,
|
||||
);
|
||||
assert!(extract_bool(py, result));
|
||||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn sync_runner_maps_a_panicked_future() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let error = run_sync::<bool, _>(
|
||||
py,
|
||||
poll_fn(|_| -> Poll<CoreResult<bool>> { panic!("route future panicked") }),
|
||||
runtime_error,
|
||||
)
|
||||
.expect_err("panicked route should become a Python exception");
|
||||
|
||||
assert!(error.is_instance_of::<PanicException>(py));
|
||||
assert_eq!(error.to_string(), "PanicException: route future panicked");
|
||||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn sync_runner_maps_a_panicked_error_mapper() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let error = run_sync::<bool, _>(
|
||||
py,
|
||||
async { Err(CoreError::InvalidRequest("invalid".to_string())) },
|
||||
panicking_error_mapper,
|
||||
)
|
||||
.expect_err("panicked mapper should become a Python exception");
|
||||
|
||||
assert!(error.is_instance_of::<PanicException>(py));
|
||||
assert_eq!(error.to_string(), "PanicException: error mapper panicked");
|
||||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn async_runner_surfaces_serializer_panics() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let module = PyModule::new(py, "runtime").expect("module should be created");
|
||||
module
|
||||
.add_function(
|
||||
wrap_pyfunction!(async_serialization_panic, &module)
|
||||
.expect("function should wrap"),
|
||||
)
|
||||
.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():
|
||||
try:
|
||||
await runtime.async_serialization_panic()
|
||||
except BaseException as error:
|
||||
assert type(error).__name__ == "PanicException"
|
||||
assert str(error) == "async serializer panicked"
|
||||
else:
|
||||
raise AssertionError("serializer panic was not raised")
|
||||
|
||||
asyncio.run(exercise())
|
||||
"#,
|
||||
)
|
||||
.expect("Python source should not contain null bytes");
|
||||
py.run(&code, Some(&locals), Some(&locals))
|
||||
.expect("serializer panic should reach the Python awaiter");
|
||||
});
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -2,4 +2,4 @@ mod gil;
|
|||
mod marshal;
|
||||
|
||||
pub use gil::{release_count, release_gil};
|
||||
pub use marshal::{from_py, to_py};
|
||||
pub use marshal::{Pythonized, from_py, panic_to_pyerr, to_py};
|
||||
|
|
|
|||
|
|
@ -1,4 +1,8 @@
|
|||
use std::any::Any;
|
||||
use std::panic::{AssertUnwindSafe, catch_unwind};
|
||||
|
||||
use pyo3::exceptions::PyValueError;
|
||||
use pyo3::panic::PanicException;
|
||||
use pyo3::prelude::*;
|
||||
use serde::Serialize;
|
||||
use serde::de::DeserializeOwned;
|
||||
|
|
@ -18,3 +22,71 @@ where
|
|||
.map(Bound::unbind)
|
||||
.map_err(|error| PyValueError::new_err(error.to_string()))
|
||||
}
|
||||
|
||||
pub struct Pythonized<T>(pub T);
|
||||
|
||||
impl<'py, T> IntoPyObject<'py> for Pythonized<T>
|
||||
where
|
||||
T: Serialize,
|
||||
{
|
||||
type Target = PyAny;
|
||||
type Output = Bound<'py, PyAny>;
|
||||
type Error = PyErr;
|
||||
|
||||
fn into_pyobject(self, py: Python<'py>) -> PyResult<Self::Output> {
|
||||
catch_unwind(AssertUnwindSafe(|| to_py(py, &self.0)))
|
||||
.map_err(panic_to_pyerr)?
|
||||
.map(|value| value.into_bound(py))
|
||||
}
|
||||
}
|
||||
|
||||
pub fn panic_to_pyerr(payload: Box<dyn Any + Send>) -> PyErr {
|
||||
let message = payload
|
||||
.downcast_ref::<String>()
|
||||
.map(String::as_str)
|
||||
.or_else(|| payload.downcast_ref::<&str>().copied())
|
||||
.unwrap_or("panic from Rust code");
|
||||
PanicException::new_err(message.to_string())
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use serde::Serializer;
|
||||
|
||||
use super::*;
|
||||
|
||||
struct PanickingSerializer;
|
||||
|
||||
impl Serialize for PanickingSerializer {
|
||||
fn serialize<S>(&self, _serializer: S) -> Result<S::Ok, S::Error>
|
||||
where
|
||||
S: Serializer,
|
||||
{
|
||||
panic!("serializer panicked")
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn pythonized_converts_on_the_attached_thread() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let value: Vec<i32> = Pythonized(vec![1, 2, 3])
|
||||
.into_pyobject(py)
|
||||
.and_then(|value| value.extract())
|
||||
.expect("value should convert");
|
||||
assert_eq!(value, vec![1, 2, 3]);
|
||||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn pythonized_maps_serializer_panics_to_a_base_exception() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let error = Pythonized(PanickingSerializer)
|
||||
.into_pyobject(py)
|
||||
.expect_err("serializer panic should become a Python exception");
|
||||
assert!(error.is_instance_of::<PanicException>(py));
|
||||
assert_eq!(error.to_string(), "PanicException: serializer panicked");
|
||||
});
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -20,6 +20,7 @@ import litellm.types.utils
|
|||
from litellm._logging import _redact_string, verbose_logger
|
||||
from litellm.anthropic_beta_headers_manager import update_headers_with_filtered_beta
|
||||
from litellm.constants import REALTIME_WEBSOCKET_MAX_MESSAGE_SIZE_BYTES
|
||||
from litellm.exceptions import APIError
|
||||
from litellm.litellm_core_utils.agentic_loop_settings import (
|
||||
DEFAULT_MAX_AGENTIC_LOOPS,
|
||||
validated_max_agentic_loops,
|
||||
|
|
@ -2400,10 +2401,27 @@ class BaseLLMHTTPHandler:
|
|||
extra_headers=headers,
|
||||
timeout=timeout,
|
||||
)
|
||||
except Exception as rust_error: # noqa: BLE001 # rollout-safety fallback: any Rust bridge failure must fall back to the Python path
|
||||
except Exception as rust_error: # noqa: BLE001
|
||||
from litellm.rust_bridge import get_native_bridge
|
||||
|
||||
native_bridge: Final = get_native_bridge()
|
||||
declined = getattr(native_bridge, "RustBridgeDeclined", None)
|
||||
upstream_failed = getattr(native_bridge, "RustUpstreamError", None)
|
||||
if isinstance(upstream_failed, type) and isinstance(rust_error, upstream_failed):
|
||||
args: Final = rust_error.args
|
||||
status: Final = args[0] if args else 0
|
||||
message: Final = args[1] if len(args) > 1 else ""
|
||||
raise APIError(
|
||||
status_code=int(status) or 500,
|
||||
message=f"litellm rust messages: {message}",
|
||||
llm_provider=custom_llm_provider,
|
||||
model=model,
|
||||
) from rust_error
|
||||
if not isinstance(declined, type) or not isinstance(rust_error, declined):
|
||||
raise
|
||||
verbose_logger.debug(
|
||||
"Rust Anthropic messages bridge raised %s; falling back to Python path",
|
||||
type(rust_error).__name__,
|
||||
"Rust Anthropic messages bridge declined before calling the provider (%s); falling back to Python path",
|
||||
rust_error,
|
||||
)
|
||||
return None
|
||||
if rust_response is None:
|
||||
|
|
|
|||
|
|
@ -1,12 +1,14 @@
|
|||
"""Tests for the optional Rust-backed Anthropic Messages path."""
|
||||
|
||||
import importlib
|
||||
from types import ModuleType
|
||||
from typing import cast
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
import litellm
|
||||
from litellm.exceptions import APIError
|
||||
from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler
|
||||
from litellm.types.llms.anthropic_messages.anthropic_response import (
|
||||
AnthropicMessagesResponse,
|
||||
|
|
@ -99,12 +101,28 @@ class ExplodingAsyncMessages:
|
|||
|
||||
|
||||
class RaisingAsyncMessages:
|
||||
def __init__(self) -> None:
|
||||
def __init__(self, error: Exception) -> None:
|
||||
self.calls = 0
|
||||
self.error = error
|
||||
|
||||
async def __call__(self, **kwargs: object) -> dict[str, object]:
|
||||
self.calls += 1
|
||||
raise RuntimeError("upstream request failed with status 400: bad request")
|
||||
raise self.error
|
||||
|
||||
|
||||
class FakeBridgeDeclined(Exception):
|
||||
pass
|
||||
|
||||
|
||||
class FakeUpstreamError(Exception):
|
||||
pass
|
||||
|
||||
|
||||
def _install_fake_bridge_exceptions(monkeypatch) -> None:
|
||||
native_bridge = ModuleType("_native")
|
||||
native_bridge.RustBridgeDeclined = FakeBridgeDeclined
|
||||
native_bridge.RustUpstreamError = FakeUpstreamError
|
||||
monkeypatch.setattr(rust_bridge_loader, "_cached_bridge", native_bridge)
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
|
|
@ -251,8 +269,9 @@ async def test_gate_invokes_rust_and_marks_response_header():
|
|||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_gate_falls_back_to_python_when_bridge_raises():
|
||||
bridge = RaisingAsyncMessages()
|
||||
async def test_gate_falls_back_only_when_bridge_declines(monkeypatch):
|
||||
_install_fake_bridge_exceptions(monkeypatch)
|
||||
bridge = RaisingAsyncMessages(FakeBridgeDeclined("unsupported request"))
|
||||
litellm.use_litellm_rust(True, amessages=bridge)
|
||||
|
||||
response = await _gate()
|
||||
|
|
@ -261,6 +280,31 @@ async def test_gate_falls_back_to_python_when_bridge_raises():
|
|||
assert bridge.calls == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_gate_surfaces_an_upstream_failure_without_fallback(monkeypatch):
|
||||
_install_fake_bridge_exceptions(monkeypatch)
|
||||
bridge = RaisingAsyncMessages(FakeUpstreamError(429, "429: rate limited"))
|
||||
litellm.use_litellm_rust(True, amessages=bridge)
|
||||
|
||||
with pytest.raises(APIError) as exc_info:
|
||||
await _gate()
|
||||
|
||||
assert exc_info.value.status_code == 429
|
||||
assert "429: rate limited" 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"))
|
||||
litellm.use_litellm_rust(True, amessages=bridge)
|
||||
|
||||
with pytest.raises(RuntimeError, match="unknown bridge failure"):
|
||||
await _gate()
|
||||
|
||||
assert bridge.calls == 1
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_gate_skips_rust_when_flag_absent():
|
||||
bridge = ExplodingAsyncMessages()
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue