mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
refactor
This commit is contained in:
parent
c6fc5e185f
commit
a67e94162a
7 changed files with 753 additions and 216 deletions
|
|
@ -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();
|
||||
|
|
|
|||
334
litellm-rust/crates/python-bridge/src/lifecycle/bindings.rs
Normal file
334
litellm-rust/crates/python-bridge/src/lifecycle/bindings.rs
Normal 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]
|
||||
);
|
||||
});
|
||||
}
|
||||
}
|
||||
|
|
@ -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")?
|
||||
|
|
|
|||
|
|
@ -10,6 +10,7 @@ mod audio_transcription;
|
|||
mod chat_completions;
|
||||
mod messages;
|
||||
mod ocr;
|
||||
mod ocr_callbacks;
|
||||
mod ocr_document;
|
||||
mod ocr_lifecycle;
|
||||
|
||||
|
|
|
|||
275
litellm-rust/crates/python-bridge/src/routes/ocr_callbacks.rs
Normal file
275
litellm-rust/crates/python-bridge/src/routes/ocr_callbacks.rs
Normal 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));
|
||||
});
|
||||
}
|
||||
}
|
||||
|
|
@ -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(),
|
||||
};
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue