mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
wip
This commit is contained in:
parent
fd48f9da75
commit
ed140cd861
8 changed files with 1106 additions and 109 deletions
|
|
@ -25,31 +25,16 @@ pub(crate) enum NativeCallStep<O> {
|
|||
Complete,
|
||||
}
|
||||
|
||||
type NativeCallFuture<'a, O> =
|
||||
Pin<Box<dyn Future<Output = Result<NativeCallStep<O>, litellm_core::Error>> + Send + 'a>>;
|
||||
|
||||
pub(crate) trait NativeCall: Send + Sync {
|
||||
type Operation: Send + 'static;
|
||||
type Result: Send + 'static;
|
||||
|
||||
fn resume(
|
||||
&mut self,
|
||||
result: Option<Self::Result>,
|
||||
) -> Pin<
|
||||
Box<
|
||||
dyn Future<Output = Result<NativeCallStep<Self::Operation>, litellm_core::Error>>
|
||||
+ Send
|
||||
+ '_,
|
||||
>,
|
||||
>;
|
||||
fn resume(&mut self, result: Option<Self::Result>) -> NativeCallFuture<'_, Self::Operation>;
|
||||
|
||||
fn interrupt(
|
||||
&mut self,
|
||||
failure: HostFailure,
|
||||
) -> Pin<
|
||||
Box<
|
||||
dyn Future<Output = Result<NativeCallStep<Self::Operation>, litellm_core::Error>>
|
||||
+ Send
|
||||
+ '_,
|
||||
>,
|
||||
>;
|
||||
fn interrupt(&mut self, failure: HostFailure) -> NativeCallFuture<'_, Self::Operation>;
|
||||
}
|
||||
|
||||
pub(crate) enum OperationClass {
|
||||
|
|
@ -74,6 +59,9 @@ pub(crate) trait PythonRoute: Send + Sync {
|
|||
fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError>;
|
||||
}
|
||||
|
||||
type HostResumeStep<R> =
|
||||
HostStep<NativeCallStep<<<R as PythonRoute>::Call as NativeCall>::Operation>, Py<PyAny>>;
|
||||
|
||||
struct NativeCallState<C: NativeCall> {
|
||||
call: C,
|
||||
result: Option<Result<NativeCallStep<C::Operation>, litellm_core::Error>>,
|
||||
|
|
@ -128,7 +116,7 @@ impl<R: PythonRoute> PythonLifecycle<R> {
|
|||
&mut self,
|
||||
py: Python<'_>,
|
||||
result: Option<Result<<R::Call as NativeCall>::Result, HostFailure>>,
|
||||
) -> PyResult<HostStep<NativeCallStep<<R::Call as NativeCall>::Operation>, Py<PyAny>>> {
|
||||
) -> PyResult<HostResumeStep<R>> {
|
||||
let call = Arc::clone(self.call.as_ref().ok_or_else(missing_state)?);
|
||||
let future = async move {
|
||||
let mut call = call.lock().await;
|
||||
|
|
|
|||
|
|
@ -2,6 +2,18 @@ use litellm_core::auth::{credential_default_fields, credential_index};
|
|||
use pyo3::prelude::*;
|
||||
use pyo3::types::{PyDict, PyList};
|
||||
|
||||
struct CredentialEntry<'py>(Bound<'py, PyAny>);
|
||||
|
||||
impl<'py> CredentialEntry<'py> {
|
||||
fn name(&self) -> PyResult<String> {
|
||||
self.0.getattr("credential_name")?.extract()
|
||||
}
|
||||
|
||||
fn values(&self) -> PyResult<Bound<'py, PyDict>> {
|
||||
Ok(self.0.getattr("credential_values")?.cast_into::<PyDict>()?)
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn prepare<'py>(
|
||||
py: Python<'py>,
|
||||
kwargs: &Bound<'py, PyDict>,
|
||||
|
|
@ -35,7 +47,7 @@ fn inherit_credentials(
|
|||
let credentials = litellm.getattr("credential_list")?.cast_into::<PyList>()?;
|
||||
let names = credentials
|
||||
.iter()
|
||||
.map(|credential| credential.getattr("credential_name")?.extract::<String>())
|
||||
.map(|credential| CredentialEntry(credential).name())
|
||||
.collect::<PyResult<Vec<_>>>()?;
|
||||
let Some(index) = credential_index(&requested, &names) else {
|
||||
py.import("litellm._logging")?.getattr("verbose_logger")?.call_method1(
|
||||
|
|
@ -44,10 +56,8 @@ fn inherit_credentials(
|
|||
)?;
|
||||
return Ok(());
|
||||
};
|
||||
let values = credentials
|
||||
.get_item(index)?
|
||||
.getattr("credential_values")?
|
||||
.cast_into::<PyDict>()?;
|
||||
let selected = CredentialEntry(credentials.get_item(index)?);
|
||||
let values = selected.values()?;
|
||||
let supplied: Vec<String> = arguments.keys().extract()?;
|
||||
let fields: Vec<String> = values.keys().extract()?;
|
||||
for name in credential_default_fields(&supplied, &fields) {
|
||||
|
|
@ -57,3 +67,248 @@ fn inherit_credentials(
|
|||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
fn eval<'py>(py: Python<'py>, source: &std::ffi::CStr) -> Bound<'py, PyDict> {
|
||||
let locals = PyDict::new(py);
|
||||
py.run(source, Some(&locals), Some(&locals)).unwrap();
|
||||
locals
|
||||
}
|
||||
|
||||
fn inherit(py: Python<'_>, locals: &Bound<'_, PyDict>) -> PyResult<()> {
|
||||
let litellm = PyModule::new(py, "credential_host")?;
|
||||
litellm.setattr(
|
||||
"credential_list",
|
||||
locals.get_item("credentials").unwrap().unwrap(),
|
||||
)?;
|
||||
inherit_credentials(
|
||||
py,
|
||||
&litellm,
|
||||
&locals
|
||||
.get_item("arguments")
|
||||
.unwrap()
|
||||
.unwrap()
|
||||
.cast_into::<PyDict>()?,
|
||||
)
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn duplicate_names_select_the_first_entry_without_reading_other_values() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let locals = eval(
|
||||
py,
|
||||
c"
|
||||
accesses = []
|
||||
class Credential:
|
||||
def __init__(self, name, values):
|
||||
self._name = name
|
||||
self._values = values
|
||||
@property
|
||||
def credential_name(self):
|
||||
accesses.append(('name', self._name))
|
||||
return self._name
|
||||
@property
|
||||
def credential_values(self):
|
||||
accesses.append(('values', self._name))
|
||||
return self._values
|
||||
credentials = [
|
||||
Credential('ocr-test', {'api_key': 'first'}),
|
||||
Credential('other', {'api_key': 'unused'}),
|
||||
Credential('ocr-test', {'api_key': 'later'}),
|
||||
]
|
||||
arguments = {'litellm_credential_name': 'ocr-test'}
|
||||
",
|
||||
);
|
||||
inherit(py, &locals).unwrap();
|
||||
let arguments = locals
|
||||
.get_item("arguments")
|
||||
.unwrap()
|
||||
.unwrap()
|
||||
.cast_into::<PyDict>()
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
arguments
|
||||
.get_item("api_key")
|
||||
.unwrap()
|
||||
.unwrap()
|
||||
.extract::<String>()
|
||||
.unwrap(),
|
||||
"first"
|
||||
);
|
||||
let accesses: Vec<(String, String)> = locals
|
||||
.get_item("accesses")
|
||||
.unwrap()
|
||||
.unwrap()
|
||||
.extract()
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
accesses,
|
||||
[
|
||||
("name".into(), "ocr-test".into()),
|
||||
("name".into(), "other".into()),
|
||||
("name".into(), "ocr-test".into()),
|
||||
("values".into(), "ocr-test".into()),
|
||||
]
|
||||
);
|
||||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn later_invalid_name_still_fails_after_an_earlier_match() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let locals = eval(
|
||||
py,
|
||||
c"
|
||||
failure = LookupError('later name')
|
||||
class Good:
|
||||
credential_name = 'ocr-test'
|
||||
credential_values = {'api_key': 'first'}
|
||||
class Bad:
|
||||
@property
|
||||
def credential_name(self):
|
||||
raise failure
|
||||
credentials = [Good(), Bad()]
|
||||
arguments = {'litellm_credential_name': 'ocr-test'}
|
||||
",
|
||||
);
|
||||
let error = inherit(py, &locals).unwrap_err();
|
||||
assert!(
|
||||
error
|
||||
.value(py)
|
||||
.is(locals.get_item("failure").unwrap().unwrap())
|
||||
);
|
||||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn selected_values_must_be_a_dictionary_and_property_errors_keep_identity() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let locals = eval(
|
||||
py,
|
||||
c"
|
||||
class Listed:
|
||||
credential_name = 'ocr-test'
|
||||
credential_values = ['not-a-dict']
|
||||
credentials = [Listed()]
|
||||
arguments = {'litellm_credential_name': 'ocr-test'}
|
||||
",
|
||||
);
|
||||
assert!(
|
||||
inherit(py, &locals)
|
||||
.unwrap_err()
|
||||
.is_instance_of::<pyo3::exceptions::PyTypeError>(py)
|
||||
);
|
||||
|
||||
let locals = eval(
|
||||
py,
|
||||
c"
|
||||
failure = RuntimeError('values failed')
|
||||
class Broken:
|
||||
credential_name = 'ocr-test'
|
||||
@property
|
||||
def credential_values(self):
|
||||
raise failure
|
||||
credentials = [Broken()]
|
||||
arguments = {'litellm_credential_name': 'ocr-test'}
|
||||
",
|
||||
);
|
||||
let error = inherit(py, &locals).unwrap_err();
|
||||
assert!(
|
||||
error
|
||||
.value(py)
|
||||
.is(locals.get_item("failure").unwrap().unwrap())
|
||||
);
|
||||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn explicit_none_is_not_overwritten_and_inherited_objects_keep_identity() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let locals = eval(
|
||||
py,
|
||||
c"
|
||||
opaque = object()
|
||||
class Credential:
|
||||
credential_name = 'ocr-test'
|
||||
credential_values = {'api_key': 'credential-key', 'opaque': opaque}
|
||||
credentials = [Credential()]
|
||||
arguments = {'litellm_credential_name': 'ocr-test', 'api_key': None}
|
||||
",
|
||||
);
|
||||
inherit(py, &locals).unwrap();
|
||||
let arguments = locals
|
||||
.get_item("arguments")
|
||||
.unwrap()
|
||||
.unwrap()
|
||||
.cast_into::<PyDict>()
|
||||
.unwrap();
|
||||
assert!(arguments.get_item("api_key").unwrap().unwrap().is_none());
|
||||
assert!(
|
||||
arguments
|
||||
.get_item("opaque")
|
||||
.unwrap()
|
||||
.unwrap()
|
||||
.is(locals.get_item("opaque").unwrap().unwrap())
|
||||
);
|
||||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn selection_rereads_the_list_after_name_properties_run() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let locals = eval(
|
||||
py,
|
||||
c"
|
||||
class First:
|
||||
@property
|
||||
def credential_name(self):
|
||||
credentials[0] = Second()
|
||||
return 'ocr-test'
|
||||
credential_values = {'api_key': 'first'}
|
||||
class Second:
|
||||
credential_name = 'ocr-test'
|
||||
credential_values = {'api_key': 'replaced'}
|
||||
credentials = [First()]
|
||||
arguments = {'litellm_credential_name': 'ocr-test'}
|
||||
",
|
||||
);
|
||||
inherit(py, &locals).unwrap();
|
||||
let arguments = locals
|
||||
.get_item("arguments")
|
||||
.unwrap()
|
||||
.unwrap()
|
||||
.cast_into::<PyDict>()
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
arguments
|
||||
.get_item("api_key")
|
||||
.unwrap()
|
||||
.unwrap()
|
||||
.extract::<String>()
|
||||
.unwrap(),
|
||||
"replaced"
|
||||
);
|
||||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn falsy_credential_names_return_before_loading_credentials() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let litellm = PyModule::new(py, "credential_host").unwrap();
|
||||
for name in [py.None(), py.eval(c"''", None, None).unwrap().unbind()] {
|
||||
let arguments = PyDict::new(py);
|
||||
arguments.set_item("litellm_credential_name", name).unwrap();
|
||||
inherit_credentials(py, &litellm, &arguments).unwrap();
|
||||
}
|
||||
});
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -3,6 +3,7 @@ use std::time::Duration;
|
|||
|
||||
use pyo3::exceptions::PyValueError;
|
||||
use pyo3::prelude::*;
|
||||
use pyo3::types::PyDict;
|
||||
use serde_json::{Map, Value};
|
||||
|
||||
use litellm_core::auth::InputSource;
|
||||
|
|
@ -39,18 +40,18 @@ impl RouteOptions {
|
|||
}
|
||||
}
|
||||
|
||||
pub(crate) fn required_value(
|
||||
name: &'static str,
|
||||
value: Value,
|
||||
expected: fn(&Value) -> bool,
|
||||
expected_name: &'static str,
|
||||
) -> PyResult<Value> {
|
||||
if expected(&value) {
|
||||
return Ok(value);
|
||||
pub(crate) fn required_array(name: &'static str, value: Value) -> PyResult<Vec<Value>> {
|
||||
match value {
|
||||
Value::Array(values) => Ok(values),
|
||||
_ => Err(PyValueError::new_err(format!("{name} must be a list"))),
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn required_object(name: &'static str, value: Value) -> PyResult<Map<String, Value>> {
|
||||
match value {
|
||||
Value::Object(values) => Ok(values),
|
||||
_ => Err(PyValueError::new_err(format!("{name} must be a dict"))),
|
||||
}
|
||||
Err(PyValueError::new_err(format!(
|
||||
"{name} must be a {expected_name}"
|
||||
)))
|
||||
}
|
||||
|
||||
pub(crate) fn object_or_empty(
|
||||
|
|
@ -58,7 +59,7 @@ pub(crate) fn object_or_empty(
|
|||
value: Option<Value>,
|
||||
) -> PyResult<Map<String, Value>> {
|
||||
match value {
|
||||
Some(value) => object(name, value),
|
||||
Some(value) => required_object(name, value),
|
||||
None => Ok(Map::new()),
|
||||
}
|
||||
}
|
||||
|
|
@ -67,14 +68,7 @@ fn optional_object(
|
|||
name: &'static str,
|
||||
value: Option<Value>,
|
||||
) -> PyResult<Option<Map<String, Value>>> {
|
||||
value.map(|value| object(name, value)).transpose()
|
||||
}
|
||||
|
||||
fn object(name: &'static str, value: Value) -> PyResult<Map<String, Value>> {
|
||||
match value {
|
||||
Value::Object(map) => Ok(map),
|
||||
_ => Err(PyValueError::new_err(format!("{name} must be a dict"))),
|
||||
}
|
||||
value.map(|value| required_object(name, value)).transpose()
|
||||
}
|
||||
|
||||
pub(crate) fn optional_timeout(timeout_seconds: Option<f64>) -> Option<Duration> {
|
||||
|
|
@ -95,7 +89,7 @@ pub(crate) fn python_timeout_seconds(py: Python<'_>, timeout: Py<PyAny>) -> PyRe
|
|||
}
|
||||
|
||||
pub(crate) fn project_optional_fields(
|
||||
kwargs: &Bound<'_, pyo3::types::PyDict>,
|
||||
kwargs: &Bound<'_, PyDict>,
|
||||
names: &[&str],
|
||||
) -> PyResult<Map<String, Value>> {
|
||||
names
|
||||
|
|
@ -108,28 +102,48 @@ pub(crate) fn project_optional_fields(
|
|||
.collect()
|
||||
}
|
||||
|
||||
struct RequestFieldSources<'py> {
|
||||
body: Option<Bound<'py, PyAny>>,
|
||||
credentials: Option<Bound<'py, PyAny>>,
|
||||
}
|
||||
|
||||
impl<'py> RequestFieldSources<'py> {
|
||||
fn extract(proxy_request: &Bound<'py, PyAny>) -> PyResult<Self> {
|
||||
let proxy_request = proxy_request.cast::<PyDict>()?;
|
||||
|
||||
let body = proxy_request
|
||||
.get_item("body_fields")?
|
||||
.or(proxy_request.get_item("body")?);
|
||||
|
||||
let credentials = proxy_request.get_item("credential_fields")?;
|
||||
|
||||
Ok(Self { body, credentials })
|
||||
}
|
||||
|
||||
fn contains(&self, name: &str) -> bool {
|
||||
self.body
|
||||
.as_ref()
|
||||
.is_some_and(|fields| fields.contains(name).unwrap_or(false))
|
||||
|| self
|
||||
.credentials
|
||||
.as_ref()
|
||||
.is_some_and(|fields| fields.contains(name).unwrap_or(false))
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn request_input_sources<'a>(
|
||||
kwargs: &Bound<'_, pyo3::types::PyDict>,
|
||||
kwargs: &Bound<'_, PyDict>,
|
||||
names: impl Iterator<Item = &'a str>,
|
||||
) -> PyResult<BTreeMap<String, InputSource>> {
|
||||
let Some(proxy_request) = kwargs.get_item("proxy_server_request")? else {
|
||||
return Ok(BTreeMap::new());
|
||||
};
|
||||
let proxy_request = proxy_request.cast_into::<pyo3::types::PyDict>()?;
|
||||
let body_fields = proxy_request
|
||||
.get_item("body_fields")?
|
||||
.or(proxy_request.get_item("body")?);
|
||||
let credential_fields = proxy_request.get_item("credential_fields")?;
|
||||
|
||||
let sources = RequestFieldSources::extract(&proxy_request)?;
|
||||
|
||||
Ok(names
|
||||
.filter_map(|name| {
|
||||
let present = body_fields
|
||||
.as_ref()
|
||||
.is_some_and(|fields| fields.contains(name).unwrap_or(false))
|
||||
|| credential_fields
|
||||
.as_ref()
|
||||
.is_some_and(|fields| fields.contains(name).unwrap_or(false));
|
||||
present.then(|| (name.to_string(), InputSource::Request))
|
||||
})
|
||||
.filter(|name| sources.contains(name))
|
||||
.map(|name| (name.to_string(), InputSource::Request))
|
||||
.collect())
|
||||
}
|
||||
|
||||
|
|
@ -151,3 +165,199 @@ pub(crate) fn marshal_headers(headers: Option<Value>) -> PyResult<HashMap<String
|
|||
})
|
||||
.collect()
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use pyo3::exceptions::PyTypeError;
|
||||
use serde_json::json;
|
||||
|
||||
fn eval<'py>(py: Python<'py>, source: &std::ffi::CStr) -> Bound<'py, PyDict> {
|
||||
let locals = PyDict::new(py);
|
||||
py.run(source, Some(&locals), Some(&locals)).unwrap();
|
||||
locals
|
||||
}
|
||||
|
||||
fn sources(
|
||||
py: Python<'_>,
|
||||
proxy: &Bound<'_, PyAny>,
|
||||
names: &[&str],
|
||||
) -> PyResult<BTreeMap<String, InputSource>> {
|
||||
let kwargs = PyDict::new(py);
|
||||
kwargs.set_item("proxy_server_request", proxy)?;
|
||||
request_input_sources(&kwargs, names.iter().copied())
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn required_shapes_preserve_nested_values_and_existing_errors() {
|
||||
let nested = json!([{"role": "user", "content": [{"type": "text", "text": "hi"}]}]);
|
||||
assert_eq!(
|
||||
Value::Array(required_array("messages", nested.clone()).unwrap()),
|
||||
nested
|
||||
);
|
||||
|
||||
let body = json!({"model": "claude", "metadata": {"user": "1"}});
|
||||
assert_eq!(
|
||||
Value::Object(required_object("body", body.clone()).unwrap()),
|
||||
body
|
||||
);
|
||||
|
||||
assert_eq!(
|
||||
required_array("messages", json!({"role": "user"}))
|
||||
.unwrap_err()
|
||||
.to_string(),
|
||||
"ValueError: messages must be a list"
|
||||
);
|
||||
assert_eq!(
|
||||
required_object("body", json!([])).unwrap_err().to_string(),
|
||||
"ValueError: body must be a dict"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn optional_parameters_treat_missing_as_empty() {
|
||||
assert_eq!(
|
||||
object_or_empty("optional_params", None).unwrap(),
|
||||
Map::new()
|
||||
);
|
||||
assert_eq!(
|
||||
object_or_empty("optional_params", Some(json!({"temperature": 0.2}))).unwrap(),
|
||||
required_object("optional_params", json!({"temperature": 0.2})).unwrap()
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn missing_none_and_empty_proxy_metadata_are_distinct() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let kwargs = PyDict::new(py);
|
||||
assert!(
|
||||
request_input_sources(&kwargs, ["api_key"].into_iter())
|
||||
.unwrap()
|
||||
.is_empty()
|
||||
);
|
||||
|
||||
kwargs.set_item("proxy_server_request", py.None()).unwrap();
|
||||
assert!(
|
||||
request_input_sources(&kwargs, ["api_key"].into_iter())
|
||||
.unwrap_err()
|
||||
.is_instance_of::<PyTypeError>(py)
|
||||
);
|
||||
|
||||
kwargs
|
||||
.set_item("proxy_server_request", PyDict::new(py))
|
||||
.unwrap();
|
||||
assert!(
|
||||
request_input_sources(&kwargs, ["api_key"].into_iter())
|
||||
.unwrap()
|
||||
.is_empty()
|
||||
);
|
||||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn body_fields_win_over_body_and_explicit_none_does_not_fall_back() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let locals = eval(
|
||||
py,
|
||||
c"
|
||||
proxy = {'body_fields': ['api_key'], 'body': ['api_base']}
|
||||
none_fields = {'body_fields': None, 'body': ['api_key']}
|
||||
body_only = {'body': ['api_base']}
|
||||
",
|
||||
);
|
||||
let named = sources(
|
||||
py,
|
||||
&locals.get_item("proxy").unwrap().unwrap(),
|
||||
&["api_key", "api_base"],
|
||||
)
|
||||
.unwrap();
|
||||
assert_eq!(named.get("api_key").copied(), Some(InputSource::Request));
|
||||
assert!(!named.contains_key("api_base"));
|
||||
|
||||
assert!(
|
||||
sources(
|
||||
py,
|
||||
&locals.get_item("none_fields").unwrap().unwrap(),
|
||||
&["api_key"],
|
||||
)
|
||||
.unwrap()
|
||||
.is_empty()
|
||||
);
|
||||
|
||||
let body_only = sources(
|
||||
py,
|
||||
&locals.get_item("body_only").unwrap().unwrap(),
|
||||
&["api_base"],
|
||||
)
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
body_only.get("api_base").copied(),
|
||||
Some(InputSource::Request)
|
||||
);
|
||||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn body_and_credential_membership_can_mark_request_fields() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let locals = eval(
|
||||
py,
|
||||
c"
|
||||
class Raising:
|
||||
def __contains__(self, item):
|
||||
raise RuntimeError('credential membership')
|
||||
proxy = {
|
||||
'body_fields': ['api_key'],
|
||||
'credential_fields': Raising(),
|
||||
}
|
||||
credentials_only = {'credential_fields': ['extra_headers']}
|
||||
erroring = {'body_fields': Raising()}
|
||||
extra = {'body_fields': ['api_key', 'unused']}
|
||||
",
|
||||
);
|
||||
let skipped = sources(
|
||||
py,
|
||||
&locals.get_item("proxy").unwrap().unwrap(),
|
||||
&["api_key"],
|
||||
)
|
||||
.unwrap();
|
||||
assert_eq!(skipped.get("api_key").copied(), Some(InputSource::Request));
|
||||
|
||||
let credentials = sources(
|
||||
py,
|
||||
&locals.get_item("credentials_only").unwrap().unwrap(),
|
||||
&["extra_headers"],
|
||||
)
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
credentials.get("extra_headers").copied(),
|
||||
Some(InputSource::Request)
|
||||
);
|
||||
|
||||
assert!(
|
||||
sources(
|
||||
py,
|
||||
&locals.get_item("erroring").unwrap().unwrap(),
|
||||
&["api_key"],
|
||||
)
|
||||
.unwrap()
|
||||
.is_empty()
|
||||
);
|
||||
|
||||
let requested = sources(
|
||||
py,
|
||||
&locals.get_item("extra").unwrap().unwrap(),
|
||||
&["api_key"],
|
||||
)
|
||||
.unwrap();
|
||||
assert_eq!(requested.len(), 1);
|
||||
assert_eq!(
|
||||
requested.get("api_key").copied(),
|
||||
Some(InputSource::Request)
|
||||
);
|
||||
});
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -9,12 +9,12 @@ use pyo3::prelude::*;
|
|||
use serde_json::Value;
|
||||
|
||||
use crate::errors::chat_completions_error_to_pyerr;
|
||||
use crate::marshal::{RouteOptions, RouteOptionsInputs, object_or_empty, required_value};
|
||||
use crate::marshal::{RouteOptions, RouteOptionsInputs, object_or_empty, required_array};
|
||||
|
||||
fn prepare_chat_completions(
|
||||
inputs: ChatCompletionsInputs,
|
||||
) -> PyResult<impl Future<Output = Result<ChatCompletionsResponse, Error>> + Send + 'static> {
|
||||
let messages = required_value("messages", inputs.messages, Value::is_array, "list")?;
|
||||
let messages = required_array("messages", inputs.messages)?;
|
||||
let optional_params = object_or_empty("optional_params", inputs.optional_params)?;
|
||||
let options = RouteOptions::from_python(RouteOptionsInputs {
|
||||
model: inputs.model,
|
||||
|
|
@ -36,7 +36,7 @@ fn prepare_chat_completions(
|
|||
} = options;
|
||||
run_chat_completions(ChatCompletionsRequest {
|
||||
model: &model,
|
||||
messages,
|
||||
messages: Value::Array(messages),
|
||||
optional_params,
|
||||
api_key: api_key.as_deref(),
|
||||
api_base: api_base.as_deref(),
|
||||
|
|
|
|||
|
|
@ -389,6 +389,82 @@ mod tests {
|
|||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn missing_and_explicit_none_optional_params_share_the_next_error() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let module = PyModule::new(py, "routes").expect("module should be created");
|
||||
crate::routes::register(&module).expect("routes should register");
|
||||
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.call(("model", &messages), Some(&omitted)))
|
||||
.expect_err("omitted optional_params should reach header validation");
|
||||
let explicit_error = module
|
||||
.getattr("chat_completions")
|
||||
.and_then(|function| function.call(("model", &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());
|
||||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn chat_completions_decline_keeps_existing_reasons() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let module = PyModule::new(py, "routes").expect("module should be created");
|
||||
crate::routes::register(&module).expect("routes should register");
|
||||
let decline = module
|
||||
.getattr("chat_completions_decline")
|
||||
.expect("decline helper should be registered");
|
||||
let empty = PyList::empty(py);
|
||||
let unreadable = py
|
||||
.eval(c"'nope'", None, None)
|
||||
.expect("string messages should convert");
|
||||
|
||||
let unknown: Option<String> = decline
|
||||
.call1(("unknown-model", &empty))
|
||||
.and_then(|value| value.extract())
|
||||
.expect("unknown providers should decline");
|
||||
assert_eq!(
|
||||
unknown.as_deref(),
|
||||
Some("provider is not on the rust chat completions path")
|
||||
);
|
||||
|
||||
let empty_reason: Option<String> = decline
|
||||
.call1(("anthropic/claude-sonnet-4-5", &empty))
|
||||
.and_then(|value| value.extract())
|
||||
.expect("empty lists should decline");
|
||||
assert_eq!(empty_reason.as_deref(), Some("empty message list"));
|
||||
|
||||
let unreadable_reason: Option<String> = decline
|
||||
.call1(("anthropic/claude-sonnet-4-5", unreadable))
|
||||
.and_then(|value| value.extract())
|
||||
.expect("non-list messages should decline");
|
||||
assert_eq!(
|
||||
unreadable_reason.as_deref(),
|
||||
Some("unreadable message list")
|
||||
);
|
||||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn generated_routes_execute_sync_and_async_contracts() {
|
||||
Python::initialize();
|
||||
|
|
|
|||
|
|
@ -6,12 +6,12 @@ use serde_json::Value;
|
|||
use std::future::Future;
|
||||
|
||||
use crate::errors::core_error_to_pyerr;
|
||||
use crate::marshal::{RouteOptions, RouteOptionsInputs, required_value};
|
||||
use crate::marshal::{RouteOptions, RouteOptionsInputs, required_object};
|
||||
|
||||
fn prepare_messages(
|
||||
inputs: MessagesInputs,
|
||||
) -> PyResult<impl Future<Output = Result<AnthropicMessagesResponse, Error>> + Send + 'static> {
|
||||
let body = required_value("body", inputs.body, Value::is_object, "dict")?;
|
||||
let body = required_object("body", inputs.body)?;
|
||||
let options = RouteOptions::from_python(RouteOptionsInputs {
|
||||
model: inputs.model,
|
||||
api_key: inputs.api_key,
|
||||
|
|
@ -32,7 +32,7 @@ fn prepare_messages(
|
|||
} = options;
|
||||
run_messages(MessagesRequest {
|
||||
model: &model,
|
||||
body,
|
||||
body: Value::Object(body),
|
||||
api_key: api_key.as_deref(),
|
||||
api_base: api_base.as_deref(),
|
||||
custom_llm_provider: custom_llm_provider.as_deref(),
|
||||
|
|
|
|||
|
|
@ -1,4 +1,4 @@
|
|||
use serde_json::Value;
|
||||
use serde_json::{Map, Value};
|
||||
use std::sync::Arc;
|
||||
|
||||
use pyo3::prelude::*;
|
||||
|
|
@ -287,23 +287,106 @@ impl PythonRoute for PythonOcrHost {
|
|||
}
|
||||
}
|
||||
|
||||
struct OcrArguments<'a, 'py> {
|
||||
request: &'a Bound<'py, PyAny>,
|
||||
kwargs: &'a Bound<'py, PyDict>,
|
||||
}
|
||||
|
||||
impl<'py> OcrArguments<'_, 'py> {
|
||||
fn lookup(&self, name: &str) -> PyResult<Bound<'py, PyAny>> {
|
||||
match self.kwargs.get_item(name)? {
|
||||
Some(value) => Ok(value),
|
||||
None => self.request.getattr(name),
|
||||
}
|
||||
}
|
||||
|
||||
fn model(&self) -> PyResult<String> {
|
||||
self.lookup("model")?.extract()
|
||||
}
|
||||
|
||||
fn custom_llm_provider(&self) -> PyResult<Option<String>> {
|
||||
self.lookup("custom_llm_provider")?.extract()
|
||||
}
|
||||
|
||||
fn document(&self) -> PyResult<CapturedDocument<'py>> {
|
||||
Ok(CapturedDocument(self.lookup("document")?))
|
||||
}
|
||||
|
||||
fn api_key(&self) -> PyResult<CapturedApiKey<'py>> {
|
||||
Ok(CapturedApiKey(self.lookup("api_key")?))
|
||||
}
|
||||
|
||||
fn api_base(&self) -> PyResult<Option<String>> {
|
||||
self.lookup("api_base")?.extract()
|
||||
}
|
||||
|
||||
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())))
|
||||
.transpose()
|
||||
}
|
||||
|
||||
fn timeout_seconds(&self) -> PyResult<Option<f64>> {
|
||||
Ok(self
|
||||
.lookup("timeout")?
|
||||
.extract::<Option<Py<PyAny>>>()?
|
||||
.map(|value| python_timeout_seconds(self.request.py(), value))
|
||||
.transpose()?
|
||||
.flatten())
|
||||
}
|
||||
}
|
||||
|
||||
struct CapturedDocument<'py>(Bound<'py, PyAny>);
|
||||
|
||||
impl<'py> CapturedDocument<'py> {
|
||||
fn as_bound(&self) -> &Bound<'py, PyAny> {
|
||||
&self.0
|
||||
}
|
||||
}
|
||||
|
||||
struct CapturedApiKey<'py>(Bound<'py, PyAny>);
|
||||
|
||||
impl CapturedApiKey<'_> {
|
||||
fn value(&self) -> PyResult<Option<String>> {
|
||||
self.0.extract()
|
||||
}
|
||||
|
||||
fn into_object(self) -> Py<PyAny> {
|
||||
self.0.unbind()
|
||||
}
|
||||
}
|
||||
|
||||
enum DocumentKind {
|
||||
File,
|
||||
Other,
|
||||
}
|
||||
|
||||
impl FromPyObject<'_, '_> for DocumentKind {
|
||||
type Error = PyErr;
|
||||
|
||||
fn extract(document: Borrowed<'_, '_, PyAny>) -> PyResult<Self> {
|
||||
let kind: String = document.get_item("type")?.extract()?;
|
||||
|
||||
Ok(match kind.as_str() {
|
||||
"file" => Self::File,
|
||||
_ => Self::Other,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
fn project_request(
|
||||
py: Python<'_>,
|
||||
request: &Bound<'_, PyAny>,
|
||||
kwargs: &Bound<'_, PyDict>,
|
||||
) -> PyResult<AdmittedOcrCall> {
|
||||
let argument = |name: &str| {
|
||||
kwargs
|
||||
.get_item(name)?
|
||||
.map(Ok)
|
||||
.unwrap_or_else(|| request.getattr(name))
|
||||
};
|
||||
let model: String = argument("model")?.extract()?;
|
||||
let custom_llm_provider: Option<String> = argument("custom_llm_provider")?.extract()?;
|
||||
let document = argument("document")?;
|
||||
let wire_document = extract_document(py, &document)?;
|
||||
let retained_document = retained_document(py, &document, &wire_document)?;
|
||||
let api_key = argument("api_key")?;
|
||||
let arguments = OcrArguments { request, kwargs };
|
||||
let model = arguments.model()?;
|
||||
let custom_llm_provider = arguments.custom_llm_provider()?;
|
||||
let document = arguments.document()?;
|
||||
let wire_document = extract_document(py, document.as_bound())?;
|
||||
let retained_document = retained_document(py, document.as_bound(), &wire_document)?;
|
||||
let api_key = arguments.api_key()?;
|
||||
let request_kwargs = kwargs;
|
||||
let consumed = consumed_optional_param_names(&model, custom_llm_provider.as_deref())
|
||||
.map_err(ocr_error_to_pyerr)?;
|
||||
|
|
@ -321,20 +404,13 @@ fn project_request(
|
|||
let wire = OcrWireRequest {
|
||||
model,
|
||||
document: wire_document,
|
||||
api_key: api_key.extract()?,
|
||||
api_base: argument("api_base")?.extract()?,
|
||||
api_key: api_key.value()?,
|
||||
api_base: arguments.api_base()?,
|
||||
custom_llm_provider,
|
||||
extra_headers: argument("extra_headers")?
|
||||
.extract::<Option<Py<PyAny>>>()?
|
||||
.map(|value| from_py(value.bind(py)))
|
||||
.transpose()?,
|
||||
extra_headers: arguments.extra_headers()?,
|
||||
optional_params,
|
||||
input_sources,
|
||||
timeout_seconds: argument("timeout")?
|
||||
.extract::<Option<Py<PyAny>>>()?
|
||||
.map(|value| python_timeout_seconds(py, value))
|
||||
.transpose()?
|
||||
.flatten(),
|
||||
timeout_seconds: arguments.timeout_seconds()?,
|
||||
};
|
||||
let request = decode_request(wire).map_err(ocr_error_to_pyerr)?;
|
||||
let provider = request.provider_name().to_string();
|
||||
|
|
@ -342,18 +418,22 @@ fn project_request(
|
|||
Ok(AdmittedOcrCall {
|
||||
request,
|
||||
document: retained_document,
|
||||
api_key: api_key.unbind(),
|
||||
api_key: api_key.into_object(),
|
||||
azure_ad_token_provider,
|
||||
provider,
|
||||
})
|
||||
}
|
||||
|
||||
fn extract_document(py: Python<'_>, document: &Bound<'_, PyAny>) -> PyResult<Value> {
|
||||
if document.get_item("type")?.extract::<String>()? != "file" {
|
||||
return from_py(document);
|
||||
match document.extract::<DocumentKind>()? {
|
||||
DocumentKind::Other => from_py(document),
|
||||
DocumentKind::File => {
|
||||
let input = document.extract()?;
|
||||
let encoded = super::document::file_document(py, input)?;
|
||||
serde_json::to_value(encoded)
|
||||
.map_err(|error| pyo3::exceptions::PyValueError::new_err(error.to_string()))
|
||||
}
|
||||
}
|
||||
serde_json::to_value(super::document::file_document(py, document.extract()?)?)
|
||||
.map_err(|error| pyo3::exceptions::PyValueError::new_err(error.to_string()))
|
||||
}
|
||||
|
||||
fn retained_document(
|
||||
|
|
@ -361,10 +441,9 @@ fn retained_document(
|
|||
document: &Bound<'_, PyAny>,
|
||||
wire_document: &Value,
|
||||
) -> PyResult<Py<PyAny>> {
|
||||
if document.get_item("type")?.extract::<String>()? == "file" {
|
||||
to_py(py, wire_document)
|
||||
} else {
|
||||
Ok(document.clone().unbind())
|
||||
match document.extract::<DocumentKind>()? {
|
||||
DocumentKind::File => to_py(py, wire_document),
|
||||
DocumentKind::Other => Ok(document.clone().unbind()),
|
||||
}
|
||||
}
|
||||
|
||||
|
|
@ -437,10 +516,38 @@ pub(super) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> {
|
|||
mod tests {
|
||||
use litellm_core::Error;
|
||||
use litellm_core::ocr::OcrDecline;
|
||||
use pyo3::exceptions::PyValueError;
|
||||
use pyo3::exceptions::{PyKeyError, PyTypeError, PyValueError};
|
||||
|
||||
use super::*;
|
||||
|
||||
fn eval<'py>(py: Python<'py>, source: &std::ffi::CStr) -> Bound<'py, PyDict> {
|
||||
let locals = PyDict::new(py);
|
||||
py.run(source, Some(&locals), Some(&locals)).unwrap();
|
||||
locals
|
||||
}
|
||||
|
||||
fn arguments<'a, 'py>(
|
||||
request: &'a Bound<'py, PyAny>,
|
||||
kwargs: &'a Bound<'py, PyDict>,
|
||||
) -> OcrArguments<'a, 'py> {
|
||||
OcrArguments { request, kwargs }
|
||||
}
|
||||
|
||||
fn stub_timeout_conversion(py: Python<'_>) {
|
||||
eval(
|
||||
py,
|
||||
c"
|
||||
import sys
|
||||
import types
|
||||
timeouts = types.ModuleType('litellm.rust_bridge.timeouts')
|
||||
timeouts.timeout_to_seconds = lambda timeout: None if timeout is None else float(timeout)
|
||||
sys.modules.setdefault('litellm', types.ModuleType('litellm'))
|
||||
sys.modules.setdefault('litellm.rust_bridge', types.ModuleType('litellm.rust_bridge'))
|
||||
sys.modules['litellm.rust_bridge.timeouts'] = timeouts
|
||||
",
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn typed_initial_decline_uses_bridge_decline_contract() {
|
||||
Python::initialize();
|
||||
|
|
@ -462,4 +569,355 @@ mod tests {
|
|||
assert!(!error.is_instance_of::<RustBridgeDeclined>(py));
|
||||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn kwargs_override_request_attributes_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'
|
||||
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);
|
||||
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);
|
||||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn document_reader_mutations_are_visible_to_later_field_reads() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
stub_timeout_conversion(py);
|
||||
let locals = eval(
|
||||
py,
|
||||
c"
|
||||
class Request:
|
||||
api_base = 'original'
|
||||
timeout = 1
|
||||
@property
|
||||
def document(self):
|
||||
return document
|
||||
class Reader:
|
||||
def read(self):
|
||||
Request.api_base = 'mutated'
|
||||
Request.timeout = 9
|
||||
return b'abc'
|
||||
document = {'type': 'file', 'file': Reader()}
|
||||
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);
|
||||
let document = arguments.document().unwrap();
|
||||
extract_document(py, document.as_bound()).unwrap();
|
||||
assert_eq!(arguments.api_base().unwrap().as_deref(), Some("mutated"));
|
||||
assert_eq!(arguments.timeout_seconds().unwrap(), Some(9.0));
|
||||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn captured_api_key_keeps_the_original_python_object() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let locals = eval(
|
||||
py,
|
||||
c"
|
||||
key = object()
|
||||
class Request:
|
||||
api_key = None
|
||||
request = Request()
|
||||
kwargs = {'api_key': key}
|
||||
",
|
||||
);
|
||||
let request = locals.get_item("request").unwrap().unwrap();
|
||||
let kwargs = locals
|
||||
.get_item("kwargs")
|
||||
.unwrap()
|
||||
.unwrap()
|
||||
.cast_into::<PyDict>()
|
||||
.unwrap();
|
||||
let captured = arguments(&request, &kwargs).api_key().unwrap();
|
||||
assert!(
|
||||
captured
|
||||
.into_object()
|
||||
.bind(py)
|
||||
.is(locals.get_item("key").unwrap().unwrap())
|
||||
);
|
||||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn file_documents_are_encoded_and_other_documents_keep_the_python_object() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let file = py
|
||||
.eval(
|
||||
c"{'type': 'file', 'file': b'%PDF-1.4', 'mime_type': 'application/pdf'}",
|
||||
None,
|
||||
None,
|
||||
)
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
extract_document(py, &file).unwrap(),
|
||||
serde_json::json!({
|
||||
"type": "document_url",
|
||||
"document_url": "data:application/pdf;base64,JVBERi0xLjQ=",
|
||||
})
|
||||
);
|
||||
|
||||
let original = py
|
||||
.eval(
|
||||
c"{'type': 'document_url', 'document_url': 'https://example.com/a.pdf'}",
|
||||
None,
|
||||
None,
|
||||
)
|
||||
.unwrap();
|
||||
let wire = extract_document(py, &original).unwrap();
|
||||
assert_eq!(
|
||||
wire,
|
||||
serde_json::json!({
|
||||
"type": "document_url",
|
||||
"document_url": "https://example.com/a.pdf",
|
||||
})
|
||||
);
|
||||
assert!(
|
||||
retained_document(py, &original, &wire)
|
||||
.unwrap()
|
||||
.bind(py)
|
||||
.is(&original)
|
||||
);
|
||||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn unknown_document_types_reach_existing_downstream_validation() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let document = py
|
||||
.eval(c"{'type': 'mystery', 'mystery': 'x'}", None, None)
|
||||
.unwrap();
|
||||
let wire_document = extract_document(py, &document).unwrap();
|
||||
assert_eq!(
|
||||
wire_document,
|
||||
serde_json::json!({"type": "mystery", "mystery": "x"})
|
||||
);
|
||||
let error = match decode_request(OcrWireRequest {
|
||||
model: "mistral/mistral-ocr-latest".into(),
|
||||
document: wire_document,
|
||||
api_key: None,
|
||||
api_base: None,
|
||||
custom_llm_provider: None,
|
||||
extra_headers: None,
|
||||
optional_params: Map::new(),
|
||||
input_sources: Default::default(),
|
||||
timeout_seconds: None,
|
||||
}) {
|
||||
Ok(_) => panic!("unknown discriminators belong to core validation"),
|
||||
Err(error) => error,
|
||||
};
|
||||
assert!(error.to_string().contains("document"));
|
||||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn document_discriminator_errors_keep_their_existing_exceptions() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let missing = py.eval(c"{}", None, None).unwrap();
|
||||
assert!(
|
||||
extract_document(py, &missing)
|
||||
.unwrap_err()
|
||||
.is_instance_of::<PyKeyError>(py)
|
||||
);
|
||||
|
||||
let non_string = py.eval(c"{'type': 1}", None, None).unwrap();
|
||||
assert!(
|
||||
extract_document(py, &non_string)
|
||||
.unwrap_err()
|
||||
.is_instance_of::<PyTypeError>(py)
|
||||
);
|
||||
|
||||
let locals = eval(
|
||||
py,
|
||||
c"
|
||||
failure = RuntimeError('type lookup failed')
|
||||
class Document:
|
||||
def __getitem__(self, key):
|
||||
raise failure
|
||||
document = Document()
|
||||
",
|
||||
);
|
||||
let error =
|
||||
extract_document(py, &locals.get_item("document").unwrap().unwrap()).unwrap_err();
|
||||
assert!(
|
||||
error
|
||||
.value(py)
|
||||
.is(locals.get_item("failure").unwrap().unwrap())
|
||||
);
|
||||
});
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn document_kind_reads_only_type_and_classification_happens_twice() {
|
||||
Python::initialize();
|
||||
Python::attach(|py| {
|
||||
let locals = eval(
|
||||
py,
|
||||
c"
|
||||
class Document(dict):
|
||||
def __init__(self):
|
||||
super().__init__({'file': b'abc'})
|
||||
self.reads = []
|
||||
def __getitem__(self, key):
|
||||
self.reads.append(key)
|
||||
if key == 'type':
|
||||
return 'file' if self.reads.count('type') == 1 else 'document_url'
|
||||
return super().__getitem__(key)
|
||||
document = Document()
|
||||
",
|
||||
);
|
||||
let document = locals.get_item("document").unwrap().unwrap();
|
||||
assert!(matches!(
|
||||
document.extract::<DocumentKind>().unwrap(),
|
||||
DocumentKind::File
|
||||
));
|
||||
let reads: Vec<String> = document.getattr("reads").unwrap().extract().unwrap();
|
||||
assert_eq!(reads, ["type"]);
|
||||
|
||||
py.run(c"document.reads = []", Some(&locals), Some(&locals))
|
||||
.unwrap();
|
||||
let wire = extract_document(py, &document).unwrap();
|
||||
let retained = retained_document(py, &document, &wire).unwrap();
|
||||
assert!(retained.bind(py).is(&document));
|
||||
let reads: Vec<String> = document.getattr("reads").unwrap().extract().unwrap();
|
||||
assert_eq!(reads, ["type", "file", "type"]);
|
||||
});
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -460,14 +460,27 @@ async def test_native_azure_ocr_rejects_coroutine_returned_by_sync_token_provide
|
|||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize("asynchronous", [False, True])
|
||||
@pytest.mark.parametrize("explicit_key", [False, True])
|
||||
@pytest.mark.parametrize(
|
||||
"override, expected_key",
|
||||
[
|
||||
({}, "credential-key"),
|
||||
({"api_key": "explicit-key"}, "explicit-key"),
|
||||
({"api_key": None}, "environment-key"),
|
||||
],
|
||||
ids=["inherit", "explicit", "explicit-none"],
|
||||
)
|
||||
async def test_native_ocr_inherits_named_credentials_without_overwriting_arguments(
|
||||
ocr_server: RecordingServer, monkeypatch: pytest.MonkeyPatch, asynchronous: bool, explicit_key: bool
|
||||
ocr_server: RecordingServer,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
asynchronous: bool,
|
||||
override: dict[str, object],
|
||||
expected_key: str,
|
||||
) -> None:
|
||||
from litellm.models.credentials import CredentialItem
|
||||
|
||||
pages: Final = [0]
|
||||
opaque: Final = object()
|
||||
monkeypatch.setenv("MISTRAL_API_KEY", "environment-key")
|
||||
monkeypatch.setattr(
|
||||
litellm,
|
||||
"credential_list",
|
||||
|
|
@ -497,14 +510,11 @@ async def test_native_ocr_inherits_named_credentials_without_overwriting_argumen
|
|||
"document": OCR_DOCUMENT,
|
||||
"litellm_credential_name": "ocr-test",
|
||||
"callbacks": [Observer()],
|
||||
**({"api_key": "explicit-key"} if explicit_key else {}),
|
||||
**override,
|
||||
}
|
||||
response: Final = await litellm.aocr(**arguments) if asynchronous else litellm.ocr(**arguments)
|
||||
assert response.pages[0].markdown == "native OCR response"
|
||||
assert (
|
||||
ocr_server.requests[0].headers["authorization"]
|
||||
== f"Bearer {'explicit-key' if explicit_key else 'credential-key'}"
|
||||
)
|
||||
assert ocr_server.requests[0].headers["authorization"] == f"Bearer {expected_key}"
|
||||
assert ocr_server.requests[0].body["pages"] == [0, 2]
|
||||
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue