This commit is contained in:
Yujong Lee 2026-09-12 08:35:34 -07:00
parent a67e94162a
commit fe8759df7b
19 changed files with 1001 additions and 621 deletions

View file

@ -9,6 +9,21 @@ use crate::AuthError;
use super::{ResolvedCredential, SecretValue, TokenProviderHandle};
pub fn credential_index(requested: &str, names: &[String]) -> Option<usize> {
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),

View file

@ -49,6 +49,7 @@ impl<T> Sourced<T> {
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};

View file

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

View file

@ -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<usize> {
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<T> {
#[serde(flatten)]

View file

@ -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<PyAny>);
pub(crate) struct PythonTokenProvider {
callback: Py<PyAny>,
contract: TokenProviderContract,
}
impl AzureAdTokenProvider {
pub(super) fn select(provider: Bound<'_, PyAny>) -> Option<Self> {
(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<Self> {
(provider.is_callable() && provider.is_truthy().unwrap_or(false)).then(|| Self {
callback: provider.unbind(),
contract,
})
}
pub(super) fn acquire(&self, py: Python<'_>) -> PyResult<ResolvedCredential> {
let provider = self.0.bind(py);
pub(crate) fn acquire(&self, py: Python<'_>) -> PyResult<ResolvedCredential> {
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::<PyString>() {
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::<PyTypeError>(py) || !error.is_instance_of::<PyException>(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<PyDict>,
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::<PyDict>()?)?;
update.set_item("model", &pre_call.model)?;
update.set_item(
"optional_params",
redact
.call1((to_py(py, &pre_call.optional_params)?,))?
.cast_into::<PyDict>()?,
)?;
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<Py<PyAny>>,
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<Py<PyDict>>,
headers: &Option<Py<PyDict>>,
) -> 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<PyAny>> {
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<PyBaseException>,
request: &Bound<'_, PyAny>,
provider: &str,
) -> PyResult<Py<PyBaseException>> {
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<PyAny>) -> PyResult<Option<f64>> {
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::<PyRuntimeError>(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::<PyRuntimeError>(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::<pyo3::exceptions::PyUnicodeEncodeError>(py));
});

View file

@ -63,77 +63,3 @@ pub(crate) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> {
module.add("RustBridgeDeclined", py.get_type::<RustBridgeDeclined>())?;
module.add("RustUpstreamError", py.get_type::<RustUpstreamError>())
}
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::<PyValueError>(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::<RustUpstreamError>(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::<PyValueError>(py));
assert_eq!(
mapped
.value(py)
.getattr("status_code")
.unwrap()
.extract::<u16>()
.unwrap(),
400
);
});
}
}

View file

@ -1,3 +1,4 @@
mod auth;
mod constants;
mod diagnostics;
mod errors;

View file

@ -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<PyAny>),
Await(Py<PyAny>),
}
pub(super) trait ExecutionBody: Send + Sync {
fn resume(&mut self, result: Option<PyResult<Py<PyAny>>>) -> PyResult<ExecutionStep>;
fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError>;
}
enum ExecutionState {
Created(Box<dyn ExecutionBody>),
Running,
Suspended(Box<dyn ExecutionBody>),
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<Py<PyAny>>>,
) -> PyResult<Py<PyAny>> {
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<Py<PyAny>> {
Self::advance(slf, py, None)
}
fn resume_value(
slf: &Bound<'_, Self>,
py: Python<'_>,
value: Py<PyAny>,
) -> PyResult<Py<PyAny>> {
Self::advance(slf, py, Some(Ok(value)))
}
fn resume_error(
slf: &Bound<'_, Self>,
py: Python<'_>,
error: Bound<'_, PyBaseException>,
) -> PyResult<Py<PyAny>> {
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);
}
}

View file

