mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
refactor(python-bridge): take NativeCall directly and fold routes into per-route folders (#45413)
* refactor(python-bridge): take NativeCall directly and fold routes into per-route folders Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com> * test(python-bridge): use rstest for updated tests and keep embedding's stub parameter name Co-Authored-By: Claude Opus 5.5 <noreply@anthropic.com> --------- Co-authored-by: Claude Opus 5.5 <noreply@anthropic.com>
This commit is contained in:
parent
85a3869dfd
commit
0721cffab2
32 changed files with 559 additions and 908 deletions
|
|
@ -936,9 +936,6 @@ mod payload_tests {
|
|||
/// The payload phases of `Logging` on top of `StubLogger`, with `pre_call` handing the
|
||||
/// payload to the case's `on_pre_call`.
|
||||
const PAYLOAD_LOGGER: &CStr = c"
|
||||
class Request:
|
||||
pass
|
||||
|
||||
class PayloadLogger(StubLogger):
|
||||
def update_from_kwargs(self, **update):
|
||||
self.update = update
|
||||
|
|
@ -954,7 +951,7 @@ class PayloadLogger(StubLogger):
|
|||
self.record('post_call', None)
|
||||
self.post = (original_response, api_key, additional_args)
|
||||
|
||||
request = Request()
|
||||
bound = {}
|
||||
kwargs = {}
|
||||
logger = PayloadLogger()
|
||||
on_pre_call = lambda additional_args: None
|
||||
|
|
@ -1225,10 +1222,10 @@ on_pre_call = lambda args: observed.append(
|
|||
def check():
|
||||
assert observed == [(True, True)], observed
|
||||
")]
|
||||
#[case::request_attribute_behind_an_omitted_keyword(c"
|
||||
#[case::bound_value_behind_an_omitted_keyword(c"
|
||||
document = {'type': 'document_url', 'document_url': 'data:application/pdf;base64,YWJj'}
|
||||
pages = [0]
|
||||
request.document = document
|
||||
bound['document'] = document
|
||||
kwargs = {'pages': pages}
|
||||
observed = []
|
||||
on_pre_call = lambda args: observed.append(
|
||||
|
|
|
|||
|
|
@ -13,21 +13,21 @@ use pyo3::{
|
|||
pub struct PublicCall {
|
||||
args: Py<PyTuple>,
|
||||
kwargs: Py<PyDict>,
|
||||
request: Py<PyAny>,
|
||||
bound: Py<PyDict>,
|
||||
}
|
||||
|
||||
impl PublicCall {
|
||||
/// Copies the keyword arguments once, so the legacy path's rewrites never reach the
|
||||
/// caller's own dict while every value keeps its identity.
|
||||
pub fn capture(
|
||||
request: &Bound<'_, PyAny>,
|
||||
bound: &Bound<'_, PyDict>,
|
||||
args: &Bound<'_, PyTuple>,
|
||||
kwargs: &Bound<'_, PyDict>,
|
||||
) -> PyResult<Self> {
|
||||
Ok(Self {
|
||||
args: args.clone().unbind(),
|
||||
kwargs: kwargs.copy()?.unbind(),
|
||||
request: request.clone().unbind(),
|
||||
bound: bound.clone().unbind(),
|
||||
})
|
||||
}
|
||||
|
||||
|
|
@ -55,31 +55,30 @@ impl PublicCall {
|
|||
py: Python<'py>,
|
||||
name: &str,
|
||||
) -> PyResult<Option<Bound<'py, PyAny>>> {
|
||||
lookup(self.kwargs.bind(py), self.request.bind(py), name)
|
||||
lookup(self.kwargs.bind(py), self.bound.bind(py), name)
|
||||
}
|
||||
|
||||
pub(crate) fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> {
|
||||
visit.call(&self.args)?;
|
||||
visit.call(&self.kwargs)?;
|
||||
visit.call(&self.request)
|
||||
visit.call(&self.bound)
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use crate::test_support::{local, local_dict};
|
||||
|
||||
fn capture<'py>(py: Python<'py>, source: &std::ffi::CStr) -> (PublicCall, Bound<'py, PyDict>) {
|
||||
let locals = PyDict::new(py);
|
||||
py.run(source, Some(&locals), Some(&locals)).unwrap();
|
||||
let request = locals.get_item("request").unwrap().unwrap();
|
||||
let kwargs = locals
|
||||
.get_item("kwargs")
|
||||
.unwrap()
|
||||
.unwrap()
|
||||
.cast_into::<PyDict>()
|
||||
.unwrap();
|
||||
let call = PublicCall::capture(&request, &PyTuple::empty(py), &kwargs).unwrap();
|
||||
let call = PublicCall::capture(
|
||||
&local_dict(&locals, "bound"),
|
||||
&PyTuple::empty(py),
|
||||
&local_dict(&locals, "kwargs"),
|
||||
)
|
||||
.unwrap();
|
||||
(call, locals)
|
||||
}
|
||||
|
||||
|
|
@ -91,25 +90,22 @@ mod tests {
|
|||
py,
|
||||
c"
|
||||
pages = [0]
|
||||
class Request:
|
||||
pass
|
||||
request = Request()
|
||||
bound = {'pages': [1]}
|
||||
kwargs = {'pages': pages}
|
||||
",
|
||||
);
|
||||
let caller = locals
|
||||
.get_item("kwargs")
|
||||
.unwrap()
|
||||
.unwrap()
|
||||
.cast_into::<PyDict>()
|
||||
.unwrap();
|
||||
let caller = local_dict(&locals, "kwargs");
|
||||
call.kwargs()
|
||||
.bind(py)
|
||||
.set_item("litellm_call_id", "call")
|
||||
.unwrap();
|
||||
assert!(!caller.contains("litellm_call_id").unwrap());
|
||||
let pages = locals.get_item("pages").unwrap().unwrap();
|
||||
assert!(call.lookup(py, "pages").unwrap().unwrap().is(&pages));
|
||||
assert!(
|
||||
call.lookup(py, "pages")
|
||||
.unwrap()
|
||||
.unwrap()
|
||||
.is(local(&locals, "pages"))
|
||||
);
|
||||
});
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -173,21 +173,23 @@ pub(crate) fn local<'py>(locals: &Bound<'py, PyDict>, name: &str) -> Bound<'py,
|
|||
locals.get_item(name).unwrap().unwrap()
|
||||
}
|
||||
|
||||
/// A legacy call over the namespace's `kwargs` (or none) and `request` (or `None`).
|
||||
pub(crate) fn local_dict<'py>(locals: &Bound<'py, PyDict>, name: &str) -> Bound<'py, PyDict> {
|
||||
local(locals, name).cast_into().unwrap()
|
||||
}
|
||||
|
||||
/// A legacy call over the namespace's `kwargs` and `bound` dicts, each empty when absent.
|
||||
pub(crate) fn legacy_call(
|
||||
py: Python<'_>,
|
||||
locals: &Bound<'_, PyDict>,
|
||||
asynchronous: bool,
|
||||
) -> LegacyLogging {
|
||||
let request = locals
|
||||
.get_item("request")
|
||||
.unwrap()
|
||||
.unwrap_or_else(|| py.None().into_bound(py));
|
||||
let kwargs = locals
|
||||
.get_item("kwargs")
|
||||
.unwrap()
|
||||
.map(|kwargs| kwargs.cast_into::<PyDict>().unwrap())
|
||||
.unwrap_or_else(|| PyDict::new(py));
|
||||
let call = PublicCall::capture(&request, &PyTuple::empty(py), &kwargs).unwrap();
|
||||
let dict = |name: &str| {
|
||||
locals
|
||||
.get_item(name)
|
||||
.unwrap()
|
||||
.map(|value| value.cast_into::<PyDict>().unwrap())
|
||||
.unwrap_or_else(|| PyDict::new(py))
|
||||
};
|
||||
let call = PublicCall::capture(&dict("bound"), &PyTuple::empty(py), &dict("kwargs")).unwrap();
|
||||
LegacyLogging::new(py, crate::LoggingOperation::Ocr, call, asynchronous)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -2,78 +2,97 @@ use pyo3::{prelude::*, types::PyDict};
|
|||
|
||||
pub fn lookup<'py>(
|
||||
kwargs: &Bound<'py, PyDict>,
|
||||
request: &Bound<'py, PyAny>,
|
||||
bound: &Bound<'py, PyDict>,
|
||||
name: &str,
|
||||
) -> PyResult<Option<Bound<'py, PyAny>>> {
|
||||
if let Some(value) = kwargs.get_item(name)? {
|
||||
return Ok(Some(value));
|
||||
match kwargs.get_item(name)? {
|
||||
Some(value) => Ok(Some(value)),
|
||||
None => bound.get_item(name),
|
||||
}
|
||||
if let Ok(bound) = request.cast::<PyDict>() {
|
||||
return bound.get_item(name);
|
||||
}
|
||||
request.getattr_opt(name)
|
||||
}
|
||||
|
||||
pub fn present<'py>(
|
||||
kwargs: &Bound<'py, PyDict>,
|
||||
bound: &Bound<'py, PyDict>,
|
||||
name: &str,
|
||||
) -> PyResult<Option<Bound<'py, PyAny>>> {
|
||||
Ok(lookup(kwargs, bound, name)?.filter(|value| !value.is_none()))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use rstest::rstest;
|
||||
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn lookup_prefers_the_keyword_even_when_none_and_falls_back_to_the_request() {
|
||||
fn dicts<'py>(
|
||||
py: Python<'py>,
|
||||
kwargs: &str,
|
||||
bound: &str,
|
||||
) -> (Bound<'py, PyDict>, Bound<'py, PyDict>) {
|
||||
let eval = |source: &str| {
|
||||
py.eval(&std::ffi::CString::new(source).unwrap(), None, None)
|
||||
.unwrap()
|
||||
.cast_into::<PyDict>()
|
||||
.unwrap()
|
||||
};
|
||||
(eval(kwargs), eval(bound))
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
#[case::keyword_wins(
|
||||
"{'api_key': 'keyword'}",
|
||||
"{'api_key': 'bound'}",
|
||||
Some(Some("keyword"))
|
||||
)]
|
||||
#[case::explicit_none_wins("{'api_key': None}", "{'api_key': 'bound'}", Some(None))]
|
||||
#[case::bound_fallback("{}", "{'api_key': 'bound'}", Some(Some("bound")))]
|
||||
#[case::missing("{}", "{}", None)]
|
||||
fn lookup_prefers_the_keyword_and_falls_back_to_bound(
|
||||
#[case] kwargs: &str,
|
||||
#[case] bound: &str,
|
||||
#[case] expected: Option<Option<&str>>,
|
||||
) {
|
||||
crate::initialize_python();
|
||||
Python::attach(|py| {
|
||||
let locals = PyDict::new(py);
|
||||
py.run(
|
||||
c"
|
||||
key = object()
|
||||
document = {'type': 'document_url'}
|
||||
class Request:
|
||||
api_key = 'from-request'
|
||||
api_base = 'from-request'
|
||||
document = document
|
||||
request = Request()
|
||||
kwargs = {'api_key': key, 'api_base': None}
|
||||
",
|
||||
Some(&locals),
|
||||
Some(&locals),
|
||||
)
|
||||
.unwrap();
|
||||
let item = |name: &str| locals.get_item(name).unwrap().unwrap();
|
||||
let kwargs = item("kwargs").cast_into::<PyDict>().unwrap();
|
||||
let request = item("request");
|
||||
let find = |name: &str| lookup(&kwargs, &request, name).unwrap();
|
||||
assert!(find("api_key").unwrap().is(item("key")));
|
||||
assert!(find("api_base").unwrap().is_none());
|
||||
assert!(find("document").unwrap().is(item("document")));
|
||||
assert!(find("model").is_none());
|
||||
let (kwargs, bound) = dicts(py, kwargs, bound);
|
||||
let value = lookup(&kwargs, &bound, "api_key")
|
||||
.unwrap()
|
||||
.map(|value| value.extract::<Option<String>>().unwrap());
|
||||
assert_eq!(value, expected.map(|value| value.map(str::to_owned)));
|
||||
});
|
||||
}
|
||||
|
||||
#[rstest::rstest]
|
||||
#[case::prepared_value("{'api_key': 'replacement'}", Some("replacement"))]
|
||||
#[case::explicit_none("{'api_key': None}", None)]
|
||||
#[case::bound_fallback("{}", Some("original"))]
|
||||
fn prepared_mapping_overrides_bound_values(
|
||||
#[case] source: &str,
|
||||
#[rstest]
|
||||
#[case::explicit_none_hides_bound("{'api_key': None}", "{'api_key': 'bound'}", None)]
|
||||
#[case::bound_none("{}", "{'api_key': None}", None)]
|
||||
#[case::bound_value("{}", "{'api_key': 'bound'}", Some("bound"))]
|
||||
fn present_treats_none_as_unset(
|
||||
#[case] kwargs: &str,
|
||||
#[case] bound: &str,
|
||||
#[case] expected: Option<&str>,
|
||||
) {
|
||||
crate::initialize_python();
|
||||
Python::attach(|py| {
|
||||
let (kwargs, bound) = dicts(py, kwargs, bound);
|
||||
let value = present(&kwargs, &bound, "api_key")
|
||||
.unwrap()
|
||||
.map(|value| value.extract::<String>().unwrap());
|
||||
assert_eq!(value.as_deref(), expected);
|
||||
});
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn lookup_returns_the_callers_object() {
|
||||
crate::initialize_python();
|
||||
Python::attach(|py| {
|
||||
let document = PyDict::new(py);
|
||||
let bound = PyDict::new(py);
|
||||
bound.set_item("api_key", "original").unwrap();
|
||||
let source = std::ffi::CString::new(source).unwrap();
|
||||
let prepared = py
|
||||
.eval(&source, None, None)
|
||||
.unwrap()
|
||||
.cast_into::<PyDict>()
|
||||
.unwrap();
|
||||
let value = lookup(&prepared, bound.as_any(), "api_key")
|
||||
bound.set_item("document", &document).unwrap();
|
||||
let found = lookup(&PyDict::new(py), &bound, "document")
|
||||
.unwrap()
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
value.extract::<Option<String>>().unwrap().as_deref(),
|
||||
expected
|
||||
);
|
||||
assert!(found.is(&document));
|
||||
});
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -19,7 +19,7 @@ mod owned;
|
|||
mod runtime;
|
||||
mod services;
|
||||
|
||||
pub use argument::lookup;
|
||||
pub use argument::{lookup, present};
|
||||
pub use binding::PythonBinding;
|
||||
pub use conversion_cache::{FromPythonCache, ToPythonCache};
|
||||
pub use driver::{CallOptions, run_call};
|
||||
|
|
|
|||
|
|
@ -17,6 +17,10 @@ mod tokenizer;
|
|||
|
||||
#[pymodule(gil_used = true)]
|
||||
mod _native {
|
||||
#[pymodule_export]
|
||||
use litellm_host_python::{ForkedAfterNativeRuntimeStarted, ProcessReservedForForking};
|
||||
use pyo3::{prelude::*, types::PyModule};
|
||||
|
||||
#[cfg(feature = "panic-test")]
|
||||
#[pymodule_export]
|
||||
use crate::diagnostics::_panic_for_test;
|
||||
|
|
@ -29,9 +33,7 @@ mod _native {
|
|||
#[pymodule_export]
|
||||
use crate::routes::audio_transcription::{atranscription, transcription};
|
||||
#[pymodule_export]
|
||||
use crate::routes::chat_completions::{
|
||||
achat_completions, acompletion, chat_completions, completion,
|
||||
};
|
||||
use crate::routes::chat_completions::{acompletion, completion};
|
||||
#[pymodule_export]
|
||||
use crate::routes::embeddings::{aembedding, embedding};
|
||||
#[pymodule_export]
|
||||
|
|
@ -51,9 +53,6 @@ mod _native {
|
|||
use crate::tokenizer::HuggingFaceEncoding;
|
||||
#[pymodule_export]
|
||||
use crate::tokenizer::Tokenizer;
|
||||
#[pymodule_export]
|
||||
use litellm_host_python::{ForkedAfterNativeRuntimeStarted, ProcessReservedForForking};
|
||||
use pyo3::{prelude::*, types::PyModule};
|
||||
|
||||
#[pymodule_init]
|
||||
fn init(module: &Bound<'_, PyModule>) -> PyResult<()> {
|
||||
|
|
@ -101,8 +100,6 @@ mod tests {
|
|||
"atranscription",
|
||||
"messages",
|
||||
"amessages",
|
||||
"chat_completions",
|
||||
"achat_completions",
|
||||
"completion",
|
||||
"acompletion",
|
||||
"responses",
|
||||
|
|
|
|||
|
|
@ -19,13 +19,6 @@ pub(crate) struct RouteOptions {
|
|||
pub(crate) timeout: Option<Duration>,
|
||||
}
|
||||
|
||||
pub(crate) fn messages_argument(value: &Bound<'_, PyAny>) -> PyResult<Vec<Value>> {
|
||||
match from_py_argument(value)? {
|
||||
Value::Array(values) => Ok(values),
|
||||
_ => Err(PyValueError::new_err("messages must be a list")),
|
||||
}
|
||||
}
|
||||
|
||||
fn required_object(name: &'static str, value: Value) -> PyResult<Map<String, Value>> {
|
||||
match value {
|
||||
Value::Object(values) => Ok(values),
|
||||
|
|
@ -449,18 +442,6 @@ module.__getattr__ = fail
|
|||
fn argument_converters_keep_nested_values_and_accept_explicit_none() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let messages = py
|
||||
.eval(
|
||||
c"[{'role': 'user', 'content': [{'type': 'text', 'text': 'hi'}]}]",
|
||||
None,
|
||||
None,
|
||||
)
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
Value::Array(messages_argument(&messages).unwrap()),
|
||||
json!([{"role": "user", "content": [{"type": "text", "text": "hi"}]}])
|
||||
);
|
||||
|
||||
let params = py.eval(c"{'temperature': 0.2}", None, None).unwrap();
|
||||
assert_eq!(
|
||||
optional_object("optional_params", ¶ms).unwrap(),
|
||||
|
|
|
|||
|
|
@ -11,3 +11,7 @@ The host driver owns sequencing and terminal events; the bridge supplies fallibl
|
|||
Use the shared `run_public_call` boundary with hooks supplied by bridge composition. `callbacks-legacy-python` owns legacy argument sharing and `Logging` dispatch behind `PublicCall` and `LegacyLogging`. Route bindings supply `callbacks-legacy-python::LoggingOperation` when composing legacy logging and may retain the request needed for projection, but must not duplicate the legacy callback contract
|
||||
|
||||
Regression tests must observe that an unstarted call does no setup, hook and preflight rewrites affect resource configuration, setup failures reach the selected failure handler once, and provider work is not replayed. Retain existing read-point and object-identity guarantees while changing setup timing
|
||||
|
||||
## Layout
|
||||
|
||||
`messages/` is the reference shape. Each lifecycle route is a folder: `mod.rs` holds `run_<route>` and the thin sync and async `#[pyfunction]` entrypoints that take `NativeCall` and call it, `host.rs` holds the `PythonBinding` host, and anything else route-specific (projection, error mapping, extra pyclasses) gets its own sibling file. Scaffolded routes without a lifecycle (embeddings, audio transcription) stay a single `<route>.rs`
|
||||
|
|
|
|||
|
|
@ -5,6 +5,8 @@ use litellm_inference_transcription::{
|
|||
use pyo3::prelude::*;
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
use super::NativeCall;
|
||||
|
||||
use crate::{
|
||||
errors::route_error_to_pyerr,
|
||||
marshal::{RouteOptions, optional_object_field, required_field, value_route_options},
|
||||
|
|
@ -40,8 +42,7 @@ async fn execute(
|
|||
}
|
||||
|
||||
#[pyfunction]
|
||||
pub(crate) fn transcription(py: Python<'_>, call: Bound<'_, PyAny>) -> PyResult<Py<PyAny>> {
|
||||
let call = super::NativeCall::extract(&call)?;
|
||||
pub(crate) fn transcription(py: Python<'_>, call: NativeCall<'_>) -> PyResult<Py<PyAny>> {
|
||||
let audio: Value =
|
||||
litellm_host_python::from_py_argument(&required_field(&call.bound, "audio")?)?;
|
||||
let options = value_route_options(&call.bound)?;
|
||||
|
|
@ -59,9 +60,8 @@ pub(crate) fn transcription(py: Python<'_>, call: Bound<'_, PyAny>) -> PyResult<
|
|||
#[pyfunction]
|
||||
pub(crate) fn atranscription<'py>(
|
||||
py: Python<'py>,
|
||||
call: Bound<'py, PyAny>,
|
||||
call: NativeCall<'py>,
|
||||
) -> PyResult<Bound<'py, PyAny>> {
|
||||
let call = super::NativeCall::extract(&call)?;
|
||||
let audio: Value =
|
||||
litellm_host_python::from_py_argument(&required_field(&call.bound, "audio")?)?;
|
||||
let options = value_route_options(&call.bound)?;
|
||||
|
|
|
|||
|
|
@ -1,150 +0,0 @@
|
|||
mod host;
|
||||
|
||||
use pyo3::types::{PyDict, PyTuple};
|
||||
|
||||
use crate::execution::{run_async, run_sync};
|
||||
use litellm_inference_chat::{ChatCompletionsRoute, Error, types::ChatCompletionsRequest};
|
||||
use litellm_llms_types::formats::chat_completions::ChatCompletionsResponse;
|
||||
use pyo3::prelude::*;
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
use crate::{
|
||||
errors::route_error_to_pyerr,
|
||||
marshal::{
|
||||
RouteOptions, messages_argument, optional_object_field, required_field, value_route_options,
|
||||
},
|
||||
};
|
||||
|
||||
async fn execute(
|
||||
http: Result<litellm_http::Client, litellm_http::Error>,
|
||||
secrets: std::sync::Arc<dyn litellm_secrets::source::SecretSource>,
|
||||
messages: Vec<Value>,
|
||||
optional_params: Map<String, Value>,
|
||||
options: RouteOptions,
|
||||
) -> Result<ChatCompletionsResponse, Error> {
|
||||
let RouteOptions {
|
||||
model,
|
||||
api_key,
|
||||
api_base,
|
||||
custom_llm_provider,
|
||||
extra_headers,
|
||||
timeout,
|
||||
} = options;
|
||||
ChatCompletionsRoute::new(http?, crate::http::resources().auth.clone(), secrets)
|
||||
.execute(
|
||||
ChatCompletionsRequest {
|
||||
model: &model,
|
||||
messages: Value::Array(messages),
|
||||
optional_params,
|
||||
api_key: api_key.as_deref(),
|
||||
api_base: api_base.as_deref(),
|
||||
custom_llm_provider: custom_llm_provider.as_deref(),
|
||||
extra_headers,
|
||||
timeout,
|
||||
},
|
||||
&(),
|
||||
None,
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
#[pyfunction]
|
||||
pub(crate) fn chat_completions(py: Python<'_>, call: Bound<'_, PyAny>) -> PyResult<Py<PyAny>> {
|
||||
let call = super::NativeCall::extract(&call)?;
|
||||
let messages: Vec<Value> = messages_argument(&required_field(&call.bound, "messages")?)?;
|
||||
let optional_params =
|
||||
optional_object_field(&call.bound, "optional_params")?.unwrap_or_default();
|
||||
let options = value_route_options(&call.bound)?;
|
||||
let http = crate::http::provider_client(py, &call.kwargs, false)?;
|
||||
let secrets = crate::secrets::source(py)?;
|
||||
run_sync(
|
||||
py,
|
||||
execute(http, secrets, messages, optional_params, options),
|
||||
route_error_to_pyerr,
|
||||
)
|
||||
}
|
||||
|
||||
#[pyfunction]
|
||||
pub(crate) fn achat_completions<'py>(
|
||||
py: Python<'py>,
|
||||
call: Bound<'py, PyAny>,
|
||||
) -> PyResult<Bound<'py, PyAny>> {
|
||||
let call = super::NativeCall::extract(&call)?;
|
||||
let messages: Vec<Value> = messages_argument(&required_field(&call.bound, "messages")?)?;
|
||||
let optional_params =
|
||||
optional_object_field(&call.bound, "optional_params")?.unwrap_or_default();
|
||||
let options = value_route_options(&call.bound)?;
|
||||
let http = crate::http::provider_client(py, &call.kwargs, true)?;
|
||||
let secrets = crate::secrets::source(py)?;
|
||||
run_async(
|
||||
py,
|
||||
execute(http, secrets, messages, optional_params, options),
|
||||
route_error_to_pyerr,
|
||||
)
|
||||
}
|
||||
|
||||
fn run_public(
|
||||
py: Python<'_>,
|
||||
request: Bound<'_, PyAny>,
|
||||
args: Bound<'_, PyTuple>,
|
||||
kwargs: Bound<'_, PyDict>,
|
||||
asynchronous: bool,
|
||||
) -> PyResult<Py<PyAny>> {
|
||||
use super::inference::InferenceHost;
|
||||
use litellm_callbacks_legacy_python::LoggingOperation;
|
||||
let host = InferenceHost::new(
|
||||
request.clone().unbind(),
|
||||
"litellm.rust_bridge.chat_completions.route_host",
|
||||
);
|
||||
let cache_call_type = if asynchronous {
|
||||
"acompletion"
|
||||
} else {
|
||||
"completion"
|
||||
};
|
||||
crate::cache::admit_native(py, &kwargs, cache_call_type)?;
|
||||
let (arguments, hooks) = crate::routes::call_hooks(
|
||||
py,
|
||||
LoggingOperation::Completion,
|
||||
&request,
|
||||
&args,
|
||||
&kwargs,
|
||||
asynchronous,
|
||||
)?;
|
||||
crate::routes::run_public_call(
|
||||
py,
|
||||
arguments,
|
||||
move |py, arguments, request| {
|
||||
let route = ChatCompletionsRoute::new(
|
||||
crate::http::provider_client(py, arguments, asynchronous)?
|
||||
.map_err(crate::http::client_error)?,
|
||||
crate::http::resources().auth.clone(),
|
||||
crate::secrets::source(py)?,
|
||||
);
|
||||
let (cache, cache_options) =
|
||||
crate::cache::configured_native(py, arguments, cache_call_type)?;
|
||||
let route = match cache {
|
||||
Some(cache) => route.with_cache(litellm_cache_response::ScopedCache::new(
|
||||
cache,
|
||||
litellm_cache_response::CacheScope::Shared,
|
||||
)),
|
||||
None => route,
|
||||
};
|
||||
Ok(route.machine(request, cache_options.policy))
|
||||
},
|
||||
host::ChatCompletionsPythonHost(host),
|
||||
hooks,
|
||||
asynchronous,
|
||||
)
|
||||
}
|
||||
|
||||
#[pyfunction]
|
||||
pub(crate) fn completion(py: Python<'_>, call: Bound<'_, PyAny>) -> PyResult<Py<PyAny>> {
|
||||
let call = super::NativeCall::extract(&call)?;
|
||||
run_public(py, call.bound.into_any(), call.args, call.kwargs, false)
|
||||
}
|
||||
|
||||
#[pyfunction]
|
||||
pub(crate) fn acompletion(py: Python<'_>, call: Bound<'_, PyAny>) -> PyResult<Py<PyAny>> {
|
||||
let call = super::NativeCall::extract(&call)?;
|
||||
run_public(py, call.bound.into_any(), call.args, call.kwargs, true)
|
||||
}
|
||||
|
|
@ -1,6 +1,6 @@
|
|||
use std::convert::Infallible;
|
||||
|
||||
use super::super::inference::InferenceHost;
|
||||
use crate::routes::inference::InferenceHost;
|
||||
use litellm_host_python::{InvokeError, PythonBinding, PythonHostCalls, PythonOwned};
|
||||
use litellm_inference_chat::{Error, route::ChatCompletions, types::ChatCompletionsCall};
|
||||
use pyo3::{
|
||||
|
|
|
|||
|
|
@ -0,0 +1,71 @@
|
|||
mod host;
|
||||
|
||||
use std::sync::Arc;
|
||||
|
||||
use host::ChatCompletionsPythonHost;
|
||||
use litellm_auth::AuthServices;
|
||||
use litellm_cache_response::{CachePolicy, ScopedCache};
|
||||
use litellm_callbacks_legacy_python::LoggingOperation;
|
||||
use litellm_host::{call::HostedMachine, protocol::Protocol};
|
||||
use litellm_inference_chat::{ChatCompletionsRoute, route::ChatCompletions};
|
||||
use litellm_secrets::source::SecretSource;
|
||||
use pyo3::prelude::*;
|
||||
|
||||
use super::{
|
||||
NativeCall,
|
||||
inference::{InferenceHost, InferenceRoute, run_inference},
|
||||
};
|
||||
|
||||
fn run_chat_completions(
|
||||
py: Python<'_>,
|
||||
call: NativeCall<'_>,
|
||||
asynchronous: bool,
|
||||
) -> PyResult<Py<PyAny>> {
|
||||
let host = InferenceHost::new(
|
||||
call.bound.clone().unbind(),
|
||||
"litellm.rust_bridge.chat_completions.route_host",
|
||||
);
|
||||
run_inference::<ChatCompletionsRoute, _>(
|
||||
py,
|
||||
call,
|
||||
asynchronous,
|
||||
ChatCompletionsPythonHost(host),
|
||||
)
|
||||
}
|
||||
|
||||
#[pyfunction]
|
||||
pub(crate) fn completion(py: Python<'_>, call: NativeCall<'_>) -> PyResult<Py<PyAny>> {
|
||||
run_chat_completions(py, call, false)
|
||||
}
|
||||
|
||||
#[pyfunction]
|
||||
pub(crate) fn acompletion(py: Python<'_>, call: NativeCall<'_>) -> PyResult<Py<PyAny>> {
|
||||
run_chat_completions(py, call, true)
|
||||
}
|
||||
|
||||
impl InferenceRoute for ChatCompletionsRoute {
|
||||
type Protocol = ChatCompletions;
|
||||
const OPERATION: LoggingOperation = LoggingOperation::Completion;
|
||||
const SYNC_CALL_TYPE: &'static str = "completion";
|
||||
const ASYNC_CALL_TYPE: &'static str = "acompletion";
|
||||
|
||||
fn new(
|
||||
http: litellm_http::Client,
|
||||
auth: Arc<AuthServices>,
|
||||
secrets: Arc<dyn SecretSource>,
|
||||
) -> Self {
|
||||
Self::new(http, auth, secrets)
|
||||
}
|
||||
|
||||
fn with_cache(self, cache: ScopedCache) -> Self {
|
||||
self.with_cache(cache)
|
||||
}
|
||||
|
||||
fn machine(
|
||||
self,
|
||||
call: <ChatCompletions as Protocol>::Request,
|
||||
policy: CachePolicy,
|
||||
) -> HostedMachine<ChatCompletions> {
|
||||
self.machine(call, policy)
|
||||
}
|
||||
}
|
||||
|
|
@ -1,17 +1,18 @@
|
|||
use pyo3::prelude::*;
|
||||
|
||||
use super::NativeCall;
|
||||
use crate::errors::RustBridgeDeclined;
|
||||
|
||||
#[pyfunction]
|
||||
pub(crate) fn embedding(call: Bound<'_, PyAny>) -> PyResult<Py<PyAny>> {
|
||||
drop(super::NativeCall::extract(&call)?);
|
||||
pub(crate) fn embedding(call: NativeCall<'_>) -> PyResult<Py<PyAny>> {
|
||||
drop(call);
|
||||
Err(RustBridgeDeclined::new_err(
|
||||
"native embeddings route is not implemented",
|
||||
))
|
||||
}
|
||||
|
||||
#[pyfunction]
|
||||
pub(crate) fn aembedding(call: Bound<'_, PyAny>) -> PyResult<Py<PyAny>> {
|
||||
pub(crate) fn aembedding(call: NativeCall<'_>) -> PyResult<Py<PyAny>> {
|
||||
embedding(call)
|
||||
}
|
||||
|
||||
|
|
@ -36,7 +37,8 @@ call = SimpleNamespace(args=(), kwargs={}, bound={'model':'test-model','input':'
|
|||
Some(&locals),
|
||||
)
|
||||
.unwrap();
|
||||
let call = locals.get_item("call").unwrap().unwrap();
|
||||
let call: super::NativeCall<'_> =
|
||||
locals.get_item("call").unwrap().unwrap().extract().unwrap();
|
||||
let error = if asynchronous {
|
||||
super::aembedding(call)
|
||||
} else {
|
||||
|
|
|
|||
|
|
@ -1,11 +1,19 @@
|
|||
use std::sync::Arc;
|
||||
|
||||
use litellm_auth::AuthServices;
|
||||
use litellm_cache_response::{CachePolicy, CacheScope, ScopedCache};
|
||||
use litellm_callbacks_legacy_python::LoggingOperation;
|
||||
use litellm_core_utils::get_llm_provider_logic::get_custom_llm_provider;
|
||||
use litellm_host_python::{from_py, lookup};
|
||||
use litellm_host::{call::HostedMachine, protocol::Protocol};
|
||||
use litellm_host_python::{PythonBinding, PythonHostCalls, from_py, present};
|
||||
use litellm_http::transport::Error as TransportError;
|
||||
use litellm_inference::RouteError;
|
||||
use litellm_secrets::source::SecretSource;
|
||||
use pyo3::{exceptions::PyValueError, prelude::*, types::PyDict};
|
||||
use serde::Serialize;
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
use super::NativeCall;
|
||||
use crate::{
|
||||
errors::{RustUpstreamError, route_error_to_pyerr},
|
||||
marshal::{
|
||||
|
|
@ -15,7 +23,7 @@ use crate::{
|
|||
};
|
||||
|
||||
pub(super) struct InferenceHost {
|
||||
pub request: Py<PyAny>,
|
||||
pub request: Py<PyDict>,
|
||||
module: &'static str,
|
||||
}
|
||||
|
||||
|
|
@ -26,7 +34,7 @@ pub(super) struct ProjectedCall {
|
|||
}
|
||||
|
||||
impl InferenceHost {
|
||||
pub fn new(request: Py<PyAny>, module: &'static str) -> Self {
|
||||
pub fn new(request: Py<PyDict>, module: &'static str) -> Self {
|
||||
Self { request, module }
|
||||
}
|
||||
|
||||
|
|
@ -85,21 +93,7 @@ impl InferenceHost {
|
|||
arguments: &Bound<'py, PyDict>,
|
||||
name: &str,
|
||||
) -> PyResult<Option<Bound<'py, PyAny>>> {
|
||||
let request = self.request.bind(py);
|
||||
if let Some(value) = lookup(arguments, request, name)? {
|
||||
return Ok((!value.is_none()).then_some(value));
|
||||
}
|
||||
if request.is_instance_of::<PyDict>() {
|
||||
return Ok(None);
|
||||
}
|
||||
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))
|
||||
present(arguments, self.request.bind(py), name)
|
||||
}
|
||||
|
||||
pub fn parameters(
|
||||
|
|
@ -140,3 +134,62 @@ impl InferenceHost {
|
|||
Ok(PyErr::from_value(mapped))
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) trait InferenceRoute: Sized + 'static {
|
||||
type Protocol: Protocol<Error = RouteError>;
|
||||
const OPERATION: LoggingOperation;
|
||||
const SYNC_CALL_TYPE: &'static str;
|
||||
const ASYNC_CALL_TYPE: &'static str;
|
||||
|
||||
fn new(
|
||||
http: litellm_http::Client,
|
||||
auth: Arc<AuthServices>,
|
||||
secrets: Arc<dyn SecretSource>,
|
||||
) -> Self;
|
||||
fn with_cache(self, cache: ScopedCache) -> Self;
|
||||
fn machine(
|
||||
self,
|
||||
call: <Self::Protocol as Protocol>::Request,
|
||||
policy: CachePolicy,
|
||||
) -> HostedMachine<Self::Protocol>;
|
||||
}
|
||||
|
||||
pub(super) fn run_inference<R, H>(
|
||||
py: Python<'_>,
|
||||
call: NativeCall<'_>,
|
||||
asynchronous: bool,
|
||||
host: H,
|
||||
) -> PyResult<Py<PyAny>>
|
||||
where
|
||||
R: InferenceRoute,
|
||||
H: PythonBinding<Protocol = R::Protocol> + PythonHostCalls<R::Protocol> + 'static,
|
||||
{
|
||||
let call_type = if asynchronous {
|
||||
R::ASYNC_CALL_TYPE
|
||||
} else {
|
||||
R::SYNC_CALL_TYPE
|
||||
};
|
||||
crate::cache::admit_native(py, &call.kwargs, call_type)?;
|
||||
let (arguments, hooks) = super::call_hooks(py, R::OPERATION, &call, asynchronous)?;
|
||||
super::run_public_call(
|
||||
py,
|
||||
arguments,
|
||||
move |py, arguments, request| {
|
||||
let route = R::new(
|
||||
crate::http::provider_client(py, arguments, asynchronous)?
|
||||
.map_err(crate::http::client_error)?,
|
||||
crate::http::resources().auth.clone(),
|
||||
crate::secrets::source(py)?,
|
||||
);
|
||||
let (cache, cache_options) = crate::cache::configured_native(py, arguments, call_type)?;
|
||||
let route = match cache {
|
||||
Some(cache) => route.with_cache(ScopedCache::new(cache, CacheScope::Shared)),
|
||||
None => route,
|
||||
};
|
||||
Ok(route.machine(request, cache_options.policy))
|
||||
},
|
||||
host,
|
||||
hooks,
|
||||
asynchronous,
|
||||
)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -2,7 +2,7 @@ use crate::cache::{CacheCall, Cached, PythonCache, Selection};
|
|||
use litellm_host_python::{PythonHostCalls, PythonOwned};
|
||||
|
||||
use bytes::Bytes;
|
||||
use litellm_host_python::{InvokeError, PythonBinding, from_py, lookup, to_py};
|
||||
use litellm_host_python::{InvokeError, PythonBinding, from_py, present, to_py};
|
||||
use litellm_http::transport::Error as TransportError;
|
||||
use litellm_inference_messages::{
|
||||
Error, MessagesCall, MessagesSettings, MessagesShaping, messages_body,
|
||||
|
|
@ -89,12 +89,12 @@ fn native_error(py: Python<'_>, error: Error) -> PyResult<PyErr> {
|
|||
/// The Python side of the Messages route: projects the prepared arguments and builds the
|
||||
/// public response, chunks and exceptions.
|
||||
pub(super) struct MessagesPythonHost {
|
||||
request: Py<PyAny>,
|
||||
request: Py<PyDict>,
|
||||
cache: PythonCache,
|
||||
}
|
||||
|
||||
impl MessagesPythonHost {
|
||||
pub(super) fn new(request: Py<PyAny>, asynchronous: bool) -> Self {
|
||||
pub(super) fn new(request: Py<PyDict>, asynchronous: bool) -> Self {
|
||||
Self {
|
||||
request,
|
||||
cache: PythonCache::new(asynchronous),
|
||||
|
|
@ -107,9 +107,7 @@ impl MessagesPythonHost {
|
|||
arguments: &Bound<'_, PyDict>,
|
||||
) -> PyResult<Result<MessagesCall, Error>> {
|
||||
let request = self.request.bind(py);
|
||||
let argument = |name: &str| -> PyResult<Option<Bound<'_, PyAny>>> {
|
||||
Ok(lookup(arguments, request, name)?.filter(|value| !value.is_none()))
|
||||
};
|
||||
let argument = |name: &str| present(arguments, request, name);
|
||||
let string = |name: &str| -> PyResult<Option<String>> {
|
||||
argument(name)?.map(|value| value.extract()).transpose()
|
||||
};
|
||||
|
|
@ -153,8 +151,7 @@ impl MessagesPythonHost {
|
|||
) -> PyResult<Option<Map<String, Value>>> {
|
||||
let request = self.request.bind(py);
|
||||
let mapping = |name: &str| -> PyResult<Option<Map<String, Value>>> {
|
||||
lookup(arguments, request, name)?
|
||||
.filter(|value| !value.is_none())
|
||||
present(arguments, request, name)?
|
||||
.map(|value| from_py(&value))
|
||||
.transpose()
|
||||
};
|
||||
|
|
@ -169,8 +166,7 @@ impl MessagesPythonHost {
|
|||
py: Python<'_>,
|
||||
arguments: &Bound<'_, PyDict>,
|
||||
) -> PyResult<Option<ProviderSpecificHeaders>> {
|
||||
lookup(arguments, self.request.bind(py), "provider_specific_header")?
|
||||
.filter(|value| !value.is_none())
|
||||
present(arguments, self.request.bind(py), "provider_specific_header")?
|
||||
.map(|value| from_py(&value))
|
||||
.transpose()
|
||||
}
|
||||
|
|
@ -200,9 +196,9 @@ impl MessagesPythonHost {
|
|||
self.request
|
||||
.bind(py)
|
||||
.get_item("custom_llm_provider")
|
||||
.and_then(|value| value.extract::<Option<String>>())
|
||||
.ok()
|
||||
.flatten()
|
||||
.and_then(|value| value.extract::<Option<String>>().ok().flatten())
|
||||
.unwrap_or_else(|| "anthropic".into())
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -2,26 +2,13 @@ mod host;
|
|||
|
||||
use host::MessagesPythonHost;
|
||||
use litellm_callbacks_legacy_python::LoggingOperation;
|
||||
use pyo3::{
|
||||
prelude::*,
|
||||
types::{PyDict, PyTuple},
|
||||
};
|
||||
use pyo3::prelude::*;
|
||||
|
||||
fn run_messages(
|
||||
py: Python<'_>,
|
||||
request: Bound<'_, PyAny>,
|
||||
args: Bound<'_, PyTuple>,
|
||||
kwargs: Bound<'_, PyDict>,
|
||||
asynchronous: bool,
|
||||
) -> PyResult<Py<PyAny>> {
|
||||
let (arguments, hooks) = crate::routes::call_hooks(
|
||||
py,
|
||||
LoggingOperation::Messages,
|
||||
&request,
|
||||
&args,
|
||||
&kwargs,
|
||||
asynchronous,
|
||||
)?;
|
||||
use super::NativeCall;
|
||||
|
||||
fn run_messages(py: Python<'_>, call: NativeCall<'_>, asynchronous: bool) -> PyResult<Py<PyAny>> {
|
||||
let (arguments, hooks) =
|
||||
crate::routes::call_hooks(py, LoggingOperation::Messages, &call, asynchronous)?;
|
||||
crate::routes::run_public_call(
|
||||
py,
|
||||
arguments,
|
||||
|
|
@ -57,20 +44,18 @@ fn run_messages(
|
|||
},
|
||||
))
|
||||
},
|
||||
MessagesPythonHost::new(request.unbind(), asynchronous),
|
||||
MessagesPythonHost::new(call.bound.unbind(), asynchronous),
|
||||
hooks,
|
||||
asynchronous,
|
||||
)
|
||||
}
|
||||
|
||||
#[pyfunction]
|
||||
pub(crate) fn messages(py: Python<'_>, call: Bound<'_, PyAny>) -> PyResult<Py<PyAny>> {
|
||||
let call = super::NativeCall::extract(&call)?;
|
||||
run_messages(py, call.bound.into_any(), call.args, call.kwargs, false)
|
||||
pub(crate) fn messages(py: Python<'_>, call: NativeCall<'_>) -> PyResult<Py<PyAny>> {
|
||||
run_messages(py, call, false)
|
||||
}
|
||||
|
||||
#[pyfunction]
|
||||
pub(crate) fn amessages(py: Python<'_>, call: Bound<'_, PyAny>) -> PyResult<Py<PyAny>> {
|
||||
let call = super::NativeCall::extract(&call)?;
|
||||
run_messages(py, call.bound.into_any(), call.args, call.kwargs, true)
|
||||
pub(crate) fn amessages(py: Python<'_>, call: NativeCall<'_>) -> PyResult<Py<PyAny>> {
|
||||
run_messages(py, call, true)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -8,8 +8,7 @@ pub(crate) mod responses;
|
|||
pub(crate) mod token_counter;
|
||||
pub(crate) mod traces;
|
||||
|
||||
use litellm_callbacks_legacy_python::LoggingOperation;
|
||||
use litellm_callbacks_legacy_python::{LegacyLogging, PublicCall};
|
||||
use litellm_callbacks_legacy_python::{LegacyLogging, LoggingOperation, PublicCall};
|
||||
use litellm_host::{call::HostedCompletion, machine::Machine, protocol::Protocol};
|
||||
use litellm_host_python::{HookChain, PythonBinding, PythonCallHooks, PythonHostCalls};
|
||||
use pyo3::{
|
||||
|
|
@ -17,14 +16,16 @@ use pyo3::{
|
|||
types::{PyDict, PyMapping, PyTuple},
|
||||
};
|
||||
|
||||
struct NativeCall<'py> {
|
||||
pub(crate) struct NativeCall<'py> {
|
||||
args: Bound<'py, PyTuple>,
|
||||
kwargs: Bound<'py, PyDict>,
|
||||
bound: Bound<'py, PyDict>,
|
||||
}
|
||||
|
||||
impl<'py> NativeCall<'py> {
|
||||
fn extract(call: &Bound<'py, PyAny>) -> PyResult<Self> {
|
||||
impl<'py> FromPyObject<'_, 'py> for NativeCall<'py> {
|
||||
type Error = PyErr;
|
||||
|
||||
fn extract(call: Borrowed<'_, 'py, PyAny>) -> PyResult<Self> {
|
||||
Ok(Self {
|
||||
args: call.getattr("args")?.cast_into()?,
|
||||
kwargs: mapping_dict(&call.getattr("kwargs")?)?,
|
||||
|
|
@ -46,12 +47,10 @@ fn mapping_dict<'py>(value: &Bound<'py, PyAny>) -> PyResult<Bound<'py, PyDict>>
|
|||
fn call_hooks(
|
||||
py: Python<'_>,
|
||||
operation: LoggingOperation,
|
||||
request: &Bound<'_, PyAny>,
|
||||
args: &Bound<'_, PyTuple>,
|
||||
kwargs: &Bound<'_, PyDict>,
|
||||
call: &NativeCall<'_>,
|
||||
asynchronous: bool,
|
||||
) -> PyResult<(Py<PyDict>, impl PythonCallHooks + use<>)> {
|
||||
let call = PublicCall::capture(request, args, kwargs)?;
|
||||
let call = PublicCall::capture(&call.bound, &call.args, &call.kwargs)?;
|
||||
let arguments = call.arguments(py);
|
||||
Ok((
|
||||
arguments,
|
||||
|
|
@ -99,6 +98,7 @@ mod tests {
|
|||
prelude::*,
|
||||
types::{PyDict, PyList},
|
||||
};
|
||||
use rstest::rstest;
|
||||
|
||||
fn value_call<'py>(
|
||||
py: Python<'py>,
|
||||
|
|
@ -126,8 +126,10 @@ mod tests {
|
|||
.unwrap()
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn route_arguments_that_fail_to_convert_raise_value_error() {
|
||||
#[rstest]
|
||||
#[case::sync("transcription")]
|
||||
#[case::asynchronous("atranscription")]
|
||||
fn route_arguments_that_fail_to_convert_raise_value_error(#[case] name: &str) {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let module = crate::native_module(py);
|
||||
|
|
@ -151,48 +153,24 @@ value = Broken()
|
|||
.expect("locals should be readable")
|
||||
.expect("helper value should exist");
|
||||
|
||||
for name in ["chat_completions", "achat_completions"] {
|
||||
let error = module
|
||||
.getattr(name)
|
||||
.and_then(|function| {
|
||||
function.call1((value_call(py, "messages", &broken, None),))
|
||||
})
|
||||
.expect_err("route should reject a value it cannot convert");
|
||||
let error = module
|
||||
.getattr(name)
|
||||
.and_then(|function| function.call1((value_call(py, "audio", &broken, None),)))
|
||||
.expect_err("route should reject a value it cannot convert");
|
||||
|
||||
assert!(
|
||||
error.is_instance_of::<pyo3::exceptions::PyValueError>(py),
|
||||
"{name} surfaced {error} instead of ValueError"
|
||||
);
|
||||
}
|
||||
assert!(
|
||||
error.is_instance_of::<pyo3::exceptions::PyValueError>(py),
|
||||
"{name} surfaced {error} instead of ValueError"
|
||||
);
|
||||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[rstest]
|
||||
fn sync_and_async_routes_apply_the_same_input_validation() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let module = crate::native_module(py);
|
||||
|
||||
let invalid_messages = PyDict::new(py);
|
||||
let sync_chat_error = module
|
||||
.getattr("chat_completions")
|
||||
.and_then(|function| {
|
||||
function.call1((value_call(py, "messages", &invalid_messages, None),))
|
||||
})
|
||||
.expect_err("sync chat should reject a non-list messages value");
|
||||
let async_chat_error = module
|
||||
.getattr("achat_completions")
|
||||
.and_then(|function| {
|
||||
function.call1((value_call(py, "messages", &invalid_messages, None),))
|
||||
})
|
||||
.expect_err("async chat should reject a non-list messages value");
|
||||
|
||||
assert_eq!(
|
||||
sync_chat_error.to_string(),
|
||||
"ValueError: messages must be a list"
|
||||
);
|
||||
assert_eq!(async_chat_error.to_string(), sync_chat_error.to_string());
|
||||
|
||||
let invalid_headers = PyList::empty(py);
|
||||
let kwargs = PyDict::new(py);
|
||||
kwargs
|
||||
|
|
@ -221,51 +199,43 @@ value = Broken()
|
|||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
#[rstest]
|
||||
fn missing_and_explicit_none_optional_params_share_the_next_error() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let module = crate::native_module(py);
|
||||
let audio = PyDict::new(py);
|
||||
let transcribe = |optional_params: Option<Bound<'_, PyAny>>| {
|
||||
let kwargs = PyDict::new(py);
|
||||
if let Some(optional_params) = optional_params {
|
||||
kwargs.set_item("optional_params", optional_params).unwrap();
|
||||
}
|
||||
module
|
||||
.getattr("transcription")
|
||||
.and_then(|function| {
|
||||
function.call1((value_call(py, "audio", &audio, Some(&kwargs)),))
|
||||
})
|
||||
.expect_err("model 'model' has no provider")
|
||||
.to_string()
|
||||
};
|
||||
|
||||
let omitted = transcribe(None);
|
||||
assert_eq!(transcribe(Some(py.None().into_bound(py))), omitted);
|
||||
assert_ne!(omitted, "ValueError: optional_params must be a dict");
|
||||
assert_eq!(
|
||||
transcribe(Some(PyList::empty(py).into_any())),
|
||||
"ValueError: optional_params must be a dict"
|
||||
);
|
||||
});
|
||||
}
|
||||
|
||||
#[rstest]
|
||||
fn route_input_validation_preserves_left_to_right_order() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let module = crate::native_module(py);
|
||||
let invalid = PyList::empty(py);
|
||||
|
||||
let chat_kwargs = PyDict::new(py);
|
||||
chat_kwargs
|
||||
.set_item("optional_params", &invalid)
|
||||
.expect("kwargs should accept optional_params");
|
||||
chat_kwargs
|
||||
.set_item("extra_headers", &invalid)
|
||||
.expect("kwargs should accept extra_headers");
|
||||
let invalid_messages = PyDict::new(py);
|
||||
let error = module
|
||||
.getattr("chat_completions")
|
||||
.and_then(|function| {
|
||||
function.call1((value_call(
|
||||
py,
|
||||
"messages",
|
||||
&invalid_messages,
|
||||
Some(&chat_kwargs),
|
||||
),))
|
||||
})
|
||||
.expect_err("messages should be validated first");
|
||||
assert_eq!(error.to_string(), "ValueError: messages must be a list");
|
||||
|
||||
let valid_messages = PyList::empty(py);
|
||||
let error = module
|
||||
.getattr("chat_completions")
|
||||
.and_then(|function| {
|
||||
function.call1((value_call(
|
||||
py,
|
||||
"messages",
|
||||
&valid_messages,
|
||||
Some(&chat_kwargs),
|
||||
),))
|
||||
})
|
||||
.expect_err("optional_params should be validated before headers");
|
||||
assert_eq!(
|
||||
error.to_string(),
|
||||
"ValueError: optional_params must be a dict"
|
||||
);
|
||||
|
||||
let headers_kwargs = PyDict::new(py);
|
||||
headers_kwargs
|
||||
.set_item("extra_headers", &invalid)
|
||||
|
|
@ -286,43 +256,4 @@ value = Broken()
|
|||
assert!(!error.to_string().contains("extra_headers"));
|
||||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn missing_and_explicit_none_optional_params_share_the_next_error() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let module = crate::native_module(py);
|
||||
let messages = PyList::empty(py);
|
||||
let headers = PyList::empty(py);
|
||||
let omitted = PyDict::new(py);
|
||||
omitted
|
||||
.set_item("extra_headers", &headers)
|
||||
.expect("kwargs should accept extra_headers");
|
||||
let explicit = PyDict::new(py);
|
||||
explicit
|
||||
.set_item("optional_params", py.None())
|
||||
.expect("kwargs should accept optional_params");
|
||||
explicit
|
||||
.set_item("extra_headers", &headers)
|
||||
.expect("kwargs should accept extra_headers");
|
||||
|
||||
let omitted_error = module
|
||||
.getattr("chat_completions")
|
||||
.and_then(|function| {
|
||||
function.call1((value_call(py, "messages", &messages, Some(&omitted)),))
|
||||
})
|
||||
.expect_err("omitted optional_params should reach header validation");
|
||||
let explicit_error = module
|
||||
.getattr("chat_completions")
|
||||
.and_then(|function| {
|
||||
function.call1((value_call(py, "messages", &messages, Some(&explicit)),))
|
||||
})
|
||||
.expect_err("None optional_params should reach header validation");
|
||||
assert_eq!(
|
||||
omitted_error.to_string(),
|
||||
"ValueError: extra_headers must be a dict"
|
||||
);
|
||||
assert_eq!(explicit_error.to_string(), omitted_error.to_string());
|
||||
});
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -27,12 +27,12 @@ enum OcrHostData {
|
|||
/// document as it goes), acquires Azure AD tokens, and builds the public response and
|
||||
/// exception.
|
||||
pub(super) struct OcrPythonHost {
|
||||
request: Py<PyAny>,
|
||||
request: Py<PyDict>,
|
||||
data: OcrHostData,
|
||||
}
|
||||
|
||||
impl OcrPythonHost {
|
||||
pub(super) fn new(request: Py<PyAny>) -> Self {
|
||||
pub(super) fn new(request: Py<PyDict>) -> Self {
|
||||
Self {
|
||||
request,
|
||||
data: OcrHostData::Unprojected,
|
||||
|
|
@ -209,7 +209,7 @@ del provider
|
|||
.unwrap()
|
||||
.cast_into::<PyDict>()
|
||||
.unwrap();
|
||||
let mut host = OcrPythonHost::new(py.None());
|
||||
let mut host = OcrPythonHost::new(PyDict::new(py).unbind());
|
||||
assert!(host.decode_request(py, &kwargs).unwrap().caller_token);
|
||||
locals.del_item("kwargs").unwrap();
|
||||
drop(kwargs);
|
||||
|
|
|
|||
|
|
@ -9,10 +9,9 @@ use litellm_core_utils::settings::ProcessEnvironment;
|
|||
use litellm_host_python::to_py;
|
||||
use litellm_inference_ocr::provider_config;
|
||||
use litellm_llms::base_llm::ocr::settings::OcrSettings;
|
||||
use pyo3::{
|
||||
prelude::*,
|
||||
types::{PyDict, PyTuple},
|
||||
};
|
||||
use pyo3::prelude::*;
|
||||
|
||||
use super::NativeCall;
|
||||
|
||||
use crate::{
|
||||
coercion::FieldSpec,
|
||||
|
|
@ -29,21 +28,9 @@ const ENABLE_AZURE_AD_TOKEN_REFRESH: FieldSpec<bool> =
|
|||
Ok(field.exact_true())
|
||||
});
|
||||
|
||||
fn run_ocr(
|
||||
py: Python<'_>,
|
||||
request: Bound<'_, PyAny>,
|
||||
args: Bound<'_, PyTuple>,
|
||||
kwargs: Bound<'_, PyDict>,
|
||||
asynchronous: bool,
|
||||
) -> PyResult<Py<PyAny>> {
|
||||
let (arguments, hooks) = crate::routes::call_hooks(
|
||||
py,
|
||||
LoggingOperation::Ocr,
|
||||
&request,
|
||||
&args,
|
||||
&kwargs,
|
||||
asynchronous,
|
||||
)?;
|
||||
fn run_ocr(py: Python<'_>, call: NativeCall<'_>, asynchronous: bool) -> PyResult<Py<PyAny>> {
|
||||
let (arguments, hooks) =
|
||||
crate::routes::call_hooks(py, LoggingOperation::Ocr, &call, asynchronous)?;
|
||||
crate::routes::run_public_call(
|
||||
py,
|
||||
arguments,
|
||||
|
|
@ -61,7 +48,7 @@ fn run_ocr(
|
|||
let route = litellm_inference_ocr::OcrRoute::new(client);
|
||||
Ok(route.machine(request, None))
|
||||
},
|
||||
OcrPythonHost::new(request.unbind()),
|
||||
OcrPythonHost::new(call.bound.unbind()),
|
||||
hooks,
|
||||
asynchronous,
|
||||
)
|
||||
|
|
@ -81,15 +68,13 @@ fn project_provider_defaults(snapshot: &Snapshot<'_>) -> PyResult<OcrSettings> {
|
|||
}
|
||||
|
||||
#[pyfunction]
|
||||
pub(crate) fn ocr(py: Python<'_>, call: Bound<'_, PyAny>) -> PyResult<Py<PyAny>> {
|
||||
let call = super::NativeCall::extract(&call)?;
|
||||
run_ocr(py, call.bound.into_any(), call.args, call.kwargs, false)
|
||||
pub(crate) fn ocr(py: Python<'_>, call: NativeCall<'_>) -> PyResult<Py<PyAny>> {
|
||||
run_ocr(py, call, false)
|
||||
}
|
||||
|
||||
#[pyfunction]
|
||||
pub(crate) fn aocr(py: Python<'_>, call: Bound<'_, PyAny>) -> PyResult<Py<PyAny>> {
|
||||
let call = super::NativeCall::extract(&call)?;
|
||||
run_ocr(py, call.bound.into_any(), call.args, call.kwargs, true)
|
||||
pub(crate) fn aocr(py: Python<'_>, call: NativeCall<'_>) -> PyResult<Py<PyAny>> {
|
||||
run_ocr(py, call, true)
|
||||
}
|
||||
|
||||
#[pyfunction]
|
||||
|
|
|
|||
|
|
@ -1,5 +1,5 @@
|
|||
use litellm_auth::SecretValue;
|
||||
use litellm_host_python::from_py;
|
||||
use litellm_host_python::{from_py, present};
|
||||
use litellm_inference_ocr::{
|
||||
types::{LiteLLMOcrRequest, OcrDocumentInput},
|
||||
wire::{OcrWireRequest, consumed_optional_params, decode_document, decode_request_input},
|
||||
|
|
@ -22,13 +22,13 @@ pub(super) struct OcrHostHandles {
|
|||
}
|
||||
|
||||
struct OcrArguments<'a, 'py> {
|
||||
request: &'a Bound<'py, PyAny>,
|
||||
bound: &'a Bound<'py, PyDict>,
|
||||
kwargs: &'a Bound<'py, PyDict>,
|
||||
}
|
||||
|
||||
impl<'py> OcrArguments<'_, 'py> {
|
||||
fn lookup(&self, name: &str) -> PyResult<Bound<'py, PyAny>> {
|
||||
litellm_host_python::lookup(self.kwargs, self.request, name)?
|
||||
litellm_host_python::lookup(self.kwargs, self.bound, name)?
|
||||
.ok_or_else(|| PyValueError::new_err(format!("missing argument: {name}")))
|
||||
}
|
||||
|
||||
|
|
@ -58,7 +58,7 @@ impl<'py> OcrArguments<'_, 'py> {
|
|||
fn extra_headers(&self) -> PyResult<Option<Map<String, Value>>> {
|
||||
self.lookup("extra_headers")?
|
||||
.extract::<Option<Py<PyAny>>>()?
|
||||
.map(|value| from_py(value.bind(self.request.py())))
|
||||
.map(|value| from_py(value.bind(self.bound.py())))
|
||||
.transpose()
|
||||
}
|
||||
|
||||
|
|
@ -66,7 +66,7 @@ impl<'py> OcrArguments<'_, 'py> {
|
|||
Ok(self
|
||||
.lookup("timeout")?
|
||||
.extract::<Option<Py<PyAny>>>()?
|
||||
.map(|value| python_timeout_seconds(self.request.py(), value))
|
||||
.map(|value| python_timeout_seconds(self.bound.py(), value))
|
||||
.transpose()?
|
||||
.flatten())
|
||||
}
|
||||
|
|
@ -110,10 +110,10 @@ impl ProjectedDocument {
|
|||
}
|
||||
|
||||
pub(super) fn project_request(
|
||||
request: &Bound<'_, PyAny>,
|
||||
bound: &Bound<'_, PyDict>,
|
||||
kwargs: &Bound<'_, PyDict>,
|
||||
) -> PyResult<(LiteLLMOcrRequest<OcrDocumentInput>, OcrHostHandles)> {
|
||||
let arguments = OcrArguments { request, kwargs };
|
||||
let arguments = OcrArguments { bound, kwargs };
|
||||
let model = arguments.model()?;
|
||||
let custom_llm_provider = arguments.custom_llm_provider()?;
|
||||
let document = ProjectedDocument::project(&arguments.document()?)?;
|
||||
|
|
@ -122,7 +122,7 @@ pub(super) fn project_request(
|
|||
.map_err(ocr_error_to_pyerr)?;
|
||||
let names = specs.iter().map(|spec| spec.name).collect::<Vec<_>>();
|
||||
let optional_params =
|
||||
project_optional_fields(names.iter().copied(), |name| kwargs.get_item(name))?;
|
||||
project_optional_fields(names.iter().copied(), |name| present(kwargs, bound, name))?;
|
||||
let input_sources = request_input_sources(
|
||||
kwargs,
|
||||
names
|
||||
|
|
@ -136,7 +136,7 @@ pub(super) fn project_request(
|
|||
let timeout_seconds = arguments.timeout_seconds()?;
|
||||
let wire = OcrWireRequest {
|
||||
model,
|
||||
document: document.resolve(request.py())?,
|
||||
document: document.resolve(bound.py())?,
|
||||
api_key,
|
||||
api_base,
|
||||
custom_llm_provider,
|
||||
|
|
@ -169,11 +169,20 @@ mod tests {
|
|||
locals
|
||||
}
|
||||
|
||||
fn dict<'py>(locals: &Bound<'py, PyDict>, name: &str) -> Bound<'py, PyDict> {
|
||||
locals
|
||||
.get_item(name)
|
||||
.unwrap()
|
||||
.unwrap()
|
||||
.cast_into::<PyDict>()
|
||||
.unwrap()
|
||||
}
|
||||
|
||||
fn arguments<'a, 'py>(
|
||||
request: &'a Bound<'py, PyAny>,
|
||||
bound: &'a Bound<'py, PyDict>,
|
||||
kwargs: &'a Bound<'py, PyDict>,
|
||||
) -> OcrArguments<'a, 'py> {
|
||||
OcrArguments { request, kwargs }
|
||||
OcrArguments { bound, kwargs }
|
||||
}
|
||||
|
||||
fn project_document(document: &Bound<'_, PyAny>) -> PyResult<OcrDocumentInput> {
|
||||
|
|
@ -204,141 +213,26 @@ sys.modules['litellm.rust_bridge.timeouts'] = timeouts
|
|||
}
|
||||
|
||||
#[test]
|
||||
fn kwargs_override_request_attributes_including_explicit_none() {
|
||||
fn kwargs_override_bound_values_including_explicit_none() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let locals = eval(
|
||||
py,
|
||||
c"
|
||||
class Request:
|
||||
def __init__(self):
|
||||
self.accesses = []
|
||||
def __getattribute__(self, name):
|
||||
if name != 'accesses':
|
||||
object.__getattribute__(self, 'accesses').append(name)
|
||||
return object.__getattribute__(self, name)
|
||||
request = Request()
|
||||
request.model = 'from-request'
|
||||
request.custom_llm_provider = 'mistral'
|
||||
bound = {'model': 'from-bound', 'custom_llm_provider': 'mistral'}
|
||||
kwargs = {'model': 'from-kwargs', 'custom_llm_provider': None}
|
||||
",
|
||||
);
|
||||
let request = locals.get_item("request").unwrap().unwrap();
|
||||
let kwargs = locals
|
||||
.get_item("kwargs")
|
||||
.unwrap()
|
||||
.unwrap()
|
||||
.cast_into::<PyDict>()
|
||||
.unwrap();
|
||||
let arguments = arguments(&request, &kwargs);
|
||||
let (bound, kwargs) = (dict(&locals, "bound"), dict(&locals, "kwargs"));
|
||||
let arguments = arguments(&bound, &kwargs);
|
||||
assert_eq!(arguments.model().unwrap(), "from-kwargs");
|
||||
assert_eq!(arguments.custom_llm_provider().unwrap(), None);
|
||||
let accesses: Vec<String> = request.getattr("accesses").unwrap().extract().unwrap();
|
||||
assert_eq!(accesses, Vec::<String>::new());
|
||||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn missing_kwargs_read_the_request_property_once() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let locals = eval(
|
||||
py,
|
||||
c"
|
||||
class Request:
|
||||
def __init__(self):
|
||||
self.reads = 0
|
||||
@property
|
||||
def model(self):
|
||||
self.reads += 1
|
||||
return 'mistral-ocr-latest'
|
||||
request = Request()
|
||||
kwargs = {}
|
||||
",
|
||||
);
|
||||
let request = locals.get_item("request").unwrap().unwrap();
|
||||
let kwargs = locals
|
||||
.get_item("kwargs")
|
||||
.unwrap()
|
||||
.unwrap()
|
||||
.cast_into::<PyDict>()
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
arguments(&request, &kwargs).model().unwrap(),
|
||||
"mistral-ocr-latest"
|
||||
);
|
||||
assert_eq!(
|
||||
request.getattr("reads").unwrap().extract::<i32>().unwrap(),
|
||||
1
|
||||
);
|
||||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn request_property_exceptions_keep_their_identity() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let locals = eval(
|
||||
py,
|
||||
c"
|
||||
failure = LookupError('model failed')
|
||||
class Request:
|
||||
@property
|
||||
def model(self):
|
||||
raise failure
|
||||
request = Request()
|
||||
kwargs = {}
|
||||
",
|
||||
);
|
||||
let request = locals.get_item("request").unwrap().unwrap();
|
||||
let kwargs = locals
|
||||
.get_item("kwargs")
|
||||
.unwrap()
|
||||
.unwrap()
|
||||
.cast_into::<PyDict>()
|
||||
.unwrap();
|
||||
let error = arguments(&request, &kwargs).model().unwrap_err();
|
||||
assert!(
|
||||
error
|
||||
.value(py)
|
||||
.is(locals.get_item("failure").unwrap().unwrap())
|
||||
);
|
||||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn unused_raising_property_is_never_inspected() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let locals = eval(
|
||||
py,
|
||||
c"
|
||||
class Request:
|
||||
@property
|
||||
def unused(self):
|
||||
raise RuntimeError('unused')
|
||||
model = 'mistral-ocr-latest'
|
||||
custom_llm_provider = None
|
||||
request = Request()
|
||||
kwargs = {}
|
||||
",
|
||||
);
|
||||
let request = locals.get_item("request").unwrap().unwrap();
|
||||
let kwargs = locals
|
||||
.get_item("kwargs")
|
||||
.unwrap()
|
||||
.unwrap()
|
||||
.cast_into::<PyDict>()
|
||||
.unwrap();
|
||||
let arguments = arguments(&request, &kwargs);
|
||||
assert_eq!(arguments.model().unwrap(), "mistral-ocr-latest");
|
||||
assert_eq!(arguments.custom_llm_provider().unwrap(), None);
|
||||
});
|
||||
}
|
||||
|
||||
/// A reader that rewrites the request while it runs shows which arguments projection
|
||||
/// read before it and which after: every other argument is read first, and the read
|
||||
/// happens exactly once.
|
||||
/// A reader that rewrites the bound arguments while it runs shows which arguments
|
||||
/// projection read before it and which after: every other argument is read first, and
|
||||
/// the read happens exactly once.
|
||||
#[test]
|
||||
fn document_readers_are_read_once_after_every_other_argument() {
|
||||
Python::initialize();
|
||||
|
|
@ -347,37 +241,28 @@ kwargs = {}
|
|||
let locals = eval(
|
||||
py,
|
||||
c"
|
||||
class Request:
|
||||
model = 'mistral/mistral-ocr-latest'
|
||||
custom_llm_provider = None
|
||||
api_key = None
|
||||
api_base = 'https://original.example.com'
|
||||
extra_headers = {'x-source': 'original'}
|
||||
timeout = 1
|
||||
@property
|
||||
def document(self):
|
||||
return document
|
||||
class Reader:
|
||||
reads = 0
|
||||
def read(self):
|
||||
Reader.reads += 1
|
||||
Request.api_base = 'https://mutated.example.com'
|
||||
Request.extra_headers = {'x-source': 'mutated'}
|
||||
Request.timeout = 9
|
||||
bound['api_base'] = 'https://mutated.example.com'
|
||||
bound['extra_headers'] = {'x-source': 'mutated'}
|
||||
bound['timeout'] = 9
|
||||
return b'abc'
|
||||
document = {'type': 'file', 'file': Reader(), 'mime_type': 'application/pdf'}
|
||||
request = Request()
|
||||
bound = {
|
||||
'model': 'mistral/mistral-ocr-latest',
|
||||
'custom_llm_provider': None,
|
||||
'api_key': None,
|
||||
'api_base': 'https://original.example.com',
|
||||
'extra_headers': {'x-source': 'original'},
|
||||
'timeout': 1,
|
||||
'document': {'type': 'file', 'file': Reader(), 'mime_type': 'application/pdf'},
|
||||
}
|
||||
kwargs = {}
|
||||
",
|
||||
);
|
||||
let request = locals.get_item("request").unwrap().unwrap();
|
||||
let kwargs = locals
|
||||
.get_item("kwargs")
|
||||
.unwrap()
|
||||
.unwrap()
|
||||
.cast_into::<PyDict>()
|
||||
.unwrap();
|
||||
let (projected, _) = project_request(&request, &kwargs).unwrap();
|
||||
let (projected, _) =
|
||||
project_request(&dict(&locals, "bound"), &dict(&locals, "kwargs")).unwrap();
|
||||
assert_eq!(
|
||||
py.eval(c"Reader.reads", Some(&locals), Some(&locals))
|
||||
.unwrap()
|
||||
|
|
@ -524,34 +409,26 @@ document = Document()
|
|||
});
|
||||
}
|
||||
|
||||
fn request_and_kwargs<'py>(
|
||||
fn bound_and_kwargs<'py>(
|
||||
py: Python<'py>,
|
||||
kwargs: &std::ffi::CStr,
|
||||
) -> (Bound<'py, PyAny>, Bound<'py, PyDict>) {
|
||||
) -> (Bound<'py, PyDict>, Bound<'py, PyDict>) {
|
||||
let locals = eval(
|
||||
py,
|
||||
c"
|
||||
class Request:
|
||||
model = 'mistral/mistral-ocr-latest'
|
||||
custom_llm_provider = 'mistral'
|
||||
document = {'type': 'document_url', 'document_url': 'https://example.com/request.pdf'}
|
||||
api_key = None
|
||||
api_base = 'https://request.example.com'
|
||||
extra_headers = {'x-source': 'request'}
|
||||
timeout = 1
|
||||
request = Request()
|
||||
bound = {
|
||||
'model': 'mistral/mistral-ocr-latest',
|
||||
'custom_llm_provider': 'mistral',
|
||||
'document': {'type': 'document_url', 'document_url': 'https://example.com/bound.pdf'},
|
||||
'api_key': None,
|
||||
'api_base': 'https://bound.example.com',
|
||||
'extra_headers': {'x-source': 'bound'},
|
||||
'timeout': 1,
|
||||
}
|
||||
",
|
||||
);
|
||||
py.run(kwargs, Some(&locals), Some(&locals)).unwrap();
|
||||
(
|
||||
locals.get_item("request").unwrap().unwrap(),
|
||||
locals
|
||||
.get_item("kwargs")
|
||||
.unwrap()
|
||||
.unwrap()
|
||||
.cast_into::<PyDict>()
|
||||
.unwrap(),
|
||||
)
|
||||
(dict(&locals, "bound"), dict(&locals, "kwargs"))
|
||||
}
|
||||
|
||||
#[test]
|
||||
|
|
@ -559,7 +436,7 @@ request = Request()
|
|||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
stub_timeout_conversion(py);
|
||||
let (request, kwargs) = request_and_kwargs(
|
||||
let (bound, kwargs) = bound_and_kwargs(
|
||||
py,
|
||||
c"
|
||||
kwargs = {
|
||||
|
|
@ -575,7 +452,7 @@ kwargs = {
|
|||
}
|
||||
",
|
||||
);
|
||||
let (projected, _) = project_request(&request, &kwargs).unwrap();
|
||||
let (projected, _) = project_request(&bound, &kwargs).unwrap();
|
||||
assert_eq!(
|
||||
projected.optional_params.keys().collect::<Vec<_>>(),
|
||||
["pages"]
|
||||
|
|
@ -584,12 +461,29 @@ kwargs = {
|
|||
});
|
||||
}
|
||||
|
||||
#[rstest::rstest]
|
||||
#[case::explicit_none_is_unset(c"bound['pages'] = [1]\nkwargs = {'pages': None}", None)]
|
||||
#[case::bound_fallback(c"bound['pages'] = [1]\nkwargs = {}", Some(serde_json::json!([1])))]
|
||||
#[case::keyword_wins(c"bound['pages'] = [1]\nkwargs = {'pages': [0]}", Some(serde_json::json!([0])))]
|
||||
fn optional_params_read_through_bound_and_drop_none(
|
||||
#[case] script: &std::ffi::CStr,
|
||||
#[case] expected: Option<Value>,
|
||||
) {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
stub_timeout_conversion(py);
|
||||
let (bound, kwargs) = bound_and_kwargs(py, script);
|
||||
let (projected, _) = project_request(&bound, &kwargs).unwrap();
|
||||
assert_eq!(projected.optional_params.get("pages").cloned(), expected);
|
||||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn replacement_kwargs_project_provider_connection_and_timeout() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
stub_timeout_conversion(py);
|
||||
let (request, kwargs) = request_and_kwargs(
|
||||
let (bound, kwargs) = bound_and_kwargs(
|
||||
py,
|
||||
c"
|
||||
kwargs = {
|
||||
|
|
@ -602,7 +496,7 @@ kwargs = {
|
|||
}
|
||||
",
|
||||
);
|
||||
let (projected, handles) = project_request(&request, &kwargs).unwrap();
|
||||
let (projected, handles) = project_request(&bound, &kwargs).unwrap();
|
||||
assert_eq!(handles.provider, "azure_ai");
|
||||
assert_eq!(projected.model, "mistral-ocr-latest");
|
||||
assert_eq!(
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
use std::convert::Infallible;
|
||||
|
||||
use super::super::inference::InferenceHost;
|
||||
use crate::routes::inference::InferenceHost;
|
||||
use litellm_host_python::{InvokeError, PythonBinding, PythonHostCalls, PythonOwned};
|
||||
use litellm_inference_responses::{Error, route::Responses, types::ResponsesCall};
|
||||
use pyo3::{
|
||||
|
|
|
|||
102
litellm-rust/crates/python-bridge/src/routes/responses/mod.rs
Normal file
102
litellm-rust/crates/python-bridge/src/routes/responses/mod.rs
Normal file
|
|
@ -0,0 +1,102 @@
|
|||
mod host;
|
||||
mod websocket;
|
||||
|
||||
use std::sync::Arc;
|
||||
|
||||
use host::ResponsesPythonHost;
|
||||
use litellm_auth::AuthServices;
|
||||
use litellm_cache_response::{CachePolicy, ScopedCache};
|
||||
use litellm_callbacks_legacy_python::LoggingOperation;
|
||||
use litellm_host::{call::HostedMachine, protocol::Protocol};
|
||||
use litellm_host_python::present;
|
||||
use litellm_inference_responses::{ResponsesRoute, route::Responses};
|
||||
use litellm_secrets::source::SecretSource;
|
||||
use pyo3::prelude::*;
|
||||
use serde_json::Value;
|
||||
pub(crate) use websocket::ResponsesWebSocketConnection;
|
||||
|
||||
use super::{
|
||||
NativeCall,
|
||||
inference::{InferenceHost, InferenceRoute, run_inference},
|
||||
};
|
||||
use crate::errors::RustBridgeDeclined;
|
||||
|
||||
const ROUTE_HOST_MODULE: &str = "litellm.rust_bridge.responses.route_host";
|
||||
|
||||
fn run_responses(py: Python<'_>, call: NativeCall<'_>, asynchronous: bool) -> PyResult<Py<PyAny>> {
|
||||
if let Some(reason) = py
|
||||
.import(ROUTE_HOST_MODULE)?
|
||||
.getattr("decline_reason")?
|
||||
.call1((&call.bound,))?
|
||||
.extract::<Option<String>>()?
|
||||
{
|
||||
return Err(RustBridgeDeclined::new_err(reason));
|
||||
}
|
||||
let argument = |name: &str| present(&call.kwargs, &call.bound, name);
|
||||
let model = argument("model")?
|
||||
.ok_or_else(|| pyo3::exceptions::PyValueError::new_err("model is required"))?
|
||||
.extract::<String>()?;
|
||||
let provider = argument("custom_llm_provider")?
|
||||
.map(|value| value.extract::<String>())
|
||||
.transpose()?;
|
||||
if provider
|
||||
.as_deref()
|
||||
.is_some_and(|provider| provider != "openai")
|
||||
|| model
|
||||
.strip_prefix("openai/")
|
||||
.unwrap_or(&model)
|
||||
.contains('/')
|
||||
{
|
||||
return Err(RustBridgeDeclined::new_err(
|
||||
"native HTTP responses provider",
|
||||
));
|
||||
}
|
||||
if argument("stream")?
|
||||
.map(|value| litellm_host_python::from_py::<Value>(&value))
|
||||
.transpose()?
|
||||
.is_some_and(|value| value == Value::Bool(true))
|
||||
{
|
||||
return Err(RustBridgeDeclined::new_err(
|
||||
"native Python responses streaming",
|
||||
));
|
||||
}
|
||||
let host = InferenceHost::new(call.bound.clone().unbind(), ROUTE_HOST_MODULE);
|
||||
run_inference::<ResponsesRoute, _>(py, call, asynchronous, ResponsesPythonHost(host))
|
||||
}
|
||||
|
||||
#[pyfunction]
|
||||
pub(crate) fn responses(py: Python<'_>, call: NativeCall<'_>) -> PyResult<Py<PyAny>> {
|
||||
run_responses(py, call, false)
|
||||
}
|
||||
|
||||
#[pyfunction]
|
||||
pub(crate) fn aresponses(py: Python<'_>, call: NativeCall<'_>) -> PyResult<Py<PyAny>> {
|
||||
run_responses(py, call, true)
|
||||
}
|
||||
|
||||
impl InferenceRoute for ResponsesRoute {
|
||||
type Protocol = Responses;
|
||||
const OPERATION: LoggingOperation = LoggingOperation::Responses;
|
||||
const SYNC_CALL_TYPE: &'static str = "responses";
|
||||
const ASYNC_CALL_TYPE: &'static str = "aresponses";
|
||||
|
||||
fn new(
|
||||
http: litellm_http::Client,
|
||||
auth: Arc<AuthServices>,
|
||||
secrets: Arc<dyn SecretSource>,
|
||||
) -> Self {
|
||||
Self::new(http, auth, secrets)
|
||||
}
|
||||
|
||||
fn with_cache(self, cache: ScopedCache) -> Self {
|
||||
self.with_cache(cache)
|
||||
}
|
||||
|
||||
fn machine(
|
||||
self,
|
||||
call: <Responses as Protocol>::Request,
|
||||
policy: CachePolicy,
|
||||
) -> HostedMachine<Responses> {
|
||||
self.machine(call, policy)
|
||||
}
|
||||
}
|
||||
|
|
@ -1,121 +1,12 @@
|
|||
mod host;
|
||||
|
||||
use litellm_inference_responses::websocket::ResponsesWebSocketConnection as RustResponsesWebSocketConnection;
|
||||
use pyo3::{
|
||||
prelude::*,
|
||||
types::{PyDict, PyTuple},
|
||||
};
|
||||
use pyo3::prelude::*;
|
||||
use serde_json::Value;
|
||||
|
||||
use crate::{
|
||||
errors::{RustBridgeDeclined, route_error_to_pyerr},
|
||||
errors::route_error_to_pyerr,
|
||||
marshal::{marshal_headers, optional_timeout},
|
||||
};
|
||||
|
||||
fn run_public(
|
||||
py: Python<'_>,
|
||||
request: Bound<'_, PyAny>,
|
||||
args: Bound<'_, PyTuple>,
|
||||
kwargs: Bound<'_, PyDict>,
|
||||
asynchronous: bool,
|
||||
) -> PyResult<Py<PyAny>> {
|
||||
use super::inference::InferenceHost;
|
||||
use litellm_callbacks_legacy_python::LoggingOperation;
|
||||
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 model = host
|
||||
.argument(py, &kwargs, "model")?
|
||||
.ok_or_else(|| pyo3::exceptions::PyValueError::new_err("model is required"))?
|
||||
.extract::<String>()?;
|
||||
let provider = host
|
||||
.argument(py, &kwargs, "custom_llm_provider")?
|
||||
.map(|value| value.extract::<String>())
|
||||
.transpose()?;
|
||||
if provider
|
||||
.as_deref()
|
||||
.is_some_and(|provider| provider != "openai")
|
||||
|| model
|
||||
.strip_prefix("openai/")
|
||||
.unwrap_or(&model)
|
||||
.contains('/')
|
||||
{
|
||||
return Err(RustBridgeDeclined::new_err(
|
||||
"native HTTP responses provider",
|
||||
));
|
||||
}
|
||||
if host
|
||||
.argument(py, &kwargs, "stream")?
|
||||
.map(|value| litellm_host_python::from_py::<Value>(&value))
|
||||
.transpose()?
|
||||
.is_some_and(|value| value == Value::Bool(true))
|
||||
{
|
||||
return Err(RustBridgeDeclined::new_err(
|
||||
"native Python responses streaming",
|
||||
));
|
||||
}
|
||||
let cache_call_type = if asynchronous {
|
||||
"aresponses"
|
||||
} else {
|
||||
"responses"
|
||||
};
|
||||
crate::cache::admit_native(py, &kwargs, cache_call_type)?;
|
||||
let (arguments, hooks) = crate::routes::call_hooks(
|
||||
py,
|
||||
LoggingOperation::Responses,
|
||||
&request,
|
||||
&args,
|
||||
&kwargs,
|
||||
asynchronous,
|
||||
)?;
|
||||
crate::routes::run_public_call(
|
||||
py,
|
||||
arguments,
|
||||
move |py, arguments, request| {
|
||||
let route = litellm_inference_responses::ResponsesRoute::new(
|
||||
crate::http::provider_client(py, arguments, asynchronous)?
|
||||
.map_err(crate::http::client_error)?,
|
||||
crate::http::resources().auth.clone(),
|
||||
crate::secrets::source(py)?,
|
||||
);
|
||||
let (cache, cache_options) =
|
||||
crate::cache::configured_native(py, arguments, cache_call_type)?;
|
||||
let route = match cache {
|
||||
Some(cache) => route.with_cache(litellm_cache_response::ScopedCache::new(
|
||||
cache,
|
||||
litellm_cache_response::CacheScope::Shared,
|
||||
)),
|
||||
None => route,
|
||||
};
|
||||
Ok(route.machine(request, cache_options.policy))
|
||||
},
|
||||
host::ResponsesPythonHost(host),
|
||||
hooks,
|
||||
asynchronous,
|
||||
)
|
||||
}
|
||||
|
||||
#[pyfunction]
|
||||
pub(crate) fn responses(py: Python<'_>, call: Bound<'_, PyAny>) -> PyResult<Py<PyAny>> {
|
||||
let call = super::NativeCall::extract(&call)?;
|
||||
run_public(py, call.bound.into_any(), call.args, call.kwargs, false)
|
||||
}
|
||||
|
||||
#[pyfunction]
|
||||
pub(crate) fn aresponses(py: Python<'_>, call: Bound<'_, PyAny>) -> PyResult<Py<PyAny>> {
|
||||
let call = super::NativeCall::extract(&call)?;
|
||||
run_public(py, call.bound.into_any(), call.args, call.kwargs, true)
|
||||
}
|
||||
|
||||
#[pyclass]
|
||||
pub(crate) struct ResponsesWebSocketConnection {
|
||||
inner: RustResponsesWebSocketConnection,
|
||||
|
|
@ -110,12 +110,6 @@ def messages(
|
|||
def amessages(
|
||||
call: NativeCall,
|
||||
) -> Coroutine[object, object, AnthropicMessagesResponse | AsyncIterator[bytes]]: ...
|
||||
def chat_completions(
|
||||
call: NativeCall,
|
||||
) -> dict[str, object]: ...
|
||||
def achat_completions(
|
||||
call: NativeCall,
|
||||
) -> Future[dict[str, object]]: ...
|
||||
|
||||
@final
|
||||
class ResponsesWebSocketConnection:
|
||||
|
|
@ -251,14 +245,12 @@ __all__ = [
|
|||
"RustUpstreamError",
|
||||
"TokenCounter",
|
||||
"Tokenizer",
|
||||
"achat_completions",
|
||||
"acompletion",
|
||||
"aembedding",
|
||||
"amessages",
|
||||
"aocr",
|
||||
"aresponses",
|
||||
"atranscription",
|
||||
"chat_completions",
|
||||
"completion",
|
||||
"embedding",
|
||||
"gil_stats",
|
||||
|
|
|
|||
|
|
@ -37,16 +37,12 @@ def response(value: Mapping[str, object]) -> ModelResponse:
|
|||
return ModelResponse(**value)
|
||||
|
||||
|
||||
def arguments(request: Mapping[str, object]) -> Mapping[str, object]:
|
||||
return request
|
||||
|
||||
|
||||
def map_failure(error: Exception, request: Mapping[str, object]) -> Exception:
|
||||
provider: Final = optional_str(request.get("custom_llm_provider")) or str(request["model"]).partition("/")[0]
|
||||
return failures.map_native_failure(
|
||||
error,
|
||||
str(request["model"]),
|
||||
provider,
|
||||
arguments(request),
|
||||
request,
|
||||
optional_str(request.get("api_base")) or optional_str(request.get("base_url")),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -39,10 +39,6 @@ def stream_hidden_params(headers: Sequence[tuple[str, str]]) -> Mapping[str, obj
|
|||
return anthropic_messages_stream_hidden_params(httpx.Headers(list(headers)))
|
||||
|
||||
|
||||
def arguments(request: Mapping[str, object]) -> Mapping[str, object]:
|
||||
return request
|
||||
|
||||
|
||||
def map_failure(error: Exception, request: Mapping[str, object], request_provider: str) -> Exception:
|
||||
if getattr(error, "messages_request_error", False):
|
||||
return litellm.BadRequestError(
|
||||
|
|
@ -51,7 +47,7 @@ def map_failure(error: Exception, request: Mapping[str, object], request_provide
|
|||
llm_provider=request_provider,
|
||||
)
|
||||
return failures.map_native_failure(
|
||||
error, str(request["model"]), request_provider, arguments(request), optional_str(request.get("api_base"))
|
||||
error, str(request["model"]), request_provider, request, optional_str(request.get("api_base"))
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -12,7 +12,7 @@ from litellm.rust_bridge import failures
|
|||
from litellm.rust_bridge.failures import UpstreamFailure
|
||||
from litellm.rust_bridge.public_call import optional_str
|
||||
|
||||
__all__ = ("UpstreamFailure", "arguments", "map_failure", "response")
|
||||
__all__ = ("UpstreamFailure", "map_failure", "response")
|
||||
|
||||
_RESPONSE_ADAPTER: Final = TypeAdapter(dict[str, object])
|
||||
|
||||
|
|
@ -27,10 +27,6 @@ def response(value: Mapping[str, object]) -> OCRResponse:
|
|||
return normalized
|
||||
|
||||
|
||||
def arguments(request: Mapping[str, object]) -> Mapping[str, object]:
|
||||
return request
|
||||
|
||||
|
||||
def map_failure(error: Exception, request: Mapping[str, object], request_provider: str) -> Exception:
|
||||
if getattr(error, "ocr_request_format_error", False):
|
||||
return litellm.UnsupportedParamsError(
|
||||
|
|
@ -39,5 +35,5 @@ def map_failure(error: Exception, request: Mapping[str, object], request_provide
|
|||
llm_provider=request_provider,
|
||||
)
|
||||
return failures.map_native_failure(
|
||||
error, str(request["model"]), request_provider, arguments(request), optional_str(request.get("api_base"))
|
||||
error, str(request["model"]), request_provider, request, optional_str(request.get("api_base"))
|
||||
)
|
||||
|
|
|
|||
|
|
@ -20,17 +20,13 @@ def response(value: Mapping[str, object]) -> ResponsesAPIResponse:
|
|||
return ResponsesAPIResponse.model_validate(value)
|
||||
|
||||
|
||||
def arguments(request: Mapping[str, object]) -> Mapping[str, object]:
|
||||
return request
|
||||
|
||||
|
||||
def map_failure(error: Exception, request: Mapping[str, object]) -> Exception:
|
||||
provider: Final = optional_str(request.get("custom_llm_provider")) or "openai"
|
||||
return failures.map_native_failure(
|
||||
error,
|
||||
str(request["model"]),
|
||||
provider,
|
||||
arguments(request),
|
||||
request,
|
||||
optional_str(request.get("api_base")) or optional_str(request.get("base_url")),
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -4,7 +4,7 @@ from typing import Final
|
|||
import pytest
|
||||
|
||||
import litellm
|
||||
from litellm.rust_bridge.chat_completions.route_host import arguments, connection_defaults, response
|
||||
from litellm.rust_bridge.chat_completions.route_host import connection_defaults, response
|
||||
from litellm.types.utils import ModelResponse
|
||||
|
||||
|
||||
|
|
@ -35,12 +35,6 @@ def test_response_builds_the_public_model_response() -> None:
|
|||
assert built.usage.total_tokens == 5
|
||||
|
||||
|
||||
def test_arguments_are_the_public_kwargs_view() -> None:
|
||||
kwargs: Final = MappingProxyType({"metadata": {"user_id": "u"}})
|
||||
|
||||
assert arguments(kwargs) is kwargs
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("provider", "global_key", "provider_key", "expected_key", "expected_base"),
|
||||
(
|
||||
|
|
|
|||
|
|
@ -1,8 +1,7 @@
|
|||
from types import MappingProxyType
|
||||
from typing import Final
|
||||
|
||||
from litellm.rust_bridge.messages.route_host import arguments, response
|
||||
from litellm.rust_bridge.public_call import NativeCall
|
||||
from litellm.rust_bridge.messages.route_host import response
|
||||
import pytest
|
||||
import litellm
|
||||
from litellm.rust_bridge.messages import route_host
|
||||
|
|
@ -29,26 +28,6 @@ def test_response_is_a_detached_public_messages_dict() -> None:
|
|||
assert "_hidden_params" not in native
|
||||
|
||||
|
||||
def test_arguments_preserve_the_bound_view() -> None:
|
||||
kwargs: Final = MappingProxyType({"litellm_metadata": {"user_id": "u"}})
|
||||
request: Final = NativeCall(
|
||||
args=(),
|
||||
kwargs=kwargs,
|
||||
bound={
|
||||
"model": "claude-sonnet-4-5",
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
"max_tokens": 16,
|
||||
"stream": None,
|
||||
"api_key": None,
|
||||
"api_base": None,
|
||||
"custom_llm_provider": "anthropic",
|
||||
**kwargs,
|
||||
},
|
||||
)
|
||||
|
||||
assert arguments(request.bound) is request.bound
|
||||
|
||||
|
||||
def test_settings_project_caller_configuration_without_resolving_a_model(monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
monkeypatch.setattr(litellm, "drop_params", False)
|
||||
monkeypatch.setattr(litellm, "reasoning_auto_summary", True)
|
||||
|
|
|
|||
|
|
@ -20,13 +20,6 @@ from typing import Final
|
|||
REQUEST_STARTED: Final = threading.Event()
|
||||
REQUEST_CANCELLED: Final = threading.Event()
|
||||
|
||||
ANTHROPIC_RESPONSE: Final = (
|
||||
b'{"id":"msg_native","type":"message","role":"assistant",'
|
||||
b'"model":"claude-sonnet-4-5","content":[{"type":"text","text":"native-message"}],'
|
||||
b'"stop_reason":"end_turn","stop_sequence":null,'
|
||||
b'"usage":{"input_tokens":2,"output_tokens":3}}'
|
||||
)
|
||||
|
||||
|
||||
class NativeRouteServer(ThreadingHTTPServer):
|
||||
request_queue_size = 64
|
||||
|
|
@ -49,7 +42,7 @@ class NativeRouteHandler(BaseHTTPRequestHandler):
|
|||
return
|
||||
|
||||
status: Final = 429 if outcome == "429" else 200
|
||||
response_body: Final = native_response(status, route)
|
||||
response_body: Final = native_response(status)
|
||||
|
||||
self.send_response(status)
|
||||
self.send_header("content-type", "application/json")
|
||||
|
|
@ -78,32 +71,23 @@ def assert_native_request(
|
|||
headers: HTTPMessage,
|
||||
body: object,
|
||||
) -> None:
|
||||
if route not in {"transcription", "chat_completions"}:
|
||||
if route != "transcription":
|
||||
raise AssertionError(f"unexpected route marker: {route!r}")
|
||||
if outcome not in {"success", "429", "hang"}:
|
||||
raise AssertionError(f"unexpected outcome marker: {outcome!r}")
|
||||
if not isinstance(body, dict):
|
||||
raise TypeError(f"{route} sent {type(body).__name__}, expected a JSON object")
|
||||
if route == "transcription":
|
||||
assert path == "/model/mistral.voxtral-mini-3b-2507/converse"
|
||||
assert headers.get("authorization", "").startswith("AWS4-HMAC-SHA256 ")
|
||||
assert headers.get("x-amz-date")
|
||||
assert body["messages"][0]["content"][0]["audio"]["source"]["bytes"] == "AQI="
|
||||
assert "The audio language is en" in body["messages"][0]["content"][1]["text"]
|
||||
return
|
||||
assert path == "/v1/messages"
|
||||
assert headers.get("x-api-key") == "sk-native"
|
||||
assert body["model"] == "claude-sonnet-4-5"
|
||||
assert body["max_tokens"] == 17
|
||||
assert body["messages"][0]["content"] == [{"type": "text", "text": "hello-from-chat"}]
|
||||
assert path == "/model/mistral.voxtral-mini-3b-2507/converse"
|
||||
assert headers.get("authorization", "").startswith("AWS4-HMAC-SHA256 ")
|
||||
assert headers.get("x-amz-date")
|
||||
assert body["messages"][0]["content"][0]["audio"]["source"]["bytes"] == "AQI="
|
||||
assert "The audio language is en" in body["messages"][0]["content"][1]["text"]
|
||||
|
||||
|
||||
def native_response(status: int, route: str | None) -> bytes:
|
||||
def native_response(status: int) -> bytes:
|
||||
if status == 429:
|
||||
return b'{"error":"native-rate-limit"}'
|
||||
if route == "transcription":
|
||||
return b'{"output":{"message":{"content":[{"text":"native-transcription"}]}}}'
|
||||
return ANTHROPIC_RESPONSE
|
||||
return b'{"output":{"message":{"content":[{"text":"native-transcription"}]}}}'
|
||||
|
||||
|
||||
def load_native(native_path: Path) -> object:
|
||||
|
|
@ -138,38 +122,23 @@ def route_kwargs(route: str, api_base: str, outcome: str) -> dict[str, object]:
|
|||
"language": "en",
|
||||
},
|
||||
}
|
||||
if route == "chat_completions":
|
||||
return common | {
|
||||
"model": "anthropic/claude-sonnet-4-5",
|
||||
"messages": [{"role": "user", "content": "hello-from-chat"}],
|
||||
"optional_params": {"max_tokens": 17},
|
||||
"api_key": "sk-native",
|
||||
}
|
||||
raise AssertionError(f"unknown route: {route}")
|
||||
|
||||
|
||||
def assert_success(route: str, response: object) -> None:
|
||||
if not isinstance(response, dict):
|
||||
raise TypeError(f"{route} returned {type(response).__name__}, expected dict")
|
||||
actual: Final = success_value(route, response)
|
||||
expected: Final = "native-transcription" if route == "transcription" else "native-message"
|
||||
if actual != expected:
|
||||
raise AssertionError(f"{route} returned {actual!r}, expected {expected!r}")
|
||||
|
||||
|
||||
def success_value(route: str, response: dict[object, object]) -> object:
|
||||
if route == "transcription":
|
||||
return response["text"]
|
||||
return response["choices"][0]["message"]["content"]
|
||||
if response["text"] != "native-transcription":
|
||||
raise AssertionError(f"{route} returned {response['text']!r}, expected 'native-transcription'")
|
||||
|
||||
|
||||
def assert_rate_limit(route: str, error: BaseException) -> None:
|
||||
if error.args != (429, native_response(429, route).decode()):
|
||||
if error.args != (429, native_response(429).decode()):
|
||||
raise AssertionError(f"{route} returned the wrong 429 error: {error!r}")
|
||||
|
||||
|
||||
def exercise_sync(native: object, api_base: str) -> None:
|
||||
for route in ("transcription", "chat_completions"):
|
||||
for route in ("transcription",):
|
||||
function: Final = getattr(native, route)
|
||||
assert_success(route, function(route_call(route, api_base, "success")))
|
||||
try:
|
||||
|
|
@ -181,7 +150,7 @@ def exercise_sync(native: object, api_base: str) -> None:
|
|||
|
||||
|
||||
async def exercise_async(native: object, api_base: str) -> None:
|
||||
for route in ("transcription", "chat_completions"):
|
||||
for route in ("transcription",):
|
||||
function: Final = getattr(native, f"a{route}")
|
||||
assert_success(route, await function(route_call(route, api_base, "success")))
|
||||
try:
|
||||
|
|
@ -192,13 +161,11 @@ async def exercise_async(native: object, api_base: str) -> None:
|
|||
raise AssertionError(f"a{route} accepted a 429 response")
|
||||
|
||||
responses: Final = await asyncio.wait_for(
|
||||
asyncio.gather(
|
||||
*(native.achat_completions(route_call("chat_completions", api_base, "success")) for _ in range(32))
|
||||
),
|
||||
asyncio.gather(*(native.atranscription(route_call("transcription", api_base, "success")) for _ in range(32))),
|
||||
timeout=15,
|
||||
)
|
||||
for response in responses:
|
||||
assert_success("chat_completions", response)
|
||||
assert_success("transcription", response)
|
||||
|
||||
|
||||
def exercise_routes(native_path: Path, api_base: str) -> object:
|
||||
|
|
@ -210,8 +177,8 @@ def exercise_routes(native_path: Path, api_base: str) -> object:
|
|||
|
||||
def exercise_signal(native: object, api_base: str) -> int:
|
||||
try:
|
||||
native.chat_completions(
|
||||
route_call("chat_completions", api_base, "hang"),
|
||||
native.transcription(
|
||||
route_call("transcription", api_base, "hang"),
|
||||
)
|
||||
except KeyboardInterrupt:
|
||||
sys.stdout.write("KeyboardInterrupt\n")
|
||||
|
|
|
|||
|
|
@ -5,8 +5,7 @@ import pytest
|
|||
from pydantic import ValidationError
|
||||
|
||||
import litellm
|
||||
from litellm.rust_bridge.responses.route_host import arguments, connection_defaults, map_failure, response
|
||||
from litellm.rust_bridge.public_call import NativeCall
|
||||
from litellm.rust_bridge.responses.route_host import connection_defaults, map_failure, response
|
||||
from litellm.types.llms.openai import ResponsesAPIResponse
|
||||
|
||||
|
||||
|
|
@ -42,26 +41,6 @@ def test_response_rejects_a_payload_missing_required_fields() -> None:
|
|||
response(MappingProxyType({"object": "response"}))
|
||||
|
||||
|
||||
def test_arguments_preserve_the_bound_view() -> None:
|
||||
kwargs: Final = MappingProxyType({"litellm_metadata": {"user_id": "u"}})
|
||||
request: Final = NativeCall(
|
||||
args=(),
|
||||
kwargs=kwargs,
|
||||
bound={
|
||||
"model": "gpt-4o",
|
||||
"input": "hi",
|
||||
"stream": None,
|
||||
"api_key": None,
|
||||
"api_base": None,
|
||||
"custom_llm_provider": "openai",
|
||||
"extra_headers": None,
|
||||
**kwargs,
|
||||
},
|
||||
)
|
||||
|
||||
assert arguments(request.bound) is request.bound
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("global_key", "provider_key", "expected"),
|
||||
(
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue