mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-29 01:42:19 +00:00
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:
parent
9b08112ed2
commit
8ab124309f
18 changed files with 963 additions and 81 deletions
|
|
@ -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() {
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
}
|
||||
}
|
||||
129
litellm-rust/crates/python-bridge/src/routes/inference.rs
Normal file
129
litellm-rust/crates/python-bridge/src/routes/inference.rs
Normal 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))
|
||||
}
|
||||
}
|
||||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -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};
|
||||
|
||||
|
|
|
|||
128
litellm-rust/crates/python-bridge/src/routes/responses/host.rs
Normal file
128
litellm-rust/crates/python-bridge/src/routes/responses/host.rs
Normal 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)
|
||||
}
|
||||
}
|
||||
|
|
@ -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"}),
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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"}),
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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})
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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})
|
||||
|
|
|
|||
281
tests/test_litellm_rust/test_inference.py
Normal file
281
tests/test_litellm_rust/test_inference.py
Normal 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
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue