From 8ab124309f92b63fcd1215d84bc04700fa5b907f Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sun, 27 Sep 2026 19:05:39 -0700 Subject: [PATCH] feat(rust): connect Python inference bindings to shared routes (#43465) * feat(rust): support the HTTP Responses API Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * feat(rust): enable native Python inference opt-in Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(rust_bridge): cover only python-only routes in the native load guard Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(rust): keep Python inference rollout disabled * fix(rust): preserve inference defaults and continuation IDs --------- Co-authored-by: Yujong Lee Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../src/routes/chat_completions.rs | 111 ++++--- .../src/routes/chat_completions/host.rs | 101 +++++++ .../python-bridge/src/routes/inference.rs | 129 ++++++++ .../crates/python-bridge/src/routes/mod.rs | 1 + .../python-bridge/src/routes/responses.rs | 115 ++++--- .../src/routes/responses/host.rs | 128 ++++++++ litellm/chat_completions/dispatch.py | 1 + litellm/responses/dispatch.py | 1 + .../chat_completions/entrypoints.py | 4 +- .../chat_completions/route_host.py | 36 ++- litellm/rust_bridge/public_call.py | 37 +++ litellm/rust_bridge/responses/entrypoints.py | 4 +- litellm/rust_bridge/responses/route_host.py | 28 +- tests/test_litellm_rust/test_inference.py | 281 ++++++++++++++++++ tests/unit/chat_completions/test_dispatch.py | 8 + tests/unit/responses/test_dispatch.py | 9 + .../chat_completions/test_route_host.py | 29 +- .../rust_bridge/responses/test_route_host.py | 21 +- 18 files changed, 963 insertions(+), 81 deletions(-) create mode 100644 litellm-rust/crates/python-bridge/src/routes/chat_completions/host.rs create mode 100644 litellm-rust/crates/python-bridge/src/routes/inference.rs create mode 100644 litellm-rust/crates/python-bridge/src/routes/responses/host.rs create mode 100644 tests/test_litellm_rust/test_inference.py diff --git a/litellm-rust/crates/python-bridge/src/routes/chat_completions.rs b/litellm-rust/crates/python-bridge/src/routes/chat_completions.rs index 1f84545d142..960eb1f4697 100644 --- a/litellm-rust/crates/python-bridge/src/routes/chat_completions.rs +++ b/litellm-rust/crates/python-bridge/src/routes/chat_completions.rs @@ -1,3 +1,5 @@ +mod host; + use pyo3::types::{PyDict, PyTuple}; use crate::errors::RustBridgeDeclined; @@ -147,58 +149,93 @@ pub(crate) fn achat_completions<'py>( ) } -#[pyfunction] -#[pyo3(signature = (request, args, kwargs))] -pub(crate) fn completion( +fn run_public( + py: Python<'_>, request: Bound<'_, PyAny>, args: Bound<'_, PyTuple>, kwargs: Bound<'_, PyDict>, + asynchronous: bool, ) -> PyResult> { - drop((request, args, kwargs)); - Err(RustBridgeDeclined::new_err( - "native chat completions route is not implemented", - )) + use super::inference::InferenceHost; + use litellm_callbacks_legacy_python::{LegacySurface, PublicCall, run_legacy_call}; + let host = InferenceHost::new( + request.clone().unbind(), + "litellm.rust_bridge.chat_completions.route_host", + ); + if let Some(reason) = py + .import("litellm.rust_bridge.chat_completions.route_host")? + .getattr("decline_reason")? + .call1((&request,))? + .extract::>()? + { + return Err(RustBridgeDeclined::new_err(reason)); + } + let admission = host::project(&host, py, &kwargs)?; + if let Some(reason) = chat_completions_decline_reason( + &admission.model, + admission.custom_llm_provider.as_deref(), + admission.messages, + &admission.optional_params, + ) { + return Err(RustBridgeDeclined::new_err(reason)); + } + if admission + .optional_params + .get("stream") + .is_some_and(|value| value == &serde_json::Value::Bool(true)) + { + return Err(RustBridgeDeclined::new_err( + "native Python chat_completions streaming", + )); + } + let route = ChatCompletionsRoute::new( + crate::http::provider_client(py, &kwargs, asynchronous)? + .map_err(crate::http::client_error)?, + crate::http::resources().auth.clone(), + crate::secrets::source(py)?, + ); + run_legacy_call( + py, + LegacySurface { + call_type: if asynchronous { + "acompletion" + } else { + "completion" + }, + input_description: "Chat completions", + stream: None, + }, + PublicCall::capture(&request, &args, &kwargs)?, + move |request| crate::logger::LoggedMachine::new(route.machine(request)), + host::ChatCompletionsPythonHost(host), + crate::preflight::sdk_preflight, + asynchronous, + ) } #[pyfunction] -#[pyo3(signature = (request, args, kwargs))] -pub(crate) fn acompletion( +pub(crate) fn completion( + py: Python<'_>, request: Bound<'_, PyAny>, args: Bound<'_, PyTuple>, kwargs: Bound<'_, PyDict>, ) -> PyResult> { - drop((request, args, kwargs)); - Err(RustBridgeDeclined::new_err( - "native chat completions route is not implemented", - )) + run_public(py, request, args, kwargs, false) +} + +#[pyfunction] +pub(crate) fn acompletion( + py: Python<'_>, + request: Bound<'_, PyAny>, + args: Bound<'_, PyTuple>, + kwargs: Bound<'_, PyDict>, +) -> PyResult> { + run_public(py, request, args, kwargs, true) } #[cfg(test)] mod tests { - use pyo3::{ - prelude::*, - types::{PyDict, PyList, PyTuple}, - }; - - use crate::errors::RustBridgeDeclined; - - #[test] - fn both_entrypoints_decline_before_provider_execution() { - Python::initialize(); - Python::attach(|py| { - let request = PyDict::new(py); - let args = PyTuple::empty(py); - let kwargs = PyDict::new(py); - - for entrypoint in [super::completion, super::acompletion] { - let error = entrypoint(request.clone().into_any(), args.clone(), kwargs.clone()) - .expect_err( - "native chat completions must decline until a route machine exists", - ); - assert!(error.is_instance_of::(py)); - } - }); - } + use pyo3::{prelude::*, types::PyList}; #[test] fn chat_completions_decline_keeps_existing_reasons() { diff --git a/litellm-rust/crates/python-bridge/src/routes/chat_completions/host.rs b/litellm-rust/crates/python-bridge/src/routes/chat_completions/host.rs new file mode 100644 index 00000000000..5011e151b48 --- /dev/null +++ b/litellm-rust/crates/python-bridge/src/routes/chat_completions/host.rs @@ -0,0 +1,101 @@ +use std::convert::Infallible; + +use super::super::inference::InferenceHost; +use litellm_core::chat_completions::{Error, route::ChatCompletions, types::ChatCompletionsCall}; +use litellm_host_python::{InvokeError, PythonBinding, PythonHostCalls, PythonOwned}; +use pyo3::{ + gc::{PyTraverseError, PyVisit}, + prelude::*, + types::PyDict, +}; + +pub(super) struct ChatCompletionsPythonHost(pub InferenceHost); + +pub(super) fn project( + host: &InferenceHost, + py: Python<'_>, + arguments: &Bound<'_, PyDict>, +) -> PyResult { + let call = host.project(py, arguments, "messages")?; + Ok(ChatCompletionsCall { + model: call.options.model, + messages: call.input, + optional_params: call.params, + api_key: call.options.api_key, + api_base: call.options.api_base, + custom_llm_provider: call.options.custom_llm_provider, + extra_headers: call.options.extra_headers, + timeout: call.options.timeout, + }) +} + +impl PythonBinding for ChatCompletionsPythonHost { + type Protocol = ChatCompletions; + type Failure = PyErr; + + fn decode_request( + &mut self, + py: Python<'_>, + arguments: &Bound<'_, PyDict>, + ) -> Result> { + let call = project(&self.0, py, arguments).map_err(InvokeError::Python)?; + if call + .optional_params + .get("stream") + .is_some_and(|value| value == &serde_json::Value::Bool(true)) + { + return Err(InvokeError::Native(Error::Unsupported( + "native Python chat_completions streaming", + ))); + } + Ok(call) + } + + fn encode_response( + &mut self, + py: Python<'_>, + response: ::Response, + ) -> PyResult> { + self.0.response(py, &response) + } + + fn encode_stream_head( + &mut self, + _py: Python<'_>, + head: ::StreamHead, + ) -> PyResult> { + match head {} + } + + fn encode_chunk( + &mut self, + _py: Python<'_>, + chunk: ::Chunk, + ) -> PyResult> { + match chunk {} + } + + fn map_error(&self, py: Python<'_>, error: Error) -> PyResult { + self.0.error(py, error) + } + fn host_error(error: &PyErr) -> Error { + Error::InvalidRequest(error.to_string().into()) + } +} + +impl PythonHostCalls for ChatCompletionsPythonHost { + fn handle_host_call( + &mut self, + _: Python<'_>, + op: Infallible, + ) -> Result<(), InvokeError> { + match op {} + } +} + +impl PythonOwned for ChatCompletionsPythonHost { + fn close(&mut self, _: Python<'_>) {} + fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> { + visit.call(&self.0.request) + } +} diff --git a/litellm-rust/crates/python-bridge/src/routes/inference.rs b/litellm-rust/crates/python-bridge/src/routes/inference.rs new file mode 100644 index 00000000000..13499bb13cb --- /dev/null +++ b/litellm-rust/crates/python-bridge/src/routes/inference.rs @@ -0,0 +1,129 @@ +use litellm_core::RouteError; +use litellm_core_utils::get_llm_provider_logic::get_custom_llm_provider; +use litellm_host_python::{from_py, lookup, to_py}; +use litellm_http::transport::Error as TransportError; +use pyo3::{exceptions::PyValueError, prelude::*, types::PyDict}; +use serde::Serialize; +use serde_json::{Map, Value}; + +use crate::{ + errors::{RustUpstreamError, route_error_to_pyerr}, + marshal::{RouteOptions, optional_timeout, python_timeout_seconds}, +}; + +pub(super) struct InferenceHost { + pub request: Py, + module: &'static str, +} + +pub(super) struct ProjectedCall { + pub options: RouteOptions, + pub input: Value, + pub params: Map, +} + +impl InferenceHost { + pub fn new(request: Py, module: &'static str) -> Self { + Self { request, module } + } + + pub fn project( + &self, + py: Python<'_>, + arguments: &Bound<'_, PyDict>, + input: &str, + ) -> PyResult { + let request = self.request.bind(py); + let argument = |name: &str| -> PyResult>> { + if let Some(value) = lookup(arguments, request, name)? { + return Ok((!value.is_none()).then_some(value)); + } + let parameter = request + .getattr("parameters")? + .call_method1("get", (name,))?; + if !parameter.is_none() { + return Ok(Some(parameter)); + } + let extra = request.getattr("kwargs")?.call_method1("get", (name,))?; + Ok((!extra.is_none()).then_some(extra)) + }; + let string = |name: &str| -> PyResult> { + argument(name)?.map(|value| value.extract()).transpose() + }; + let names: Vec = py.import(self.module)?.getattr("PARAMETERS")?.extract()?; + let params = names + .iter() + .filter_map(|name| match argument(name) { + Ok(Some(value)) => Some(from_py(&value).map(|value| (name.clone(), value))), + Ok(None) => None, + Err(error) => Some(Err(error)), + }) + .collect::>>()?; + let timeout = argument("timeout")? + .or(argument("request_timeout")?) + .map(|value| python_timeout_seconds(py, value.unbind())) + .transpose()? + .flatten(); + let model = string("model")?.ok_or_else(|| PyValueError::new_err("model is required"))?; + let custom_llm_provider = string("custom_llm_provider")?; + let provider = get_custom_llm_provider(&model, custom_llm_provider.as_deref()) + .map_or("", |resolved| resolved.custom_llm_provider); + let (default_key, default_base): (Option, Option) = py + .import(self.module)? + .getattr("connection_defaults")? + .call1((provider,))? + .extract()?; + Ok(ProjectedCall { + options: RouteOptions { + model, + api_key: string("api_key")? + .filter(|key| !key.is_empty()) + .or(default_key), + api_base: string("api_base")? + .filter(|base| !base.is_empty()) + .or(string("base_url")?.filter(|base| !base.is_empty())) + .or(default_base), + custom_llm_provider, + extra_headers: argument("extra_headers")? + .map(|value| from_py(&value)) + .transpose()?, + timeout: optional_timeout(timeout), + }, + input: from_py( + &argument(input)? + .ok_or_else(|| PyValueError::new_err(format!("{input} is required")))?, + )?, + params, + }) + } + + pub fn response(&self, py: Python<'_>, response: &impl Serialize) -> PyResult> { + py.import(self.module)? + .getattr("response")? + .call1((to_py(py, response)?,)) + .map(Bound::unbind) + } + + pub fn error(&self, py: Python<'_>, error: RouteError) -> PyResult { + if let RouteError::Secret(source) = &error + && let Some(original) = crate::secrets::python_error(py, source.source_error()) + { + return Ok(original); + } + let native = match error { + RouteError::Transport(TransportError::Http { status, body }) => { + let error = RustUpstreamError::new_err((status, body)); + error + .value(py) + .setattr("headers", Vec::<(String, String)>::new())?; + error + } + other => route_error_to_pyerr(other), + }; + let mapped = py + .import(self.module)? + .getattr("map_failure")? + .call1((native.value(py), self.request.bind(py)))?; + Ok(PyErr::from_value(mapped)) + } +} diff --git a/litellm-rust/crates/python-bridge/src/routes/mod.rs b/litellm-rust/crates/python-bridge/src/routes/mod.rs index dd694fa589f..313bbb3945c 100644 --- a/litellm-rust/crates/python-bridge/src/routes/mod.rs +++ b/litellm-rust/crates/python-bridge/src/routes/mod.rs @@ -1,6 +1,7 @@ pub(crate) mod audio_transcription; pub(crate) mod chat_completions; pub(crate) mod embeddings; +mod inference; pub(crate) mod messages; pub(crate) mod ocr; pub(crate) mod responses; diff --git a/litellm-rust/crates/python-bridge/src/routes/responses.rs b/litellm-rust/crates/python-bridge/src/routes/responses.rs index bf17ef6edde..4cacb850f81 100644 --- a/litellm-rust/crates/python-bridge/src/routes/responses.rs +++ b/litellm-rust/crates/python-bridge/src/routes/responses.rs @@ -1,3 +1,5 @@ +mod host; + use litellm_core::responses::websocket::ResponsesWebSocketConnection as RustResponsesWebSocketConnection; use pyo3::{ prelude::*, @@ -10,30 +12,94 @@ use crate::{ marshal::{marshal_headers, optional_timeout}, }; -#[pyfunction] -#[pyo3(signature = (request, args, kwargs))] -pub(crate) fn responses( +fn run_public( + py: Python<'_>, request: Bound<'_, PyAny>, args: Bound<'_, PyTuple>, kwargs: Bound<'_, PyDict>, + asynchronous: bool, ) -> PyResult> { - drop((request, args, kwargs)); - Err(RustBridgeDeclined::new_err( - "native responses route is not implemented", - )) + use super::inference::InferenceHost; + use litellm_callbacks_legacy_python::{LegacySurface, PublicCall, run_legacy_call}; + let host = InferenceHost::new( + request.clone().unbind(), + "litellm.rust_bridge.responses.route_host", + ); + if let Some(reason) = py + .import("litellm.rust_bridge.responses.route_host")? + .getattr("decline_reason")? + .call1((&request,))? + .extract::>()? + { + return Err(RustBridgeDeclined::new_err(reason)); + } + let admission = host::project(&host, py, &kwargs)?; + if admission + .custom_llm_provider + .as_deref() + .is_some_and(|provider| provider != "openai") + || admission + .model + .strip_prefix("openai/") + .unwrap_or(&admission.model) + .contains('/') + { + return Err(RustBridgeDeclined::new_err( + "native HTTP responses provider", + )); + } + if admission + .optional_params + .get("stream") + .is_some_and(|value| value == &serde_json::Value::Bool(true)) + { + return Err(RustBridgeDeclined::new_err( + "native Python responses streaming", + )); + } + let route = litellm_core::responses::ResponsesRoute::new( + crate::http::provider_client(py, &kwargs, asynchronous)? + .map_err(crate::http::client_error)?, + crate::http::resources().auth.clone(), + crate::secrets::source(py)?, + ); + run_legacy_call( + py, + LegacySurface { + call_type: if asynchronous { + "aresponses" + } else { + "responses" + }, + input_description: "Responses", + stream: None, + }, + PublicCall::capture(&request, &args, &kwargs)?, + move |request| crate::logger::LoggedMachine::new(route.machine(request)), + host::ResponsesPythonHost(host), + crate::preflight::sdk_preflight, + asynchronous, + ) } #[pyfunction] -#[pyo3(signature = (request, args, kwargs))] -pub(crate) fn aresponses( +pub(crate) fn responses( + py: Python<'_>, request: Bound<'_, PyAny>, args: Bound<'_, PyTuple>, kwargs: Bound<'_, PyDict>, ) -> PyResult> { - drop((request, args, kwargs)); - Err(RustBridgeDeclined::new_err( - "native responses route is not implemented", - )) + run_public(py, request, args, kwargs, false) +} + +#[pyfunction] +pub(crate) fn aresponses( + py: Python<'_>, + request: Bound<'_, PyAny>, + args: Bound<'_, PyTuple>, + kwargs: Bound<'_, PyDict>, +) -> PyResult> { + run_public(py, request, args, kwargs, true) } #[pyclass] @@ -89,28 +155,7 @@ mod tests { use std::{ffi::CString, time::Duration}; use futures_util::{SinkExt, StreamExt}; - use pyo3::{ - prelude::*, - types::{PyDict, PyTuple}, - }; - - use crate::errors::RustBridgeDeclined; - - #[test] - fn both_entrypoints_decline_before_provider_execution() { - Python::initialize(); - Python::attach(|py| { - let request = PyDict::new(py); - let args = PyTuple::empty(py); - let kwargs = PyDict::new(py); - - for entrypoint in [super::responses, super::aresponses] { - let error = entrypoint(request.clone().into_any(), args.clone(), kwargs.clone()) - .expect_err("native responses must decline until a route machine exists"); - assert!(error.is_instance_of::(py)); - } - }); - } + use pyo3::{prelude::*, types::PyDict}; use tokio::net::TcpListener; use tokio_tungstenite::{accept_async, tungstenite::Message}; diff --git a/litellm-rust/crates/python-bridge/src/routes/responses/host.rs b/litellm-rust/crates/python-bridge/src/routes/responses/host.rs new file mode 100644 index 00000000000..b054cbe5f38 --- /dev/null +++ b/litellm-rust/crates/python-bridge/src/routes/responses/host.rs @@ -0,0 +1,128 @@ +use std::convert::Infallible; + +use super::super::inference::InferenceHost; +use litellm_core::responses::{Error, route::Responses, types::ResponsesCall}; +use litellm_host_python::{InvokeError, PythonBinding, PythonHostCalls, PythonOwned}; +use pyo3::{ + gc::{PyTraverseError, PyVisit}, + prelude::*, + types::PyDict, +}; + +pub(super) struct ResponsesPythonHost(pub InferenceHost); + +pub(super) fn project( + host: &InferenceHost, + py: Python<'_>, + arguments: &Bound<'_, PyDict>, +) -> PyResult { + let call = host.project(py, arguments, "input")?; + let optional_params = call + .params + .into_iter() + .map(|(name, value)| { + let value = match (name.as_str(), value) { + ("previous_response_id", serde_json::Value::String(id)) => { + let decoded: String = py + .import("litellm.responses.utils")? + .getattr("ResponsesAPIRequestUtils")? + .call_method1( + "decode_previous_response_id_to_original_previous_response_id", + (id,), + )? + .extract()?; + serde_json::Value::String(decoded) + } + (_, value) => value, + }; + Ok((name, value)) + }) + .collect::>()?; + Ok(ResponsesCall { + model: call.options.model, + input: call.input, + optional_params, + api_key: call.options.api_key, + api_base: call.options.api_base, + custom_llm_provider: call.options.custom_llm_provider, + extra_headers: call.options.extra_headers, + timeout: call.options.timeout, + }) +} + +impl PythonBinding for ResponsesPythonHost { + type Protocol = Responses; + type Failure = PyErr; + + fn decode_request( + &mut self, + py: Python<'_>, + arguments: &Bound<'_, PyDict>, + ) -> Result> { + let call = project(&self.0, py, arguments).map_err(InvokeError::Python)?; + if call + .optional_params + .get("stream") + .is_some_and(|value| value == &serde_json::Value::Bool(true)) + { + return Err(InvokeError::Native(Error::Unsupported( + "native Python responses streaming", + ))); + } + Ok(call) + } + + fn encode_response( + &mut self, + py: Python<'_>, + response: ::Response, + ) -> PyResult> { + self.0.response(py, &response) + } + + fn encode_stream_head( + &mut self, + _py: Python<'_>, + head: ::StreamHead, + ) -> PyResult> { + let _ = head; + Err(pyo3::exceptions::PyRuntimeError::new_err( + "native Python Responses streaming is not supported", + )) + } + + fn encode_chunk( + &mut self, + _py: Python<'_>, + chunk: ::Chunk, + ) -> PyResult> { + let _ = chunk; + Err(pyo3::exceptions::PyRuntimeError::new_err( + "native Python Responses streaming is not supported", + )) + } + + fn map_error(&self, py: Python<'_>, error: Error) -> PyResult { + self.0.error(py, error) + } + fn host_error(error: &PyErr) -> Error { + Error::InvalidRequest(error.to_string().into()) + } +} + +impl PythonHostCalls for ResponsesPythonHost { + fn handle_host_call( + &mut self, + _: Python<'_>, + op: Infallible, + ) -> Result<(), InvokeError> { + match op {} + } +} + +impl PythonOwned for ResponsesPythonHost { + fn close(&mut self, _: Python<'_>) {} + fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> { + visit.call(&self.0.request) + } +} diff --git a/litellm/chat_completions/dispatch.py b/litellm/chat_completions/dispatch.py index 8e4b7f82eb8..e0fce35bb7a 100644 --- a/litellm/chat_completions/dispatch.py +++ b/litellm/chat_completions/dispatch.py @@ -69,6 +69,7 @@ def _public_request( custom_llm_provider=optional_str(extra.get("custom_llm_provider")), extra_headers=optional_mapping(fields.get("extra_headers")), kwargs=extra, + parameters=MappingProxyType({name: value for name, value in fields.items() if name != "kwargs"}), ) diff --git a/litellm/responses/dispatch.py b/litellm/responses/dispatch.py index d240356805c..52f9219b3fd 100644 --- a/litellm/responses/dispatch.py +++ b/litellm/responses/dispatch.py @@ -61,6 +61,7 @@ def _public_request( custom_llm_provider=optional_str(fields.get("custom_llm_provider")), extra_headers=optional_mapping(fields.get("extra_headers")), kwargs=extra, + parameters=MappingProxyType({name: value for name, value in fields.items() if name != "kwargs"}), ) diff --git a/litellm/rust_bridge/chat_completions/entrypoints.py b/litellm/rust_bridge/chat_completions/entrypoints.py index 6e41600c42e..d8cde8d0c66 100644 --- a/litellm/rust_bridge/chat_completions/entrypoints.py +++ b/litellm/rust_bridge/chat_completions/entrypoints.py @@ -1,7 +1,8 @@ from __future__ import annotations from collections.abc import Awaitable, Mapping, Sequence -from dataclasses import dataclass +from dataclasses import dataclass, field +from types import MappingProxyType from typing import Final, Protocol, cast # noqa: TID251 # validates dynamically loaded native callables from litellm.rust_bridge.bindings import NativeBinding @@ -18,6 +19,7 @@ class LiteLLMChatCompletionsRequest: custom_llm_provider: str | None extra_headers: Mapping[str, object] | None kwargs: Mapping[str, object] + parameters: Mapping[str, object] = field(default_factory=lambda: MappingProxyType({})) class NativeCompletion(Protocol): diff --git a/litellm/rust_bridge/chat_completions/route_host.py b/litellm/rust_bridge/chat_completions/route_host.py index 9a00ce340ba..9902e8e9946 100644 --- a/litellm/rust_bridge/chat_completions/route_host.py +++ b/litellm/rust_bridge/chat_completions/route_host.py @@ -1,11 +1,38 @@ from __future__ import annotations from collections.abc import Mapping +from typing import Final +import litellm +from litellm.constants import OPENAI_CHAT_COMPLETION_PARAMS from litellm.rust_bridge import failures from litellm.rust_bridge.chat_completions.entrypoints import LiteLLMChatCompletionsRequest +from litellm.rust_bridge.public_call import inference_decline_reason from litellm.types.utils import ModelResponse +_TRANSPORT_PARAMETERS: Final = frozenset( + { + "api_base", + "api_key", + "api_version", + "deployment_id", + "organization", + "base_url", + "default_headers", + "timeout", + "request_timeout", + "max_retries", + "extra_headers", + } +) +PARAMETERS: Final = tuple(name for name in OPENAI_CHAT_COMPLETION_PARAMS if name not in _TRANSPORT_PARAMETERS) + + +def connection_defaults(provider: str) -> tuple[str | None, str | None]: + if provider == "anthropic": + return litellm.anthropic_key or litellm.api_key, litellm.api_base + return None, None + def response(value: Mapping[str, object]) -> ModelResponse: return ModelResponse(**value) @@ -15,5 +42,10 @@ def arguments(request: LiteLLMChatCompletionsRequest) -> Mapping[str, object]: return request.kwargs -def map_failure(error: Exception, request: LiteLLMChatCompletionsRequest, request_provider: str) -> Exception: - return failures.map_failure(error, request.model, request_provider, arguments(request)) +def map_failure(error: Exception, request: LiteLLMChatCompletionsRequest) -> Exception: + provider: Final = request.custom_llm_provider or request.model.partition("/")[0] + return failures.map_native_failure(error, request.model, provider, arguments(request), request.api_base) + + +def decline_reason(request: LiteLLMChatCompletionsRequest) -> str | None: + return inference_decline_reason(PARAMETERS, {**request.parameters, **request.kwargs}) diff --git a/litellm/rust_bridge/public_call.py b/litellm/rust_bridge/public_call.py index 2a41926a802..17ff7e52165 100644 --- a/litellm/rust_bridge/public_call.py +++ b/litellm/rust_bridge/public_call.py @@ -40,3 +40,40 @@ def optional_sequence(value: object) -> Sequence[object] | None: if isinstance(value, str | bytes) or not isinstance(value, Sequence): return None return cast("Sequence[object]", value) # cast-ok: the same caller-owned object is handed on unchanged + + +def inference_decline_reason(parameters: tuple[str, ...], kwargs: Mapping[str, object]) -> str | None: + import litellm + + if litellm.cache is not None or litellm.drop_params or litellm.modify_params: + return "native inference does not implement the configured cache or parameter rewrites" + context: Final = frozenset( + { + "model", + "messages", + "input", + "api_key", + "api_base", + "base_url", + "custom_llm_provider", + "extra_headers", + "timeout", + "request_timeout", + "callbacks", + "success_callback", + "failure_callback", + "metadata", + "litellm_metadata", + "litellm_call_id", + "litellm_trace_id", + "litellm_logging_obj", + "litellm_credential_name", + "proxy_server_request", + } + ) + for name, value in kwargs.items(): + if value is None: + continue + if name not in parameters and name not in context: + return f"native inference does not implement {name}" + return None diff --git a/litellm/rust_bridge/responses/entrypoints.py b/litellm/rust_bridge/responses/entrypoints.py index 9bba7406b6d..d9f9489ac22 100644 --- a/litellm/rust_bridge/responses/entrypoints.py +++ b/litellm/rust_bridge/responses/entrypoints.py @@ -1,7 +1,8 @@ from __future__ import annotations from collections.abc import Awaitable, Mapping -from dataclasses import dataclass +from dataclasses import dataclass, field +from types import MappingProxyType from typing import Final, Protocol, cast # noqa: TID251 # validates dynamically loaded native callables from litellm.rust_bridge.bindings import NativeBinding @@ -18,6 +19,7 @@ class LiteLLMResponsesRequest: custom_llm_provider: str | None extra_headers: Mapping[str, object] | None kwargs: Mapping[str, object] + parameters: Mapping[str, object] = field(default_factory=lambda: MappingProxyType({})) class NativeResponses(Protocol): diff --git a/litellm/rust_bridge/responses/route_host.py b/litellm/rust_bridge/responses/route_host.py index 180b89c4412..110dd02ab22 100644 --- a/litellm/rust_bridge/responses/route_host.py +++ b/litellm/rust_bridge/responses/route_host.py @@ -1,10 +1,20 @@ from __future__ import annotations from collections.abc import Mapping +from typing import Final +import litellm +from litellm import get_llm_provider from litellm.rust_bridge import failures +from litellm.rust_bridge.public_call import inference_decline_reason from litellm.rust_bridge.responses.entrypoints import LiteLLMResponsesRequest -from litellm.types.llms.openai import ResponsesAPIResponse +from litellm.types.llms.openai import ResponsesAPIOptionalRequestParams, ResponsesAPIResponse + +PARAMETERS: Final = tuple(ResponsesAPIOptionalRequestParams.__annotations__) + + +def connection_defaults(_provider: str) -> tuple[str | None, str | None]: + return litellm.api_key or litellm.openai_key, litellm.api_base def response(value: Mapping[str, object]) -> ResponsesAPIResponse: @@ -15,5 +25,17 @@ def arguments(request: LiteLLMResponsesRequest) -> Mapping[str, object]: return request.kwargs -def map_failure(error: Exception, request: LiteLLMResponsesRequest, request_provider: str) -> Exception: - return failures.map_failure(error, request.model, request_provider, arguments(request)) +def map_failure(error: Exception, request: LiteLLMResponsesRequest) -> Exception: + provider: Final = request.custom_llm_provider or "openai" + return failures.map_native_failure(error, request.model, provider, arguments(request), request.api_base) + + +def decline_reason(request: LiteLLMResponsesRequest) -> str | None: + if request.custom_llm_provider is None and "/" not in request.model: + try: + _, provider, _, _ = get_llm_provider(model=request.model) + except Exception: # noqa: BLE001 # unresolved models stay on the existing Python dispatch path + return "native Responses could not resolve the provider" + if provider != "openai": + return "native HTTP responses provider" + return inference_decline_reason(PARAMETERS, {**request.parameters, **request.kwargs}) diff --git a/tests/test_litellm_rust/test_inference.py b/tests/test_litellm_rust/test_inference.py new file mode 100644 index 00000000000..863741b01a6 --- /dev/null +++ b/tests/test_litellm_rust/test_inference.py @@ -0,0 +1,281 @@ +import asyncio +from collections.abc import Awaitable, Coroutine, Mapping +from typing import Final, Literal, TypeAlias + +import pytest +from pydantic import JsonValue, TypeAdapter + +import litellm +from litellm import RateLimitError +from litellm.integrations.custom_logger import CustomLogger +from litellm.responses.utils import ResponsesAPIRequestUtils +from litellm.rust_bridge import _native +from litellm.rust_bridge.chat_completions.entrypoints import LiteLLMChatCompletionsRequest +from litellm.rust_bridge.responses.entrypoints import LiteLLMResponsesRequest +from litellm.types.llms.openai import ResponsesAPIResponse +from litellm.types.utils import ModelResponse +from tests.test_litellm_rust.support.callback_recorder import RecordingLogger +from tests.test_litellm_rust.support.recording_server import RecordingServer, ResponseSpec +from tests.test_litellm_rust.support.requests import MESSAGES, MESSAGES_MODEL, MESSAGES_RESPONSE, request_body + +pytestmark = pytest.mark.requires_rust_extension +Route: TypeAlias = Literal["chat", "responses"] +_OBJECT: Final = TypeAdapter(dict[str, object]) +NativeResult: TypeAlias = ( + ModelResponse + | ResponsesAPIResponse + | Coroutine[object, object, ModelResponse] + | Coroutine[object, object, ResponsesAPIResponse] +) +RESPONSES_MODEL: Final = "openai/gpt-6-sol" +RESPONSES_RESPONSE: Final[dict[str, JsonValue]] = { + "id": "resp_native", + "object": "response", + "created_at": 1, + "model": RESPONSES_MODEL.removeprefix("openai/"), + "status": "completed", + "output": [ + { + "type": "message", + "id": "msg_native", + "role": "assistant", + "status": "completed", + "content": [{"type": "output_text", "text": "native response", "annotations": []}], + } + ], + "usage": {"input_tokens": 5, "output_tokens": 4, "total_tokens": 9}, +} + + +@pytest.fixture(params=("chat", "responses")) +def route(request: pytest.FixtureRequest) -> Route: + return TypeAdapter(Route).validate_python(request.param) + + +def native_call( + route: Route, asynchronous: bool, server: RecordingServer, options: Mapping[str, object] +) -> NativeResult: + server.default_response = ResponseSpec(body=MESSAGES_RESPONSE if route == "chat" else RESPONSES_RESPONSE) + if route == "chat": + kwargs: Final = { + "model": MESSAGES_MODEL, + "messages": list(MESSAGES), + "api_key": "test-key", + "api_base": server.base_url, + "max_tokens": 32, + **options, + } + request: Final = LiteLLMChatCompletionsRequest( + MESSAGES_MODEL, list(MESSAGES), None, "test-key", server.base_url, None, None, kwargs + ) + return (_native.acompletion if asynchronous else _native.completion)(request, (), kwargs) + response_kwargs: Final = { + "model": RESPONSES_MODEL, + "input": "hello", + "api_key": "test-key", + "api_base": server.base_url, + "max_output_tokens": 32, + **options, + } + response_request: Final = LiteLLMResponsesRequest( + RESPONSES_MODEL, "hello", None, "test-key", server.base_url, "openai", None, response_kwargs + ) + return (_native.aresponses if asynchronous else _native.responses)(response_request, (), response_kwargs) + + +async def execute(route: Route, asynchronous: bool, server: RecordingServer, options: Mapping[str, object]) -> object: + if not asynchronous: + return await asyncio.to_thread(native_call, route, False, server, options) + result: Final = native_call(route, True, server, options) + assert isinstance(result, Awaitable) + return await result + + +@pytest.mark.asyncio +@pytest.mark.parametrize("asynchronous", (False, True)) +async def test_native_inference_returns_public_models_and_logs_once( + route: Route, + asynchronous: bool, + recording_server: RecordingServer, +) -> None: + recorder: Final = RecordingLogger() + result: Final = await execute(route, asynchronous, recording_server, {"callbacks": [recorder], "temperature": 0.25}) + assert len(recording_server.requests) == 1 + sent: Final = recording_server.requests[0] + body: Final = _OBJECT.validate_python(sent.body) + assert body["temperature"] == 0.25 + assert request_body(_OBJECT.validate_python(recorder.wait_for("log_pre_api_call")[0].kwargs)) == body + if route == "chat": + assert isinstance(result, ModelResponse) + assert result.choices[0].message.content == "Hello from native Messages" + assert sent.path == "/v1/messages" + else: + assert isinstance(result, ResponsesAPIResponse) + assert result.output_text == "native response" + assert sent.path == "/responses" + success: Final = await recorder.wait_for_async("async_log_success_event" if asynchronous else "log_success_event") + assert len(success) == 1 + if isinstance(result, ModelResponse): + assert success[0].response is result + else: + logged: Final = success[0].response + assert isinstance(logged, ResponsesAPIResponse) + assert isinstance(result, ResponsesAPIResponse) + assert logged.id == result.id + assert logged.output_text == result.output_text + + +@pytest.mark.asyncio +async def test_native_inference_pre_call_edits_reach_the_provider( + route: Route, recording_server: RecordingServer +) -> None: + class Edit(CustomLogger): + def log_pre_api_call(self, model: object, messages: object, kwargs: dict[str, object]) -> None: + request_body(kwargs)["temperature"] = 0.75 + + await execute(route, True, recording_server, {"callbacks": [Edit()]}) + assert _OBJECT.validate_python(recording_server.requests[0].body)["temperature"] == 0.75 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("asynchronous", (False, True)) +async def test_native_inference_provider_failure_is_terminal_and_shared_with_callbacks( + route: Route, + asynchronous: bool, + recording_server: RecordingServer, +) -> None: + recorder: Final = RecordingLogger() + recording_server.enqueue( + ResponseSpec(body={"error": {"message": "slow down", "type": "rate_limit_error"}}, status=429) + ) + with pytest.raises(RateLimitError) as caught: + await execute(route, asynchronous, recording_server, {"callbacks": [recorder]}) + assert getattr(caught.value, "status_code", None) == 429 + assert len(recording_server.requests) == 1 + failure: Final = await recorder.wait_for_async("async_log_failure_event" if asynchronous else "log_failure_event") + assert len(failure) == 1 + assert _OBJECT.validate_python(failure[0].kwargs)["exception"] is caught.value + assert not any("success" in name for name in recorder.names) + + +@pytest.mark.asyncio +async def test_unstarted_native_inference_has_no_provider_or_callback_effects( + route: Route, + recording_server: RecordingServer, +) -> None: + recorder: Final = RecordingLogger() + recording_server.expected_requests = 0 + pending: Final = native_call(route, True, recording_server, {"callbacks": [recorder]}) + assert asyncio.iscoroutine(pending) + pending.close() + assert not recording_server.requests + assert not recorder.events + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "options", + ( + {"stream": True}, + {"extra_body": {"provider_option": True}}, + {"mock_response": "mock"}, + {"num_retries": 1}, + {"use_chat_completions_api": True}, + {"model_list": []}, + ), +) +async def test_native_inference_declines_unsupported_requests_before_callbacks( + route: Route, + recording_server: RecordingServer, + options: Mapping[str, object], +) -> None: + recorder: Final = RecordingLogger() + recording_server.expected_requests = 0 + with pytest.raises(_native.RustBridgeDeclined): + native_call(route, True, recording_server, {**options, "callbacks": [recorder]}) + assert not recording_server.requests + assert not recorder.events + + +@pytest.mark.asyncio +async def test_native_projection_reads_positional_parameters(route: Route, recording_server: RecordingServer) -> None: + from litellm.chat_completions.dispatch import ( + _DISPATCH as chat_dispatch, # pyright: ignore[reportPrivateUsage] # exercise the request passed to the native boundary + ) + from litellm.responses.dispatch import ( + _DISPATCH as responses_dispatch, # pyright: ignore[reportPrivateUsage] # exercise the request passed to the native boundary + ) + + recording_server.default_response = ResponseSpec(body=MESSAGES_RESPONSE if route == "chat" else RESPONSES_RESPONSE) + kwargs: Final = {"api_key": "test-key", "api_base": recording_server.base_url} + if route == "chat": + args: Final = (MESSAGES_MODEL, list(MESSAGES), 12.0, 0.25) + request: Final = chat_dispatch.request(args, kwargs) + assert request is not None + await asyncio.to_thread(_native.completion, request, args, kwargs) + assert _OBJECT.validate_python(recording_server.requests[0].body)["temperature"] == 0.25 + else: + response_args: Final = ("hello", RESPONSES_MODEL, None, "Be brief", 16) + response_request: Final = responses_dispatch.request(response_args, kwargs) + assert response_request is not None + await asyncio.to_thread(_native.responses, response_request, response_args, kwargs) + body: Final = _OBJECT.validate_python(recording_server.requests[0].body) + assert body["instructions"] == "Be brief" + assert body["max_output_tokens"] == 16 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("asynchronous", (False, True)) +@pytest.mark.parametrize("source", ("explicit", "base_url", "global", "provider", "environment", "empty")) +async def test_native_connection_settings_reach_the_provider( + route: Route, + asynchronous: bool, + source: str, + recording_server: RecordingServer, + monkeypatch: pytest.MonkeyPatch, +) -> None: + key: Final = "selected-key" + explicit: Final = source in ("explicit", "base_url") + monkeypatch.setattr(litellm, "api_key", key if source in ("global", "empty") else ("unused" if explicit else None)) + monkeypatch.setattr(litellm, "openai_key", key if source == "provider" else ("unused" if explicit else None)) + monkeypatch.setattr(litellm, "anthropic_key", key if source == "provider" else ("unused" if explicit else None)) + monkeypatch.setattr( + litellm, + "api_base", + None if source == "environment" else ("http://127.0.0.1:1" if explicit else recording_server.base_url), + ) + monkeypatch.setenv( + "OPENAI_API_KEY" if route == "responses" else "ANTHROPIC_API_KEY", key if source == "environment" else "unused" + ) + for name in ("OPENAI_BASE_URL", "OPENAI_API_BASE", "ANTHROPIC_BASE_URL", "ANTHROPIC_API_BASE"): + monkeypatch.setenv(name, recording_server.base_url if source == "environment" else "http://127.0.0.1:1") + result: Final = await execute( + route, + asynchronous, + recording_server, + { + "api_key": key if explicit else ("" if source == "empty" else None), + "api_base": recording_server.base_url if source == "explicit" else ("" if source == "empty" else None), + **({"base_url": recording_server.base_url} if source == "base_url" else {}), + }, + ) + assert isinstance(result, ModelResponse | ResponsesAPIResponse) + assert len(recording_server.requests) == 1 + headers: Final = recording_server.requests[0].headers + assert headers["x-api-key" if route == "chat" else "authorization"] == (key if route == "chat" else f"Bearer {key}") + + +@pytest.mark.asyncio +@pytest.mark.parametrize("asynchronous", (False, True)) +@pytest.mark.parametrize("encoded", (False, True)) +async def test_native_responses_decode_continuation_ids( + asynchronous: bool, encoded: bool, recording_server: RecordingServer +) -> None: + original: Final = "resp_upstream" + previous: Final = ( + ResponsesAPIRequestUtils._build_responses_api_response_id("openai", "deployment", original) + if encoded + else original + ) + await execute("responses", asynchronous, recording_server, {"previous_response_id": previous}) + assert _OBJECT.validate_python(recording_server.requests[0].body)["previous_response_id"] == original diff --git a/tests/unit/chat_completions/test_dispatch.py b/tests/unit/chat_completions/test_dispatch.py index 40b1c0ef019..c9274321a0f 100644 --- a/tests/unit/chat_completions/test_dispatch.py +++ b/tests/unit/chat_completions/test_dispatch.py @@ -355,3 +355,11 @@ def test_internal_acompletion_marker_bypasses_native() -> None: ) assert response is expected + + +def test_positional_parameters_remain_available_to_native_projection() -> None: + request: Final = _DISPATCH.request(("anthropic/test-model", MESSAGES, 12.0, 0.25), {}) + assert request is not None + assert request.parameters["timeout"] == 12.0 + assert request.parameters["temperature"] == 0.25 + assert request.messages is MESSAGES diff --git a/tests/unit/responses/test_dispatch.py b/tests/unit/responses/test_dispatch.py index 45cb5c4f1ad..637d4bc0a1e 100644 --- a/tests/unit/responses/test_dispatch.py +++ b/tests/unit/responses/test_dispatch.py @@ -317,3 +317,12 @@ def test_responses_with_retries_uses_the_dispatch_entrypoint(monkeypatch: pytest assert result is expected assert calls[0]["num_retries"] == 0 assert calls[0]["max_retries"] == 0 + + +def test_positional_parameters_remain_available_to_native_projection() -> None: + include: Final = ["reasoning.encrypted_content"] + request: Final = _DISPATCH.request((INPUT, "openai/test-model", include, "Be brief", 16), {}) + assert request is not None + assert request.parameters["include"] is include + assert request.parameters["instructions"] == "Be brief" + assert request.parameters["max_output_tokens"] == 16 diff --git a/tests/unit/rust_bridge/chat_completions/test_route_host.py b/tests/unit/rust_bridge/chat_completions/test_route_host.py index 848f5a00eb3..7f9295e93b3 100644 --- a/tests/unit/rust_bridge/chat_completions/test_route_host.py +++ b/tests/unit/rust_bridge/chat_completions/test_route_host.py @@ -1,7 +1,10 @@ from types import MappingProxyType from typing import Final -from litellm.rust_bridge.chat_completions.route_host import arguments, response +import pytest + +import litellm +from litellm.rust_bridge.chat_completions.route_host import arguments, connection_defaults, response from litellm.rust_bridge.chat_completions.entrypoints import LiteLLMChatCompletionsRequest from litellm.types.utils import ModelResponse @@ -47,3 +50,27 @@ def test_arguments_are_the_public_kwargs_view() -> None: ) assert arguments(request) is kwargs + + +@pytest.mark.parametrize( + ("provider", "global_key", "provider_key", "expected_key", "expected_base"), + ( + ("anthropic", "global", "provider", "provider", "https://configured.invalid"), + ("anthropic", "global", None, "global", "https://configured.invalid"), + ("anthropic", "global", "", "global", "https://configured.invalid"), + ("anthropic", None, None, None, "https://configured.invalid"), + ("bedrock", "global", "provider", None, None), + ), +) +def test_connection_defaults_preserve_provider_precedence( + monkeypatch: pytest.MonkeyPatch, + provider: str, + global_key: str | None, + provider_key: str | None, + expected_key: str | None, + expected_base: str | None, +) -> None: + monkeypatch.setattr(litellm, "api_key", global_key) + monkeypatch.setattr(litellm, "anthropic_key", provider_key) + monkeypatch.setattr(litellm, "api_base", "https://configured.invalid") + assert connection_defaults(provider) == (expected_key, expected_base) diff --git a/tests/unit/rust_bridge/responses/test_route_host.py b/tests/unit/rust_bridge/responses/test_route_host.py index 49bf19e7d8a..d04e02b0dda 100644 --- a/tests/unit/rust_bridge/responses/test_route_host.py +++ b/tests/unit/rust_bridge/responses/test_route_host.py @@ -4,7 +4,8 @@ from typing import Final import pytest from pydantic import ValidationError -from litellm.rust_bridge.responses.route_host import arguments, response +import litellm +from litellm.rust_bridge.responses.route_host import arguments, connection_defaults, response from litellm.rust_bridge.responses.entrypoints import LiteLLMResponsesRequest from litellm.types.llms.openai import ResponsesAPIResponse @@ -55,3 +56,21 @@ def test_arguments_are_the_public_kwargs_view() -> None: ) assert arguments(request) is kwargs + + +@pytest.mark.parametrize( + ("global_key", "provider_key", "expected"), + ( + ("global", "provider", "global"), + (None, "provider", "provider"), + ("", "provider", "provider"), + (None, None, None), + ), +) +def test_connection_defaults_preserve_openai_precedence( + monkeypatch: pytest.MonkeyPatch, global_key: str | None, provider_key: str | None, expected: str | None +) -> None: + monkeypatch.setattr(litellm, "api_key", global_key) + monkeypatch.setattr(litellm, "openai_key", provider_key) + monkeypatch.setattr(litellm, "api_base", "https://configured.invalid/v1") + assert connection_defaults("openai") == (expected, litellm.api_base)