diff --git a/litellm-rust/crates/core/src/llms/azure_ai/ocr/cohere_parse_transformation.rs b/litellm-rust/crates/core/src/llms/azure_ai/ocr/cohere_parse_transformation.rs index 4af36cd9be9..0dacb02e383 100644 --- a/litellm-rust/crates/core/src/llms/azure_ai/ocr/cohere_parse_transformation.rs +++ b/litellm-rust/crates/core/src/llms/azure_ai/ocr/cohere_parse_transformation.rs @@ -99,9 +99,6 @@ impl BaseOcrConfig for AzureAICohereParseConfig { CohereParseConfig.transform_ocr_response(model, raw_response, request_format) } - fn retains_document(&self, document: &OcrDocument) -> bool { - !document.is_remote() - } fn validate_request_body(&self, body: &Value) -> Result<(), crate::ocr::Error> { let document = crate::ocr::prepare::body_document(body)?; diff --git a/litellm-rust/crates/core/src/llms/azure_ai/ocr/document_intelligence/transformation.rs b/litellm-rust/crates/core/src/llms/azure_ai/ocr/document_intelligence/transformation.rs index 23e0349496f..f53a4d43680 100644 --- a/litellm-rust/crates/core/src/llms/azure_ai/ocr/document_intelligence/transformation.rs +++ b/litellm-rust/crates/core/src/llms/azure_ai/ocr/document_intelligence/transformation.rs @@ -542,10 +542,6 @@ impl BaseOcrConfig for AzureDocumentIntelligenceOCRConfig { build_request(document) } - /// The body is `urlSource`/`base64Source`, not a `document` field. - fn retains_document(&self, _document: &OcrDocument) -> bool { - false - } } impl AzureDocumentIntelligenceOCRConfig { diff --git a/litellm-rust/crates/core/src/llms/azure_ai/ocr/transformation.rs b/litellm-rust/crates/core/src/llms/azure_ai/ocr/transformation.rs index a99a549cf84..c7e5c13f930 100644 --- a/litellm-rust/crates/core/src/llms/azure_ai/ocr/transformation.rs +++ b/litellm-rust/crates/core/src/llms/azure_ai/ocr/transformation.rs @@ -100,9 +100,6 @@ impl BaseOcrConfig for AzureAIOCRConfig { MistralOCRConfig.transform_ocr_response(model, raw_response, request_format) } - fn retains_document(&self, document: &OcrDocument) -> bool { - !document.is_remote() - } fn validate_request_body(&self, body: &Value) -> Result<(), crate::ocr::Error> { validate_inline_document(&crate::ocr::prepare::body_document(body)?) diff --git a/litellm-rust/crates/core/src/llms/base_llm/ocr/transformation.rs b/litellm-rust/crates/core/src/llms/base_llm/ocr/transformation.rs index 33f3523f220..8af304b7d8d 100644 --- a/litellm-rust/crates/core/src/llms/base_llm/ocr/transformation.rs +++ b/litellm-rust/crates/core/src/llms/base_llm/ocr/transformation.rs @@ -135,14 +135,6 @@ pub(crate) trait BaseOcrConfig: Send + Sync + Sized + 'static { } } - /// Whether the `document` field in the outgoing body is owned by the - /// provider transform and must survive guardrail body rewrites. - /// Providers that inline remote URLs return `false` for remote documents - /// so a hook may still replace the fetched payload. - fn retains_document(&self, _document: &OcrDocument) -> bool { - true - } - /// Provider-specific check applied to the composed body, both before and /// after guardrail hooks. Defaults to accepting any body. fn validate_request_body(&self, _body: &Value) -> Result<(), crate::ocr::Error> { @@ -178,7 +170,6 @@ pub(crate) trait BaseOcrConfig: Send + Sync + Sized + 'static { request, &url, headers, - self.retains_document(&request.document), body, |body| self.validate_request_body(body), ) diff --git a/litellm-rust/crates/core/src/llms/vertex_ai/ocr/deepseek_transformation.rs b/litellm-rust/crates/core/src/llms/vertex_ai/ocr/deepseek_transformation.rs index 80f4c0cd700..87af0125710 100644 --- a/litellm-rust/crates/core/src/llms/vertex_ai/ocr/deepseek_transformation.rs +++ b/litellm-rust/crates/core/src/llms/vertex_ai/ocr/deepseek_transformation.rs @@ -193,11 +193,6 @@ impl BaseOcrConfig for VertexAIDeepSeekOCRConfig { }) } - /// The body carries the document inside `messages`, not a top-level - /// `document` field, so there is nothing for guardrails to retain. - fn retains_document(&self, _document: &OcrDocument) -> bool { - false - } } pub(crate) fn normalize_response( diff --git a/litellm-rust/crates/core/src/llms/vertex_ai/ocr/transformation.rs b/litellm-rust/crates/core/src/llms/vertex_ai/ocr/transformation.rs index 9d7774e9b65..4329a4f05f0 100644 --- a/litellm-rust/crates/core/src/llms/vertex_ai/ocr/transformation.rs +++ b/litellm-rust/crates/core/src/llms/vertex_ai/ocr/transformation.rs @@ -109,9 +109,6 @@ impl BaseOcrConfig for VertexAIOCRConfig { MistralOCRConfig.transform_ocr_response(model, raw_response, request_format) } - fn retains_document(&self, document: &OcrDocument) -> bool { - !document.is_remote() - } fn validate_request_body(&self, body: &Value) -> Result<(), crate::ocr::Error> { validate_inline_document(&crate::ocr::prepare::body_document(body)?) diff --git a/litellm-rust/crates/core/src/ocr/hooks.rs b/litellm-rust/crates/core/src/ocr/hooks.rs index 2860470b68b..0f7ef898281 100644 --- a/litellm-rust/crates/core/src/ocr/hooks.rs +++ b/litellm-rust/crates/core/src/ocr/hooks.rs @@ -22,11 +22,10 @@ pub struct OcrPreCallRequest { pub struct OcrDuringCallRequest { pub model: String, pub custom_llm_provider: String, + pub api_key: Option, pub url: String, pub headers: Vec<(String, String)>, pub body: Value, - #[serde(skip)] - pub retained_fields: Vec, } #[derive(Clone, Debug, Serialize)] diff --git a/litellm-rust/crates/core/src/ocr/lifecycle.rs b/litellm-rust/crates/core/src/ocr/lifecycle.rs index 35289eaebde..9c8a058d3d6 100644 --- a/litellm-rust/crates/core/src/ocr/lifecycle.rs +++ b/litellm-rust/crates/core/src/ocr/lifecycle.rs @@ -863,12 +863,6 @@ mod tests { Box::pin(async move { assert_eq!(request.body["pages"], json!([2])); assert_eq!(request.body.get("future"), Some(&Value::Null)); - assert!( - !request - .retained_fields - .iter() - .any(|field| field == "pages" || field == "document") - ); request.body.as_object_mut().unwrap().remove("future"); request.body["hook_option"] = json!({"nested":[null,false,0]}); Ok(request) diff --git a/litellm-rust/crates/core/src/ocr/prepare.rs b/litellm-rust/crates/core/src/ocr/prepare.rs index ba173585cf2..2acbffcb3be 100644 --- a/litellm-rust/crates/core/src/ocr/prepare.rs +++ b/litellm-rust/crates/core/src/ocr/prepare.rs @@ -10,7 +10,6 @@ pub(crate) async fn transform_request_body( request: &PreparedOcrRequest, url: &str, headers: &[(String, String)], - retains_document: bool, body: B, validate: impl Fn(&Value) -> Result<(), super::Error>, ) -> Result @@ -24,30 +23,15 @@ where )?; validate(&composed)?; let (body, headers) = if request.hooks.intercepts_requests() { - let body = composed; - let retained_fields = request - .optional_params - .keys() - .filter(|name| body.get(*name).is_some()) - .cloned() - .chain(retains_document.then(|| "document".to_string())) - .filter(|name| { - request - .optional_params - .get("extra_body") - .and_then(Value::as_object) - .is_none_or(|overrides| !overrides.contains_key(name)) - }) - .collect(); let changed = request .hooks .during_call(OcrDuringCallRequest { model: request.model.clone(), custom_llm_provider: request.provider_name().into(), + api_key: request.connection.api_key.clone(), url: url.into(), headers: headers.to_vec(), - body, - retained_fields, + body: composed, }) .await?; if !changed.body.is_object() { @@ -94,6 +78,7 @@ pub(crate) async fn guardrail_document( .during_call(OcrDuringCallRequest { model: request.model.clone(), custom_llm_provider: request.provider_name().into(), + api_key: request.connection.api_key.clone(), url: url.into(), headers: headers.to_vec(), body: serde_json::to_value(&request.document).map_err(|_| { @@ -101,7 +86,6 @@ pub(crate) async fn guardrail_document( path: "document".into(), } })?, - retained_fields: Vec::new(), }) .await?; let document = super::json::decode_request_value(changed.body, "guardrail.document")?; diff --git a/litellm-rust/crates/python-bridge/src/lifecycle/arguments.rs b/litellm-rust/crates/python-bridge/src/lifecycle/arguments.rs new file mode 100644 index 00000000000..f890e17a7ad --- /dev/null +++ b/litellm-rust/crates/python-bridge/src/lifecycle/arguments.rs @@ -0,0 +1,186 @@ +use pyo3::exceptions::PyTypeError; +use pyo3::prelude::*; +use pyo3::types::{PyDict, PyTuple}; + +pub(crate) struct Signature { + pub name: &'static str, + pub parameters: &'static [&'static str], + pub required: usize, +} + +#[derive(Debug)] +pub(crate) struct BoundArguments<'py> { + kwargs: Bound<'py, PyDict>, +} + +impl Signature { + pub(crate) fn bind<'py>( + &self, + args: &Bound<'py, PyTuple>, + kwargs: &Bound<'py, PyDict>, + ) -> PyResult> { + if args.len() > self.parameters.len() { + return Err(PyTypeError::new_err(format!( + "{}() takes {} positional arguments but {} were given", + self.name, + self.parameters.len(), + args.len() + ))); + } + let bound = kwargs.copy()?; + for (name, value) in self.parameters.iter().zip(args.iter()) { + if kwargs.contains(name)? { + return Err(PyTypeError::new_err(format!( + "{}() got multiple values for argument '{name}'", + self.name + ))); + } + bound.set_item(name, value)?; + } + let missing: Vec<&str> = self.parameters[..self.required] + .iter() + .copied() + .filter(|name| !bound.contains(name).unwrap_or(false)) + .collect(); + match missing.as_slice() { + [] => {} + [name] => { + return Err(PyTypeError::new_err(format!( + "{}() missing 1 required positional argument: '{name}'", + self.name + ))); + } + names => { + let quoted: Vec = names.iter().map(|name| format!("'{name}'")).collect(); + let (last, rest) = quoted.split_last().expect("at least two names"); + return Err(PyTypeError::new_err(format!( + "{}() missing {} required positional arguments: {} and {last}", + self.name, + names.len(), + rest.join(", ") + ))); + } + } + Ok(BoundArguments { kwargs: bound }) + } +} + +impl<'py> BoundArguments<'py> { + pub(crate) fn get(&self, name: &str) -> PyResult>> { + self.kwargs.get_item(name) + } + + pub(crate) fn required(&self, name: &str) -> PyResult> { + self.get(name)? + .ok_or_else(|| PyTypeError::new_err(format!("missing required argument '{name}'"))) + } + + pub(crate) fn extract(&self, name: &str) -> PyResult + where + T: for<'a> FromPyObject<'a, 'py>, + for<'a> >::Error: Into, + { + self.required(name)?.extract().map_err(Into::into) + } + + pub(crate) fn optional(&self, name: &str) -> PyResult> + where + T: for<'a> FromPyObject<'a, 'py>, + for<'a> >::Error: Into, + { + self.get(name)? + .filter(|value| !value.is_none()) + .map(|value| value.extract().map_err(Into::into)) + .transpose() + } + + pub(crate) fn kwargs(&self) -> &Bound<'py, PyDict> { + &self.kwargs + } +} + +#[cfg(test)] +mod tests { + use super::*; + + const OCR: Signature = Signature { + name: "ocr", + parameters: &["model", "document", "api_key"], + required: 2, + }; + + fn call<'py>( + py: Python<'py>, + args: &[&str], + kwargs: &[(&str, &str)], + ) -> PyResult> { + let args = PyTuple::new(py, args).unwrap(); + let dict = PyDict::new(py); + for (name, value) in kwargs { + dict.set_item(name, value).unwrap(); + } + OCR.bind(&args, &dict) + } + + #[test] + fn positional_and_keyword_arguments_bind_like_python() { + Python::initialize(); + Python::attach(|py| { + let bound = call(py, &["m", "d"], &[("api_key", "k"), ("extra", "x")]).unwrap(); + assert_eq!(bound.extract::("model").unwrap(), "m"); + assert_eq!(bound.extract::("document").unwrap(), "d"); + assert_eq!( + bound.optional::("api_key").unwrap().as_deref(), + Some("k") + ); + assert_eq!(bound.kwargs().len(), 4); + + let bound = call(py, &[], &[("document", "d"), ("model", "m")]).unwrap(); + assert_eq!(bound.extract::("model").unwrap(), "m"); + }); + } + + #[test] + fn binding_errors_match_python_messages() { + Python::initialize(); + Python::attach(|py| { + let error = call(py, &["m", "d"], &[("model", "dup")]).unwrap_err(); + assert_eq!( + error.to_string(), + "TypeError: ocr() got multiple values for argument 'model'" + ); + + let error = call(py, &["m"], &[]).unwrap_err(); + assert_eq!( + error.to_string(), + "TypeError: ocr() missing 1 required positional argument: 'document'" + ); + + let error = call(py, &[], &[]).unwrap_err(); + assert_eq!( + error.to_string(), + "TypeError: ocr() missing 2 required positional arguments: 'model' and 'document'" + ); + + let error = call(py, &["m", "d", "k", "extra"], &[]).unwrap_err(); + assert_eq!( + error.to_string(), + "TypeError: ocr() takes 3 positional arguments but 4 were given" + ); + }); + } + + #[test] + fn explicit_none_is_absent_for_optional_and_present_for_required() { + Python::initialize(); + Python::attach(|py| { + let args = PyTuple::new(py, ["m"]).unwrap(); + let dict = PyDict::new(py); + dict.set_item("document", py.None()).unwrap(); + dict.set_item("api_key", py.None()).unwrap(); + let bound = OCR.bind(&args, &dict).unwrap(); + assert!(bound.required("document").unwrap().is_none()); + assert_eq!(bound.optional::("api_key").unwrap(), None); + }); + } +} diff --git a/litellm-rust/crates/python-bridge/src/lifecycle/mod.rs b/litellm-rust/crates/python-bridge/src/lifecycle/mod.rs index 7ae22b2a5be..2256a6888e2 100644 --- a/litellm-rust/crates/python-bridge/src/lifecycle/mod.rs +++ b/litellm-rust/crates/python-bridge/src/lifecycle/mod.rs @@ -15,10 +15,12 @@ use tokio::sync::Mutex; use crate::execution::{poll_async_value, run_async_value, run_sync_value}; +mod arguments; mod bindings; mod handle; mod preparation; +pub(crate) use arguments::{BoundArguments, Signature}; use bindings::DeploymentHooks; pub(crate) use bindings::PythonLogger; use handle::{Execution, ExecutionBody, ExecutionStep}; diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/callbacks.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/callbacks.rs index 2dbb968b764..d4caa008052 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/callbacks.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/callbacks.rs @@ -75,7 +75,7 @@ impl PythonLogger { pub(crate) fn pre_ocr( &self, py: Python<'_>, - api_key: &Option>, + api_key: Option<&str>, body: &Bound<'_, PyDict>, headers: &Bound<'_, PyDict>, url: &str, @@ -152,12 +152,13 @@ pub(super) fn response(py: Python<'_>, response: &LiteLLMOcrResponse) -> PyResul pub(super) fn map_failure( py: Python<'_>, error: &Py, - request: &Bound<'_, PyAny>, + model: &str, provider: &str, + kwargs: &Py, ) -> PyResult> { Ok(py .import("litellm.rust_bridge.ocr")? .getattr("map_failure")? - .call1((error, request, provider))? + .call1((error, model, provider, kwargs))? .extract()?) } diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/lifecycle.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/lifecycle.rs index e74228778eb..6d5f40c5aa9 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/lifecycle.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/lifecycle.rs @@ -2,7 +2,7 @@ use pyo3::prelude::*; use pyo3::types::{PyDict, PyTuple}; use litellm_auth::ResolvedCredential; -use litellm_core::ocr::hooks::{OcrDuringCallRequest, OcrPostCallRequest, OcrPreCallRequest}; +use litellm_core::ocr::hooks::{OcrDuringCallRequest, OcrPostCallRequest}; use litellm_core::ocr::{OcrAdmission, OcrCall, OcrClient, OcrHostOperation, OcrHostResult}; use litellm_python_interop::{ from_py_preserving_errors as from_py, to_py_preserving_errors as to_py, @@ -10,81 +10,100 @@ use litellm_python_interop::{ use super::callbacks; use super::errors::to_pyerr as ocr_error_to_pyerr; -use super::project::{ProjectedOcrCall, ProjectedOcrFields, PythonOcrInput, admitted_call}; +use super::project::{admitted_call, project}; +use crate::auth::PythonTokenProvider; use crate::lifecycle::{ - OperationClass, PythonCallState, PythonRoute, missing_state, now, run_call, + OperationClass, PythonCallState, PythonRoute, Signature, missing_state, now, run_call, +}; + +const SIGNATURE: Signature = Signature { + name: "ocr", + parameters: &[ + "model", + "document", + "api_key", + "api_base", + "timeout", + "custom_llm_provider", + "extra_headers", + ], + required: 2, +}; + +const ASYNC_SIGNATURE: Signature = Signature { + name: "aocr", + ..SIGNATURE }; struct PythonOcrHost { state: PythonCallState, - data: OcrHostData, -} - -enum OcrHostData { - Unprojected { request: Py }, - Projected(Box), - Released, + projected: Option, } struct ProjectedOcrHost { - fields: ProjectedOcrFields, + model: String, + provider: &'static str, + secret_fields: Vec<&'static str>, + azure_ad_token_provider: Option, pre_call: Option, - callback_inputs: Option>, body: Option>, headers: Option>, } impl PythonOcrHost { fn projected(&self) -> PyResult<&ProjectedOcrHost> { - match &self.data { - OcrHostData::Projected(projected) => Ok(projected), - _ => Err(missing_state()), - } + self.projected.as_ref().ok_or_else(missing_state) } fn projected_mut(&mut self) -> PyResult<&mut ProjectedOcrHost> { - match &mut self.data { - OcrHostData::Projected(projected) => Ok(projected), - _ => Err(missing_state()), - } + self.projected.as_mut().ok_or_else(missing_state) } - fn pre_call( - &mut self, - py: Python<'_>, - request: OcrPreCallRequest, - ) -> PyResult { - let kwargs = self.state.kwargs.bind(py); - let callback_inputs = kwargs.copy()?; - callback_inputs.set_item("document", &self.projected()?.fields.document)?; - let projected = self.projected_mut()?; - projected.callback_inputs = Some(callback_inputs.unbind()); - projected.pre_call = Some((&request).into()); - Ok(request) + fn project(&mut self, py: Python<'_>) -> PyResult { + let signature = if self.state.asynchronous { + &ASYNC_SIGNATURE + } else { + &SIGNATURE + }; + let arguments = signature.bind(self.state.args.bind(py), self.state.kwargs.bind(py))?; + let projected = project(py, &arguments)?; + let has_token_provider = projected.azure_ad_token_provider.is_some(); + self.projected = Some(ProjectedOcrHost { + model: projected.request.model.clone(), + provider: projected.request.provider_name(), + secret_fields: projected.secret_fields, + azure_ad_token_provider: projected.azure_ad_token_provider, + pre_call: None, + body: None, + headers: None, + }); + Ok(OcrHostResult::Request(Ok(( + Box::new(projected.request), + has_token_provider, + )))) } fn acquire_azure_ad_token(&self, py: Python<'_>) -> PyResult { - let provider = self - .projected()? - .fields + self.projected()? .azure_ad_token_provider .as_ref() - .ok_or_else(missing_state)?; - provider.acquire(py) + .ok_or_else(missing_state)? + .acquire(py) } - fn python_pre_call( + fn during_call( &mut self, py: Python<'_>, mut request: OcrDuringCallRequest, ) -> PyResult { let projected = self.projected()?; let pre_call = projected.pre_call.as_ref().ok_or_else(missing_state)?; - self.state.logger()?.update_ocr( + let logger = self.state.logger()?; + logger.update_ocr( py, &self.state.kwargs, pre_call, - &projected.fields.secret_fields, + &projected.secret_fields, &request.url, )?; let body = to_py(py, &request.body)? @@ -94,23 +113,19 @@ impl PythonOcrHost { for (name, value) in &request.headers { headers.set_item(name, value)?; } - let api_key = self.projected()?.fields.api_key.clone_ref(py); - let projected = self.projected_mut()?; - projected.body = Some(body.clone().unbind()); - projected.headers = Some(headers.clone().unbind()); - self.state - .logger()? - .pre_ocr(py, &Some(api_key), &body, &headers, &request.url)?; - let headers = headers + logger.pre_ocr(py, request.api_key.as_deref(), &body, &headers, &request.url)?; + request.body = from_py(&body)?; + request.headers = headers .iter() .map(|(name, value)| Ok((name.extract::()?, value.extract::()?))) .collect::>>()?; - request.body = from_py(&body)?; - request.headers = headers; + let projected = self.projected_mut()?; + projected.body = Some(body.unbind()); + projected.headers = Some(headers.unbind()); Ok(request) } - fn python_post_call( + fn post_call( &mut self, py: Python<'_>, request: OcrPostCallRequest, @@ -124,6 +139,24 @@ impl PythonOcrHost { )?; Ok(request) } + + fn map_failure(&mut self, py: Python<'_>, error: litellm_core::ocr::Error) -> PyResult<()> { + if self.state.error.is_none() { + self.state.retain_error(py, ocr_error_to_pyerr(error)); + } + if self.state.end.is_none() { + self.state.end = Some(now(py)?); + } + let error = self.state.error.as_ref().ok_or_else(missing_state)?; + let (model, provider) = match &self.projected { + Some(projected) => (projected.model.as_str(), projected.provider), + None => ("", ""), + }; + let mapped = callbacks::map_failure(py, error, model, provider, &self.state.kwargs)?; + self.state + .retain_error(py, PyErr::from_value(mapped.into_bound(py).into_any())); + Ok(()) + } } impl PythonRoute for PythonOcrHost { @@ -157,36 +190,19 @@ impl PythonRoute for PythonOcrHost { fn invoke(&mut self, py: Python<'_>, operation: OcrHostOperation) -> PyResult { Ok(match operation { - OcrHostOperation::ProjectRequest => { - let OcrHostData::Unprojected { request } = &self.data else { - return Err(missing_state()); - }; - let projected = ProjectedOcrCall::try_from(PythonOcrInput { - request: request.bind(py), - kwargs: self.state.kwargs.bind(py), - })?; - let has_token_provider = projected.fields.azure_ad_token_provider.is_some(); - let request = projected.request; - self.data = OcrHostData::Projected(Box::new(ProjectedOcrHost { - fields: projected.fields, - pre_call: None, - callback_inputs: None, - body: None, - headers: None, - })); - OcrHostResult::Request(Ok((Box::new(request), has_token_provider))) - } + OcrHostOperation::ProjectRequest => self.project(py)?, OcrHostOperation::AcquireAzureAdToken => { OcrHostResult::AzureAdToken(Ok(self.acquire_azure_ad_token(py)?)) } OcrHostOperation::PreCall(request) => { - OcrHostResult::PreCall(Ok(self.pre_call(py, request)?)) + self.projected_mut()?.pre_call = Some((&request).into()); + OcrHostResult::PreCall(Ok(request)) } OcrHostOperation::DuringCall(request) => { - OcrHostResult::DuringCall(Ok(self.python_pre_call(py, request)?)) + OcrHostResult::DuringCall(Ok(self.during_call(py, request)?)) } OcrHostOperation::PostCall(request) => { - OcrHostResult::PostCall(Ok(self.python_post_call(py, request)?)) + OcrHostResult::PostCall(Ok(self.post_call(py, request)?)) } OcrHostOperation::ConstructResponse(response) => { self.state.end = Some(now(py)?); @@ -194,24 +210,7 @@ impl PythonRoute for PythonOcrHost { OcrHostResult::Lifecycle(Ok(())) } OcrHostOperation::MapFailure(error) => { - if self.state.error.is_none() { - self.state.retain_error(py, ocr_error_to_pyerr(error)); - } - if self.state.end.is_none() { - self.state.end = Some(now(py)?); - } - let error = self.state.error.as_ref().ok_or_else(missing_state)?; - let (request, provider) = match &self.data { - OcrHostData::Unprojected { request } => (request.bind(py), ""), - OcrHostData::Projected(projected) => ( - projected.fields.boundary_request.bind(py), - projected.fields.provider, - ), - OcrHostData::Released => return Err(missing_state()), - }; - let mapped = callbacks::map_failure(py, error, request, provider)?; - self.state - .retain_error(py, PyErr::from_value(mapped.into_bound(py).into_any())); + self.map_failure(py, error)?; OcrHostResult::Lifecycle(Ok(())) } OcrHostOperation::Lifecycle(_) @@ -221,24 +220,18 @@ impl PythonRoute for PythonOcrHost { } fn cleanup(&mut self) { - self.data = OcrHostData::Released; + self.projected = None; } + fn traverse(&self, visit: &pyo3::gc::PyVisit<'_>) -> Result<(), pyo3::gc::PyTraverseError> { - match &self.data { - OcrHostData::Unprojected { request } => visit.call(request), - OcrHostData::Projected(projected) => { - visit.call(&projected.fields.boundary_request)?; - visit.call(&projected.fields.document)?; - visit.call(&projected.fields.api_key)?; - if let Some(provider) = &projected.fields.azure_ad_token_provider { - provider.traverse(visit)?; - } - visit.call(&projected.callback_inputs)?; - visit.call(&projected.body)?; - visit.call(&projected.headers) - } - OcrHostData::Released => Ok(()), + let Some(projected) = &self.projected else { + return Ok(()); + }; + if let Some(provider) = &projected.azure_ad_token_provider { + provider.traverse(visit)?; } + visit.call(&projected.body)?; + visit.call(&projected.headers) } } @@ -257,11 +250,12 @@ fn call( asynchronous: bool, ) -> PyResult> { let kwargs = kwargs.unwrap_or_else(|| PyDict::new(py)); - let name = if asynchronous { "aocr" } else { "ocr" }; - let request = py - .import("litellm.rust_bridge.ocr")? - .getattr("bind_request")? - .call1((name, &args, &kwargs))?; + let signature = if asynchronous { + &ASYNC_SIGNATURE + } else { + &SIGNATURE + }; + signature.bind(&args, &kwargs)?; let client = OcrClient::shared().map_err(ocr_error_to_pyerr)?; let call = admitted_call(OcrCall::admit( client, @@ -276,11 +270,9 @@ fn call( args.unbind(), kwargs.copy()?.unbind(), asynchronous, - name, + signature.name, )?, - data: OcrHostData::Unprojected { - request: request.unbind(), - }, + projected: None, }; run_call(py, call, host) } diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs index f404c5ed293..b7f9613a5a0 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs @@ -3,7 +3,6 @@ mod document; mod errors; mod lifecycle; mod project; -mod request; use pyo3::prelude::*; diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/project.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/project.rs index 36450b1360b..accc8000cdb 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/project.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/project.rs @@ -1,154 +1,137 @@ +use std::collections::BTreeMap; use std::sync::Arc; +use std::time::Duration; -use litellm_core::ocr::{LiteLLMOcrRequest, NativeOutcome, OcrCall, consumed_optional_params}; -use litellm_python_interop::{ - from_py_preserving_errors as from_py, to_py_preserving_errors as to_py, -}; use pyo3::prelude::*; use pyo3::types::PyDict; use serde_json::{Map, Value}; +use litellm_auth::InputSource; +use litellm_core::ocr::{ + LiteLLMOcrRequest, NativeOutcome, OcrCall, OcrCredentialInputs, OcrDocument, + consumed_optional_params, +}; +use litellm_python_interop::from_py_preserving_errors as from_py; + use super::errors::to_pyerr as ocr_error_to_pyerr; use super::lifecycle::BridgeOcrHooks; -use super::request::BridgeOcrRequest; use crate::auth::{AZURE_AD_TOKEN_PROVIDER, PythonTokenProvider}; use crate::errors::RustBridgeDeclined; +use crate::lifecycle::BoundArguments; use crate::marshal::{project_optional_fields, python_timeout_seconds, request_input_sources}; const BOUND_FIELDS: &[&str] = &["model", "document", "timeout", "input_sources"]; -pub(super) struct ProjectedOcrFields { - pub boundary_request: Py, - pub document: Py, - pub api_key: Py, +pub(super) struct ProjectedOcrCall { + pub request: LiteLLMOcrRequest, pub azure_ad_token_provider: Option, - pub provider: &'static str, pub secret_fields: Vec<&'static str>, } -pub(super) struct ProjectedOcrCall { - pub request: LiteLLMOcrRequest, - pub fields: ProjectedOcrFields, -} - -pub(super) struct PythonOcrInput<'a, 'py> { - pub request: &'a Bound<'py, PyAny>, - pub kwargs: &'a Bound<'py, PyDict>, -} - -impl<'py> PythonOcrInput<'_, 'py> { - fn lookup(&self, name: &str) -> PyResult> { - match self.kwargs.get_item(name)? { - Some(value) => Ok(value), - None => self.request.getattr(name), - } - } - - fn model(&self) -> PyResult { - self.lookup("model")?.extract() - } - - fn custom_llm_provider(&self) -> PyResult> { - self.lookup("custom_llm_provider")?.extract() - } - - fn document(&self) -> PyResult> { - self.lookup("document") - } - - fn api_key(&self) -> PyResult> { - self.lookup("api_key") - } - - fn api_base(&self) -> PyResult> { - self.lookup("api_base")?.extract() - } - - fn extra_headers(&self) -> PyResult>> { - self.lookup("extra_headers")? - .extract::>>()? - .map(|value| from_py(value.bind(self.request.py()))) - .transpose() - } - - fn timeout_seconds(&self) -> PyResult> { - Ok(self - .lookup("timeout")? - .extract::>>()? - .map(|value| python_timeout_seconds(self.request.py(), value)) - .transpose()? - .flatten()) - } -} - -fn project_document(py: Python<'_>, document: &Bound<'_, PyAny>) -> PyResult<(Value, Py)> { +fn project_document(py: Python<'_>, document: &Bound<'_, PyAny>) -> PyResult { let kind: String = document.get_item("type")?.extract()?; if kind != "file" { - return Ok((from_py(document)?, document.clone().unbind())); + return from_py(document); } let encoded = super::document::file_document(py, document.extract()?)?; - let wire = serde_json::to_value(encoded) - .map_err(|error| pyo3::exceptions::PyValueError::new_err(error.to_string()))?; - let retained = to_py(py, &wire)?; - Ok((wire, retained)) + serde_json::to_value(encoded) + .map_err(|error| pyo3::exceptions::PyValueError::new_err(error.to_string())) } -impl TryFrom> for ProjectedOcrCall { - type Error = PyErr; +fn header_pairs(headers: Option>) -> Result, litellm_core::ocr::Error> { + headers + .unwrap_or_default() + .into_iter() + .map(|(name, value)| { + value + .as_str() + .map(|value| (name.clone(), value.to_string())) + .ok_or_else(|| litellm_core::ocr::Error::RequestField { + path: format!("extra_headers.{name}"), + }) + }) + .collect() +} - fn try_from(input: PythonOcrInput<'_, '_>) -> PyResult { - let request = input.request; - let kwargs = input.kwargs; - let py = request.py(); - let boundary_request = request.clone().unbind(); - let arguments = input; - let model = arguments.model()?; - let custom_llm_provider = arguments.custom_llm_provider()?; - let (wire_document, retained_document) = project_document(py, &arguments.document()?)?; - let api_key = arguments.api_key()?; - let specs = consumed_optional_params(&model, custom_llm_provider.as_deref()) - .map_err(ocr_error_to_pyerr)?; - let optional_params = project_optional_fields(kwargs, &specs, BOUND_FIELDS)?; - let input_sources = request_input_sources( - kwargs, - optional_params.keys().map(String::as_str).chain([ - "api_key", - "api_base", - "extra_headers", - ]), - )?; - let azure_ad_token_provider = kwargs - .get_item("azure_ad_token_provider")? - .and_then(|provider| PythonTokenProvider::select(provider, AZURE_AD_TOKEN_PROVIDER)); - let request = LiteLLMOcrRequest::try_from(BridgeOcrRequest { - model, - document: wire_document, - api_key: api_key.extract()?, - api_base: arguments.api_base()?, - custom_llm_provider, - extra_headers: arguments.extra_headers()?, - optional_params: optional_params.into(), - input_sources, - timeout_seconds: arguments.timeout_seconds()?, - }) +fn source_for(sources: &BTreeMap, name: &str) -> InputSource { + sources.get(name).copied().unwrap_or_default() +} + +pub(super) fn project(py: Python<'_>, arguments: &BoundArguments<'_>) -> PyResult { + let kwargs: &Bound<'_, PyDict> = arguments.kwargs(); + let model: String = arguments.extract("model")?; + let custom_llm_provider: Option = arguments.optional("custom_llm_provider")?; + let document = project_document(py, &arguments.required("document")?)?; + let api_key: Option = arguments.optional("api_key")?; + let specs = consumed_optional_params(&model, custom_llm_provider.as_deref()) .map_err(ocr_error_to_pyerr)?; - let provider = request.provider_name(); - Ok(Self { - request: request.with_host_hooks(Arc::new(BridgeOcrHooks), None), - fields: ProjectedOcrFields { - boundary_request, - document: retained_document, - api_key: api_key.unbind(), - azure_ad_token_provider, - provider, - secret_fields: specs - .into_iter() - .filter(|spec| spec.secret) - .map(|spec| spec.name) - .collect(), - }, + let optional_params = project_optional_fields(kwargs, &specs, BOUND_FIELDS)?; + let input_sources = request_input_sources( + kwargs, + optional_params + .keys() + .map(String::as_str) + .chain(["api_key", "api_base", "extra_headers"]), + )?; + let azure_ad_token_provider = kwargs + .get_item("azure_ad_token_provider")? + .and_then(|provider| PythonTokenProvider::select(provider, AZURE_AD_TOKEN_PROVIDER)); + let api_base: Option = arguments.optional("api_base")?; + let extra_headers: Option> = arguments + .get("extra_headers")? + .filter(|value| !value.is_none()) + .map(|value| from_py(&value)) + .transpose()?; + let timeout = arguments + .get("timeout")? + .filter(|value| !value.is_none()) + .map(|value| python_timeout_seconds(py, value.unbind())) + .transpose()? + .flatten() + .map(|seconds| { + Duration::try_from_secs_f64(seconds).map_err(|_| { + litellm_core::ocr::Error::RequestField { + path: "timeout_seconds".into(), + } + }) }) - } + .transpose() + .map_err(ocr_error_to_pyerr)?; + + let request = (|| { + let core = LiteLLMOcrRequest::new( + model, + OcrDocument::try_from(document)?, + custom_llm_provider.as_deref(), + optional_params.into(), + )?; + let transport = core.transport.clone().with_overrides( + header_pairs(extra_headers)?, + source_for(&input_sources, "extra_headers"), + timeout, + ); + let credentials = OcrCredentialInputs::new( + api_key, + source_for(&input_sources, "api_key"), + api_base, + source_for(&input_sources, "api_base"), + ); + Ok::<_, litellm_core::ocr::Error>( + core.with_connection_inputs(credentials, transport, input_sources) + .with_host_hooks(Arc::new(BridgeOcrHooks), None), + ) + })() + .map_err(ocr_error_to_pyerr)?; + + Ok(ProjectedOcrCall { + request, + azure_ad_token_provider, + secret_fields: specs + .into_iter() + .filter(|spec| spec.secret) + .map(|spec| spec.name) + .collect(), + }) } pub(super) fn admitted_call(outcome: NativeOutcome) -> PyResult { @@ -187,28 +170,6 @@ mod tests { locals } - fn arguments<'a, 'py>( - request: &'a Bound<'py, PyAny>, - kwargs: &'a Bound<'py, PyDict>, - ) -> PythonOcrInput<'a, 'py> { - PythonOcrInput { 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(); @@ -232,210 +193,7 @@ sys.modules['litellm.rust_bridge.timeouts'] = timeouts } #[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::() - .unwrap(); - let arguments = arguments(&request, &kwargs); - assert_eq!(arguments.model().unwrap(), "from-kwargs"); - assert_eq!(arguments.custom_llm_provider().unwrap(), None); - let accesses: Vec = request.getattr("accesses").unwrap().extract().unwrap(); - assert_eq!(accesses, Vec::::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::() - .unwrap(); - assert_eq!( - arguments(&request, &kwargs).model().unwrap(), - "mistral-ocr-latest" - ); - assert_eq!( - request.getattr("reads").unwrap().extract::().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::() - .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::() - .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::() - .unwrap(); - let arguments = arguments(&request, &kwargs); - let document = arguments.document().unwrap(); - project_document(py, &document).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::() - .unwrap(); - let captured = arguments(&request, &kwargs).api_key().unwrap(); - assert!( - captured - .unbind() - .bind(py) - .is(locals.get_item("key").unwrap().unwrap()) - ); - }); - } - - #[test] - fn file_documents_are_encoded_and_other_documents_keep_the_python_object() { + fn file_documents_are_encoded_and_other_documents_pass_through() { Python::initialize(); Python::attach(|py| { let file = py @@ -446,7 +204,7 @@ kwargs = {'api_key': key} ) .unwrap(); assert_eq!( - project_document(py, &file).unwrap().0, + project_document(py, &file).unwrap(), serde_json::json!({ "type": "document_url", "document_url": "data:application/pdf;base64,JVBERi0xLjQ=", @@ -460,15 +218,13 @@ kwargs = {'api_key': key} None, ) .unwrap(); - let (wire, retained) = project_document(py, &original).unwrap(); assert_eq!( - wire, + project_document(py, &original).unwrap(), serde_json::json!({ "type": "document_url", "document_url": "https://example.com/a.pdf", }) ); - assert!(retained.bind(py).is(&original)); }); } @@ -479,25 +235,12 @@ kwargs = {'api_key': key} let document = py .eval(c"{'type': 'mystery', 'mystery': 'x'}", None, None) .unwrap(); - let wire_document = project_document(py, &document).unwrap().0; + let wire_document = project_document(py, &document).unwrap(); assert_eq!( wire_document, serde_json::json!({"type": "mystery", "mystery": "x"}) ); - let error = match LiteLLMOcrRequest::try_from(BridgeOcrRequest { - model: "mistral/mistral-ocr-latest".into(), - document: wire_document, - api_key: None, - api_base: None, - custom_llm_provider: None, - extra_headers: None, - optional_params: Default::default(), - input_sources: Default::default(), - timeout_seconds: None, - }) { - Ok(_) => panic!("unknown discriminators belong to core validation"), - Err(error) => error, - }; + let error = OcrDocument::try_from(wire_document).unwrap_err(); assert!(error.to_string().contains("document")); }); } @@ -560,9 +303,8 @@ document = Document() ", ); let document = locals.get_item("document").unwrap().unwrap(); - let (wire, retained) = project_document(py, &document).unwrap(); + let wire = project_document(py, &document).unwrap(); assert_eq!(wire["type"], "document_url"); - assert!(!retained.bind(py).is(&document)); let reads: Vec = document.getattr("reads").unwrap().extract().unwrap(); assert_eq!(reads, ["type", "mime_type", "file"]); }); diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/request.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/request.rs deleted file mode 100644 index 051540bdee3..00000000000 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/request.rs +++ /dev/null @@ -1,75 +0,0 @@ -use std::collections::BTreeMap; -use std::time::Duration; - -use litellm_auth::InputSource; -use litellm_core::call_arguments::CallArguments; -use litellm_core::ocr::{LiteLLMOcrRequest, OcrCredentialInputs, OcrDocument}; -use serde_json::{Map, Value}; - -pub(super) struct BridgeOcrRequest { - pub model: String, - pub document: Value, - pub api_key: Option, - pub api_base: Option, - pub custom_llm_provider: Option, - pub extra_headers: Option>, - pub optional_params: CallArguments, - pub input_sources: BTreeMap, - pub timeout_seconds: Option, -} - -impl TryFrom for LiteLLMOcrRequest { - type Error = litellm_core::ocr::Error; - - fn try_from(request: BridgeOcrRequest) -> Result { - let api_key_source = source_for(&request.input_sources, "api_key"); - let api_base_source = source_for(&request.input_sources, "api_base"); - let extra_headers_source = source_for(&request.input_sources, "extra_headers"); - let timeout = request - .timeout_seconds - .map(|seconds| { - Duration::try_from_secs_f64(seconds).map_err(|_| Self::Error::RequestField { - path: "timeout_seconds".into(), - }) - }) - .transpose()?; - let headers = request - .extra_headers - .unwrap_or_default() - .into_iter() - .map(|(name, value)| { - value - .as_str() - .map(|value| (name.clone(), value.to_string())) - .ok_or_else(|| Self::Error::RequestField { - path: format!("extra_headers.{name}"), - }) - }) - .collect::, _>>()?; - let core_request = LiteLLMOcrRequest::new( - request.model, - OcrDocument::try_from(request.document)?, - request.custom_llm_provider.as_deref(), - request.optional_params, - )?; - let transport = - core_request - .transport - .clone() - .with_overrides(headers, extra_headers_source, timeout); - Ok(core_request.with_connection_inputs( - OcrCredentialInputs::new( - request.api_key, - api_key_source, - request.api_base, - api_base_source, - ), - transport, - request.input_sources, - )) - } -} - -fn source_for(sources: &BTreeMap, name: &str) -> InputSource { - sources.get(name).copied().unwrap_or_default() -} diff --git a/litellm/ocr/main.py b/litellm/ocr/main.py index 36b7fa2a35c..c6fca6078e5 100644 --- a/litellm/ocr/main.py +++ b/litellm/ocr/main.py @@ -6,7 +6,7 @@ from litellm.ocr import legacy from litellm.ocr.legacy import convert_file_document_to_url_document, get_mime_type from litellm.rust_bridge.bindings import native_exception_types from litellm.rust_bridge.configuration import rust_ocr_enabled -from litellm.rust_bridge.ocr import NATIVE_AOCR, NATIVE_OCR, bind_request +from litellm.rust_bridge.ocr import NATIVE_AOCR, NATIVE_OCR __all__ = ("aocr", "convert_file_document_to_url_document", "get_mime_type", "ocr") @@ -15,7 +15,6 @@ def ocr( *args: object, **kwargs: object, # kwargs-ok: preserve the public OCR call shape ) -> OCRResponse | Coroutine[object, object, OCRResponse]: - bind_request("ocr", args, kwargs) native: Final = NATIVE_OCR.load() if rust_ocr_enabled() and not kwargs.get("aocr") else None if native is not None: try: @@ -29,7 +28,6 @@ def ocr( async def aocr(*args: object, **kwargs: object) -> OCRResponse: # kwargs-ok: preserve the public OCR call shape - bind_request("aocr", args, kwargs) native: Final = NATIVE_AOCR.load() if rust_ocr_enabled() and not kwargs.get("aocr") else None if native is not None: try: diff --git a/litellm/rust_bridge/ocr.py b/litellm/rust_bridge/ocr.py index 5ed9f3e2d59..66ec39b5280 100644 --- a/litellm/rust_bridge/ocr.py +++ b/litellm/rust_bridge/ocr.py @@ -1,11 +1,9 @@ from __future__ import annotations from collections.abc import Coroutine, Mapping -from dataclasses import dataclass from types import MappingProxyType from typing import Final, Protocol, cast # noqa: TID251 # validates dynamically loaded native callables -import httpx from pydantic import TypeAdapter import litellm @@ -13,38 +11,6 @@ from litellm.llms.base_llm.ocr.transformation import PROVIDER_NATIVE_RESPONSE_KE from litellm.rust_bridge.bindings import NativeBinding -@dataclass(frozen=True, slots=True) -class LiteLLMOcrRequest: - model: str - document: Mapping[str, object] - api_key: str | None - api_base: str | None - timeout: float | httpx.Timeout | None - custom_llm_provider: str | None - extra_headers: dict[str, object] | None - kwargs: Mapping[str, object] - - -def _bind_request( - model: str, - document: Mapping[str, object], - api_key: str | None = None, - api_base: str | None = None, - timeout: float | httpx.Timeout | None = None, - custom_llm_provider: str | None = None, - extra_headers: dict[str, object] | None = None, - **kwargs: object, # kwargs-ok: public OCR accepts provider-specific options -) -> LiteLLMOcrRequest: - return LiteLLMOcrRequest(model, document, api_key, api_base, timeout, custom_llm_provider, extra_headers, kwargs) - - -def bind_request(name: str, args: tuple[object, ...], kwargs: dict[str, object]) -> LiteLLMOcrRequest: - try: - return _bind_request(*args, **kwargs) # pyright: ignore[reportArgumentType] # Python binds arguments before native validation - except TypeError as error: - raise TypeError(str(error).replace("_bind_request()", f"{name}()")) from None - - class RustOcr(Protocol): def __call__(self, *args: object, **kwargs: object) -> OCRResponse: ... @@ -88,17 +54,17 @@ class ExceptionMapper(Protocol): ) -> Exception: ... -def map_failure(error: Exception, request: LiteLLMOcrRequest, request_provider: str) -> Exception: +def map_failure(error: Exception, model: str, provider: str, kwargs: Mapping[str, object]) -> Exception: mapper: Final = cast( ExceptionMapper, litellm.exception_type ) # cast-ok: bounded adapter for the public exception mapper try: return mapper( - model=request.model.removeprefix(f"{request_provider}/"), - custom_llm_provider=request_provider, + model=model, + custom_llm_provider=provider, original_exception=error, - completion_kwargs=dict(request.kwargs), # mutable-ok: exception mapper requires owned kwargs - extra_kwargs=dict(request.kwargs), # mutable-ok: exception mapper requires owned kwargs + completion_kwargs=dict(kwargs), # mutable-ok: exception mapper requires owned kwargs + extra_kwargs=dict(kwargs), # mutable-ok: exception mapper requires owned kwargs ) except Exception as public_error: public_error.__context__ = error diff --git a/tests/test_litellm/rust_bridge/test_ocr_lifecycle.py b/tests/test_litellm/rust_bridge/test_ocr_lifecycle.py index eb76447bffa..baf5295d571 100644 --- a/tests/test_litellm/rust_bridge/test_ocr_lifecycle.py +++ b/tests/test_litellm/rust_bridge/test_ocr_lifecycle.py @@ -101,32 +101,33 @@ def test_public_binding_keeps_keyword_model_and_document_in_native_hook_kwargs() @pytest.mark.parametrize("enabled", [False, True], ids=["flag-disabled", "flag-enabled"]) -def test_public_duplicate_argument_error_does_not_depend_on_native_selection(enabled: bool) -> None: - native: Final = Mock(side_effect=AssertionError("binding errors precede admission")) - document: Final = {"type": "document_url", "document_url": "https://example.com"} +@pytest.mark.parametrize( + "call, message", + [ + ( + lambda: litellm.ocr("mistral/mistral-ocr-latest", {"type": "document_url"}, model="duplicate"), + r"ocr\(\) got multiple values for argument 'model'", + ), + ( + lambda: litellm.ocr("mistral/mistral-ocr-latest"), + r"ocr\(\) missing 1 required positional argument: 'document'", + ), + ], + ids=["duplicate", "missing"], +) +def test_public_binding_errors_do_not_depend_on_native_selection( + monkeypatch: pytest.MonkeyPatch, enabled: bool, call, message: str +) -> None: + fallback: Final = Mock(side_effect=AssertionError("legacy must not run after a native binding error")) + if enabled and NATIVE_OCR.load() is not None: + monkeypatch.setattr(legacy, "ocr", fallback) litellm.rust(enabled) - NATIVE_OCR.override(native) try: - with pytest.raises(TypeError, match=r"ocr\(\) got multiple values for argument 'model'"): - litellm.ocr("mistral/mistral-ocr-latest", document, model="duplicate") + with pytest.raises(TypeError, match=message): + call() finally: - NATIVE_OCR.reset() litellm.rust(None) - assert native.call_count == 0 - - -@pytest.mark.parametrize("enabled", [False, True], ids=["flag-disabled", "flag-enabled"]) -def test_public_missing_required_argument_error_does_not_depend_on_native_selection(enabled: bool) -> None: - native: Final = Mock(side_effect=AssertionError("binding errors precede admission")) - litellm.rust(enabled) - NATIVE_OCR.override(native) - try: - with pytest.raises(TypeError, match=r"ocr\(\) missing 1 required positional argument: 'document'"): - litellm.ocr("mistral/mistral-ocr-latest") - finally: - NATIVE_OCR.reset() - litellm.rust(None) - assert native.call_count == 0 + assert fallback.call_count == 0 @pytest.mark.asyncio