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 <yujong@berri.ai>
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-09-27 19:05:39 -07:00 • committed by GitHub
parent 9b08112ed2
commit 8ab124309f
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
18 changed files with 963 additions and 81 deletions

View file

@ -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<Py<PyAny>> {
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::<Option<String>>()?
{
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<Py<PyAny>> {
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<Py<PyAny>> {
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::<RustBridgeDeclined>(py));
}
});
}
use pyo3::{prelude::*, types::PyList};
#[test]
fn chat_completions_decline_keeps_existing_reasons() {

View file

@ -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<ChatCompletionsCall> {
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<ChatCompletionsCall, InvokeError<Error>> {
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: <ChatCompletions as litellm_host::protocol::Protocol>::Response,
) -> PyResult<Py<PyAny>> {
self.0.response(py, &response)
}
fn encode_stream_head(
&mut self,
_py: Python<'_>,
head: <ChatCompletions as litellm_host::protocol::Protocol>::StreamHead,
) -> PyResult<Py<PyAny>> {
match head {}
}
fn encode_chunk(
&mut self,
_py: Python<'_>,
chunk: <ChatCompletions as litellm_host::protocol::Protocol>::Chunk,
) -> PyResult<Py<PyAny>> {
match chunk {}
}
fn map_error(&self, py: Python<'_>, error: Error) -> PyResult<PyErr> {
self.0.error(py, error)
}
fn host_error(error: &PyErr) -> Error {
Error::InvalidRequest(error.to_string().into())
}
}
impl PythonHostCalls<ChatCompletions> for ChatCompletionsPythonHost {
fn handle_host_call(
&mut self,
_: Python<'_>,
op: Infallible,
) -> Result<(), InvokeError<Error>> {
match op {}
}
}
impl PythonOwned for ChatCompletionsPythonHost {
fn close(&mut self, _: Python<'_>) {}
fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> {
visit.call(&self.0.request)
}
}

View file

@ -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<PyAny>,
module: &'static str,
}
pub(super) struct ProjectedCall {
pub options: RouteOptions,
pub input: Value,
pub params: Map<String, Value>,
}
impl InferenceHost {
pub fn new(request: Py<PyAny>, module: &'static str) -> Self {
Self { request, module }
}
pub fn project(
&self,
py: Python<'_>,
arguments: &Bound<'_, PyDict>,
input: &str,
) -> PyResult<ProjectedCall> {
let request = self.request.bind(py);
let argument = |name: &str| -> PyResult<Option<Bound<'_, PyAny>>> {
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<Option<String>> {
argument(name)?.map(|value| value.extract()).transpose()
};
let names: Vec<String> = 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::<PyResult<Map<String, Value>>>()?;
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<String>, Option<String>) = 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<PyAny>> {
py.import(self.module)?
.getattr("response")?
.call1((to_py(py, response)?,))
.map(Bound::unbind)
}
pub fn error(&self, py: Python<'_>, error: RouteError) -> PyResult<PyErr> {
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))
}
}

View file

@ -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;

View file

@ -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<Py<PyAny>> {
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::<Option<String>>()?
{
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<Py<PyAny>> {
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<Py<PyAny>> {
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::<RustBridgeDeclined>(py));
}
});
}
use pyo3::{prelude::*, types::PyDict};
use tokio::net::TcpListener;
use tokio_tungstenite::{accept_async, tungstenite::Message};

View file

@ -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<ResponsesCall> {
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::<PyResult<_>>()?;
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<ResponsesCall, InvokeError<Error>> {
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: <Responses as litellm_host::protocol::Protocol>::Response,
) -> PyResult<Py<PyAny>> {
self.0.response(py, &response)
}
fn encode_stream_head(
&mut self,
_py: Python<'_>,
head: <Responses as litellm_host::protocol::Protocol>::StreamHead,
) -> PyResult<Py<PyAny>> {
let _ = head;
Err(pyo3::exceptions::PyRuntimeError::new_err(
"native Python Responses streaming is not supported",
))
}
fn encode_chunk(
&mut self,
_py: Python<'_>,
chunk: <Responses as litellm_host::protocol::Protocol>::Chunk,
) -> PyResult<Py<PyAny>> {
let _ = chunk;
Err(pyo3::exceptions::PyRuntimeError::new_err(
"native Python Responses streaming is not supported",
))
}
fn map_error(&self, py: Python<'_>, error: Error) -> PyResult<PyErr> {
self.0.error(py, error)
}
fn host_error(error: &PyErr) -> Error {
Error::InvalidRequest(error.to_string().into())
}
}
impl PythonHostCalls<Responses> for ResponsesPythonHost {
fn handle_host_call(
&mut self,
_: Python<'_>,
op: Infallible,
) -> Result<(), InvokeError<Error>> {
match op {}
}
}
impl PythonOwned for ResponsesPythonHost {
fn close(&mut self, _: Python<'_>) {}
fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> {
visit.call(&self.0.request)
}
}

View file

@ -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"}),
)

View file

@ -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"}),
)

View file

@ -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):

View file

@ -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})

View file

@ -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

View file

@ -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):

View file

@ -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})

View file

@ -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

View file

@ -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

View file

@ -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

View file

@ -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)

View file

@ -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)