This commit is contained in:
Yujong Lee 2026-09-12 08:21:05 -07:00
parent c6fc5e185f
commit a67e94162a
7 changed files with 753 additions and 216 deletions

View file

@ -14,8 +14,12 @@ use tokio::sync::Mutex;
use crate::errors::ocr_error_to_pyerr;
use crate::execution::{run_async_value, run_sync_value};
mod bindings;
mod preparation;
use bindings::DeploymentHooks;
pub(crate) use bindings::PythonLogger;
pub(crate) trait PythonRoute: Send + Sync {
fn state(&self) -> &PythonCallState;
fn state_mut(&mut self) -> &mut PythonCallState;
@ -401,7 +405,7 @@ impl<R: PythonRoute> Drop for PythonLifecycle<R> {
pub(crate) struct PythonCallState {
pub args: Py<PyTuple>,
pub kwargs: Py<PyDict>,
pub logger: Option<Py<PyAny>>,
pub logger: Option<PythonLogger>,
pub start: Py<PyAny>,
pub end: Option<Py<PyAny>>,
pub response: Option<Py<PyAny>>,
@ -427,32 +431,31 @@ impl PythonCallState {
match phase {
HostPhase::Setup => self.setup(py)?,
HostPhase::DeploymentPreCall => {
return Ok(HostStep::Suspend(
py.import("litellm.utils")?
.getattr("async_pre_call_deployment_hook")?
.call1((&self.kwargs, self.call_type))?
.unbind(),
));
return Ok(HostStep::Suspend(DeploymentHooks::before_call(
py,
&self.kwargs,
self.call_type,
)?));
}
HostPhase::Prepare => self.prepare(py)?,
HostPhase::DeploymentPostCall => {
return Ok(HostStep::Suspend(
py.import("litellm.utils")?
.getattr("async_post_call_success_deployment_hook")?
.call1((&self.kwargs, &self.response, self.call_type))?
.unbind(),
));
return Ok(HostStep::Suspend(DeploymentHooks::after_success(
py,
&self.kwargs,
&self.response,
self.call_type,
)?));
}
HostPhase::Finalize => self.finalize(py)?,
HostPhase::Success => self.dispatch_success(py)?,
HostPhase::DeploymentFailure => {
if let Some(error) = &self.error {
return Ok(HostStep::Suspend(
py.import("litellm.utils")?
.getattr("async_post_call_failure_deployment_hook")?
.call1((&self.kwargs, error, self.call_type))?
.unbind(),
));
return Ok(HostStep::Suspend(DeploymentHooks::after_failure(
py,
&self.kwargs,
error,
self.call_type,
)?));
}
}
HostPhase::Failure | HostPhase::AsyncFailure => {
@ -502,59 +505,48 @@ impl PythonCallState {
})
}
pub fn logger<'py>(&self, py: Python<'py>) -> PyResult<Bound<'py, PyAny>> {
self.logger
.as_ref()
.map(|value| value.bind(py).clone())
.ok_or_else(|| {
pyo3::exceptions::PyRuntimeError::new_err("call logging is not initialized")
})
pub fn logger(&self) -> PyResult<&PythonLogger> {
self.logger.as_ref().ok_or_else(|| {
pyo3::exceptions::PyRuntimeError::new_err("call logging is not initialized")
})
}
pub fn setup(&mut self, py: Python<'_>) -> PyResult<()> {
self.start = now(py)?;
self.internal = py
.import("litellm._internal_context")?
.getattr("is_internal_call")?
.call_method0("get")?
.extract()?;
let result = py
.import("litellm.rust_bridge.lifecycle")?
.getattr("setup")?
.call1((
self.call_type,
&self.args,
&self.kwargs,
&self.start,
self.asynchronous,
))?;
self.logger = Some(result.getattr("logger")?.unbind());
self.kwargs = result.getattr("kwargs")?.cast_into::<PyDict>()?.unbind();
self.internal = bindings::is_internal_call(py)?;
let result = bindings::setup(
py,
self.call_type,
&self.args,
&self.kwargs,
&self.start,
self.asynchronous,
)?;
self.logger = Some(result.logger()?);
self.kwargs = result.kwargs()?;
Ok(())
}
pub fn prepare(&mut self, py: Python<'_>) -> PyResult<()> {
self.kwargs = preparation::prepare(py, self.kwargs.bind(py), &self.logger(py)?)?.unbind();
self.kwargs = preparation::prepare(py, self.kwargs.bind(py), self.logger()?)?.unbind();
Ok(())
}
pub fn finalize(&self, py: Python<'_>) -> PyResult<()> {
py.import("litellm.rust_bridge.lifecycle")?
.getattr("finalize")?
.call1((
&self.response,
self.logger(py)?,
&self.kwargs,
&self.start,
&self.end,
))?;
Ok(())
bindings::finalize(
py,
&self.response,
self.logger()?,
&self.kwargs,
&self.start,
&self.end,
)
}
pub fn dispatch_success(&self, py: Python<'_>) -> PyResult<()> {
match self.try_dispatch_success(py) {
Err(error) if error.is_instance_of::<PyException>(py) => {
error.write_unraisable(py, self.logger.as_ref().map(|logger| logger.bind(py)));
error.write_unraisable(py, self.logger.as_ref().map(|logger| logger.object(py)));
Ok(())
}
result => result,
@ -562,9 +554,9 @@ impl PythonCallState {
}
fn try_dispatch_success(&self, py: Python<'_>) -> PyResult<()> {
let logger = self.logger(py)?;
let logger = self.logger()?;
let pending = PendingSuccess {
logger: logger.clone().unbind(),
logger: logger.clone_ref(py),
response: self.response.as_ref().map(|value| value.clone_ref(py)),
start: self.start.clone_ref(py),
end: self.end.as_ref().map(|value| value.clone_ref(py)),
@ -579,12 +571,9 @@ impl PythonCallState {
.get_item("fallbacks")?
.is_none_or(|value| value.is_none())
{
if logger
.getattr("_defer_async_logging")
.is_ok_and(|value| value.is_truthy().unwrap_or(false))
{
logger.setattr(
"_native_pending_logging",
if logger.defers_async_logging(py) {
logger.defer_success(
py,
Py::new(
py,
PendingLogging {
@ -596,12 +585,7 @@ impl PythonCallState {
pending.asynchronous(py)?;
}
}
logger
.call_method1(
"handle_sync_success_callbacks_for_async_calls",
(&self.response, &self.start, &self.end),
)
.map(|_| ())
logger.sync_success_for_async_call(py, &self.response, &self.start, &self.end)
}
}
@ -616,29 +600,13 @@ impl PythonCallState {
let Some(error) = &self.error else {
return Ok(None);
};
let trace = py
.import("traceback")?
.getattr("format_exception")?
.call1((error,))?;
let trace = pyo3::types::PyString::new(py, "").call_method1("join", (trace,))?;
let value = self.logger(py)?.call_method1(
if asynchronous {
"async_failure_handler"
} else {
"failure_handler"
},
(error, trace, &self.start, &self.end),
)?;
Ok(asynchronous.then(|| value.unbind()))
self.logger()?
.failure(py, error, &self.start, &self.end, asynchronous)
}
pub fn cleanup(&mut self, py: Python<'_>) {
if let Some(logger) = self.logger.take()
&& let Err(error) = py.import("litellm.utils").and_then(|utils| {
utils
.getattr("_restore_correlation_context_if_supported")?
.call1((logger,))
})
&& let Err(error) = logger.restore_context(py)
{
error.write_unraisable(py, None);
}
@ -651,7 +619,9 @@ impl PythonCallState {
pub fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> {
visit.call(&self.args)?;
visit.call(&self.kwargs)?;
visit.call(&self.logger)?;
if let Some(logger) = &self.logger {
logger.traverse(visit)?;
}
visit.call(&self.start)?;
visit.call(&self.end)?;
visit.call(&self.response)?;
@ -660,54 +630,21 @@ impl PythonCallState {
}
struct PendingSuccess {
logger: Py<PyAny>,
logger: PythonLogger,
response: Option<Py<PyAny>>,
start: Py<PyAny>,
end: Option<Py<PyAny>>,
}
impl PendingSuccess {
fn context(py: Python<'_>) -> PyResult<Py<PyAny>> {
py.import("contextvars")?
.call_method0("copy_context")
.map(Bound::unbind)
}
fn sync(&self, py: Python<'_>) -> PyResult<()> {
let context = Self::context(py)?;
py.import("litellm.litellm_core_utils.litellm_logging")?
.getattr("executor")?
.call_method1(
"submit",
(
context.getattr(py, "run")?,
self.logger.getattr(py, "success_handler")?,
&self.response,
&self.start,
&self.end,
),
)?;
Ok(())
self.logger
.submit_success(py, &self.response, &self.start, &self.end)
}
fn asynchronous(&self, py: Python<'_>) -> PyResult<()> {
let context = Self::context(py)?;
let worker = py
.import("litellm.litellm_core_utils.logging_worker")?
.getattr("GLOBAL_LOGGING_WORKER")?
.getattr("ensure_initialized_and_enqueue")?;
let coroutine = self.logger.call_method1(
py,
"async_success_handler",
(&self.response, &self.start, &self.end),
)?;
let enqueue = context.call_method1(py, "run", (worker, &coroutine));
if enqueue.is_err()
&& let Err(error) = coroutine.call_method0(py, "close")
{
error.write_unraisable(py, Some(coroutine.bind(py)));
}
enqueue.map(|_| ())
self.logger
.enqueue_success(py, &self.response, &self.start, &self.end)
}
}
@ -725,7 +662,7 @@ impl PendingLogging {
{
match pending.asynchronous(py) {
Err(error) if error.is_instance_of::<PyException>(py) => {
error.write_unraisable(py, Some(pending.logger.bind(py)));
error.write_unraisable(py, Some(pending.logger.object(py)));
}
result => return result,
}
@ -735,7 +672,7 @@ impl PendingLogging {
fn __traverse__(&self, visit: pyo3::gc::PyVisit<'_>) -> Result<(), pyo3::gc::PyTraverseError> {
if let Some(pending) = &self.pending {
visit.call(&pending.logger)?;
pending.logger.traverse(&visit)?;
visit.call(&pending.response)?;
visit.call(&pending.start)?;
visit.call(&pending.end)?;
@ -944,7 +881,7 @@ assert reference() is None
PythonCallState {
args: PyTuple::empty(py).unbind(),
kwargs: PyDict::new(py).unbind(),
logger: Some(logger),
logger: Some(logger.extract(py).unwrap()),
start: py.None(),
end: Some(py.None()),
response: Some(response),
@ -1079,7 +1016,12 @@ logger = Logger()
py,
PendingLogging {
pending: Some(PendingSuccess {
logger: locals.get_item("logger").unwrap().unwrap().unbind(),
logger: locals
.get_item("logger")
.unwrap()
.unwrap()
.extract()
.unwrap(),
response: Some(py.None()),
start: py.None(),
end: Some(py.None()),
@ -1104,6 +1046,54 @@ assert observed == ['created', 'release', 'closed']
});
}
#[test]
fn deferred_logging_collects_cycles_through_typed_logger() {
Python::initialize();
Python::attach(|py| {
let locals = PyDict::new(py);
py.run(
pyo3::ffi::c_str!("class Logger: pass\nlogger = Logger()"),
Some(&locals),
Some(&locals),
)
.unwrap();
let pending = Py::new(
py,
PendingLogging {
pending: Some(PendingSuccess {
logger: locals
.get_item("logger")
.unwrap()
.unwrap()
.extract()
.unwrap(),
response: None,
start: py.None(),
end: None,
}),
},
)
.unwrap();
locals.set_item("pending", pending).unwrap();
py.run(
pyo3::ffi::c_str!(
r#"
import gc
import weakref
logger.pending = pending
reference = weakref.ref(logger)
del logger, pending
gc.collect()
assert reference() is None
"#
),
Some(&locals),
Some(&locals),
)
.unwrap();
});
}
#[test]
fn coroutine_collects_cycles_retained_by_bridge_host() {
Python::initialize();

View file

@ -0,0 +1,334 @@
use pyo3::exceptions::PyBaseException;
use pyo3::gc::{PyTraverseError, PyVisit};
use pyo3::prelude::*;
use pyo3::types::{PyDict, PyTuple};
#[derive(FromPyObject)]
pub(crate) struct PythonLogger(Py<PyAny>);
impl PythonLogger {
pub(crate) fn object<'py>(&self, py: Python<'py>) -> &Bound<'py, PyAny> {
self.0.bind(py)
}
pub(crate) fn clone_ref(&self, py: Python<'_>) -> Self {
Self(self.0.clone_ref(py))
}
pub(crate) fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> {
visit.call(&self.0)
}
pub(super) fn defers_async_logging(&self, py: Python<'_>) -> bool {
self.object(py)
.getattr("_defer_async_logging")
.is_ok_and(|value| value.is_truthy().unwrap_or(false))
}
pub(super) fn defer_success(
&self,
py: Python<'_>,
pending: Py<super::PendingLogging>,
) -> PyResult<()> {
self.object(py).setattr("_native_pending_logging", pending)
}
pub(super) fn sync_success_for_async_call(
&self,
py: Python<'_>,
response: &Option<Py<PyAny>>,
start: &Py<PyAny>,
end: &Option<Py<PyAny>>,
) -> PyResult<()> {
self.object(py).call_method1(
"handle_sync_success_callbacks_for_async_calls",
(response, start, end),
)?;
Ok(())
}
pub(super) fn failure(
&self,
py: Python<'_>,
error: &Py<PyBaseException>,
start: &Py<PyAny>,
end: &Option<Py<PyAny>>,
asynchronous: bool,
) -> PyResult<Option<Py<PyAny>>> {
let trace = py
.import("traceback")?
.getattr("format_exception")?
.call1((error,))?;
let trace = pyo3::types::PyString::new(py, "").call_method1("join", (trace,))?;
let value = self.object(py).call_method1(
if asynchronous {
"async_failure_handler"
} else {
"failure_handler"
},
(error, trace, start, end),
)?;
Ok(asynchronous.then(|| value.unbind()))
}
pub(super) fn restore_context(&self, py: Python<'_>) -> PyResult<()> {
py.import("litellm.utils")?
.getattr("_restore_correlation_context_if_supported")?
.call1((self.object(py),))?;
Ok(())
}
pub(super) fn submit_success(
&self,
py: Python<'_>,
response: &Option<Py<PyAny>>,
start: &Py<PyAny>,
end: &Option<Py<PyAny>>,
) -> PyResult<()> {
let context = py.import("contextvars")?.call_method0("copy_context")?;
py.import("litellm.litellm_core_utils.litellm_logging")?
.getattr("executor")?
.call_method1(
"submit",
(
context.getattr("run")?,
self.object(py).getattr("success_handler")?,
response,
start,
end,
),
)?;
Ok(())
}
pub(super) fn enqueue_success(
&self,
py: Python<'_>,
response: &Option<Py<PyAny>>,
start: &Py<PyAny>,
end: &Option<Py<PyAny>>,
) -> PyResult<()> {
let context = py.import("contextvars")?.call_method0("copy_context")?;
let worker = py
.import("litellm.litellm_core_utils.logging_worker")?
.getattr("GLOBAL_LOGGING_WORKER")?
.getattr("ensure_initialized_and_enqueue")?;
let coroutine = self
.object(py)
.call_method1("async_success_handler", (response, start, end))?;
let enqueue = context.call_method1("run", (worker, &coroutine));
if enqueue.is_err()
&& let Err(error) = coroutine.call_method0("close")
{
error.write_unraisable(py, Some(&coroutine));
}
enqueue.map(|_| ())
}
}
pub(super) struct SetupResult<'py>(Bound<'py, PyAny>);
impl SetupResult<'_> {
pub(super) fn logger(&self) -> PyResult<PythonLogger> {
self.0.getattr("logger")?.extract()
}
pub(super) fn kwargs(&self) -> PyResult<Py<PyDict>> {
Ok(self.0.getattr("kwargs")?.extract()?)
}
}
pub(super) fn setup<'py>(
py: Python<'py>,
call_type: &str,
args: &Py<PyTuple>,
kwargs: &Py<PyDict>,
start: &Py<PyAny>,
asynchronous: bool,
) -> PyResult<SetupResult<'py>> {
py.import("litellm.rust_bridge.lifecycle")?
.getattr("setup")?
.call1((call_type, args, kwargs, start, asynchronous))
.map(SetupResult)
}
pub(super) fn finalize(
py: Python<'_>,
response: &Option<Py<PyAny>>,
logger: &PythonLogger,
kwargs: &Py<PyDict>,
start: &Py<PyAny>,
end: &Option<Py<PyAny>>,
) -> PyResult<()> {
py.import("litellm.rust_bridge.lifecycle")?
.getattr("finalize")?
.call1((response, logger.object(py), kwargs, start, end))?;
Ok(())
}
pub(super) fn is_internal_call(py: Python<'_>) -> PyResult<bool> {
py.import("litellm._internal_context")?
.getattr("is_internal_call")?
.call_method0("get")?
.extract()
}
pub(super) struct DeploymentHooks;
impl DeploymentHooks {
pub(super) fn before_call(
py: Python<'_>,
kwargs: &Py<PyDict>,
call_type: &str,
) -> PyResult<Py<PyAny>> {
py.import("litellm.utils")?
.getattr("async_pre_call_deployment_hook")?
.call1((kwargs, call_type))
.map(Bound::unbind)
}
pub(super) fn after_success(
py: Python<'_>,
kwargs: &Py<PyDict>,
response: &Option<Py<PyAny>>,
call_type: &str,
) -> PyResult<Py<PyAny>> {
py.import("litellm.utils")?
.getattr("async_post_call_success_deployment_hook")?
.call1((kwargs, response, call_type))
.map(Bound::unbind)
}
pub(super) fn after_failure(
py: Python<'_>,
kwargs: &Py<PyDict>,
error: &Py<PyBaseException>,
call_type: &str,
) -> PyResult<Py<PyAny>> {
py.import("litellm.utils")?
.getattr("async_post_call_failure_deployment_hook")?
.call1((kwargs, error, call_type))
.map(Bound::unbind)
}
}
#[cfg(test)]
mod tests {
use super::*;
use pyo3::exceptions::PyTypeError;
#[test]
fn setup_fields_are_checked_in_order_without_eager_logger_method_reads() {
Python::initialize();
Python::attach(|py| {
let locals = PyDict::new(py);
py.run(
pyo3::ffi::c_str!(
r#"
reads = []
class Logger:
def __getattribute__(self, name):
reads.append(name)
raise AssertionError('logger methods must remain lazy')
logger = Logger()
class Setup:
@property
def logger(self):
reads.append('logger')
return logger
@property
def kwargs(self):
reads.append('kwargs')
return []
result = Setup()
"#
),
Some(&locals),
Some(&locals),
)
.unwrap();
let result = SetupResult(locals.get_item("result").unwrap().unwrap());
let logger = result.logger().unwrap();
assert!(
logger
.object(py)
.is(locals.get_item("logger").unwrap().unwrap())
);
assert_eq!(
locals
.get_item("reads")
.unwrap()
.unwrap()
.extract::<Vec<String>>()
.unwrap(),
["logger"]
);
assert!(
result
.kwargs()
.unwrap_err()
.is_instance_of::<PyTypeError>(py)
);
assert_eq!(
locals
.get_item("reads")
.unwrap()
.unwrap()
.extract::<Vec<String>>()
.unwrap(),
["logger", "kwargs"]
);
});
}
#[test]
fn logger_resolves_each_callback_at_invocation_and_preserves_arguments() {
Python::initialize();
Python::attach(|py| {
let locals = PyDict::new(py);
py.run(
pyo3::ffi::c_str!(
r#"
calls = []
response, start, end = object(), object(), object()
class Logger:
@property
def handle_sync_success_callbacks_for_async_calls(self):
generation = len(calls)
def callback(*args):
assert args == (response, start, end)
calls.append(generation)
return callback
logger = Logger()
"#
),
Some(&locals),
Some(&locals),
)
.unwrap();
let logger: PythonLogger = locals
.get_item("logger")
.unwrap()
.unwrap()
.extract()
.unwrap();
let response = Some(locals.get_item("response").unwrap().unwrap().unbind());
let start = locals.get_item("start").unwrap().unwrap().unbind();
let end = Some(locals.get_item("end").unwrap().unwrap().unbind());
for _ in 0..2 {
logger
.sync_success_for_async_call(py, &response, &start, &end)
.unwrap();
}
assert_eq!(
locals
.get_item("calls")
.unwrap()
.unwrap()
.extract::<Vec<usize>>()
.unwrap(),
[0, 1]
);
});
}
}

View file

@ -5,10 +5,10 @@ use pyo3::types::{PyDict, PyList};
pub(super) fn prepare<'py>(
py: Python<'py>,
kwargs: &Bound<'py, PyDict>,
logger: &Bound<'py, PyAny>,
logger: &super::PythonLogger,
) -> PyResult<Bound<'py, PyDict>> {
let arguments = kwargs.copy()?;
arguments.set_item("litellm_logging_obj", logger)?;
arguments.set_item("litellm_logging_obj", logger.object(py))?;
let litellm = py.import("litellm")?;
inherit_credentials(py, &litellm, &arguments)?;
py.import("litellm.rust_bridge.lifecycle")?

View file

@ -10,6 +10,7 @@ mod audio_transcription;
mod chat_completions;
mod messages;
mod ocr;
mod ocr_callbacks;
mod ocr_document;
mod ocr_lifecycle;

View file

@ -0,0 +1,275 @@
use pyo3::exceptions::{PyBaseException, PyException, PyRuntimeError, PyTypeError};
use pyo3::gc::{PyTraverseError, PyVisit};
use pyo3::prelude::*;
use pyo3::types::{PyDict, PyString};
use serde_json::Value;
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;
use crate::lifecycle::PythonLogger;
pub(super) struct AzureAdTokenProvider(Py<PyAny>);
impl AzureAdTokenProvider {
pub(super) fn select(provider: Bound<'_, PyAny>) -> Option<Self> {
(provider.is_callable() && provider.is_truthy().unwrap_or(false))
.then(|| Self(provider.unbind()))
}
pub(super) fn acquire(&self, py: Python<'_>) -> PyResult<ResolvedCredential> {
let provider = self.0.bind(py);
if !provider.is_callable() {
return Err(PyTypeError::new_err(
"Azure AD token provider must be callable",
));
}
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 {}")
.call_method1("format", (token.get_type(),))?;
return Err(PyTypeError::new_err(message.unbind()));
}
Ok(token)
})()
.map_err(|error| {
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: {}")
.call_method1("format", (error.value(py),))
{
Ok(message) => {
let wrapped = PyRuntimeError::new_err(message.unbind());
wrapped.set_context(py, Some(error.clone_ref(py)));
wrapped.set_cause(py, Some(error));
wrapped
}
Err(format_error) => {
format_error.set_context(py, Some(error));
format_error
}
}
})?;
Ok(ResolvedCredential::AccessToken {
token: SecretValue::new(token.extract::<String>()?),
expires_on: None,
})
}
pub(super) fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> {
visit.call(&self.0)
}
}
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 super::*;
#[test]
fn token_callback_preserves_exception_identity_and_explicit_chaining() {
Python::initialize();
Python::attach(|py| {
let locals = PyDict::new(py);
py.run(
pyo3::ffi::c_str!(
r#"
class ProviderError(Exception):
def __format__(self, specification):
return 'unavailable'
ordinary = ProviderError('must use __format__')
type_error = TypeError('signature')
abort = KeyboardInterrupt('cancelled')
def provider(error):
def acquire():
raise error
return acquire
"#
),
Some(&locals),
Some(&locals),
)
.unwrap();
for name in ["ordinary", "type_error", "abort"] {
let original = locals.get_item(name).unwrap().unwrap();
let callback = locals
.get_item("provider")
.unwrap()
.unwrap()
.call1((&original,))
.unwrap();
let provider = AzureAdTokenProvider::select(callback).unwrap();
let error = provider.acquire(py).unwrap_err();
if name == "ordinary" {
assert!(error.is_instance_of::<PyRuntimeError>(py));
assert!(error.cause(py).unwrap().value(py).is(&original));
assert!(
error
.value(py)
.getattr("__context__")
.unwrap()
.is(&original)
);
assert_eq!(
error.value(py).str().unwrap().to_str().unwrap(),
"Failed to get Azure AD token: unavailable"
);
} else {
assert!(error.value(py).is(&original));
}
}
});
}
#[test]
fn invalid_token_type_formatting_preserves_python_failure_semantics() {
Python::initialize();
Python::attach(|py| {
let locals = PyDict::new(py);
py.run(
pyo3::ffi::c_str!(
r#"
failure = ValueError('formatting failed')
class TokenType(type):
def __format__(cls, specification):
raise failure
class Token(metaclass=TokenType):
pass
def provider():
return Token()
"#
),
Some(&locals),
Some(&locals),
)
.unwrap();
let provider =
AzureAdTokenProvider::select(locals.get_item("provider").unwrap().unwrap())
.unwrap();
let error = provider.acquire(py).unwrap_err();
assert!(error.is_instance_of::<PyRuntimeError>(py));
assert!(
error
.cause(py)
.unwrap()
.value(py)
.is(locals.get_item("failure").unwrap().unwrap())
);
});
}
#[test]
fn token_string_extraction_errors_are_not_wrapped_as_callback_failures() {
Python::initialize();
Python::attach(|py| {
let callback = py
.eval(pyo3::ffi::c_str!("lambda: '\\ud800'"), None, None)
.unwrap();
let provider = AzureAdTokenProvider::select(callback).unwrap();
let error = provider.acquire(py).unwrap_err();
assert!(error.is_instance_of::<pyo3::exceptions::PyUnicodeEncodeError>(py));
});
}
}

View file

@ -4,7 +4,7 @@ use std::sync::Arc;
use pyo3::prelude::*;
use pyo3::types::{PyDict, PyTuple};
use litellm_core::auth::{ResolvedCredential, SecretValue};
use litellm_core::auth::ResolvedCredential;
use litellm_core::ocr::hooks::{OcrDuringCallRequest, OcrPostCallRequest, OcrPreCallRequest};
use litellm_core::ocr::wire::{OcrWireRequest, consumed_optional_param_names, decode_request};
use litellm_core::ocr::{
@ -14,6 +14,7 @@ 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};
@ -23,7 +24,7 @@ struct PythonOcrHost {
pre_call: Option<OcrPreCallRequest>,
document: Option<Py<PyAny>>,
api_key: Option<Py<PyAny>>,
azure_ad_token_provider: Option<Py<PyAny>>,
azure_ad_token_provider: Option<AzureAdTokenProvider>,
provider: String,
retained_fields: Option<Py<PyDict>>,
body: Option<Py<PyDict>>,
@ -34,7 +35,7 @@ struct AdmittedOcrCall {
request: litellm_core::ocr::LiteLLMOcrRequest,
document: Py<PyAny>,
api_key: Py<PyAny>,
azure_ad_token_provider: Option<Py<PyAny>>,
azure_ad_token_provider: Option<AzureAdTokenProvider>,
provider: String,
}
@ -70,15 +71,7 @@ impl PythonOcrHost {
.azure_ad_token_provider
.as_ref()
.ok_or_else(missing_state)?;
let token: String = py
.import("litellm.rust_bridge.ocr_lifecycle")?
.getattr("call_azure_ad_token_provider")?
.call1((provider,))?
.extract()?;
Ok(ResolvedCredential::AccessToken {
token: SecretValue::new(token),
expires_on: None,
})
provider.acquire(py)
}
fn python_pre_call(
@ -108,35 +101,9 @@ impl PythonOcrHost {
}
self.body = Some(body.clone().unbind());
self.headers = Some(headers.clone().unbind());
let logger = self.state.logger(py)?;
let redact = py
.import("litellm.rust_bridge.ocr")?
.getattr("redact_logging_params")?;
let update = PyDict::new(py);
update.set_item("kwargs", redact.call1((&self.state.kwargs,))?)?;
update.set_item("model", &pre_call.model)?;
update.set_item(
"optional_params",
redact.call1((to_py(py, &pre_call.optional_params)?,))?,
)?;
let params = PyDict::new(py);
params.set_item(
"litellm_call_id",
self.state.kwargs.bind(py).get_item("litellm_call_id")?,
)?;
params.set_item("api_base", &request.url)?;
update.set_item("litellm_params", params)?;
update.set_item("custom_llm_provider", &pre_call.custom_llm_provider)?;
logger.call_method("update_from_kwargs", (), Some(&update))?;
let additional = PyDict::new(py);
additional.set_item("complete_input_dict", &body)?;
additional.set_item("headers", &headers)?;
additional.set_item("api_base", &request.url)?;
let kwargs = PyDict::new(py);
kwargs.set_item("input", "OCR document processing")?;
kwargs.set_item("api_key", &self.api_key)?;
kwargs.set_item("additional_args", additional)?;
logger.call_method("pre_call", (), Some(&kwargs))?;
let logger = self.state.logger()?;
logger.update_ocr(py, &self.state.kwargs, pre_call, &request.url)?;
logger.pre_ocr(py, &self.api_key, &body, &headers, &request.url)?;
let headers = headers
.iter()
.map(|(name, value)| Ok((name.extract::<String>()?, value.extract::<String>()?)))
@ -151,14 +118,9 @@ impl PythonOcrHost {
py: Python<'_>,
request: OcrPostCallRequest,
) -> PyResult<OcrPostCallRequest> {
let logger = self.state.logger(py)?;
let kwargs = PyDict::new(py);
kwargs.set_item("original_response", to_py(py, &request.original_response)?)?;
let additional = PyDict::new(py);
additional.set_item("complete_input_dict", &self.body)?;
additional.set_item("headers", &self.headers)?;
kwargs.set_item("additional_args", additional)?;
logger.call_method("post_call", (), Some(&kwargs))?;
self.state
.logger()?
.post_ocr(py, &request.original_response, &self.body, &self.headers)?;
Ok(request)
}
}
@ -203,12 +165,7 @@ impl PythonRoute for PythonOcrHost {
}
OcrHostOperation::ConstructResponse(response) => {
self.state.end = Some(now(py)?);
self.state.response = Some(
py.import("litellm.rust_bridge.ocr")?
.getattr("_response")?
.call1((to_py(py, response.as_ref())?,))?
.unbind(),
);
self.state.response = Some(ocr_callbacks::response(py, response.as_ref())?);
OcrHostResult::Lifecycle(Ok(()))
}
OcrHostOperation::MapFailure(error) => {
@ -220,11 +177,9 @@ 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 = py
.import("litellm.rust_bridge.ocr_lifecycle")?
.getattr("map_failure")?
.call1((error, request, &self.provider))?;
self.state.retain_error(py, PyErr::from_value(mapped));
let mapped = ocr_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(()))
}
OcrHostOperation::Lifecycle(_)
@ -247,7 +202,9 @@ impl PythonRoute for PythonOcrHost {
visit.call(&self.request)?;
visit.call(&self.document)?;
visit.call(&self.api_key)?;
visit.call(&self.azure_ad_token_provider)?;
if let Some(provider) = &self.azure_ad_token_provider {
provider.traverse(visit)?;
}
visit.call(&self.retained_fields)?;
visit.call(&self.body)?;
visit.call(&self.headers)
@ -278,8 +235,7 @@ fn project_request(
let input_sources = extract_input_sources(request_kwargs, &consumed)?;
let azure_ad_token_provider = request_kwargs
.get_item("azure_ad_token_provider")?
.filter(|provider| provider.is_callable() && provider.is_truthy().unwrap_or(false))
.map(Bound::unbind);
.and_then(AzureAdTokenProvider::select);
let wire = OcrWireRequest {
model,
document: wire_document,
@ -294,12 +250,7 @@ fn project_request(
input_sources,
timeout_seconds: argument("timeout")?
.extract::<Option<Py<PyAny>>>()?
.map(|value| {
py.import("litellm.rust_bridge.timeouts")?
.getattr("timeout_to_seconds")?
.call1((value,))?
.extract()
})
.map(|value| ocr_callbacks::timeout_seconds(py, value))
.transpose()?
.flatten(),
};

View file

@ -50,20 +50,6 @@ def arguments(request: LiteLLMOcrRequest) -> Mapping[str, object]:
return request.kwargs
def call_azure_ad_token_provider(provider: object) -> str:
if not callable(provider):
raise TypeError("Azure AD token provider must be callable")
try:
token: Final = provider()
if not isinstance(token, str):
raise TypeError(f"Azure AD token must be a string, got {type(token)}")
return token
except TypeError:
raise
except Exception as error:
raise RuntimeError(f"Failed to get Azure AD token: {error}") from error
def map_failure(error: Exception, request: LiteLLMOcrRequest, request_provider: str) -> Exception:
mapper: Final = cast( # cast-ok: bounded adapter for the legacy public exception mapper
ExceptionMapper, litellm.exception_type