@ -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<O> {
Host(O),
Complete,
}
pub(crate) trait NativeCall: Send + Sync {
type Operation: Send + 'static;
type Result: Send + 'static;
fn resume(
&mut self,
result: Option<Self::Result>,
) -> Pin<
Box<
dyn Future<Output = Result<NativeCallStep<Self::Operation>, litellm_core::Error>>
+ Send
+ '_,
>,
>;
fn interrupt(
&mut self,
failure: HostFailure,
) -> Pin<
Box<
dyn Future<Output = Result<NativeCallStep<Self::Operation>, 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<OcrHostResult>;
fn classify(operation: &<Self::Call as NativeCall>::Operation) -> OperationClass;
fn lifecycle_result() -> <Self::Call as NativeCall>::Result;
fn map_error(error: litellm_core::Error) -> PyErr;
fn invoke(
&mut self,
py: Python<'_>,
operation: <Self::Call as NativeCall>::Operation,
) -> PyResult<<Self::Call as NativeCall>::Result>;
fn cleanup(&mut self);
fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError>;
}
enum ExecutionStep {
Return(Py<PyAny>),
Await(Py<PyAny>),
}
trait ExecutionBody: Send + Sync {
fn resume(&mut self, result: Option<PyResult<Py<PyAny>>>) -> PyResult<ExecutionStep>;
fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError>;
}
enum ExecutionState {
Created(Box<dyn ExecutionBody>),
Running,
Suspended(Box<dyn ExecutionBody>),
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<Py<PyAny>>>,
) -> PyResult<Py<PyAny>> {
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<Py<PyAny>> {
Self::advance(slf, py, None)
}
fn resume_value(
slf: &Bound<'_, Self>,
py: Python<'_>,
value: Py<PyAny>,
) -> PyResult<Py<PyAny>> {
Self::advance(slf, py, Some(Ok(value)))
}
fn resume_error(
slf: &Bound<'_, Self>,
py: Python<'_>,
error: Bound<'_, PyBaseException>,
) -> PyResult<Py<PyAny>> {
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<Result<OcrCallStep, litellm_core::Error>>,
struct NativeCallState<C: NativeCall> {
call: C,
result: Option<Result<NativeCallStep<C::Operation>, litellm_core::Error>>,
}
enum PendingOperation {
@ -173,20 +86,20 @@ enum PendingOperation {
struct PythonLifecycle<R: PythonRoute> {
route: R,
call: Option<Arc<Mutex<NativeCall>>>,
call: Option<Arc<Mutex<NativeCallState<R::Call>>>>,
pending: Option<PendingOperation>,
native_abort: Option<AbortHandle>,
}
pub(crate) fn run_call<R: PythonRoute + 'static>(
py: Python<'_>,
call: OcrCall,
call: R::Call,
route: R,
) -> PyResult<Py<PyAny>> {
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<R: PythonRoute> PythonLifecycle<R> {
fn resume_core(
&mut self,
py: Python<'_>,
result: Option<OcrHostResult>,
) -> PyResult<HostStep<OcrCallStep, Py<PyAny>>> {
result: Option<Result<<R::Call as NativeCall>::Result, HostFailure>>,
) -> PyResult<HostStep<NativeCallStep<<R::Call as NativeCall>::Operation>, Py<PyAny>>> {
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<R: PythonRoute> PythonLifecycle<R> {
}
}
fn take_native_result(&self) -> PyResult<OcrCallStep> {
fn take_native_result(&self) -> PyResult<NativeCallStep<<R::Call as NativeCall>::Operation>> {
self.call
.as_ref()
.ok_or_else(missing_state)?
@ -253,7 +167,7 @@ impl<R: PythonRoute> PythonLifecycle<R> {
.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<R: PythonRoute> PythonLifecycle<R> {
py: Python<'_>,
error: PyErr,
phase: Option<HostPhase>,
) -> OcrHostResult {
) -> HostFailure {
let native = litellm_core::Error::InvalidRequest(error.to_string());
let cancelled = !error.is_instance_of::<PyException>(py);
let failure = if !cancelled {
@ -276,7 +190,7 @@ impl<R: PythonRoute> PythonLifecycle<R> {
if state.end.is_none() {
state.end = now(py).ok();
}
OcrHostResult::Lifecycle(Err(failure))
failure
}
fn drive(
@ -289,16 +203,16 @@ impl<R: PythonRoute> PythonLifecycle<R> {
(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<R: PythonRoute> PythonLifecycle<R> {
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<R: PythonRoute> PythonLifecycle<R> {
.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<Self::Result>,
) -> Pin<
Box<
dyn Future<Output = Result<NativeCallStep<Self::Operation>, 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<Output = Result<NativeCallStep<Self::Operation>, 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();
});
}

View file

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

View file

@ -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<String>,
@ -84,6 +87,52 @@ pub(crate) fn optional_timeout(timeout_seconds: Option<f64>) -> Option<Duration>
})
}
pub(crate) fn python_timeout_seconds(py: Python<'_>, timeout: Py<PyAny>) -> PyResult<Option<f64>> {
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<Map<String, Value>> {
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<Item = &'a str>,
) -> PyResult<BTreeMap<String, InputSource>> {
let Some(proxy_request) = kwargs.get_item("proxy_server_request")? else {
return Ok(BTreeMap::new());
};
let proxy_request = proxy_request.cast_into::<pyo3::types::PyDict>()?;
let body_fields = proxy_request
.get_item("body_fields")?
.or(proxy_request.get_item("body")?);
let credential_fields = proxy_request.get_item("credential_fields")?;
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<Value>) -> PyResult<HashMap<String, String>> {
let value = match headers {
Some(headers) => headers,

View file

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

View file

@ -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<PyDict>,
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::<PyDict>()?)?;
update.set_item("model", &pre_call.model)?;
update.set_item(
"optional_params",
redact
.call1((to_py(py, &pre_call.optional_params)?,))?
.cast_into::<PyDict>()?,
)?;
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<Py<PyAny>>,
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<Py<PyDict>>,
headers: &Option<Py<PyDict>>,
) -> 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<PyAny>> {
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<PyBaseException>,
request: &Bound<'_, PyAny>,
provider: &str,
) -> PyResult<Py<PyBaseException>> {
Ok(py
.import("litellm.rust_bridge.ocr_lifecycle")?
.getattr("map_failure")?
.call1((error, request, provider))?
.extract()?)
}

View file

@ -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<u8>),
}
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<String>)> {
if file.is_instance_of::<PyString>() {
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::<PyBytes>() {
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::<String>())
.transpose()?;
let value = reader.call0()?;
let bytes = if value.is_instance_of::<PyString>() {
FileBytes::Native(value.extract::<String>()?.into_bytes())
} else if value.is_instance_of::<PyBytes>() {
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<String>,
mime_type: Option<String>,
}
impl FromPyObject<'_, '_> for FileDocumentInput {
type Error = PyErr;
fn extract(document: Borrowed<'_, '_, PyAny>) -> PyResult<Self> {
let py = document.py();
let file = document.get_item("file").map_err(|error| {
if error.is_instance_of::<pyo3::exceptions::PyKeyError>(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::<PyDict>()?
.get_item("mime_type")?
.map(|value| value.extract::<String>())
.transpose()?;
Ok(Self {
bytes,
name,
mime_type,
})
}
}
pub(super) fn file_document(py: Python<'_>, document: FileDocumentInput) -> PyResult<OcrDocument> {
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<Py<PyAny>> {
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<Py<PyAny>> {
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::<FileDocumentInput>().err().unwrap();
assert!(error.is_instance_of::<PyValueError>(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::<FileDocumentInput>().err().unwrap();
assert!(error.is_instance_of::<PyTypeError>(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::<FileDocumentInput>().err().unwrap();
assert!(
error
.value(py)
.is(locals.get_item("failure").unwrap().unwrap())
);
});
}
}

View file

@ -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::<PyValueError>(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::<RustUpstreamError>(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::<PyValueError>(py));
assert_eq!(
mapped
.value(py)
.getattr("status_code")
.unwrap()
.extract::<u16>()
.unwrap(),
400
);
});
}
}

View file

@ -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<OcrPreCallRequest>,
document: Option<Py<PyAny>>,
api_key: Option<Py<PyAny>>,
azure_ad_token_provider: Option<AzureAdTokenProvider>,
azure_ad_token_provider: Option<PythonTokenProvider>,
provider: String,
retained_fields: Option<Py<PyDict>>,
body: Option<Py<PyDict>>,
@ -35,7 +41,7 @@ struct AdmittedOcrCall {
request: litellm_core::ocr::LiteLLMOcrRequest,
document: Py<PyAny>,
api_key: Py<PyAny>,
azure_ad_token_provider: Option<AzureAdTokenProvider>,
azure_ad_token_provider: Option<PythonTokenProvider>,
provider: String,
}
@ -125,7 +131,56 @@ impl PythonOcrHost {
}
}
impl NativeCall for OcrCall {
type Operation = OcrHostOperation;
type Result = OcrHostResult;
fn resume(
&mut self,
result: Option<Self::Result>,
) -> std::pin::Pin<
Box<
dyn std::future::Future<
Output = Result<NativeCallStep<Self::Operation>, 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<NativeCallStep<Self::Operation>, 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<OcrHostResult> {
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::<Option<Py<PyAny>>>()?
.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<Map<String, Value>> {
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<std::collections::BTreeMap<String, litellm_core::auth::InputSource>> {
let Some(proxy_request) = kwargs.get_item("proxy_server_request")? else {
return Ok(Default::default());
};
let proxy_request = proxy_request.cast_into::<PyDict>()?;
let body_fields = proxy_request
.get_item("body_fields")?
.or(proxy_request.get_item("body")?);
let credential_fields = proxy_request.get_item("credential_fields")?;
let 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<Value> {
if document.get_item("type")?.extract::<String>()? != "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()))
}

View file

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

View file

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

View file

@ -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<u8>),
}
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<String>)> {
if file.is_instance_of::<PyString>() {
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::<PyBytes>() {
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::<String>())
.transpose()?;
let value = reader.call0()?;
let bytes = if value.is_instance_of::<PyString>() {
FileBytes::Native(value.extract::<String>()?.into_bytes())
} else if value.is_instance_of::<PyBytes>() {
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<OcrDocument> {
let file = document.get_item("file").map_err(|error| {
if error.is_instance_of::<pyo3::exceptions::PyKeyError>(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::<PyDict>()?
.get_item("mime_type")?
.map(|value| value.extract::<String>())
.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<Py<PyAny>> {
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<Py<PyAny>> {
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)?)
}