This commit is contained in:
Yujong Lee 2026-09-12 08:48:13 -07:00
parent fd48f9da75
commit ed140cd861
8 changed files with 1106 additions and 109 deletions

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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