mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
wip
This commit is contained in:
parent
a67e94162a
commit
fe8759df7b
19 changed files with 1001 additions and 621 deletions
|
|
@ -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),
|
||||
|
|
|
|||
|
|
@ -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};
|
||||
|
|
|
|||
|
|
@ -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)]
|
||||
|
|
|
|||
|
|
@ -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)]
|
||||
|
|
|
|||
|
|
@ -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));
|
||||
});
|
||||
|
|
@ -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
|
||||
);
|
||||
});
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,3 +1,4 @@
|
|||
mod auth;
|
||||
mod constants;
|
||||
mod diagnostics;
|
||||
mod errors;
|
||||
|
|
|
|||
139
litellm-rust/crates/python-bridge/src/lifecycle/handle.rs
Normal file
139
litellm-rust/crates/python-bridge/src/lifecycle/handle.rs
Normal 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);
|
||||
}
|
||||
}
|
||||
|
|
@ -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();
|
||||
});
|
||||
}
|
||||
|
|
@ -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};
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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)?;
|
||||
|
|
|
|||
102
litellm-rust/crates/python-bridge/src/routes/ocr/callbacks.rs
Normal file
102
litellm-rust/crates/python-bridge/src/routes/ocr/callbacks.rs
Normal 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()?)
|
||||
}
|
||||
259
litellm-rust/crates/python-bridge/src/routes/ocr/document.rs
Normal file
259
litellm-rust/crates/python-bridge/src/routes/ocr/document.rs
Normal 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())
|
||||
);
|
||||
});
|
||||
}
|
||||
}
|
||||
77
litellm-rust/crates/python-bridge/src/routes/ocr/errors.rs
Normal file
77
litellm-rust/crates/python-bridge/src/routes/ocr/errors.rs
Normal 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
|
||||
);
|
||||
});
|
||||
}
|
||||
}
|
||||
|
|
@ -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()))
|
||||
}
|
||||
|
||||
18
litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs
Normal file
18
litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs
Normal 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)
|
||||
}
|
||||
|
|
@ -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(
|
||||
|
|
@ -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)?)
|
||||
}
|
||||
Loading…
Add table
Reference in a new issue