From fe8759df7bbd590f315e03bf45698420873aae6c Mon Sep 17 00:00:00 2001 From: Yujong Lee Date: Sat, 12 Sep 2026 08:35:34 -0700 Subject: [PATCH] wip --- .../crates/core/src/auth/credential.rs | 15 + litellm-rust/crates/core/src/auth/mod.rs | 1 + litellm-rust/crates/core/src/ocr/mod.rs | 1 - litellm-rust/crates/core/src/ocr/prepare.rs | 15 - .../src/{routes/ocr_callbacks.rs => auth.rs} | 171 ++------ .../crates/python-bridge/src/errors.rs | 74 ---- litellm-rust/crates/python-bridge/src/lib.rs | 1 + .../python-bridge/src/lifecycle/handle.rs | 139 +++++++ .../src/{lifecycle.rs => lifecycle/mod.rs} | 389 +++++++++--------- .../src/lifecycle/preparation.rs | 2 +- .../crates/python-bridge/src/marshal.rs | 51 ++- .../crates/python-bridge/src/routes/mod.rs | 5 - .../python-bridge/src/routes/ocr/callbacks.rs | 102 +++++ .../python-bridge/src/routes/ocr/document.rs | 259 ++++++++++++ .../python-bridge/src/routes/ocr/errors.rs | 77 ++++ .../{ocr_lifecycle.rs => ocr/lifecycle.rs} | 152 ++++--- .../python-bridge/src/routes/ocr/mod.rs | 18 + .../src/routes/{ocr.rs => ocr/value.rs} | 2 +- .../python-bridge/src/routes/ocr_document.rs | 148 ------- 19 files changed, 1001 insertions(+), 621 deletions(-) rename litellm-rust/crates/python-bridge/src/{routes/ocr_callbacks.rs => auth.rs} (51%) create mode 100644 litellm-rust/crates/python-bridge/src/lifecycle/handle.rs rename litellm-rust/crates/python-bridge/src/{lifecycle.rs => lifecycle/mod.rs} (79%) create mode 100644 litellm-rust/crates/python-bridge/src/routes/ocr/callbacks.rs create mode 100644 litellm-rust/crates/python-bridge/src/routes/ocr/document.rs create mode 100644 litellm-rust/crates/python-bridge/src/routes/ocr/errors.rs rename litellm-rust/crates/python-bridge/src/routes/{ocr_lifecycle.rs => ocr/lifecycle.rs} (78%) create mode 100644 litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs rename litellm-rust/crates/python-bridge/src/routes/{ocr.rs => ocr/value.rs} (98%) delete mode 100644 litellm-rust/crates/python-bridge/src/routes/ocr_document.rs diff --git a/litellm-rust/crates/core/src/auth/credential.rs b/litellm-rust/crates/core/src/auth/credential.rs index b5235b6780c..c64d331b877 100644 --- a/litellm-rust/crates/core/src/auth/credential.rs +++ b/litellm-rust/crates/core/src/auth/credential.rs @@ -9,6 +9,21 @@ use crate::AuthError; use super::{ResolvedCredential, SecretValue, TokenProviderHandle}; +pub fn credential_index(requested: &str, names: &[String]) -> Option { + names.iter().position(|name| name == requested) +} + +pub fn credential_default_fields<'a>( + supplied: &[String], + credential_fields: &'a [String], +) -> Vec<&'a str> { + credential_fields + .iter() + .filter(|name| !supplied.contains(name)) + .map(String::as_str) + .collect() +} + #[derive(Clone, Debug, PartialEq, Eq)] pub enum CredentialFileRef { Path(PathBuf), diff --git a/litellm-rust/crates/core/src/auth/mod.rs b/litellm-rust/crates/core/src/auth/mod.rs index 35d9c676f65..2940a983fb9 100644 --- a/litellm-rust/crates/core/src/auth/mod.rs +++ b/litellm-rust/crates/core/src/auth/mod.rs @@ -49,6 +49,7 @@ impl Sourced { pub use credential::{ CredentialFileRef, CredentialLookup, CredentialLookupFuture, CredentialPlan, CredentialPlanResolution, CredentialRef, CredentialResolver, CredentialResolverHandle, + credential_default_fields, credential_index, }; pub use http::{CredentialPlacement, RequestAuth}; pub use policy::{CredentialPlanKind, CredentialRule, ExistingHeaderBehavior, ProviderAuthPolicy}; diff --git a/litellm-rust/crates/core/src/ocr/mod.rs b/litellm-rust/crates/core/src/ocr/mod.rs index a6e84a7aff3..e29fd6ac572 100644 --- a/litellm-rust/crates/core/src/ocr/mod.rs +++ b/litellm-rust/crates/core/src/ocr/mod.rs @@ -17,7 +17,6 @@ pub use lifecycle::{ NativeOutcome, NativeResult, NoopOcrHost, OcrAdmission, OcrCall, OcrCallStep, OcrDecline, OcrHookHost, OcrHost, OcrHostOperation, OcrHostResult, }; -pub use prepare::{credential_default_fields, credential_index}; pub use types::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrConnection, OcrDocument}; #[cfg(test)] diff --git a/litellm-rust/crates/core/src/ocr/prepare.rs b/litellm-rust/crates/core/src/ocr/prepare.rs index aef6e233930..ee3ae2e78cb 100644 --- a/litellm-rust/crates/core/src/ocr/prepare.rs +++ b/litellm-rust/crates/core/src/ocr/prepare.rs @@ -6,21 +6,6 @@ use super::error::{OcrError, OcrRequestError}; use super::hooks::OcrDuringCallRequest; use super::types::{LiteLLMOcrRequest, OcrDocument}; -pub fn credential_index(requested: &str, names: &[String]) -> Option { - names.iter().position(|name| name == requested) -} - -pub fn credential_default_fields<'a>( - supplied: &[String], - credential_fields: &'a [String], -) -> Vec<&'a str> { - credential_fields - .iter() - .filter(|name| !supplied.contains(name)) - .map(String::as_str) - .collect() -} - #[derive(Debug, Deserialize)] pub(crate) struct ParsedProviderParams { #[serde(flatten)] diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr_callbacks.rs b/litellm-rust/crates/python-bridge/src/auth.rs similarity index 51% rename from litellm-rust/crates/python-bridge/src/routes/ocr_callbacks.rs rename to litellm-rust/crates/python-bridge/src/auth.rs index b29a50c17e6..8dc0b7aabf0 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr_callbacks.rs +++ b/litellm-rust/crates/python-bridge/src/auth.rs @@ -1,35 +1,47 @@ -use pyo3::exceptions::{PyBaseException, PyException, PyRuntimeError, PyTypeError}; +use litellm_core::auth::{ResolvedCredential, SecretValue}; +use pyo3::exceptions::{PyException, PyRuntimeError, PyTypeError}; use pyo3::gc::{PyTraverseError, PyVisit}; use pyo3::prelude::*; -use pyo3::types::{PyDict, PyString}; -use serde_json::Value; +use pyo3::types::PyString; -use litellm_core::auth::{ResolvedCredential, SecretValue}; -use litellm_core::ocr::LiteLLMOcrResponse; -use litellm_core::ocr::hooks::OcrPreCallRequest; -use litellm_python_interop::to_py_preserving_errors as to_py; +#[derive(Clone, Copy)] +pub(crate) struct TokenProviderContract { + callable_error: &'static str, + token_type_error: &'static str, + callback_error: &'static str, +} -use crate::lifecycle::PythonLogger; +pub(crate) const AZURE_AD_TOKEN_PROVIDER: TokenProviderContract = TokenProviderContract { + callable_error: "Azure AD token provider must be callable", + token_type_error: "Azure AD token must be a string, got {}", + callback_error: "Failed to get Azure AD token: {}", +}; -pub(super) struct AzureAdTokenProvider(Py); +pub(crate) struct PythonTokenProvider { + callback: Py, + contract: TokenProviderContract, +} -impl AzureAdTokenProvider { - pub(super) fn select(provider: Bound<'_, PyAny>) -> Option { - (provider.is_callable() && provider.is_truthy().unwrap_or(false)) - .then(|| Self(provider.unbind())) +impl PythonTokenProvider { + pub(crate) fn select( + provider: Bound<'_, PyAny>, + contract: TokenProviderContract, + ) -> Option { + (provider.is_callable() && provider.is_truthy().unwrap_or(false)).then(|| Self { + callback: provider.unbind(), + contract, + }) } - pub(super) fn acquire(&self, py: Python<'_>) -> PyResult { - let provider = self.0.bind(py); + pub(crate) fn acquire(&self, py: Python<'_>) -> PyResult { + let provider = self.callback.bind(py); if !provider.is_callable() { - return Err(PyTypeError::new_err( - "Azure AD token provider must be callable", - )); + return Err(PyTypeError::new_err(self.contract.callable_error)); } let token = (|| { let token = provider.call0()?; if !token.is_instance_of::() { - let message = PyString::new(py, "Azure AD token must be a string, got {}") + let message = PyString::new(py, self.contract.token_type_error) .call_method1("format", (token.get_type(),))?; return Err(PyTypeError::new_err(message.unbind())); } @@ -39,7 +51,7 @@ impl AzureAdTokenProvider { if error.is_instance_of::(py) || !error.is_instance_of::(py) { return error; } - match PyString::new(py, "Failed to get Azure AD token: {}") + match PyString::new(py, self.contract.callback_error) .call_method1("format", (error.value(py),)) { Ok(message) => { @@ -60,112 +72,16 @@ impl AzureAdTokenProvider { }) } - pub(super) fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> { - visit.call(&self.0) + pub(crate) fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> { + visit.call(&self.callback) } } -impl PythonLogger { - pub(crate) fn update_ocr( - &self, - py: Python<'_>, - kwargs: &Py, - pre_call: &OcrPreCallRequest, - url: &str, - ) -> PyResult<()> { - let redact = py - .import("litellm.rust_bridge.ocr")? - .getattr("redact_logging_params")?; - let update = PyDict::new(py); - update.set_item("kwargs", redact.call1((kwargs,))?.cast_into::()?)?; - update.set_item("model", &pre_call.model)?; - update.set_item( - "optional_params", - redact - .call1((to_py(py, &pre_call.optional_params)?,))? - .cast_into::()?, - )?; - let params = PyDict::new(py); - params.set_item( - "litellm_call_id", - kwargs.bind(py).get_item("litellm_call_id")?, - )?; - params.set_item("api_base", url)?; - update.set_item("litellm_params", params)?; - update.set_item("custom_llm_provider", &pre_call.custom_llm_provider)?; - self.object(py) - .call_method("update_from_kwargs", (), Some(&update))?; - Ok(()) - } - - pub(crate) fn pre_ocr( - &self, - py: Python<'_>, - api_key: &Option>, - body: &Bound<'_, PyDict>, - headers: &Bound<'_, PyDict>, - url: &str, - ) -> PyResult<()> { - let additional = PyDict::new(py); - additional.set_item("complete_input_dict", body)?; - additional.set_item("headers", headers)?; - additional.set_item("api_base", url)?; - let kwargs = PyDict::new(py); - kwargs.set_item("input", "OCR document processing")?; - kwargs.set_item("api_key", api_key)?; - kwargs.set_item("additional_args", additional)?; - self.object(py).call_method("pre_call", (), Some(&kwargs))?; - Ok(()) - } - - pub(crate) fn post_ocr( - &self, - py: Python<'_>, - original_response: &Value, - body: &Option>, - headers: &Option>, - ) -> PyResult<()> { - let kwargs = PyDict::new(py); - kwargs.set_item("original_response", to_py(py, original_response)?)?; - let additional = PyDict::new(py); - additional.set_item("complete_input_dict", body)?; - additional.set_item("headers", headers)?; - kwargs.set_item("additional_args", additional)?; - self.object(py) - .call_method("post_call", (), Some(&kwargs))?; - Ok(()) - } -} - -pub(super) fn response(py: Python<'_>, response: &LiteLLMOcrResponse) -> PyResult> { - py.import("litellm.rust_bridge.ocr")? - .getattr("_response")? - .call1((to_py(py, response)?,)) - .map(Bound::unbind) -} - -pub(super) fn map_failure( - py: Python<'_>, - error: &Py, - request: &Bound<'_, PyAny>, - provider: &str, -) -> PyResult> { - Ok(py - .import("litellm.rust_bridge.ocr_lifecycle")? - .getattr("map_failure")? - .call1((error, request, provider))? - .extract()?) -} - -pub(super) fn timeout_seconds(py: Python<'_>, timeout: Py) -> PyResult> { - py.import("litellm.rust_bridge.timeouts")? - .getattr("timeout_to_seconds")? - .call1((timeout,))? - .extract() -} - #[cfg(test)] mod tests { + use pyo3::exceptions::PyRuntimeError; + use pyo3::types::PyDict; + use super::*; #[test] @@ -200,7 +116,8 @@ def provider(error): .unwrap() .call1((&original,)) .unwrap(); - let provider = AzureAdTokenProvider::select(callback).unwrap(); + let provider = + PythonTokenProvider::select(callback, AZURE_AD_TOKEN_PROVIDER).unwrap(); let error = provider.acquire(py).unwrap_err(); if name == "ordinary" { assert!(error.is_instance_of::(py)); @@ -245,9 +162,11 @@ def provider(): Some(&locals), ) .unwrap(); - let provider = - AzureAdTokenProvider::select(locals.get_item("provider").unwrap().unwrap()) - .unwrap(); + let provider = PythonTokenProvider::select( + locals.get_item("provider").unwrap().unwrap(), + AZURE_AD_TOKEN_PROVIDER, + ) + .unwrap(); let error = provider.acquire(py).unwrap_err(); assert!(error.is_instance_of::(py)); assert!( @@ -267,7 +186,7 @@ def provider(): let callback = py .eval(pyo3::ffi::c_str!("lambda: '\\ud800'"), None, None) .unwrap(); - let provider = AzureAdTokenProvider::select(callback).unwrap(); + let provider = PythonTokenProvider::select(callback, AZURE_AD_TOKEN_PROVIDER).unwrap(); let error = provider.acquire(py).unwrap_err(); assert!(error.is_instance_of::(py)); }); diff --git a/litellm-rust/crates/python-bridge/src/errors.rs b/litellm-rust/crates/python-bridge/src/errors.rs index 74006dae117..0ef04bbcb7b 100644 --- a/litellm-rust/crates/python-bridge/src/errors.rs +++ b/litellm-rust/crates/python-bridge/src/errors.rs @@ -63,77 +63,3 @@ pub(crate) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> { module.add("RustBridgeDeclined", py.get_type::())?; module.add("RustUpstreamError", py.get_type::()) } - -pub(crate) fn ocr_error_to_pyerr(err: Error) -> PyErr { - match err { - Error::MissingField("document_url" | "image_url") => { - PyValueError::new_err("Document URL is required") - } - Error::Http { status, body } => ocr_upstream_error(status, body), - Error::Network(message) if message.contains("timed out") => { - ocr_upstream_error(408, message) - } - other => { - let status = other.http_status_code(); - let error = core_error_to_pyerr(other); - if let Some(status) = status { - Python::attach(|py| { - let value = error.value(py); - value.setattr("status_code", status).ok(); - value.setattr("message", value.to_string()).ok(); - }); - } - error - } - } -} - -fn ocr_upstream_error(status: u16, message: String) -> PyErr { - let error = RustUpstreamError::new_err((status, message.clone())); - Python::attach(|py| { - let value = error.value(py); - value.setattr("status_code", status).ok(); - value.setattr("message", message).ok(); - }); - error -} - -#[cfg(test)] -mod ocr_error_tests { - use super::*; - - #[test] - fn ocr_errors_preserve_python_validation_and_provider_details() { - Python::initialize(); - Python::attach(|py| { - for field in ["document_url", "image_url"] { - let mapped = ocr_error_to_pyerr(Error::MissingField(field)); - assert!(mapped.is_instance_of::(py)); - assert_eq!(mapped.value(py).to_string(), "Document URL is required"); - } - let mapped = ocr_error_to_pyerr(Error::Http { - status: 429, - body: r#"{"message":"rate limited"}"#.to_string(), - }); - assert!(mapped.is_instance_of::(py)); - let args: (u16, String) = mapped - .value(py) - .getattr("args") - .and_then(|args| args.extract()) - .expect("OCR failures retain status and unprefixed provider message"); - assert_eq!(args, (429, r#"{"message":"rate limited"}"#.to_string())); - - let mapped = ocr_error_to_pyerr(Error::InvalidRequest("invalid format".into())); - assert!(mapped.is_instance_of::(py)); - assert_eq!( - mapped - .value(py) - .getattr("status_code") - .unwrap() - .extract::() - .unwrap(), - 400 - ); - }); - } -} diff --git a/litellm-rust/crates/python-bridge/src/lib.rs b/litellm-rust/crates/python-bridge/src/lib.rs index 2700ff207d4..5a5aec72362 100644 --- a/litellm-rust/crates/python-bridge/src/lib.rs +++ b/litellm-rust/crates/python-bridge/src/lib.rs @@ -1,3 +1,4 @@ +mod auth; mod constants; mod diagnostics; mod errors; diff --git a/litellm-rust/crates/python-bridge/src/lifecycle/handle.rs b/litellm-rust/crates/python-bridge/src/lifecycle/handle.rs new file mode 100644 index 00000000000..17a480a7225 --- /dev/null +++ b/litellm-rust/crates/python-bridge/src/lifecycle/handle.rs @@ -0,0 +1,139 @@ +use std::panic::{AssertUnwindSafe, catch_unwind}; + +use litellm_python_interop::panic_to_pyerr; +use pyo3::exceptions::{PyBaseException, PyRuntimeError}; +use pyo3::gc::{PyTraverseError, PyVisit}; +use pyo3::prelude::*; + +pub(super) enum ExecutionStep { + Return(Py), + Await(Py), +} + +pub(super) trait ExecutionBody: Send + Sync { + fn resume(&mut self, result: Option>>) -> PyResult; + fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError>; +} + +enum ExecutionState { + Created(Box), + Running, + Suspended(Box), + Closed, +} + +#[pyclass] +pub(super) struct Execution { + state: ExecutionState, +} + +impl Execution { + pub(super) fn new(body: impl ExecutionBody + 'static) -> Self { + Self { + state: ExecutionState::Created(Box::new(body)), + } + } + + fn advance( + slf: &Bound<'_, Self>, + py: Python<'_>, + result: Option>>, + ) -> PyResult> { + let mut body = { + let mut execution = slf.borrow_mut(); + match (&execution.state, result.is_some()) { + (ExecutionState::Created(_), false) | (ExecutionState::Suspended(_), true) => {} + (ExecutionState::Running, _) => { + return Err(PyRuntimeError::new_err("execution is already running")); + } + (ExecutionState::Closed, _) => { + return Err(PyRuntimeError::new_err("execution is closed")); + } + _ => { + return Err(PyRuntimeError::new_err( + "execution requires start before resume and can only start once", + )); + } + } + match std::mem::replace(&mut execution.state, ExecutionState::Running) { + ExecutionState::Created(body) | ExecutionState::Suspended(body) => body, + _ => unreachable!(), + } + }; + let outcome = catch_unwind(AssertUnwindSafe(|| { + let step = body.resume(result)?; + let (tag, value, suspended) = match step { + ExecutionStep::Await(value) => ("Await", value, true), + ExecutionStep::Return(value) => ("Complete", value, false), + }; + let step = py + .import("litellm.rust_bridge.lifecycle")? + .getattr(tag)? + .call1((value,))? + .unbind(); + Ok((step, suspended)) + })) + .map_err(panic_to_pyerr) + .and_then(|result| result); + match outcome { + Ok((step, true)) if matches!(slf.borrow().state, ExecutionState::Running) => { + slf.borrow_mut().state = ExecutionState::Suspended(body); + Ok(step) + } + outcome => { + slf.borrow_mut().state = ExecutionState::Closed; + drop(body); + outcome.and_then(|(step, suspended)| { + if suspended { + Err(PyRuntimeError::new_err( + "execution was closed while running", + )) + } else { + Ok(step) + } + }) + } + } + } +} + +#[pymethods] +impl Execution { + fn start(slf: &Bound<'_, Self>, py: Python<'_>) -> PyResult> { + Self::advance(slf, py, None) + } + + fn resume_value( + slf: &Bound<'_, Self>, + py: Python<'_>, + value: Py, + ) -> PyResult> { + Self::advance(slf, py, Some(Ok(value))) + } + + fn resume_error( + slf: &Bound<'_, Self>, + py: Python<'_>, + error: Bound<'_, PyBaseException>, + ) -> PyResult> { + Self::advance(slf, py, Some(Err(PyErr::from_value(error.into_any())))) + } + + fn close(slf: &Bound<'_, Self>) { + let state = std::mem::replace(&mut slf.borrow_mut().state, ExecutionState::Closed); + drop(state); + } + + fn __traverse__(&self, visit: PyVisit<'_>) -> Result<(), PyTraverseError> { + match &self.state { + ExecutionState::Created(body) | ExecutionState::Suspended(body) => { + body.traverse(&visit) + } + _ => Ok(()), + } + } + + fn __clear__(slf: &Bound<'_, Self>) { + Self::close(slf); + } +} diff --git a/litellm-rust/crates/python-bridge/src/lifecycle.rs b/litellm-rust/crates/python-bridge/src/lifecycle/mod.rs similarity index 79% rename from litellm-rust/crates/python-bridge/src/lifecycle.rs rename to litellm-rust/crates/python-bridge/src/lifecycle/mod.rs index 6dd7b9f32d4..48f6e3faac4 100644 --- a/litellm-rust/crates/python-bridge/src/lifecycle.rs +++ b/litellm-rust/crates/python-bridge/src/lifecycle/mod.rs @@ -1,169 +1,82 @@ -use std::panic::{AssertUnwindSafe, catch_unwind}; +use std::future::Future; +use std::pin::Pin; use std::sync::Arc; use futures_util::future::{AbortHandle, Abortable}; use litellm_core::call_lifecycle::host::{HostFailure, HostPhase, HostStep}; -use litellm_core::ocr::{OcrCall, OcrCallStep, OcrHostOperation, OcrHostResult}; -use litellm_python_interop::panic_to_pyerr; use pyo3::exceptions::{PyBaseException, PyException, PyRuntimeError}; use pyo3::gc::{PyTraverseError, PyVisit}; use pyo3::prelude::*; use pyo3::types::{PyDict, PyTuple}; use tokio::sync::Mutex; -use crate::errors::ocr_error_to_pyerr; use crate::execution::{run_async_value, run_sync_value}; mod bindings; +mod handle; mod preparation; use bindings::DeploymentHooks; pub(crate) use bindings::PythonLogger; +use handle::{Execution, ExecutionBody, ExecutionStep}; + +pub(crate) enum NativeCallStep { + Host(O), + Complete, +} + +pub(crate) trait NativeCall: Send + Sync { + type Operation: Send + 'static; + type Result: Send + 'static; + + fn resume( + &mut self, + result: Option, + ) -> Pin< + Box< + dyn Future, litellm_core::Error>> + + Send + + '_, + >, + >; + + fn interrupt( + &mut self, + failure: HostFailure, + ) -> Pin< + Box< + dyn Future, litellm_core::Error>> + + Send + + '_, + >, + >; +} + +pub(crate) enum OperationClass { + Phase(HostPhase), + Route, +} pub(crate) trait PythonRoute: Send + Sync { + type Call: NativeCall + 'static; + fn state(&self) -> &PythonCallState; fn state_mut(&mut self) -> &mut PythonCallState; - fn invoke(&mut self, py: Python<'_>, operation: OcrHostOperation) -> PyResult; + fn classify(operation: &::Operation) -> OperationClass; + fn lifecycle_result() -> ::Result; + fn map_error(error: litellm_core::Error) -> PyErr; + fn invoke( + &mut self, + py: Python<'_>, + operation: ::Operation, + ) -> PyResult<::Result>; fn cleanup(&mut self); fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError>; } -enum ExecutionStep { - Return(Py), - Await(Py), -} - -trait ExecutionBody: Send + Sync { - fn resume(&mut self, result: Option>>) -> PyResult; - fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError>; -} - -enum ExecutionState { - Created(Box), - Running, - Suspended(Box), - Closed, -} - -#[pyclass] -struct Execution { - state: ExecutionState, -} - -impl Execution { - fn new(body: impl ExecutionBody + 'static) -> Self { - Self { - state: ExecutionState::Created(Box::new(body)), - } - } - - fn advance( - slf: &Bound<'_, Self>, - py: Python<'_>, - result: Option>>, - ) -> PyResult> { - let mut body = { - let mut execution = slf.borrow_mut(); - match (&execution.state, result.is_some()) { - (ExecutionState::Created(_), false) | (ExecutionState::Suspended(_), true) => {} - (ExecutionState::Running, _) => { - return Err(PyRuntimeError::new_err("execution is already running")); - } - (ExecutionState::Closed, _) => { - return Err(PyRuntimeError::new_err("execution is closed")); - } - _ => { - return Err(PyRuntimeError::new_err( - "execution requires start before resume and can only start once", - )); - } - } - match std::mem::replace(&mut execution.state, ExecutionState::Running) { - ExecutionState::Created(body) | ExecutionState::Suspended(body) => body, - _ => unreachable!(), - } - }; - let outcome = catch_unwind(AssertUnwindSafe(|| { - let step = body.resume(result)?; - let (tag, value, suspended) = match step { - ExecutionStep::Await(value) => ("Await", value, true), - ExecutionStep::Return(value) => ("Complete", value, false), - }; - let step = py - .import("litellm.rust_bridge.lifecycle")? - .getattr(tag)? - .call1((value,))? - .unbind(); - Ok((step, suspended)) - })) - .map_err(panic_to_pyerr) - .and_then(|result| result); - match outcome { - Ok((step, true)) if matches!(slf.borrow().state, ExecutionState::Running) => { - slf.borrow_mut().state = ExecutionState::Suspended(body); - Ok(step) - } - outcome => { - slf.borrow_mut().state = ExecutionState::Closed; - drop(body); - outcome.and_then(|(step, suspended)| { - if suspended { - Err(PyRuntimeError::new_err( - "execution was closed while running", - )) - } else { - Ok(step) - } - }) - } - } - } -} - -#[pymethods] -impl Execution { - fn start(slf: &Bound<'_, Self>, py: Python<'_>) -> PyResult> { - Self::advance(slf, py, None) - } - - fn resume_value( - slf: &Bound<'_, Self>, - py: Python<'_>, - value: Py, - ) -> PyResult> { - Self::advance(slf, py, Some(Ok(value))) - } - - fn resume_error( - slf: &Bound<'_, Self>, - py: Python<'_>, - error: Bound<'_, PyBaseException>, - ) -> PyResult> { - Self::advance(slf, py, Some(Err(PyErr::from_value(error.into_any())))) - } - - fn close(slf: &Bound<'_, Self>) { - let state = std::mem::replace(&mut slf.borrow_mut().state, ExecutionState::Closed); - drop(state); - } - - fn __traverse__(&self, visit: PyVisit<'_>) -> Result<(), PyTraverseError> { - match &self.state { - ExecutionState::Created(body) | ExecutionState::Suspended(body) => { - body.traverse(&visit) - } - _ => Ok(()), - } - } - - fn __clear__(slf: &Bound<'_, Self>) { - Self::close(slf); - } -} - -struct NativeCall { - call: OcrCall, - result: Option>, +struct NativeCallState { + call: C, + result: Option, litellm_core::Error>>, } enum PendingOperation { @@ -173,20 +86,20 @@ enum PendingOperation { struct PythonLifecycle { route: R, - call: Option>>, + call: Option>>>, pending: Option, native_abort: Option, } pub(crate) fn run_call( py: Python<'_>, - call: OcrCall, + call: R::Call, route: R, ) -> PyResult> { let asynchronous = route.state().asynchronous; let mut lifecycle = PythonLifecycle { route, - call: Some(Arc::new(Mutex::new(NativeCall { call, result: None }))), + call: Some(Arc::new(Mutex::new(NativeCallState { call, result: None }))), pending: None, native_abort: None, }; @@ -214,14 +127,15 @@ impl PythonLifecycle { fn resume_core( &mut self, py: Python<'_>, - result: Option, - ) -> PyResult>> { + result: Option::Result, HostFailure>>, + ) -> PyResult::Operation>, Py>> { let call = Arc::clone(self.call.as_ref().ok_or_else(missing_state)?); let future = async move { let mut call = call.lock().await; let result = match result { - Some(OcrHostResult::Lifecycle(Err(failure))) => call.call.interrupt(failure).await, - result => call.call.resume(result).await, + Some(Err(failure)) => call.call.interrupt(failure).await, + Some(Ok(result)) => call.call.resume(Some(result)).await, + None => call.call.resume(None).await, }; call.result = Some(result); Ok(()) @@ -244,7 +158,7 @@ impl PythonLifecycle { } } - fn take_native_result(&self) -> PyResult { + fn take_native_result(&self) -> PyResult::Operation>> { self.call .as_ref() .ok_or_else(missing_state)? @@ -253,7 +167,7 @@ impl PythonLifecycle { .result .take() .ok_or_else(missing_state)? - .map_err(ocr_error_to_pyerr) + .map_err(R::map_error) } fn host_failure( @@ -261,7 +175,7 @@ impl PythonLifecycle { py: Python<'_>, error: PyErr, phase: Option, - ) -> OcrHostResult { + ) -> HostFailure { let native = litellm_core::Error::InvalidRequest(error.to_string()); let cancelled = !error.is_instance_of::(py); let failure = if !cancelled { @@ -276,7 +190,7 @@ impl PythonLifecycle { if state.end.is_none() { state.end = now(py).ok(); } - OcrHostResult::Lifecycle(Err(failure)) + failure } fn drive( @@ -289,16 +203,16 @@ impl PythonLifecycle { (Some(PendingOperation::Native), Some(result)) => match result { Ok(_) => HostStep::Ready(self.take_native_result()?), Err(error) => { - let result = self.host_failure(py, error, None); - self.resume_core(py, Some(result))? + let failure = self.host_failure(py, error, None); + self.resume_core(py, Some(Err(failure)))? } }, (Some(PendingOperation::Host(phase)), Some(result)) => { let result = result.and_then(|value| self.route.state_mut().accept(py, phase, value)); let result = match result { - Ok(()) => OcrHostResult::Lifecycle(Ok(())), - Err(error) => self.host_failure(py, error, Some(phase)), + Ok(()) => Ok(R::lifecycle_result()), + Err(error) => Err(self.host_failure(py, error, Some(phase))), }; self.resume_core(py, Some(result))? } @@ -307,7 +221,7 @@ impl PythonLifecycle { loop { let operation = match step { HostStep::Suspend(awaitable) => return Ok(ExecutionStep::Await(awaitable)), - HostStep::Ready(OcrCallStep::Complete(_)) => { + HostStep::Ready(NativeCallStep::Complete) => { return self .route .state_mut() @@ -316,44 +230,30 @@ impl PythonLifecycle { .map(ExecutionStep::Return) .ok_or_else(missing_state); } - HostStep::Ready(OcrCallStep::Host(operation)) => operation, + HostStep::Ready(NativeCallStep::Host(operation)) => operation, }; - let phase = match &operation { - OcrHostOperation::Lifecycle(phase) => Some(*phase), - OcrHostOperation::Failure { .. } => Some(HostPhase::Failure), - OcrHostOperation::Success { .. } => Some(HostPhase::Success), - _ => None, + let phase = match R::classify(&operation) { + OperationClass::Phase(phase) => Some(phase), + OperationClass::Route => None, }; - let result = match operation { - OcrHostOperation::Success { .. } => self - .route - .state_mut() - .invoke(py, HostPhase::Success) - .map(|_| OcrHostResult::Lifecycle(Ok(()))), - OcrHostOperation::Failure { .. } => self - .route - .state_mut() - .invoke(py, HostPhase::Failure) - .map(|_| OcrHostResult::Lifecycle(Ok(()))), - OcrHostOperation::Lifecycle(phase) => { - match self.route.state_mut().invoke(py, phase) { - Ok(HostStep::Suspend(awaitable)) => { - self.pending = Some(PendingOperation::Host(phase)); - return Ok(ExecutionStep::Await(awaitable)); - } - Ok(HostStep::Ready(value)) => self - .route - .state_mut() - .accept(py, phase, value) - .map(|()| OcrHostResult::Lifecycle(Ok(()))), - Err(error) => Err(error), + let result = match phase { + Some(phase) => match self.route.state_mut().invoke(py, phase) { + Ok(HostStep::Suspend(awaitable)) => { + self.pending = Some(PendingOperation::Host(phase)); + return Ok(ExecutionStep::Await(awaitable)); } - } - operation => self.route.invoke(py, operation), + Ok(HostStep::Ready(value)) => self + .route + .state_mut() + .accept(py, phase, value) + .map(|()| R::lifecycle_result()), + Err(error) => Err(error), + }, + None => self.route.invoke(py, operation), }; let result = match result { - Ok(result) => result, - Err(error) => self.host_failure(py, error, phase), + Ok(result) => Ok(result), + Err(error) => Err(self.host_failure(py, error, phase)), }; step = self.resume_core(py, Some(result))?; } @@ -762,6 +662,111 @@ mod tests { Execution::new(CallingBody(callback)) } + struct SyntheticCall(bool); + + impl NativeCall for SyntheticCall { + type Operation = (); + type Result = (); + + fn resume( + &mut self, + result: Option, + ) -> Pin< + Box< + dyn Future, litellm_core::Error>> + + Send + + '_, + >, + > { + Box::pin(async move { + match (self.0, result) { + (false, None) => { + self.0 = true; + Ok(NativeCallStep::Host(())) + } + (true, Some(())) => Ok(NativeCallStep::Complete), + _ => Err(litellm_core::Error::InvalidRequest( + "invalid synthetic lifecycle state".into(), + )), + } + }) + } + + fn interrupt( + &mut self, + _: HostFailure, + ) -> Pin< + Box< + dyn Future, litellm_core::Error>> + + Send + + '_, + >, + > { + Box::pin(async { Ok(NativeCallStep::Complete) }) + } + } + + struct SyntheticRoute(PythonCallState); + + impl PythonRoute for SyntheticRoute { + type Call = SyntheticCall; + + fn state(&self) -> &PythonCallState { + &self.0 + } + + fn state_mut(&mut self) -> &mut PythonCallState { + &mut self.0 + } + + fn classify(_: &()) -> OperationClass { + OperationClass::Route + } + + fn lifecycle_result() {} + + fn map_error(error: litellm_core::Error) -> PyErr { + crate::errors::core_error_to_pyerr(error) + } + + fn invoke(&mut self, py: Python<'_>, _: ()) -> PyResult<()> { + self.0.response = Some( + pyo3::types::PyString::new(py, "shared lifecycle") + .into_any() + .unbind(), + ); + Ok(()) + } + + fn cleanup(&mut self) {} + + fn traverse(&self, _: &PyVisit<'_>) -> Result<(), PyTraverseError> { + Ok(()) + } + } + + #[test] + fn shared_runner_executes_a_non_ocr_adapter() { + Python::initialize(); + Python::attach(|py| { + let route = SyntheticRoute( + PythonCallState::new( + py, + PyTuple::empty(py).unbind(), + PyDict::new(py).unbind(), + false, + "synthetic", + ) + .unwrap(), + ); + let value: String = run_call(py, SyntheticCall(false), route) + .unwrap() + .extract(py) + .unwrap(); + assert_eq!(value, "shared lifecycle"); + }); + } + #[test] fn python_driver_preserves_inline_await_and_native_ownership() { let _guard = PYTHON_GLOBALS @@ -771,7 +776,7 @@ mod tests { Python::attach(|py| { py.import("asyncio").unwrap(); let source = std::ffi::CString::new(include_str!( - "../../../../litellm/rust_bridge/lifecycle.py" + "../../../../../litellm/rust_bridge/lifecycle.py" )) .unwrap(); let module = PyModule::from_code( @@ -797,7 +802,7 @@ mod tests { wrap_pyfunction!(calling_execution, py).unwrap(), ) .unwrap(); - let probe = std::ffi::CString::new(include_str!("../tests/lifecycle.py")).unwrap(); + let probe = std::ffi::CString::new(include_str!("../../tests/lifecycle.py")).unwrap(); py.run(&probe, Some(&locals), Some(&locals)).unwrap(); }); } diff --git a/litellm-rust/crates/python-bridge/src/lifecycle/preparation.rs b/litellm-rust/crates/python-bridge/src/lifecycle/preparation.rs index dd07b46ca3b..97098d175da 100644 --- a/litellm-rust/crates/python-bridge/src/lifecycle/preparation.rs +++ b/litellm-rust/crates/python-bridge/src/lifecycle/preparation.rs @@ -1,4 +1,4 @@ -use litellm_core::ocr::{credential_default_fields, credential_index}; +use litellm_core::auth::{credential_default_fields, credential_index}; use pyo3::prelude::*; use pyo3::types::{PyDict, PyList}; diff --git a/litellm-rust/crates/python-bridge/src/marshal.rs b/litellm-rust/crates/python-bridge/src/marshal.rs index a14e4b55d82..7f4fdfcd02d 100644 --- a/litellm-rust/crates/python-bridge/src/marshal.rs +++ b/litellm-rust/crates/python-bridge/src/marshal.rs @@ -1,10 +1,13 @@ -use std::collections::HashMap; +use std::collections::{BTreeMap, HashMap}; use std::time::Duration; use pyo3::exceptions::PyValueError; use pyo3::prelude::*; use serde_json::{Map, Value}; +use litellm_core::auth::InputSource; +use litellm_python_interop::from_py_preserving_errors as from_py; + pub(crate) struct RouteOptions { pub(crate) model: String, pub(crate) api_key: Option, @@ -84,6 +87,52 @@ pub(crate) fn optional_timeout(timeout_seconds: Option) -> Option }) } +pub(crate) fn python_timeout_seconds(py: Python<'_>, timeout: Py) -> PyResult> { + py.import("litellm.rust_bridge.timeouts")? + .getattr("timeout_to_seconds")? + .call1((timeout,))? + .extract() +} + +pub(crate) fn project_optional_fields( + kwargs: &Bound<'_, pyo3::types::PyDict>, + names: &[&str], +) -> PyResult> { + names + .iter() + .filter_map(|name| match kwargs.get_item(name) { + Ok(Some(value)) => Some(from_py(&value).map(|value| ((*name).to_string(), value))), + Ok(None) => None, + Err(error) => Some(Err(error)), + }) + .collect() +} + +pub(crate) fn request_input_sources<'a>( + kwargs: &Bound<'_, pyo3::types::PyDict>, + names: impl Iterator, +) -> PyResult> { + let Some(proxy_request) = kwargs.get_item("proxy_server_request")? else { + return Ok(BTreeMap::new()); + }; + let proxy_request = proxy_request.cast_into::()?; + let body_fields = proxy_request + .get_item("body_fields")? + .or(proxy_request.get_item("body")?); + let credential_fields = proxy_request.get_item("credential_fields")?; + 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)) + }) + .collect()) +} + pub(crate) fn marshal_headers(headers: Option) -> PyResult> { let value = match headers { Some(headers) => headers, diff --git a/litellm-rust/crates/python-bridge/src/routes/mod.rs b/litellm-rust/crates/python-bridge/src/routes/mod.rs index 29d48aaf824..d1c26e1d9cd 100644 --- a/litellm-rust/crates/python-bridge/src/routes/mod.rs +++ b/litellm-rust/crates/python-bridge/src/routes/mod.rs @@ -10,14 +10,9 @@ mod audio_transcription; mod chat_completions; mod messages; mod ocr; -mod ocr_callbacks; -mod ocr_document; -mod ocr_lifecycle; pub(crate) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> { ocr::register(module)?; - ocr_document::register(module)?; - ocr_lifecycle::register(module)?; audio_transcription::register(module)?; messages::register(module)?; chat_completions::register(module)?; diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/callbacks.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/callbacks.rs new file mode 100644 index 00000000000..48747847489 --- /dev/null +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/callbacks.rs @@ -0,0 +1,102 @@ +use pyo3::exceptions::PyBaseException; +use pyo3::prelude::*; +use pyo3::types::PyDict; +use serde_json::Value; + +use litellm_core::ocr::LiteLLMOcrResponse; +use litellm_core::ocr::hooks::OcrPreCallRequest; +use litellm_python_interop::to_py_preserving_errors as to_py; + +use crate::lifecycle::PythonLogger; + +impl PythonLogger { + pub(crate) fn update_ocr( + &self, + py: Python<'_>, + kwargs: &Py, + pre_call: &OcrPreCallRequest, + url: &str, + ) -> PyResult<()> { + let redact = py + .import("litellm.rust_bridge.ocr")? + .getattr("redact_logging_params")?; + let update = PyDict::new(py); + update.set_item("kwargs", redact.call1((kwargs,))?.cast_into::()?)?; + update.set_item("model", &pre_call.model)?; + update.set_item( + "optional_params", + redact + .call1((to_py(py, &pre_call.optional_params)?,))? + .cast_into::()?, + )?; + let params = PyDict::new(py); + params.set_item( + "litellm_call_id", + kwargs.bind(py).get_item("litellm_call_id")?, + )?; + params.set_item("api_base", url)?; + update.set_item("litellm_params", params)?; + update.set_item("custom_llm_provider", &pre_call.custom_llm_provider)?; + self.object(py) + .call_method("update_from_kwargs", (), Some(&update))?; + Ok(()) + } + + pub(crate) fn pre_ocr( + &self, + py: Python<'_>, + api_key: &Option>, + body: &Bound<'_, PyDict>, + headers: &Bound<'_, PyDict>, + url: &str, + ) -> PyResult<()> { + let additional = PyDict::new(py); + additional.set_item("complete_input_dict", body)?; + additional.set_item("headers", headers)?; + additional.set_item("api_base", url)?; + let kwargs = PyDict::new(py); + kwargs.set_item("input", "OCR document processing")?; + kwargs.set_item("api_key", api_key)?; + kwargs.set_item("additional_args", additional)?; + self.object(py).call_method("pre_call", (), Some(&kwargs))?; + Ok(()) + } + + pub(crate) fn post_ocr( + &self, + py: Python<'_>, + original_response: &Value, + body: &Option>, + headers: &Option>, + ) -> PyResult<()> { + let kwargs = PyDict::new(py); + kwargs.set_item("original_response", to_py(py, original_response)?)?; + let additional = PyDict::new(py); + additional.set_item("complete_input_dict", body)?; + additional.set_item("headers", headers)?; + kwargs.set_item("additional_args", additional)?; + self.object(py) + .call_method("post_call", (), Some(&kwargs))?; + Ok(()) + } +} + +pub(super) fn response(py: Python<'_>, response: &LiteLLMOcrResponse) -> PyResult> { + py.import("litellm.rust_bridge.ocr")? + .getattr("_response")? + .call1((to_py(py, response)?,)) + .map(Bound::unbind) +} + +pub(super) fn map_failure( + py: Python<'_>, + error: &Py, + request: &Bound<'_, PyAny>, + provider: &str, +) -> PyResult> { + Ok(py + .import("litellm.rust_bridge.ocr_lifecycle")? + .getattr("map_failure")? + .call1((error, request, provider))? + .extract()?) +} diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/document.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/document.rs new file mode 100644 index 00000000000..11e9c78251f --- /dev/null +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/document.rs @@ -0,0 +1,259 @@ +use std::io::Read; +use std::path::PathBuf; + +use pyo3::exceptions::{PyFileNotFoundError, PyTypeError, PyValueError}; +use pyo3::prelude::*; +use pyo3::pybacked::PyBackedBytes; +use pyo3::types::{PyBytes, PyDict, PyString}; + +use litellm_core::constants::OCR_INLINE_MAX_BYTES; +use litellm_core::ocr::{OcrDocument, encode_file_document, mime_type_for_name, upload_mime_type}; +use litellm_python_interop::to_py_preserving_errors; + +enum FileBytes { + Python(PyBackedBytes), + Native(Vec), +} + +impl AsRef<[u8]> for FileBytes { + fn as_ref(&self) -> &[u8] { + match self { + Self::Python(bytes) => bytes, + Self::Native(bytes) => bytes, + } + } +} + +fn read_file_input( + py: Python<'_>, + file: &Bound<'_, PyAny>, +) -> PyResult<(FileBytes, Option)> { + if file.is_instance_of::() { + return Err(PyValueError::new_err( + "OCR file input does not accept bare str values. Pass bytes, a pathlib.Path, or a file-like object.", + )); + } + if file.is_instance(&py.import("os")?.getattr("PathLike")?)? { + let path: PathBuf = file.extract()?; + let name = path + .file_name() + .map(|value| value.to_string_lossy().into_owned()); + let bytes = py + .detach(|| { + let mut bytes = Vec::new(); + std::fs::File::open(&path)? + .take(OCR_INLINE_MAX_BYTES as u64 + 1) + .read_to_end(&mut bytes)?; + Ok::<_, std::io::Error>(bytes) + }) + .map_err(|error| { + if error.kind() == std::io::ErrorKind::NotFound { + PyFileNotFoundError::new_err(format!("File not found: {}", path.display())) + } else { + error.into() + } + })?; + return Ok((FileBytes::Native(bytes), name)); + } + if file.is_instance_of::() { + return Ok((FileBytes::Python(file.extract()?), None)); + } + let reader = file + .getattr_opt("read")? + .filter(|value| value.is_callable()); + let Some(reader) = reader else { + return Err(PyValueError::new_err(format!( + "Unsupported file input type: {}. Expected pathlib.Path, bytes, or a file-like object.", + file.get_type(), + ))); + }; + let name = file + .getattr_opt("name")? + .filter(|value| !value.is_none()) + .map(|value| value.extract::()) + .transpose()?; + let value = reader.call0()?; + let bytes = if value.is_instance_of::() { + FileBytes::Native(value.extract::()?.into_bytes()) + } else if value.is_instance_of::() { + FileBytes::Python(value.extract()?) + } else { + return Err(PyTypeError::new_err(format!( + "OCR file read must return bytes or str, got {}", + value.get_type(), + ))); + }; + Ok((bytes, name)) +} + +pub(super) struct FileDocumentInput { + bytes: FileBytes, + name: Option, + mime_type: Option, +} + +impl FromPyObject<'_, '_> for FileDocumentInput { + type Error = PyErr; + + fn extract(document: Borrowed<'_, '_, PyAny>) -> PyResult { + let py = document.py(); + let file = document.get_item("file").map_err(|error| { + if error.is_instance_of::(py) { + PyValueError::new_err("document with type='file' must include a 'file' field containing a pathlib.Path, file-like object, or bytes") + } else { + error + } + })?; + if file.is_none() { + return Err(PyValueError::new_err( + "document with type='file' must include a 'file' field containing a pathlib.Path, file-like object, or bytes", + )); + } + let (bytes, name) = read_file_input(py, &file)?; + let mime_type = document + .cast::()? + .get_item("mime_type")? + .map(|value| value.extract::()) + .transpose()?; + Ok(Self { + bytes, + name, + mime_type, + }) + } +} + +pub(super) fn file_document(py: Python<'_>, document: FileDocumentInput) -> PyResult { + py.detach(|| { + encode_file_document( + document.bytes.as_ref(), + document.name.as_deref(), + document.mime_type.as_deref(), + ) + }) + .map_err(|error| PyValueError::new_err(error.to_string())) +} + +#[pyfunction] +fn _ocr_file_document(py: Python<'_>, document: Bound<'_, PyAny>) -> PyResult> { + to_py_preserving_errors(py, &file_document(py, document.extract()?)?) +} + +#[pyfunction] +fn _ocr_mime_type(file_name: &str) -> String { + mime_type_for_name(file_name).into() +} + +#[pyfunction] +#[pyo3(signature = (file_content, file_name=None, content_type=None))] +fn _ocr_upload_document( + py: Python<'_>, + file_content: &Bound<'_, PyBytes>, + file_name: Option<&str>, + content_type: Option<&str>, +) -> PyResult> { + let bytes: PyBackedBytes = file_content.extract()?; + let document = py + .detach(|| { + encode_file_document( + &bytes, + None, + Some(upload_mime_type(file_name, content_type)), + ) + }) + .map_err(|error| PyValueError::new_err(error.to_string()))?; + to_py_preserving_errors(py, &document) +} + +pub(super) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> { + module.add("_OCR_MAX_FILE_BYTES", OCR_INLINE_MAX_BYTES)?; + module.add_function(wrap_pyfunction!(_ocr_upload_document, module)?)?; + module.add_function(wrap_pyfunction!(_ocr_file_document, module)?)?; + module.add_function(wrap_pyfunction!(_ocr_mime_type, module)?) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn extraction_validates_required_file_and_optional_mime_type() { + Python::initialize(); + Python::attach(|py| { + for expression in [c"{}", c"{'file': None}"] { + let document = py.eval(expression, None, None).unwrap(); + let error = document.extract::().err().unwrap(); + assert!(error.is_instance_of::(py)); + assert!(error.to_string().contains("must include a 'file' field")); + } + for expression in [ + c"{'file': b'abc', 'mime_type': None}", + c"{'file': b'abc', 'mime_type': 7}", + ] { + let document = py.eval(expression, None, None).unwrap(); + let error = document.extract::().err().unwrap(); + assert!(error.is_instance_of::(py)); + } + let document = py.eval(c"{'file': b'abc'}", None, None).unwrap(); + let input: FileDocumentInput = document.extract().unwrap(); + assert_eq!(input.bytes.as_ref(), b"abc"); + assert_eq!(input.name, None); + assert_eq!(input.mime_type, None); + }); + } + + #[test] + fn extraction_reads_mime_type_after_consuming_file_once() { + Python::initialize(); + Python::attach(|py| { + let locals = PyDict::new(py); + py.run( + c"class Reader: + def read(self): + assert document['mime_type'] == 7 + document['mime_type'] = 'image/png' + return b'abc' +document = {'file': Reader(), 'mime_type': 7}", + Some(&locals), + Some(&locals), + ) + .unwrap(); + let document = locals.get_item("document").unwrap().unwrap(); + let input: FileDocumentInput = document.extract().unwrap(); + assert_eq!(input.bytes.as_ref(), b"abc"); + assert_eq!(input.mime_type.as_deref(), Some("image/png")); + let result = file_document(py, input).unwrap(); + assert_eq!( + serde_json::to_value(result).unwrap(), + serde_json::json!({ + "type": "image_url", "image_url": "data:image/png;base64,YWJj" + }) + ); + }); + } + + #[test] + fn extraction_preserves_reader_key_error_identity() { + Python::initialize(); + Python::attach(|py| { + let locals = PyDict::new(py); + py.run( + c"failure = KeyError('reader failed') +class Reader: + def read(self): + raise failure +document = {'file': Reader()}", + Some(&locals), + Some(&locals), + ) + .unwrap(); + let document = locals.get_item("document").unwrap().unwrap(); + let error = document.extract::().err().unwrap(); + assert!( + error + .value(py) + .is(locals.get_item("failure").unwrap().unwrap()) + ); + }); + } +} diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/errors.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/errors.rs new file mode 100644 index 00000000000..f5fb7bead7d --- /dev/null +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/errors.rs @@ -0,0 +1,77 @@ +use litellm_core::error::Error; +use pyo3::exceptions::PyValueError; +use pyo3::prelude::*; + +use crate::errors::{RustUpstreamError, core_error_to_pyerr}; + +pub(super) fn to_pyerr(error: Error) -> PyErr { + match error { + Error::MissingField("document_url" | "image_url") => { + PyValueError::new_err("Document URL is required") + } + Error::Http { status, body } => upstream_error(status, body), + Error::Network(message) if message.contains("timed out") => upstream_error(408, message), + other => { + let status = other.http_status_code(); + let error = core_error_to_pyerr(other); + if let Some(status) = status { + Python::attach(|py| { + let value = error.value(py); + value.setattr("status_code", status).ok(); + value.setattr("message", value.to_string()).ok(); + }); + } + error + } + } +} + +fn upstream_error(status: u16, message: String) -> PyErr { + let error = RustUpstreamError::new_err((status, message.clone())); + Python::attach(|py| { + let value = error.value(py); + value.setattr("status_code", status).ok(); + value.setattr("message", message).ok(); + }); + error +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn preserves_python_validation_and_provider_details() { + Python::initialize(); + Python::attach(|py| { + for field in ["document_url", "image_url"] { + let mapped = to_pyerr(Error::MissingField(field)); + assert!(mapped.is_instance_of::(py)); + assert_eq!(mapped.value(py).to_string(), "Document URL is required"); + } + let mapped = to_pyerr(Error::Http { + status: 429, + body: r#"{"message":"rate limited"}"#.to_string(), + }); + assert!(mapped.is_instance_of::(py)); + let args: (u16, String) = mapped + .value(py) + .getattr("args") + .and_then(|args| args.extract()) + .expect("OCR failures retain status and unprefixed provider message"); + assert_eq!(args, (429, r#"{"message":"rate limited"}"#.to_string())); + + let mapped = to_pyerr(Error::InvalidRequest("invalid format".into())); + assert!(mapped.is_instance_of::(py)); + assert_eq!( + mapped + .value(py) + .getattr("status_code") + .unwrap() + .extract::() + .unwrap(), + 400 + ); + }); + } +} diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr_lifecycle.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/lifecycle.rs similarity index 78% rename from litellm-rust/crates/python-bridge/src/routes/ocr_lifecycle.rs rename to litellm-rust/crates/python-bridge/src/routes/ocr/lifecycle.rs index 0b03f626010..10252d09ec5 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr_lifecycle.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/lifecycle.rs @@ -1,4 +1,4 @@ -use serde_json::{Map, Value}; +use serde_json::Value; use std::sync::Arc; use pyo3::prelude::*; @@ -14,9 +14,15 @@ use litellm_python_interop::{ from_py_preserving_errors as from_py, to_py_preserving_errors as to_py, }; -use super::ocr_callbacks::{self, AzureAdTokenProvider}; -use crate::errors::{RustBridgeDeclined, ocr_error_to_pyerr}; -use crate::lifecycle::{PythonCallState, PythonRoute, missing_state, now, run_call}; +use super::callbacks; +use super::errors::to_pyerr as ocr_error_to_pyerr; +use crate::auth::{AZURE_AD_TOKEN_PROVIDER, PythonTokenProvider}; +use crate::errors::RustBridgeDeclined; +use crate::lifecycle::{ + NativeCall, NativeCallStep, OperationClass, PythonCallState, PythonRoute, missing_state, now, + run_call, +}; +use crate::marshal::{project_optional_fields, python_timeout_seconds, request_input_sources}; struct PythonOcrHost { state: PythonCallState, @@ -24,7 +30,7 @@ struct PythonOcrHost { pre_call: Option, document: Option>, api_key: Option>, - azure_ad_token_provider: Option, + azure_ad_token_provider: Option, provider: String, retained_fields: Option>, body: Option>, @@ -35,7 +41,7 @@ struct AdmittedOcrCall { request: litellm_core::ocr::LiteLLMOcrRequest, document: Py, api_key: Py, - azure_ad_token_provider: Option, + azure_ad_token_provider: Option, provider: String, } @@ -125,7 +131,56 @@ impl PythonOcrHost { } } +impl NativeCall for OcrCall { + type Operation = OcrHostOperation; + type Result = OcrHostResult; + + fn resume( + &mut self, + result: Option, + ) -> std::pin::Pin< + Box< + dyn std::future::Future< + Output = Result, litellm_core::Error>, + > + Send + + '_, + >, + > { + Box::pin(async move { + OcrCall::resume(self, result).await.map(|step| match step { + litellm_core::ocr::OcrCallStep::Host(operation) => NativeCallStep::Host(operation), + litellm_core::ocr::OcrCallStep::Complete(_) => NativeCallStep::Complete, + }) + }) + } + + fn interrupt( + &mut self, + failure: litellm_core::call_lifecycle::host::HostFailure, + ) -> std::pin::Pin< + Box< + dyn std::future::Future< + Output = Result, litellm_core::Error>, + > + Send + + '_, + >, + > { + Box::pin(async move { + OcrCall::interrupt(self, failure) + .await + .map(|step| match step { + litellm_core::ocr::OcrCallStep::Host(operation) => { + NativeCallStep::Host(operation) + } + litellm_core::ocr::OcrCallStep::Complete(_) => NativeCallStep::Complete, + }) + }) + } +} + impl PythonRoute for PythonOcrHost { + type Call = OcrCall; + fn state(&self) -> &PythonCallState { &self.state } @@ -134,6 +189,27 @@ impl PythonRoute for PythonOcrHost { &mut self.state } + fn classify(operation: &OcrHostOperation) -> OperationClass { + match operation { + OcrHostOperation::Lifecycle(phase) => OperationClass::Phase(*phase), + OcrHostOperation::Success { .. } => { + OperationClass::Phase(litellm_core::call_lifecycle::host::HostPhase::Success) + } + OcrHostOperation::Failure { .. } => { + OperationClass::Phase(litellm_core::call_lifecycle::host::HostPhase::Failure) + } + _ => OperationClass::Route, + } + } + + fn lifecycle_result() -> OcrHostResult { + OcrHostResult::Lifecycle(Ok(())) + } + + fn map_error(error: litellm_core::Error) -> PyErr { + ocr_error_to_pyerr(error) + } + fn invoke(&mut self, py: Python<'_>, operation: OcrHostOperation) -> PyResult { Ok(match operation { OcrHostOperation::ProjectRequest => { @@ -165,7 +241,7 @@ impl PythonRoute for PythonOcrHost { } OcrHostOperation::ConstructResponse(response) => { self.state.end = Some(now(py)?); - self.state.response = Some(ocr_callbacks::response(py, response.as_ref())?); + self.state.response = Some(callbacks::response(py, response.as_ref())?); OcrHostResult::Lifecycle(Ok(())) } OcrHostOperation::MapFailure(error) => { @@ -177,7 +253,7 @@ impl PythonRoute for PythonOcrHost { } let error = self.state.error.as_ref().ok_or_else(missing_state)?; let request = self.request.as_ref().ok_or_else(missing_state)?.bind(py); - let mapped = ocr_callbacks::map_failure(py, error, request, &self.provider)?; + let mapped = callbacks::map_failure(py, error, request, &self.provider)?; self.state .retain_error(py, PyErr::from_value(mapped.into_bound(py).into_any())); OcrHostResult::Lifecycle(Ok(())) @@ -231,11 +307,17 @@ fn project_request( let request_kwargs = kwargs; let consumed = consumed_optional_param_names(&model, custom_llm_provider.as_deref()) .map_err(ocr_error_to_pyerr)?; - let optional_params = extract_optional_params(request_kwargs, &consumed)?; - let input_sources = extract_input_sources(request_kwargs, &consumed)?; + let optional_params = project_optional_fields(request_kwargs, &consumed)?; + let input_sources = request_input_sources( + request_kwargs, + consumed + .iter() + .copied() + .chain(["api_key", "api_base", "extra_headers"]), + )?; let azure_ad_token_provider = request_kwargs .get_item("azure_ad_token_provider")? - .and_then(AzureAdTokenProvider::select); + .and_then(|provider| PythonTokenProvider::select(provider, AZURE_AD_TOKEN_PROVIDER)); let wire = OcrWireRequest { model, document: wire_document, @@ -250,7 +332,7 @@ fn project_request( input_sources, timeout_seconds: argument("timeout")? .extract::>>()? - .map(|value| ocr_callbacks::timeout_seconds(py, value)) + .map(|value| python_timeout_seconds(py, value)) .transpose()? .flatten(), }; @@ -266,55 +348,11 @@ fn project_request( }) } -fn extract_optional_params( - kwargs: &Bound<'_, PyDict>, - consumed: &[&str], -) -> PyResult> { - let mut optional_params = Map::new(); - for name in consumed { - if let Some(value) = kwargs.get_item(name)? { - optional_params.insert((*name).to_string(), from_py(&value)?); - } - } - Ok(optional_params) -} - -fn extract_input_sources( - kwargs: &Bound<'_, PyDict>, - consumed: &[&str], -) -> PyResult> { - let Some(proxy_request) = kwargs.get_item("proxy_server_request")? else { - return Ok(Default::default()); - }; - let proxy_request = proxy_request.cast_into::()?; - 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 mut sources = std::collections::BTreeMap::new(); - for name in consumed - .iter() - .copied() - .chain(["api_key", "api_base", "extra_headers"]) - { - 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)); - if present { - sources.insert(name.to_string(), litellm_core::auth::InputSource::Request); - } - } - Ok(sources) -} - fn extract_document(py: Python<'_>, document: &Bound<'_, PyAny>) -> PyResult { if document.get_item("type")?.extract::()? != "file" { return from_py(document); } - serde_json::to_value(super::ocr_document::file_document(py, document)?) + serde_json::to_value(super::document::file_document(py, document.extract()?)?) .map_err(|error| pyo3::exceptions::PyValueError::new_err(error.to_string())) } diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs new file mode 100644 index 00000000000..8c6469a4bf7 --- /dev/null +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs @@ -0,0 +1,18 @@ +mod callbacks; +mod document; +mod errors; +mod lifecycle; +mod value; + +use pyo3::prelude::*; + +pub(super) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> { + value::register(module)?; + document::register(module)?; + lifecycle::register(module) +} + +#[cfg(feature = "trace-parity")] +pub(super) fn register_trace(module: &Bound<'_, PyModule>) -> PyResult<()> { + value::register_trace(module) +} diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/value.rs similarity index 98% rename from litellm-rust/crates/python-bridge/src/routes/ocr.rs rename to litellm-rust/crates/python-bridge/src/routes/ocr/value.rs index c5def64c2f1..94f760e1563 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/value.rs @@ -6,7 +6,7 @@ use litellm_core::ocr::wire::{OcrWireRequest, decode_request, is_supported_reque use pyo3::prelude::*; use serde_json::Value; -use crate::errors::ocr_error_to_pyerr; +use super::errors::to_pyerr as ocr_error_to_pyerr; use crate::marshal::{RouteOptions, RouteOptionsInputs, object_or_empty}; fn prepare_ocr( diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr_document.rs b/litellm-rust/crates/python-bridge/src/routes/ocr_document.rs deleted file mode 100644 index 8a77616f8d5..00000000000 --- a/litellm-rust/crates/python-bridge/src/routes/ocr_document.rs +++ /dev/null @@ -1,148 +0,0 @@ -use std::io::Read; -use std::path::PathBuf; - -use pyo3::exceptions::{PyFileNotFoundError, PyTypeError, PyValueError}; -use pyo3::prelude::*; -use pyo3::pybacked::PyBackedBytes; -use pyo3::types::{PyBytes, PyDict, PyString}; - -use litellm_core::constants::OCR_INLINE_MAX_BYTES; -use litellm_core::ocr::{OcrDocument, encode_file_document, mime_type_for_name, upload_mime_type}; -use litellm_python_interop::to_py_preserving_errors; - -enum FileBytes { - Python(PyBackedBytes), - Native(Vec), -} - -impl AsRef<[u8]> for FileBytes { - fn as_ref(&self) -> &[u8] { - match self { - Self::Python(bytes) => bytes, - Self::Native(bytes) => bytes, - } - } -} - -fn read_file_input( - py: Python<'_>, - file: &Bound<'_, PyAny>, -) -> PyResult<(FileBytes, Option)> { - if file.is_instance_of::() { - return Err(PyValueError::new_err( - "OCR file input does not accept bare str values. Pass bytes, a pathlib.Path, or a file-like object.", - )); - } - if file.is_instance(&py.import("os")?.getattr("PathLike")?)? { - let path: PathBuf = file.extract()?; - let name = path - .file_name() - .map(|value| value.to_string_lossy().into_owned()); - let bytes = py - .detach(|| { - let mut bytes = Vec::new(); - std::fs::File::open(&path)? - .take(OCR_INLINE_MAX_BYTES as u64 + 1) - .read_to_end(&mut bytes)?; - Ok::<_, std::io::Error>(bytes) - }) - .map_err(|error| { - if error.kind() == std::io::ErrorKind::NotFound { - PyFileNotFoundError::new_err(format!("File not found: {}", path.display())) - } else { - error.into() - } - })?; - return Ok((FileBytes::Native(bytes), name)); - } - if file.is_instance_of::() { - return Ok((FileBytes::Python(file.extract()?), None)); - } - let reader = file - .getattr_opt("read")? - .filter(|value| value.is_callable()); - let Some(reader) = reader else { - return Err(PyValueError::new_err(format!( - "Unsupported file input type: {}. Expected pathlib.Path, bytes, or a file-like object.", - file.get_type(), - ))); - }; - let name = file - .getattr_opt("name")? - .filter(|value| !value.is_none()) - .map(|value| value.extract::()) - .transpose()?; - let value = reader.call0()?; - let bytes = if value.is_instance_of::() { - FileBytes::Native(value.extract::()?.into_bytes()) - } else if value.is_instance_of::() { - FileBytes::Python(value.extract()?) - } else { - return Err(PyTypeError::new_err(format!( - "OCR file read must return bytes or str, got {}", - value.get_type(), - ))); - }; - Ok((bytes, name)) -} - -pub(super) fn file_document(py: Python<'_>, document: &Bound<'_, PyAny>) -> PyResult { - let file = document.get_item("file").map_err(|error| { - if error.is_instance_of::(py) { - PyValueError::new_err("document with type='file' must include a 'file' field containing a pathlib.Path, file-like object, or bytes") - } else { - error - } - })?; - if file.is_none() { - return Err(PyValueError::new_err( - "document with type='file' must include a 'file' field containing a pathlib.Path, file-like object, or bytes", - )); - } - let (bytes, name) = read_file_input(py, &file)?; - let mime = document - .cast::()? - .get_item("mime_type")? - .map(|value| value.extract::()) - .transpose()?; - py.detach(|| encode_file_document(bytes.as_ref(), name.as_deref(), mime.as_deref())) - .map_err(|error| PyValueError::new_err(error.to_string())) -} - -#[pyfunction] -fn _ocr_file_document(py: Python<'_>, document: Bound<'_, PyAny>) -> PyResult> { - to_py_preserving_errors(py, &file_document(py, &document)?) -} - -#[pyfunction] -fn _ocr_mime_type(file_name: &str) -> String { - mime_type_for_name(file_name).into() -} - -#[pyfunction] -#[pyo3(signature = (file_content, file_name=None, content_type=None))] -fn _ocr_upload_document( - py: Python<'_>, - file_content: &Bound<'_, PyBytes>, - file_name: Option<&str>, - content_type: Option<&str>, -) -> PyResult> { - let bytes: PyBackedBytes = file_content.extract()?; - let document = py - .detach(|| { - encode_file_document( - &bytes, - None, - Some(upload_mime_type(file_name, content_type)), - ) - }) - .map_err(|error| PyValueError::new_err(error.to_string()))?; - to_py_preserving_errors(py, &document) -} - -pub(super) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> { - module.add("_OCR_MAX_FILE_BYTES", OCR_INLINE_MAX_BYTES)?; - module.add_function(wrap_pyfunction!(_ocr_upload_document, module)?)?; - module.add_function(wrap_pyfunction!(_ocr_file_document, module)?)?; - module.add_function(wrap_pyfunction!(_ocr_mime_type, module)?) -}