From 60d046ccc43c4a2989e72763629a906145fd68e2 Mon Sep 17 00:00:00 2001 From: Yujong Lee Date: Wed, 16 Sep 2026 19:07:43 -0700 Subject: [PATCH] Revert "refactor callback" This reverts commit 1771e5bb68e56ed4a9afafd66b9986eb422a4add. --- .../core/src/call_lifecycle/callbacks.rs | 747 ------------------ .../crates/core/src/call_lifecycle/host.rs | 8 +- .../crates/core/src/call_lifecycle/mod.rs | 6 - .../document_intelligence/transformation.rs | 3 +- .../src/llms/cohere/ocr/transformation.rs | 16 +- .../src/llms/reducto/ocr/transformation.rs | 14 +- .../src/llms/vertex_ai/ocr/transformation.rs | 8 +- litellm-rust/crates/core/src/ocr/client.rs | 12 +- litellm-rust/crates/core/src/ocr/document.rs | 27 +- litellm-rust/crates/core/src/ocr/error.rs | 2 - litellm-rust/crates/core/src/ocr/handler.rs | 2 +- litellm-rust/crates/core/src/ocr/hooks.rs | 1 - litellm-rust/crates/core/src/ocr/lifecycle.rs | 270 ++----- litellm-rust/crates/core/src/ocr/mod.rs | 5 +- litellm-rust/crates/core/src/ocr/prepare.rs | 4 +- litellm-rust/crates/core/src/ocr/types.rs | 9 +- .../crates/python-bridge/src/lifecycle/mod.rs | 37 +- .../crates/python-bridge/src/marshal.rs | 5 - .../python-bridge/src/routes/ocr/callbacks.rs | 173 +++- .../python-bridge/src/routes/ocr/document.rs | 63 +- .../python-bridge/src/routes/ocr/errors.rs | 10 +- .../python-bridge/src/routes/ocr/host.rs | 209 ++--- .../python-bridge/src/routes/ocr/mod.rs | 2 +- .../python-bridge/src/routes/ocr/project.rs | 90 ++- litellm/proxy/ocr_endpoints/endpoints.py | 10 +- litellm/rust_bridge/bindings.py | 5 - litellm/rust_bridge/ocr.py | 82 +- .../test_litellm/rust_bridge/test_bindings.py | 18 - .../rust_bridge/test_ocr_lifecycle.py | 68 +- tests/test_litellm_rust/test_ocr.py | 110 +-- 30 files changed, 465 insertions(+), 1551 deletions(-) delete mode 100644 litellm-rust/crates/core/src/call_lifecycle/callbacks.rs diff --git a/litellm-rust/crates/core/src/call_lifecycle/callbacks.rs b/litellm-rust/crates/core/src/call_lifecycle/callbacks.rs deleted file mode 100644 index 82d33898bd0..00000000000 --- a/litellm-rust/crates/core/src/call_lifecycle/callbacks.rs +++ /dev/null @@ -1,747 +0,0 @@ -use super::host::HostPhase; - -#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash, PartialOrd, Ord)] -pub struct CallbackId(pub u64); - -#[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub enum Delivery { - Inline, - Await, - Worker, - Background, -} - -#[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub enum CallbackKind { - CustomLogger, - Callable { internal: bool }, - Named { known: bool }, - Opaque, -} - -impl CallbackKind { - fn runs_sync_handler_for_async_call(self) -> bool { - match self { - Self::Callable { internal } => !internal, - Self::Named { known } => !known, - Self::CustomLogger | Self::Opaque => false, - } - } -} - -#[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub enum CallbackMethod { - LogPreApiCall, - LogPostApiCall, - PreCallDeploymentHook, - PostCallSuccessDeploymentHook, - PostCallFailureDeploymentHook, - LoggingHook, - AsyncLoggingHook, - LogSuccessEvent, - AsyncLogSuccessEvent, - LogFailureEvent, - AsyncLogFailureEvent, -} - -#[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub struct CallbackInvocation { - pub target: CallbackId, - pub method: CallbackMethod, - pub delivery: Delivery, -} - -#[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub enum TargetErrorPolicy { - Contain, - Propagate, -} - -#[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub enum LoggedMarker { - SyncSuccess, - AsyncSuccess, - SyncFailure, - AsyncFailure, -} - -impl LoggedMarker { - pub const fn key(self) -> &'static str { - match self { - Self::SyncSuccess => "has_logged_sync_success", - Self::AsyncSuccess => "has_logged_async_success", - Self::SyncFailure => "has_logged_sync_failure", - Self::AsyncFailure => "has_logged_async_failure", - } - } -} - -#[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub enum CallbackFamily { - RequestPreCall, - RequestPostCall, - DeploymentPreCall, - DeploymentPostCall, - DeploymentFailure, - SyncSuccess, - AsyncSuccess, - SyncFailure, - AsyncFailure, -} - -impl CallbackFamily { - pub const fn delivery(self) -> Delivery { - match self { - Self::RequestPreCall | Self::RequestPostCall | Self::SyncFailure => Delivery::Inline, - Self::DeploymentPreCall - | Self::DeploymentPostCall - | Self::DeploymentFailure - | Self::AsyncFailure => Delivery::Await, - Self::SyncSuccess => Delivery::Worker, - Self::AsyncSuccess => Delivery::Background, - } - } - - pub const fn dispatch_method(self) -> CallbackMethod { - match self { - Self::RequestPreCall => CallbackMethod::LogPreApiCall, - Self::RequestPostCall => CallbackMethod::LogPostApiCall, - Self::DeploymentPreCall => CallbackMethod::PreCallDeploymentHook, - Self::DeploymentPostCall => CallbackMethod::PostCallSuccessDeploymentHook, - Self::DeploymentFailure => CallbackMethod::PostCallFailureDeploymentHook, - Self::SyncSuccess => CallbackMethod::LogSuccessEvent, - Self::AsyncSuccess => CallbackMethod::AsyncLogSuccessEvent, - Self::SyncFailure => CallbackMethod::LogFailureEvent, - Self::AsyncFailure => CallbackMethod::AsyncLogFailureEvent, - } - } - - pub const fn hook_method(self) -> Option { - match self { - Self::SyncSuccess => Some(CallbackMethod::LoggingHook), - Self::AsyncSuccess => Some(CallbackMethod::AsyncLoggingHook), - _ => None, - } - } - - pub const fn marker(self) -> Option { - match self { - Self::SyncSuccess => Some(LoggedMarker::SyncSuccess), - Self::AsyncSuccess => Some(LoggedMarker::AsyncSuccess), - Self::SyncFailure => Some(LoggedMarker::SyncFailure), - Self::AsyncFailure => Some(LoggedMarker::AsyncFailure), - _ => None, - } - } - - pub const fn prepares_logging(self) -> bool { - self.marker().is_some() - } - - pub const fn error_policy(self) -> TargetErrorPolicy { - match self { - Self::DeploymentPreCall | Self::DeploymentPostCall => TargetErrorPolicy::Propagate, - _ => TargetErrorPolicy::Contain, - } - } - - pub fn targets(self, global: &[CallbackId], dynamic: Option<&[CallbackId]>) -> Vec { - match self { - Self::RequestPreCall | Self::RequestPostCall => global - .iter() - .chain(dynamic.unwrap_or_default()) - .copied() - .collect(), - Self::DeploymentPreCall | Self::DeploymentPostCall | Self::DeploymentFailure => { - global.to_vec() - } - Self::SyncSuccess | Self::AsyncSuccess | Self::SyncFailure | Self::AsyncFailure => { - let Some(dynamic) = dynamic else { - return global.to_vec(); - }; - let mut seen = std::collections::HashSet::new(); - dynamic - .iter() - .chain(global) - .copied() - .filter(|id| seen.insert(*id)) - .collect() - } - } - } -} - -#[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub enum ReleaseGate { - Immediate, - Deferred, -} - -#[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub struct SuccessDispatch { - pub family: CallbackFamily, - pub delivery: Delivery, - pub gate: ReleaseGate, -} - -#[derive(Clone, Debug, PartialEq, Eq)] -pub struct SuccessFacts { - pub asynchronous: bool, - pub internal: bool, - pub fallbacks: bool, - pub deferred: bool, - pub sync_target_kinds: Vec, -} - -pub fn plan_success(facts: &SuccessFacts) -> Vec { - if !facts.asynchronous { - return vec![SuccessDispatch { - family: CallbackFamily::SyncSuccess, - delivery: Delivery::Worker, - gate: ReleaseGate::Immediate, - }]; - } - let background = (!facts.internal && !facts.fallbacks).then_some(SuccessDispatch { - family: CallbackFamily::AsyncSuccess, - delivery: Delivery::Background, - gate: if facts.deferred { - ReleaseGate::Deferred - } else { - ReleaseGate::Immediate - }, - }); - let worker = facts - .sync_target_kinds - .iter() - .any(|kind| kind.runs_sync_handler_for_async_call()) - .then_some(SuccessDispatch { - family: CallbackFamily::SyncSuccess, - delivery: Delivery::Worker, - gate: ReleaseGate::Immediate, - }); - background.into_iter().chain(worker).collect() -} - -pub fn plan_failure( - phase: HostPhase, - asynchronous: bool, - internal: bool, -) -> Option { - if asynchronous && internal { - return None; - } - match phase { - HostPhase::Failure => Some(CallbackFamily::SyncFailure), - HostPhase::AsyncFailure => Some(CallbackFamily::AsyncFailure), - _ => None, - } -} - -pub trait DispatchFacts { - fn eligible(&mut self, target: CallbackId, method: CallbackMethod) -> bool; -} - -#[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub enum DispatchStep { - PrepareLogging, - Invoke(CallbackInvocation), - MarkLogged(LoggedMarker), - Complete { aborted: bool }, -} - -#[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub enum InvocationOutcome { - Completed, - Failed, -} - -#[derive(Clone, Copy, Debug, PartialEq, Eq)] -enum Position { - Prepare, - Hook(usize), - Mark, - Dispatch(usize), - Complete { aborted: bool }, -} - -pub struct DispatchCursor { - family: CallbackFamily, - targets: Vec, - stream: bool, - position: Position, -} - -impl DispatchCursor { - pub fn start( - family: CallbackFamily, - targets: Vec, - already_logged: bool, - stream: bool, - ) -> Self { - let position = if family.marker().is_some() && already_logged { - Position::Complete { aborted: false } - } else if family.prepares_logging() { - Position::Prepare - } else { - Position::Dispatch(0) - }; - Self { - family, - targets, - stream, - position, - } - } - - pub const fn family(&self) -> CallbackFamily { - self.family - } - - pub fn targets(&self) -> &[CallbackId] { - &self.targets - } - - pub fn accept(&mut self, outcome: InvocationOutcome) { - if outcome == InvocationOutcome::Failed - && self.family.error_policy() == TargetErrorPolicy::Propagate - { - self.position = Position::Complete { aborted: true }; - } - } - - pub fn next(&mut self, facts: &mut dyn DispatchFacts) -> DispatchStep { - loop { - match self.position { - Position::Prepare => { - self.position = self.after_prepare(); - return DispatchStep::PrepareLogging; - } - Position::Hook(index) => { - let Some(method) = self.family.hook_method() else { - self.position = Position::Mark; - continue; - }; - let Some(target) = self.targets.get(index).copied() else { - self.position = Position::Mark; - continue; - }; - self.position = Position::Hook(index + 1); - if facts.eligible(target, method) { - return DispatchStep::Invoke(self.invocation(target, method)); - } - } - Position::Mark => { - self.position = Position::Dispatch(0); - match self.family.marker() { - Some(marker) if !self.stream => return DispatchStep::MarkLogged(marker), - _ => continue, - } - } - Position::Dispatch(index) => { - let Some(target) = self.targets.get(index).copied() else { - self.position = Position::Complete { aborted: false }; - continue; - }; - self.position = Position::Dispatch(index + 1); - let method = self.family.dispatch_method(); - if facts.eligible(target, method) { - return DispatchStep::Invoke(self.invocation(target, method)); - } - } - Position::Complete { aborted } => return DispatchStep::Complete { aborted }, - } - } - } - - fn after_prepare(&self) -> Position { - if self.family.hook_method().is_some() { - Position::Hook(0) - } else { - Position::Mark - } - } - - const fn invocation(&self, target: CallbackId, method: CallbackMethod) -> CallbackInvocation { - CallbackInvocation { - target, - method, - delivery: self.family.delivery(), - } - } -} - -#[cfg(test)] -mod tests { - use super::*; - - struct AllEligible; - - impl DispatchFacts for AllEligible { - fn eligible(&mut self, _: CallbackId, _: CallbackMethod) -> bool { - true - } - } - - struct Gate bool>(F); - - impl bool> DispatchFacts for Gate { - fn eligible(&mut self, target: CallbackId, method: CallbackMethod) -> bool { - (self.0)(target, method) - } - } - - fn drain(cursor: &mut DispatchCursor, facts: &mut dyn DispatchFacts) -> Vec { - let mut steps = Vec::new(); - loop { - let step = cursor.next(facts); - steps.push(step); - match step { - DispatchStep::Complete { .. } => return steps, - DispatchStep::Invoke(_) => cursor.accept(InvocationOutcome::Completed), - _ => {} - } - } - } - - fn ids(values: &[u64]) -> Vec { - values.iter().copied().map(CallbackId).collect() - } - - #[test] - fn terminal_families_order_dynamic_before_global_and_keep_first_duplicate() { - let combined = CallbackFamily::SyncSuccess.targets(&ids(&[3, 1, 4]), Some(&ids(&[1, 2]))); - assert_eq!(combined, ids(&[1, 2, 3, 4])); - } - - #[test] - fn terminal_families_copy_global_without_dedup_when_dynamic_is_absent() { - let combined = CallbackFamily::AsyncFailure.targets(&ids(&[3, 3, 1]), None); - assert_eq!(combined, ids(&[3, 3, 1])); - } - - #[test] - fn request_families_order_global_before_dynamic_without_dedup() { - let combined = CallbackFamily::RequestPreCall.targets(&ids(&[1, 2]), Some(&ids(&[2, 3]))); - assert_eq!(combined, ids(&[1, 2, 2, 3])); - } - - #[test] - fn success_runs_every_hook_before_any_dispatch_and_marks_between_passes() { - let mut cursor = - DispatchCursor::start(CallbackFamily::SyncSuccess, ids(&[1, 2]), false, false); - let steps = drain(&mut cursor, &mut AllEligible); - let invocation = |target, method| { - DispatchStep::Invoke(CallbackInvocation { - target: CallbackId(target), - method, - delivery: Delivery::Worker, - }) - }; - assert_eq!( - steps, - vec![ - DispatchStep::PrepareLogging, - invocation(1, CallbackMethod::LoggingHook), - invocation(2, CallbackMethod::LoggingHook), - DispatchStep::MarkLogged(LoggedMarker::SyncSuccess), - invocation(1, CallbackMethod::LogSuccessEvent), - invocation(2, CallbackMethod::LogSuccessEvent), - DispatchStep::Complete { aborted: false }, - ] - ); - } - - #[test] - fn async_success_uses_async_leaf_methods_and_background_delivery() { - let mut cursor = - DispatchCursor::start(CallbackFamily::AsyncSuccess, ids(&[7]), false, false); - let steps = drain(&mut cursor, &mut AllEligible); - let methods: Vec<_> = steps - .iter() - .filter_map(|step| match step { - DispatchStep::Invoke(invocation) => { - assert_eq!(invocation.delivery, Delivery::Background); - Some(invocation.method) - } - _ => None, - }) - .collect(); - assert_eq!( - methods, - [ - CallbackMethod::AsyncLoggingHook, - CallbackMethod::AsyncLogSuccessEvent - ] - ); - } - - #[test] - fn failure_and_request_families_have_no_hook_pass() { - for (family, delivery, method) in [ - ( - CallbackFamily::SyncFailure, - Delivery::Inline, - CallbackMethod::LogFailureEvent, - ), - ( - CallbackFamily::AsyncFailure, - Delivery::Await, - CallbackMethod::AsyncLogFailureEvent, - ), - ] { - let mut cursor = DispatchCursor::start(family, ids(&[1, 2]), false, false); - let steps = drain(&mut cursor, &mut AllEligible); - assert_eq!(steps[0], DispatchStep::PrepareLogging); - assert!(matches!(steps[1], DispatchStep::MarkLogged(_))); - assert_eq!( - &steps[2..], - &[ - DispatchStep::Invoke(CallbackInvocation { - target: CallbackId(1), - method, - delivery - }), - DispatchStep::Invoke(CallbackInvocation { - target: CallbackId(2), - method, - delivery - }), - DispatchStep::Complete { aborted: false }, - ] - ); - } - let mut cursor = - DispatchCursor::start(CallbackFamily::RequestPreCall, ids(&[1]), false, false); - let steps = drain(&mut cursor, &mut AllEligible); - assert_eq!( - steps, - vec![ - DispatchStep::Invoke(CallbackInvocation { - target: CallbackId(1), - method: CallbackMethod::LogPreApiCall, - delivery: Delivery::Inline, - }), - DispatchStep::Complete { aborted: false }, - ] - ); - } - - #[test] - fn already_logged_marker_skips_the_whole_terminal_family_but_not_request_families() { - let mut cursor = - DispatchCursor::start(CallbackFamily::AsyncSuccess, ids(&[1]), true, false); - assert_eq!( - cursor.next(&mut AllEligible), - DispatchStep::Complete { aborted: false } - ); - let mut cursor = - DispatchCursor::start(CallbackFamily::RequestPostCall, ids(&[1]), true, false); - assert!(matches!( - cursor.next(&mut AllEligible), - DispatchStep::Invoke(_) - )); - } - - #[test] - fn streaming_skips_the_marker_write_but_still_dispatches() { - let mut cursor = DispatchCursor::start(CallbackFamily::SyncSuccess, ids(&[1]), false, true); - let steps = drain(&mut cursor, &mut AllEligible); - assert!( - !steps - .iter() - .any(|step| matches!(step, DispatchStep::MarkLogged(_))) - ); - assert_eq!( - steps - .iter() - .filter(|step| matches!(step, DispatchStep::Invoke(_))) - .count(), - 2 - ); - } - - #[test] - fn ineligible_targets_are_skipped_per_method_without_affecting_others() { - let mut cursor = - DispatchCursor::start(CallbackFamily::SyncSuccess, ids(&[1, 2]), false, false); - let mut facts = Gate(|target, method| { - !(target == CallbackId(1) && method == CallbackMethod::LoggingHook) - && !(target == CallbackId(2) && method == CallbackMethod::LogSuccessEvent) - }); - let invoked: Vec<_> = drain(&mut cursor, &mut facts) - .into_iter() - .filter_map(|step| match step { - DispatchStep::Invoke(invocation) => Some((invocation.target.0, invocation.method)), - _ => None, - }) - .collect(); - assert_eq!( - invoked, - [ - (2, CallbackMethod::LoggingHook), - (1, CallbackMethod::LogSuccessEvent) - ] - ); - } - - #[test] - fn contained_failures_continue_and_propagating_failures_abort() { - let mut cursor = - DispatchCursor::start(CallbackFamily::SyncFailure, ids(&[1, 2]), false, false); - assert_eq!(cursor.next(&mut AllEligible), DispatchStep::PrepareLogging); - assert!(matches!( - cursor.next(&mut AllEligible), - DispatchStep::MarkLogged(_) - )); - assert!(matches!( - cursor.next(&mut AllEligible), - DispatchStep::Invoke(_) - )); - cursor.accept(InvocationOutcome::Failed); - assert!(matches!( - cursor.next(&mut AllEligible), - DispatchStep::Invoke(CallbackInvocation { - target: CallbackId(2), - .. - }) - )); - - let mut cursor = DispatchCursor::start( - CallbackFamily::DeploymentPreCall, - ids(&[1, 2]), - false, - false, - ); - assert!(matches!( - cursor.next(&mut AllEligible), - DispatchStep::Invoke(_) - )); - cursor.accept(InvocationOutcome::Failed); - assert_eq!( - cursor.next(&mut AllEligible), - DispatchStep::Complete { aborted: true } - ); - } - - #[test] - fn sync_sdk_success_selects_one_worker_dispatch() { - let plan = plan_success(&SuccessFacts { - asynchronous: false, - internal: true, - fallbacks: true, - deferred: true, - sync_target_kinds: vec![], - }); - assert_eq!( - plan, - [SuccessDispatch { - family: CallbackFamily::SyncSuccess, - delivery: Delivery::Worker, - gate: ReleaseGate::Immediate, - }] - ); - } - - #[test] - fn async_sdk_success_enqueues_background_then_worker_only_for_external_sync_targets() { - let base = SuccessFacts { - asynchronous: true, - internal: false, - fallbacks: false, - deferred: false, - sync_target_kinds: vec![ - CallbackKind::CustomLogger, - CallbackKind::Named { known: true }, - ], - }; - assert_eq!( - plan_success(&base), - [SuccessDispatch { - family: CallbackFamily::AsyncSuccess, - delivery: Delivery::Background, - gate: ReleaseGate::Immediate, - }] - ); - let with_external = SuccessFacts { - sync_target_kinds: vec![ - CallbackKind::CustomLogger, - CallbackKind::Callable { internal: false }, - ], - deferred: true, - ..base.clone() - }; - assert_eq!( - plan_success(&with_external), - [ - SuccessDispatch { - family: CallbackFamily::AsyncSuccess, - delivery: Delivery::Background, - gate: ReleaseGate::Deferred, - }, - SuccessDispatch { - family: CallbackFamily::SyncSuccess, - delivery: Delivery::Worker, - gate: ReleaseGate::Immediate, - }, - ] - ); - let internal_or_fallback = SuccessFacts { - internal: true, - sync_target_kinds: vec![CallbackKind::Named { known: false }], - ..base - }; - assert_eq!( - plan_success(&internal_or_fallback), - [SuccessDispatch { - family: CallbackFamily::SyncSuccess, - delivery: Delivery::Worker, - gate: ReleaseGate::Immediate, - }] - ); - } - - #[test] - fn opaque_and_internal_targets_never_trigger_the_worker_for_async_calls() { - let plan = plan_success(&SuccessFacts { - asynchronous: true, - internal: false, - fallbacks: true, - deferred: false, - sync_target_kinds: vec![ - CallbackKind::Opaque, - CallbackKind::Callable { internal: true }, - ], - }); - assert!(plan.is_empty()); - } - - #[test] - fn failure_families_follow_the_phase_and_skip_internal_async_calls() { - assert_eq!( - plan_failure(HostPhase::Failure, false, true), - Some(CallbackFamily::SyncFailure) - ); - assert_eq!( - plan_failure(HostPhase::AsyncFailure, true, false), - Some(CallbackFamily::AsyncFailure) - ); - assert_eq!(plan_failure(HostPhase::Failure, true, true), None); - assert_eq!(plan_failure(HostPhase::Success, false, false), None); - } - - #[test] - fn delivery_is_a_property_of_the_family_not_of_the_callable() { - assert_eq!(CallbackFamily::RequestPreCall.delivery(), Delivery::Inline); - assert_eq!( - CallbackFamily::DeploymentPreCall.delivery(), - Delivery::Await - ); - assert_eq!(CallbackFamily::SyncSuccess.delivery(), Delivery::Worker); - assert_eq!( - CallbackFamily::AsyncSuccess.delivery(), - Delivery::Background - ); - assert_eq!(CallbackFamily::SyncFailure.delivery(), Delivery::Inline); - assert_eq!(CallbackFamily::AsyncFailure.delivery(), Delivery::Await); - } -} diff --git a/litellm-rust/crates/core/src/call_lifecycle/host.rs b/litellm-rust/crates/core/src/call_lifecycle/host.rs index 0e93b95b947..f9e18b9f116 100644 --- a/litellm-rust/crates/core/src/call_lifecycle/host.rs +++ b/litellm-rust/crates/core/src/call_lifecycle/host.rs @@ -223,11 +223,9 @@ mod tests { ] { assert_eq!(lifecycle.phase(), phase); assert!( - lifecycle - .accept(Err(HostFailure::Error(crate::ocr::Error::InvalidRequest( - "callback".into() - )))) - .is_none() + lifecycle.accept(Err(HostFailure::Error(crate::ocr::Error::InvalidRequest( + "callback".into() + )))).is_none() ); } assert_eq!(lifecycle.phase(), HostPhase::Complete); diff --git a/litellm-rust/crates/core/src/call_lifecycle/mod.rs b/litellm-rust/crates/core/src/call_lifecycle/mod.rs index 992b29ad63e..fcd9cbd2ab8 100644 --- a/litellm-rust/crates/core/src/call_lifecycle/mod.rs +++ b/litellm-rust/crates/core/src/call_lifecycle/mod.rs @@ -1,15 +1,9 @@ use std::future::Future; use std::time::{Instant, SystemTime, UNIX_EPOCH}; -pub mod callbacks; pub mod host; pub mod types; -pub use callbacks::{ - CallbackFamily, CallbackId, CallbackInvocation, CallbackKind, CallbackMethod, Delivery, - DispatchCursor, DispatchFacts, DispatchStep, InvocationOutcome, LoggedMarker, ReleaseGate, - SuccessDispatch, SuccessFacts, TargetErrorPolicy, plan_failure, plan_success, -}; pub use types::{ CallLifecycleContext, CallLifecyclePhase, CallLifecyclePhaseTiming, CallLifecycleRequest, CallLifecycleTiming, diff --git a/litellm-rust/crates/core/src/llms/azure_ai/ocr/document_intelligence/transformation.rs b/litellm-rust/crates/core/src/llms/azure_ai/ocr/document_intelligence/transformation.rs index 9529a726ac3..50d5b703152 100644 --- a/litellm-rust/crates/core/src/llms/azure_ai/ocr/document_intelligence/transformation.rs +++ b/litellm-rust/crates/core/src/llms/azure_ai/ocr/document_intelligence/transformation.rs @@ -816,8 +816,7 @@ mod tests { "type":"document_url", "document_url":"https://example.com/document.pdf" })) - .unwrap() - .into(); + .unwrap().into(); perform_ocr(request).await.unwrap(); server.await.unwrap(); diff --git a/litellm-rust/crates/core/src/llms/cohere/ocr/transformation.rs b/litellm-rust/crates/core/src/llms/cohere/ocr/transformation.rs index fc11f62833c..fb16445e11d 100644 --- a/litellm-rust/crates/core/src/llms/cohere/ocr/transformation.rs +++ b/litellm-rust/crates/core/src/llms/cohere/ocr/transformation.rs @@ -361,12 +361,10 @@ mod tests { } }), ); - let request = request.with_document( - serde_json::from_value(json!({ - "type":"image_url","image_url":"https://example.com/original.png" - })) - .unwrap(), - ); + let request = request.with_document(serde_json::from_value(json!({ + "type":"image_url","image_url":"https://example.com/original.png" + })) + .unwrap()); let request = crate::ocr::prepare::prepare_request(request); let http = CohereParseConfig .prepare_request(&request, &crate::ocr::test_support::ocr_client()) @@ -505,12 +503,10 @@ mod tests { "https://example.com", json!({"output_format":null,"req_format":null}), ); - let request = request.with_document( - serde_json::from_value( + let request = request.with_document(serde_json::from_value( json!({"type":"image_url","image_url":"https://example.com/a.png"}), ) - .unwrap(), - ); + .unwrap()); assert_eq!( request.response_format().unwrap(), crate::ocr::types::OcrResponseFormat::Litellm diff --git a/litellm-rust/crates/core/src/llms/reducto/ocr/transformation.rs b/litellm-rust/crates/core/src/llms/reducto/ocr/transformation.rs index 4b05f1ee168..7492a2ef6ec 100644 --- a/litellm-rust/crates/core/src/llms/reducto/ocr/transformation.rs +++ b/litellm-rust/crates/core/src/llms/reducto/ocr/transformation.rs @@ -738,8 +738,7 @@ mod tests { "result":{"chunks":[]} }))]) .await; - let request = - crate::ocr::test_support::with_source(wire_request(model, &base, options), source); + let request = crate::ocr::test_support::with_source(wire_request(model, &base, options), source); perform_ocr(request).await.unwrap(); server.await.unwrap(); @@ -858,9 +857,7 @@ mod tests { #[tokio::test] async fn rejects_invalid_document_sources_before_network(#[case] source: &str) { let request = crate::ocr::test_support::with_source( - wire_request("reducto/parse-v3", "http://127.0.0.1:1", json!({})), - source, - ); + wire_request("reducto/parse-v3", "http://127.0.0.1:1", json!({})), source); assert!(perform_ocr(request).await.is_err()); } @@ -916,9 +913,7 @@ mod tests { let raw = json!({"job_id":"job-1","result":{"chunks":[]}}); let (base, seen, server) = mock_server(vec![MockResponse::json(raw)]).await; let mut request = crate::ocr::test_support::with_source( - wire_request("reducto/parse-v3", &base, json!({})), - "reducto://ready.pdf", - ); + wire_request("reducto/parse-v3", &base, json!({})), "reducto://ready.pdf"); request.transport.extra_headers = vec![("authorization".into(), "Bearer existing".into())]; let response = perform_ocr(request).await.unwrap(); @@ -989,7 +984,6 @@ mod tests { request: OcrDuringCallRequest, ) -> OcrHookFuture<'_, OcrDuringCallRequest> { Box::pin(async move { - assert_eq!(request.optional_params["use_cache"], json!(true)); assert_eq!( request.body["document_url"], "data:application/pdf;base64,YWJj" @@ -1006,7 +1000,7 @@ mod tests { async fn guardrail_rewrites_document_before_upload() { let (base, seen, server) = mock_server(vec![MockResponse::json(json!({"result":{"chunks":[]}}))]).await; - let mut request = wire_request("reducto/parse-v3", &base, json!({"use_cache":true})); + let mut request = wire_request("reducto/parse-v3", &base, json!({})); request.hooks = Arc::new(RewriteDocument); perform_ocr(request).await.unwrap(); diff --git a/litellm-rust/crates/core/src/llms/vertex_ai/ocr/transformation.rs b/litellm-rust/crates/core/src/llms/vertex_ai/ocr/transformation.rs index a3eb35f16f2..beca4dc141d 100644 --- a/litellm-rust/crates/core/src/llms/vertex_ai/ocr/transformation.rs +++ b/litellm-rust/crates/core/src/llms/vertex_ai/ocr/transformation.rs @@ -338,12 +338,8 @@ mod tests { options.clone(), ); let vertex = wire_request("vertex_ai/mistral-ocr-maas", "https://vertex.test", options); - let direct = crate::ocr::prepare::prepare_request( - crate::ocr::test_support::resolved_request(direct), - ); - let vertex = crate::ocr::prepare::prepare_request( - crate::ocr::test_support::resolved_request(vertex), - ); + let direct = crate::ocr::prepare::prepare_request(crate::ocr::test_support::resolved_request(direct)); + let vertex = crate::ocr::prepare::prepare_request(crate::ocr::test_support::resolved_request(vertex)); let direct_http = MistralOCRConfig .prepare_request(&direct, &client) .await diff --git a/litellm-rust/crates/core/src/ocr/client.rs b/litellm-rust/crates/core/src/ocr/client.rs index 6ceb4eedbd2..5881519855c 100644 --- a/litellm-rust/crates/core/src/ocr/client.rs +++ b/litellm-rust/crates/core/src/ocr/client.rs @@ -39,10 +39,9 @@ impl OcrClient { ) -> Result { use super::{ NativeOutcome, OcrAdmission, OcrCall, OcrCallStep, OcrHookHost, OcrHost, - OcrHostOperation, OcrHostResult, OcrProjectedRequest, + OcrHostOperation, OcrHostResult, }; - let intercepts_requests = request.hooks.intercepts_requests(); let host = OcrHookHost::new(request.hooks.clone()); let mut request = Some(request); let NativeOutcome::Completed(mut call) = OcrCall::admit(self.clone(), OcrAdmission::all()) @@ -55,15 +54,14 @@ impl OcrClient { loop { match call.resume(result.take()).await? { OcrCallStep::Host(OcrHostOperation::ProjectRequest) => { - result = Some(OcrHostResult::Request(Ok(OcrProjectedRequest { - request: Box::new(request.take().ok_or_else(|| { + result = Some(OcrHostResult::Request(Ok(( + Box::new(request.take().ok_or_else(|| { crate::ocr::Error::InvalidRequest( "OCR request was already projected".into(), ) })?), - intercepts_requests, - host_token_provider: false, - }))) + false, + )))) } OcrCallStep::Host(operation) => result = Some(host.invoke(operation).await), OcrCallStep::Complete(response) => return Ok(response), diff --git a/litellm-rust/crates/core/src/ocr/document.rs b/litellm-rust/crates/core/src/ocr/document.rs index e9a13f0fd11..0f737e381ef 100644 --- a/litellm-rust/crates/core/src/ocr/document.rs +++ b/litellm-rust/crates/core/src/ocr/document.rs @@ -15,10 +15,7 @@ pub(crate) fn read_path_document( ) -> Result { let mut bytes = Vec::new(); std::fs::File::open(path) - .and_then(|file| { - file.take(OCR_INLINE_MAX_BYTES as u64 + 1) - .read_to_end(&mut bytes) - }) + .and_then(|file| file.take(OCR_INLINE_MAX_BYTES as u64 + 1).read_to_end(&mut bytes)) .map_err(|source| super::Error::FileRead { path: path.to_owned(), source: std::sync::Arc::new(source), @@ -204,14 +201,9 @@ mod tests { #[test] fn path_preparation_preserves_io_causes_and_enforces_the_inline_limit() { - let path = - std::env::temp_dir().join(format!("ocr-document-{:032x}.png", rand::random::())); + let path = std::env::temp_dir().join(format!("ocr-document-{:032x}.png", rand::random::())); let error = read_path_document(&path, None).unwrap_err(); - let super::super::Error::FileRead { - path: failed_path, - source, - } = error - else { + let super::super::Error::FileRead { path: failed_path, source } = error else { panic!("missing path must produce a typed file error"); }; assert_eq!(failed_path, path); @@ -223,16 +215,9 @@ mod tests { let image = read_path_document(&path, None).unwrap(); let overridden = read_path_document(&path, Some("application/pdf")).unwrap(); std::fs::remove_file(&path).unwrap(); - assert!(matches!( - oversized, - Err(super::super::Error::InlineDocumentTooLarge) - )); - assert!( - matches!(image, OcrDocument::ImageUrl { image_url, .. } if image_url == "data:image/png;base64,aW1hZ2UgYnl0ZXM=") - ); - assert!( - matches!(overridden, OcrDocument::DocumentUrl { document_url, .. } if document_url == "data:application/pdf;base64,aW1hZ2UgYnl0ZXM=") - ); + assert!(matches!(oversized, Err(super::super::Error::InlineDocumentTooLarge))); + assert!(matches!(image, OcrDocument::ImageUrl { image_url, .. } if image_url == "data:image/png;base64,aW1hZ2UgYnl0ZXM=")); + assert!(matches!(overridden, OcrDocument::DocumentUrl { document_url, .. } if document_url == "data:application/pdf;base64,aW1hZ2UgYnl0ZXM=")); } fn document(source: &str) -> OcrDocument { diff --git a/litellm-rust/crates/core/src/ocr/error.rs b/litellm-rust/crates/core/src/ocr/error.rs index a8ac969577f..0c70c478305 100644 --- a/litellm-rust/crates/core/src/ocr/error.rs +++ b/litellm-rust/crates/core/src/ocr/error.rs @@ -8,8 +8,6 @@ pub enum Error { }, #[error("File is empty or could not be read")] EmptyFile, - #[error("Host OCR document read failed")] - HostDocumentRead, #[error("Failed to read OCR file {}: {source}", path.display())] FileRead { path: std::path::PathBuf, diff --git a/litellm-rust/crates/core/src/ocr/handler.rs b/litellm-rust/crates/core/src/ocr/handler.rs index 7e42111da0a..d66da5c81d5 100644 --- a/litellm-rust/crates/core/src/ocr/handler.rs +++ b/litellm-rust/crates/core/src/ocr/handler.rs @@ -2,7 +2,7 @@ use std::sync::Arc; use super::OcrClient; use super::hooks::{OcrHooks, OcrLifecycleHooks, OcrPostCallRequest}; -use super::types::{LiteLLMOcrResponse, PreparedOcrRequest, ResolvedOcrRequest}; +use super::types::{ResolvedOcrRequest, LiteLLMOcrResponse, PreparedOcrRequest}; use crate::call_lifecycle::{CallLifecycle, CallLifecycleContext}; use crate::llms::base_llm::ocr::transformation::OcrResponseContext; diff --git a/litellm-rust/crates/core/src/ocr/hooks.rs b/litellm-rust/crates/core/src/ocr/hooks.rs index ee9b0f776d9..20cb3841c87 100644 --- a/litellm-rust/crates/core/src/ocr/hooks.rs +++ b/litellm-rust/crates/core/src/ocr/hooks.rs @@ -22,7 +22,6 @@ pub struct OcrPreCallRequest { pub struct OcrDuringCallRequest { pub model: String, pub custom_llm_provider: String, - pub optional_params: Value, pub api_key: Option, pub url: String, pub headers: Vec<(String, String)>, diff --git a/litellm-rust/crates/core/src/ocr/lifecycle.rs b/litellm-rust/crates/core/src/ocr/lifecycle.rs index dd88e11dca6..e013c9b9c62 100644 --- a/litellm-rust/crates/core/src/ocr/lifecycle.rs +++ b/litellm-rust/crates/core/src/ocr/lifecycle.rs @@ -1,7 +1,6 @@ use std::future::Future; use std::pin::Pin; use std::sync::Arc; -use std::time::{SystemTime, UNIX_EPOCH}; use tokio::sync::{mpsc, oneshot}; @@ -10,8 +9,8 @@ use super::hooks::{ OcrDuringCallRequest, OcrHookFuture, OcrHooks, OcrLogFuture, OcrPostCallRequest, OcrPreCallRequest, }; -use super::types::ResolvedOcrRequest; use super::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrClient, OcrDocumentInput, OcrFileContent}; +use super::types::ResolvedOcrRequest; use crate::call_lifecycle::host::{ HostCall, HostCallFuture, HostCallStep, HostFailure, HostLifecycle, HostPhase, }; @@ -83,14 +82,8 @@ impl OcrHostOperation { } } -pub struct OcrProjectedRequest { - pub request: Box, - pub intercepts_requests: bool, - pub host_token_provider: bool, -} - pub enum OcrHostResult { - Request(Result), + Request(Result<(Box, bool), super::Error>), Document(Result), Lifecycle(Result<(), HostFailure>), AzureAdToken(Result), @@ -167,10 +160,9 @@ impl OcrCall { Some(OcrHostResult::Request(result)) if self.projecting => { self.projecting = false; match result { - Ok(projected) => { - self.execution.set_request(*projected.request); - self.execution.intercepts_requests = projected.intercepts_requests; - self.execution.host_token_provider = projected.host_token_provider; + Ok((request, azure_ad_token_provider)) => { + self.execution.request = Some(*request); + self.execution.azure_ad_token_provider = azure_ad_token_provider; } Err(error) => self.accept(Err(HostFailure::Error(error))), } @@ -276,15 +268,6 @@ impl OcrCall { fn accept(&mut self, result: Result<(), HostFailure>) { let cancelled = matches!(&result, Err(HostFailure::Cancelled(_))); if let Some(error) = self.lifecycle.accept(result) { - if let Some((_, timing)) = self - .execution - .terminal - .lock() - .unwrap_or_else(|error| error.into_inner()) - .as_mut() - { - timing.end_time = epoch_seconds(); - } if cancelled { self.error = Some(error); } else { @@ -350,33 +333,11 @@ struct OcrExecution { pending_result: Option>, execution: Option>>, completed: bool, - intercepts_requests: bool, - host_token_provider: bool, + azure_ad_token_provider: bool, terminal: Arc>>, } impl OcrExecution { - fn set_request(&mut self, mut request: LiteLLMOcrRequest) { - let call_id = request - .litellm_call_id - .clone() - .unwrap_or_else(|| format!("ocr-{:032x}", rand::random::())); - let context = CallLifecycleContext::new( - "ocr", - request.model.clone(), - request.provider_name(), - call_id.clone(), - ); - request.litellm_call_id = Some(call_id); - let start = epoch_seconds(); - *self - .terminal - .lock() - .unwrap_or_else(|error| error.into_inner()) = - Some((context, CallLifecycleTiming::new(start, start, Vec::new()))); - self.request = Some(request); - } - fn new(client: OcrClient) -> Self { let (operations_tx, operations_rx) = mpsc::unbounded_channel(); Self { @@ -389,8 +350,7 @@ impl OcrExecution { pending_result: None, execution: None, completed: false, - intercepts_requests: false, - host_token_provider: false, + azure_ad_token_provider: false, terminal: Arc::default(), } } @@ -406,16 +366,11 @@ impl OcrExecution { } let result = if self.reading { let Some(OcrHostResult::Document(content)) = result else { - return Err(super::Error::InvalidRequest( - "OCR document read result is required".into(), - )); + return Err(super::Error::InvalidRequest("OCR document read result is required".into())); }; self.reading = false; let content = content?; - let request = self - .request - .take() - .expect("pending document read has a request"); + let request = self.request.take().expect("pending document read has a request"); let OcrDocumentInput::HostReader { mime_type } = &request.document else { unreachable!("only host readers request document reads"); }; @@ -473,10 +428,7 @@ impl OcrExecution { async fn prepare(&mut self) -> Result, super::Error> { if self.preparation.is_none() { - let request = self - .request - .take() - .expect("admitted OCR call has a request"); + let request = self.request.take().expect("admitted OCR call has a request"); if matches!(request.document, OcrDocumentInput::HostReader { .. }) { self.request = Some(request); self.reading = true; @@ -492,25 +444,15 @@ impl OcrExecution { OcrDocumentInput::Path { path, mime_type } => { super::document::read_path_document(path, mime_type.as_deref())? } - OcrDocumentInput::Bytes { - bytes, - file_name, - mime_type, - } => super::document::encode_file_document( - bytes, - file_name.as_deref(), - mime_type.as_deref(), - )?, + OcrDocumentInput::Bytes { bytes, file_name, mime_type } => { + super::document::encode_file_document(bytes, file_name.as_deref(), mime_type.as_deref())? + } _ => unreachable!("only native file inputs require preparation"), }; Ok(request.with_document(document)) })); } - let result = self - .preparation - .as_mut() - .expect("document preparation started") - .await; + let result = self.preparation.as_mut().expect("document preparation started").await; self.preparation = None; let request = result.map_err(|error| super::Error::DocumentTask(Arc::new(error)))??; self.start(request); @@ -519,7 +461,8 @@ impl OcrExecution { fn start(&mut self, mut request: ResolvedOcrRequest) { let client = self.client.take().expect("admitted OCR call has a client"); - if self.host_token_provider { + let intercepts_requests = request.hooks.intercepts_requests(); + if self.azure_ad_token_provider { request.azure_ad_token_provider = Some(TokenProviderHandle::new(Arc::new( OcrAzureAdTokenProvider { operations: self.operations_tx.clone(), @@ -528,7 +471,7 @@ impl OcrExecution { } request.hooks = Arc::new(ProtocolHooks { operations: self.operations_tx.clone(), - intercepts_requests: self.intercepts_requests, + intercepts_requests, terminal: self.terminal.clone(), }); self.execution = Some(tokio::spawn(async move { @@ -574,13 +517,6 @@ struct ProtocolHooks { terminal: Arc>>, } -fn epoch_seconds() -> f64 { - SystemTime::now() - .duration_since(UNIX_EPOCH) - .unwrap_or_default() - .as_secs_f64() -} - #[derive(Debug)] struct OcrAzureAdTokenProvider { operations: mpsc::UnboundedSender, @@ -807,60 +743,36 @@ mod tests { use crate::ocr::test_support::{MockResponse, mock_server, perform_ocr, wire_request}; use crate::ocr::{ LiteLLMOcrRequest, NativeOutcome, NoopOcrHost, OcrAdmission, OcrCall, OcrCallStep, - OcrDecline, OcrDocument, OcrHost, OcrHostOperation, OcrHostResult, OcrProjectedRequest, + OcrDecline, OcrDocument, OcrHost, OcrHostOperation, OcrHostResult, }; - fn projected_request(request: LiteLLMOcrRequest) -> OcrHostResult { - OcrHostResult::Request(Ok(OcrProjectedRequest { - intercepts_requests: request.hooks.intercepts_requests(), - request: Box::new(request), - host_token_provider: false, - })) - } - #[tokio::test] async fn sdk_paths_are_read_during_execution_and_hooks_observe_normalized_documents() { struct CaptureDocument(Arc>>); impl OcrHooks for CaptureDocument { - fn intercepts_requests(&self) -> bool { - true - } + fn intercepts_requests(&self) -> bool { true } fn pre_call(&self, request: OcrPreCallRequest) -> OcrHookFuture<'_, OcrPreCallRequest> { self.0.lock().unwrap().push(request.document.clone()); Box::pin(async { Ok(request) }) } } let (base, seen, server) = mock_server(vec![MockResponse::json(json!({"pages":[]}))]).await; - let path = - std::env::temp_dir().join(format!("ocr-sdk-{:032x}.pdf", rand::random::())); + let path = std::env::temp_dir().join(format!("ocr-sdk-{:032x}.pdf", rand::random::())); let documents = Arc::new(Mutex::new(Vec::new())); let request = LiteLLMOcrRequest::from_inputs( - "mistral/model".into(), - path.clone(), - None, - Default::default(), - crate::ocr::OcrConnectionInputs { - api_base: Some(base), - api_key: Some("test-key".into()), - ..Default::default() - }, - ) - .unwrap() - .with_host_hooks(Arc::new(CaptureDocument(documents.clone())), None); + "mistral/model".into(), path.clone(), None, Default::default(), + crate::ocr::OcrConnectionInputs { api_base: Some(base), api_key: Some("test-key".into()), ..Default::default() }, + ).unwrap().with_host_hooks(Arc::new(CaptureDocument(documents.clone())), None); std::fs::write(&path, b"sdk document").unwrap(); let result = perform_ocr(request).await; std::fs::remove_file(&path).unwrap(); result.unwrap(); server.await.unwrap(); let expected = json!({"type":"document_url","document_url":"data:application/pdf;base64,c2RrIGRvY3VtZW50"}); - assert_eq!( - serde_json::to_value(&documents.lock().unwrap()[0]).unwrap(), - expected - ); + assert_eq!(serde_json::to_value(&documents.lock().unwrap()[0]).unwrap(), expected); let requests = seen.lock().unwrap(); assert_eq!(requests.len(), 1); - let body: Value = - serde_json::from_str(requests[0].split_once("\r\n\r\n").unwrap().1).unwrap(); + let body: Value = serde_json::from_str(requests[0].split_once("\r\n\r\n").unwrap().1).unwrap(); assert_eq!(body["document"], expected); } @@ -869,80 +781,42 @@ mod tests { let (base, seen, server) = mock_server(vec![MockResponse::json(json!({"pages":[]}))]).await; let request = wire_request("mistral/model", &base, json!({})) .with_document(crate::ocr::OcrDocumentInput::HostReader { mime_type: None }) - .with_host_hooks( - Arc::new(AdmissionSpy { - effects: Arc::new(Mutex::new(0)), - }), - None, - ); + .with_host_hooks(Arc::new(AdmissionSpy { effects: Arc::new(Mutex::new(0)) }), None); let mut request = Some(request); - let NativeOutcome::Completed(mut call) = - OcrCall::admit(crate::ocr::test_support::ocr_client(), OcrAdmission::all()) - else { - panic!("admission declined"); - }; + let NativeOutcome::Completed(mut call) = OcrCall::admit(crate::ocr::test_support::ocr_client(), OcrAdmission::all()) else { panic!("admission declined"); }; let mut result = None; let mut reads = 0; let mut pre_calls = 0; - while let OcrCallStep::Host(operation) = call.resume(result.take()).await.unwrap() { - result = Some(match operation { - OcrHostOperation::ProjectRequest => { - assert_eq!(reads, 0); - projected_request(request.take().unwrap()) + loop { + match call.resume(result.take()).await.unwrap() { + OcrCallStep::Host(operation) => { + result = Some(match operation { + OcrHostOperation::ProjectRequest => { + assert_eq!(reads, 0); + OcrHostResult::Request(Ok((Box::new(request.take().unwrap()), false))) + } + OcrHostOperation::ReadDocument => { + reads += 1; + OcrHostResult::Document(Ok(crate::ocr::OcrFileContent { + bytes: bytes::Bytes::from_static(b"image"), file_name: Some("scan.png".into()), + })) + } + OcrHostOperation::PreCall(request) => { + assert_eq!(reads, 1); + pre_calls += 1; + assert!(matches!(&request.document, OcrDocument::ImageUrl { image_url, .. } if image_url == "data:image/png;base64,aW1hZ2U=")); + OcrHostResult::PreCall(Ok(request)) + } + operation => NoopOcrHost.invoke(operation).await, + }); } - OcrHostOperation::ReadDocument => { - reads += 1; - OcrHostResult::Document(Ok(crate::ocr::OcrFileContent { - bytes: bytes::Bytes::from_static(b"image"), - file_name: Some("scan.png".into()), - })) - } - OcrHostOperation::PreCall(request) => { - assert_eq!(reads, 1); - pre_calls += 1; - assert!( - matches!(&request.document, OcrDocument::ImageUrl { image_url, .. } if image_url == "data:image/png;base64,aW1hZ2U=") - ); - OcrHostResult::PreCall(Ok(request)) - } - operation => NoopOcrHost.invoke(operation).await, - }); + OcrCallStep::Complete(_) => break, + } } server.await.unwrap(); assert_eq!((reads, pre_calls, seen.lock().unwrap().len()), (1, 1, 1)); } - #[tokio::test] - async fn sdk_file_failure_dispatches_the_typed_error_once() { - struct CaptureFailure(Arc>>); - impl OcrHooks for CaptureFailure { - fn failure<'a>( - &'a self, - _: &'a CallLifecycleContext, - error: &'a crate::ocr::Error, - _: &'a CallLifecycleTiming, - ) -> OcrLogFuture<'a> { - self.0.lock().unwrap().push(error.clone()); - Box::pin(async {}) - } - } - let path = - std::env::temp_dir().join(format!("missing-ocr-{:032x}.pdf", rand::random::())); - let failures = Arc::new(Mutex::new(Vec::new())); - let request = wire_request("mistral/model", "http://127.0.0.1:1", json!({})) - .with_document(path.clone().into()) - .with_host_hooks(Arc::new(CaptureFailure(failures.clone())), None); - let error = perform_ocr(request).await.unwrap_err(); - assert!( - matches!(error, crate::ocr::Error::FileRead { path: actual, .. } if actual == path) - ); - let failures = failures.lock().unwrap(); - assert_eq!(failures.len(), 1); - assert!( - matches!(&failures[0], crate::ocr::Error::FileRead { source, .. } if source.kind() == std::io::ErrorKind::NotFound) - ); - } - #[tokio::test] async fn cancellation_acknowledges_blocking_preparation_completion() { use std::sync::atomic::{AtomicBool, Ordering}; @@ -962,8 +836,7 @@ mod tests { std::future::poll_fn(|cx| { assert!(stop.as_mut().poll(cx).is_pending()); std::task::Poll::Ready(()) - }) - .await; + }).await; assert!(!finished.load(Ordering::SeqCst)); release_tx.send(()).unwrap(); stop.await; @@ -1161,8 +1034,6 @@ mod tests { mut request: OcrDuringCallRequest, ) -> OcrHookFuture<'_, OcrDuringCallRequest> { Box::pin(async move { - assert_eq!(request.optional_params["pages"], json!([0])); - assert_eq!(request.optional_params["extra_body"]["pages"], json!([2])); assert_eq!(request.body["pages"], json!([2])); assert_eq!(request.body.get("future"), Some(&Value::Null)); request.body.as_object_mut().unwrap().remove("future"); @@ -1416,7 +1287,10 @@ mod tests { result = Some(OcrHostResult::Lifecycle(Ok(()))) } OcrHostOperation::ProjectRequest => { - result = Some(projected_request(request.take().unwrap())) + result = Some(OcrHostResult::Request(Ok(( + Box::new(request.take().unwrap()), + false, + )))) } OcrHostOperation::AcquireAzureAdToken => { panic!("test request has no token provider") @@ -1438,9 +1312,7 @@ mod tests { Ok(request) })); } - OcrHostOperation::PostCall(_) | OcrHostOperation::ReadDocument => { - panic!("transport should not be reached") - } + OcrHostOperation::PostCall(_) | OcrHostOperation::ReadDocument => panic!("transport should not be reached"), }, Err(error) => break error, Ok(OcrCallStep::Complete(_)) => panic!("failed call completed"), @@ -1473,7 +1345,10 @@ mod tests { let error = loop { match call.resume(result.take()).await { Ok(OcrCallStep::Host(OcrHostOperation::ProjectRequest)) => { - result = Some(projected_request(request.take().unwrap())); + result = Some(OcrHostResult::Request(Ok(( + Box::new(request.take().unwrap()), + false, + )))); } Ok(OcrCallStep::Host(operation)) => { if let OcrHostOperation::PostCall(request) = &operation { @@ -1537,7 +1412,7 @@ mod tests { }); result = Some(match operation { OcrHostOperation::ProjectRequest => { - projected_request(request.take().unwrap()) + OcrHostResult::Request(Ok((Box::new(request.take().unwrap()), false))) } operation => host.invoke(operation).await, }); @@ -1597,9 +1472,7 @@ mod tests { OcrHostResult::Lifecycle(Err(HostFailure::Error(selected.clone()))) } OcrHostOperation::Failure { error, .. } => { - assert!( - matches!(error, crate::ocr::Error::InvalidRequest(message) if message == "public metadata failed") - ); + assert!(matches!(error, crate::ocr::Error::InvalidRequest(message) if message == "public metadata failed")); failures.push("sync"); OcrHostResult::Lifecycle(Err(HostFailure::Error( crate::ocr::Error::InvalidRequest("failure callback failed".into()), @@ -1615,7 +1488,7 @@ mod tests { panic!("finalization failure used provider/success dispatch") } OcrHostOperation::ProjectRequest => { - projected_request(request.take().unwrap()) + OcrHostResult::Request(Ok((Box::new(request.take().unwrap()), false))) } operation => host.invoke(operation).await, }); @@ -1625,9 +1498,7 @@ mod tests { } }; server.await.unwrap(); - assert!( - matches!(error, crate::ocr::Error::InvalidRequest(message) if message == "public metadata failed") - ); + assert!(matches!(error, crate::ocr::Error::InvalidRequest(message) if message == "public metadata failed")); assert_eq!(failures, ["sync", "async"]); assert_eq!(seen.lock().unwrap().len(), 1); } @@ -1654,7 +1525,10 @@ mod tests { match call.resume(result.take()).await.unwrap() { OcrCallStep::Host(OcrHostOperation::PreCall(_)) => break, OcrCallStep::Host(OcrHostOperation::ProjectRequest) => { - result = Some(projected_request(request.take().unwrap())) + result = Some(OcrHostResult::Request(Ok(( + Box::new(request.take().unwrap()), + false, + )))) } OcrCallStep::Host(operation) => result = Some(host.invoke(operation).await), OcrCallStep::Complete(_) => panic!("provider executed before pre-call result"), @@ -1877,7 +1751,7 @@ mod tests { _ = entered.notified() => break, step = call.resume(result.take()) => { result = Some(match step.unwrap() { - OcrCallStep::Host(OcrHostOperation::ProjectRequest) => projected_request(request.take().unwrap()), + OcrCallStep::Host(OcrHostOperation::ProjectRequest) => OcrHostResult::Request(Ok((Box::new(request.take().unwrap()), false))), OcrCallStep::Host(operation) => NoopOcrHost.invoke(operation).await, OcrCallStep::Complete(_) => panic!("pending provider completed"), }); @@ -1904,9 +1778,7 @@ mod tests { ) .await .unwrap(); - assert!( - matches!(result, Err(crate::ocr::Error::InvalidRequest(message)) if message == "cancelled") - ); + assert!(matches!(result, Err(crate::ocr::Error::InvalidRequest(message)) if message == "cancelled")); assert!( dropped.load(Ordering::SeqCst), "cancellation returned while provider captures were still alive" diff --git a/litellm-rust/crates/core/src/ocr/mod.rs b/litellm-rust/crates/core/src/ocr/mod.rs index 95bb827c104..eb3162cc79e 100644 --- a/litellm-rust/crates/core/src/ocr/mod.rs +++ b/litellm-rust/crates/core/src/ocr/mod.rs @@ -18,13 +18,12 @@ pub use client::{OcrClient, ocr}; pub use document::{encode_file_document, mime_type_for_name, upload_mime_type}; pub use lifecycle::{ NativeOutcome, NativeResult, NoopOcrHost, OcrAdmission, OcrCall, OcrCallStep, OcrDecline, - OcrHookHost, OcrHost, OcrHostOperation, OcrHostResult, OcrProjectedRequest, + OcrHookHost, OcrHost, OcrHostOperation, OcrHostResult, }; pub use provider_config::{get_api_key_env_var, get_health_check_document}; pub use types::{ LiteLLMOcrRequest, LiteLLMOcrResponse, OcrConnection, OcrConnectionInputs, OcrCredentialInputs, - OcrDocument, OcrDocumentInput, OcrFileContent, OcrPage, OcrPageDimensions, OcrPageImage, - OcrTransportConfig, OcrUsageInfo, + OcrDocument, OcrDocumentInput, OcrFileContent, OcrPage, OcrPageDimensions, OcrPageImage, OcrTransportConfig, OcrUsageInfo, }; #[cfg(test)] diff --git a/litellm-rust/crates/core/src/ocr/prepare.rs b/litellm-rust/crates/core/src/ocr/prepare.rs index 19e8912ee7d..785f4a761b6 100644 --- a/litellm-rust/crates/core/src/ocr/prepare.rs +++ b/litellm-rust/crates/core/src/ocr/prepare.rs @@ -3,7 +3,7 @@ use serde_json::Value; use super::OcrClient; use super::hooks::OcrDuringCallRequest; -use super::types::{OcrConnection, OcrDocument, PreparedOcrRequest, ResolvedOcrRequest}; +use super::types::{ResolvedOcrRequest, OcrConnection, OcrDocument, PreparedOcrRequest}; pub(crate) async fn transform_request_body( client: &OcrClient, @@ -28,7 +28,6 @@ where .during_call(OcrDuringCallRequest { model: request.model.clone(), custom_llm_provider: request.provider_name().into(), - optional_params: Value::Object(request.optional_params.clone().into()), api_key: request.connection.api_key.clone(), url: url.into(), headers: headers.to_vec(), @@ -79,7 +78,6 @@ pub(crate) async fn guardrail_document( .during_call(OcrDuringCallRequest { model: request.model.clone(), custom_llm_provider: request.provider_name().into(), - optional_params: Value::Object(request.optional_params.clone().into()), api_key: request.connection.api_key.clone(), url: url.into(), headers: headers.to_vec(), diff --git a/litellm-rust/crates/core/src/ocr/types.rs b/litellm-rust/crates/core/src/ocr/types.rs index d92cd546be4..b7483224a32 100644 --- a/litellm-rust/crates/core/src/ocr/types.rs +++ b/litellm-rust/crates/core/src/ocr/types.rs @@ -3,8 +3,8 @@ use std::path::PathBuf; use std::sync::Arc; use std::time::Duration; -use bytes::Bytes; use serde::{Deserialize, Serialize}; +use bytes::Bytes; use serde_json::{Map, Value}; use serde_with::serde_as; @@ -93,10 +93,7 @@ impl From for OcrDocumentInput { impl From for OcrDocumentInput { fn from(path: PathBuf) -> Self { - Self::Path { - path, - mime_type: None, - } + Self::Path { path, mime_type: None } } } @@ -327,6 +324,7 @@ impl LiteLLMOcrRequest { config, }) } + } impl LiteLLMOcrRequest { @@ -385,6 +383,7 @@ impl LiteLLMOcrRequest { ..self } } + } impl LiteLLMOcrRequest { diff --git a/litellm-rust/crates/python-bridge/src/lifecycle/mod.rs b/litellm-rust/crates/python-bridge/src/lifecycle/mod.rs index 3a9e305c9f6..a11162f1c55 100644 --- a/litellm-rust/crates/python-bridge/src/lifecycle/mod.rs +++ b/litellm-rust/crates/python-bridge/src/lifecycle/mod.rs @@ -25,12 +25,17 @@ use bindings::DeploymentHooks; pub(crate) use bindings::PythonLogger; use handle::{Execution, ExecutionBody, ExecutionStep}; +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 phase(operation: &::Operation) -> Option; + fn classify(operation: &::Operation) -> OperationClass; fn lifecycle_result() -> ::Result; fn map_error(error: ::Error) -> PyErr; fn host_error(message: String) -> ::Error; @@ -162,8 +167,12 @@ impl PythonLifecycle { HostFailure::Cancelled(native) }; let state = self.route.state_mut(); - state.retain_first_error(py, error, cancelled && phase != Some(HostPhase::DeploymentFailure)); - let _ = state.finish(py); + if state.error.is_none() || (cancelled && phase != Some(HostPhase::DeploymentFailure)) { + state.retain_error(py, error); + } + if state.end.is_none() { + state.end = now(py).ok(); + } failure } @@ -206,7 +215,10 @@ impl PythonLifecycle { } HostStep::Ready(NativeCallStep::Host(operation)) => operation, }; - let phase = R::phase(&operation); + let phase = match R::classify(&operation) { + OperationClass::Phase(phase) => Some(phase), + OperationClass::Route => None, + }; let result = match phase { Some(phase) => match self.route.state_mut().invoke(py, phase) { Ok(HostStep::Suspend(awaitable)) => { @@ -487,19 +499,6 @@ impl PythonCallState { self.error = Some(error.into_value(py)); } - pub fn retain_first_error(&mut self, py: Python<'_>, error: PyErr, replace: bool) { - if self.error.is_none() || replace { - self.retain_error(py, error); - } - } - - pub fn finish(&mut self, py: Python<'_>) -> PyResult<()> { - if self.end.is_none() { - self.end = Some(now(py)?); - } - Ok(()) - } - pub fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> { visit.call(&self.args)?; visit.call(&self.kwargs)?; @@ -695,8 +694,8 @@ mod tests { &mut self.0 } - fn phase(_: &()) -> Option { - None + fn classify(_: &()) -> OperationClass { + OperationClass::Route } fn lifecycle_result() {} diff --git a/litellm-rust/crates/python-bridge/src/marshal.rs b/litellm-rust/crates/python-bridge/src/marshal.rs index 4dac89b1b23..c96052aed97 100644 --- a/litellm-rust/crates/python-bridge/src/marshal.rs +++ b/litellm-rust/crates/python-bridge/src/marshal.rs @@ -13,11 +13,6 @@ use litellm_python_interop::from_py_preserving_errors as from_py; use crate::auth::{AZURE_AD_TOKEN_PROVIDER, PythonTokenProvider}; use crate::lifecycle::BoundArguments; -pub(crate) struct Projection { - pub native: Native, - pub retained: Retained, -} - /// Fields every lifecycle route reads from its bound `*args, **kwargs` before /// asking core to build the typed request. Route-specific inputs (for example /// the OCR `document`) are read separately by the route. diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/callbacks.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/callbacks.rs index de70f7ce6f0..65896279715 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/callbacks.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/callbacks.rs @@ -3,55 +3,142 @@ use pyo3::types::PyDict; use serde_json::Value; use litellm_core::ocr::LiteLLMOcrResponse; -use litellm_core::ocr::hooks::OcrDuringCallRequest; +use litellm_core::ocr::hooks::OcrPreCallRequest; use litellm_python_interop::to_py_preserving_errors as to_py; -use super::host::PythonPayload; use crate::lifecycle::PythonLogger; -pub(super) fn update_logging( +pub(super) struct OcrLoggingFields { + model: String, + custom_llm_provider: String, + optional_params: Value, +} + +impl From<&OcrPreCallRequest> for OcrLoggingFields { + fn from(request: &OcrPreCallRequest) -> Self { + Self { + model: request.model.clone(), + custom_llm_provider: request.custom_llm_provider.clone(), + optional_params: request.optional_params.clone(), + } + } +} + +impl PythonLogger { + pub(super) fn update_ocr( + &self, + py: Python<'_>, + kwargs: &Py, + pre_call: &OcrLoggingFields, + secret_fields: &[&str], + url: &str, + ) -> PyResult<()> { + let update = PyDict::new(py); + update.set_item("kwargs", redact(py, kwargs.bind(py), secret_fields)?)?; + update.set_item("model", &pre_call.model)?; + update.set_item( + "optional_params", + redact( + py, + &to_py(py, &pre_call.optional_params)? + .into_bound(py) + .cast_into::()?, + secret_fields, + )?, + )?; + 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)?; + for name in ["logger_fn", "litellm_request_debug"] { + if let Some(value) = kwargs.bind(py).get_item(name)? { + params.set_item(name, value)?; + } + } + for name in custom_pricing_fields(py)? { + if let Some(value) = kwargs.bind(py).get_item(&name)? + && !value.is_none() + { + params.set_item(name, value)?; + } + } + 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<&str>, + 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>, + headers: Option<&Py>, + ) -> PyResult<()> { + let additional = PyDict::new(py); + additional.set_item("complete_input_dict", body)?; + additional.set_item("headers", headers)?; + let kwargs = PyDict::new(py); + kwargs.set_item("original_response", to_py(py, original_response)?)?; + kwargs.set_item("additional_args", &additional)?; + self.object(py) + .call_method("post_call", (), Some(&kwargs))?; + Ok(()) + } +} + +fn custom_pricing_fields(py: Python<'_>) -> PyResult> { + py.import("litellm.types.utils")? + .getattr("CustomPricingLiteLLMParams")? + .getattr("model_fields")? + .cast_into::()? + .keys() + .iter() + .map(|name| name.extract::()) + .collect() +} + +fn redact( py: Python<'_>, - logger: &PythonLogger, - kwargs: &Py, - request: &OcrDuringCallRequest, + params: &Bound<'_, PyDict>, secret_fields: &[&str], -) -> PyResult<()> { - py.import("litellm.rust_bridge.ocr")? - .getattr("update_logging")? - .call1(( - logger.object(py), - kwargs, - &request.model, - &request.custom_llm_provider, - to_py(py, &request.optional_params)?, - secret_fields, - &request.url, - ))?; - Ok(()) -} - -pub(super) fn pre_call( - py: Python<'_>, - logger: &PythonLogger, - request: &OcrDuringCallRequest, - payload: &PythonPayload, -) -> PyResult<()> { - py.import("litellm.rust_bridge.ocr")? - .getattr("pre_call")? - .call1((logger.object(py), request.api_key.as_deref(), &payload.body, &payload.headers, &request.url))?; - Ok(()) -} - -pub(super) fn post_call( - py: Python<'_>, - logger: &PythonLogger, - original_response: &Value, - payload: &PythonPayload, -) -> PyResult<()> { - py.import("litellm.rust_bridge.ocr")? - .getattr("post_call")? - .call1((logger.object(py), to_py(py, original_response)?, &payload.body, &payload.headers))?; - Ok(()) +) -> PyResult> { + let redacted = PyDict::new(py); + for (name, value) in params { + let name = name.extract::()?; + if name == "proxy_server_request" { + continue; + } + if secret_fields.contains(&name.as_str()) { + redacted.set_item(name, "****")?; + } else { + redacted.set_item(name, value)?; + } + } + Ok(redacted.unbind()) } pub(super) fn response(py: Python<'_>, response: &LiteLLMOcrResponse) -> PyResult> { diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/document.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/document.rs index 8d414435e9d..6df86e0dc54 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/document.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/document.rs @@ -36,11 +36,11 @@ fn extract_bytes(value: &Bound<'_, PyAny>) -> PyResult { return Ok(Bytes::from_owner(value.extract::()?)); } // Bytes subclasses may retain GC edges that a native Bytes owner cannot traverse. - Ok(Bytes::copy_from_slice( - value.extract::()?.as_ref(), - )) + Ok(Bytes::copy_from_slice(value.extract::()?.as_ref())) } + + pub(super) struct FileDocumentInput { pub input: OcrDocumentInput, pub reader: Option, @@ -56,11 +56,9 @@ impl FromPyObject<'_, '_> for FileDocumentInput { Err(error) if error.is_instance_of::(py) => None, Err(error) => return Err(error), }; - let missing = || { - PyValueError::new_err( - "document with type='file' must include a 'file' field containing a pathlib.Path, file-like object, or bytes", - ) - }; + let missing = || PyValueError::new_err( + "document with type='file' must include a 'file' field containing a pathlib.Path, file-like object, or bytes", + ); let file = document.get_item("file").map_err(|error| { if error.is_instance_of::(py) { missing() @@ -95,9 +93,7 @@ impl FromPyObject<'_, '_> for FileDocumentInput { reader: None, }); } - let reader = file - .getattr_opt("read")? - .filter(|value| value.is_callable()); + 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.", @@ -111,10 +107,7 @@ impl FromPyObject<'_, '_> for FileDocumentInput { .transpose()?; Ok(Self { input: OcrDocumentInput::HostReader { mime_type }, - reader: Some(PythonFileReader { - reader: reader.unbind(), - name, - }), + reader: Some(PythonFileReader { reader: reader.unbind(), name }), }) } } @@ -134,21 +127,12 @@ mod tests { let error = document.extract::().err().unwrap(); assert!(error.is_instance_of::(py)); } - for expression in [ - c"{'file': b'abc', 'mime_type': None}", - c"{'file': b'abc', 'mime_type': 7}", - ] { - let error = py - .eval(expression, None, None) - .unwrap() - .extract::() - .err() - .unwrap(); + for expression in [c"{'file': b'abc', 'mime_type': None}", c"{'file': b'abc', 'mime_type': 7}"] { + let error = py.eval(expression, None, None).unwrap().extract::().err().unwrap(); assert!(error.is_instance_of::(py)); } let locals = PyDict::new(py); - py.run( - c"from pathlib import Path + py.run(c"from pathlib import Path failure = KeyError('reader failed') class Reader: def __init__(self): @@ -159,31 +143,12 @@ class Reader: reader = Reader() document = {'file': reader} path_document = {'file': Path('/nonexistent/ocr-projection-test.pdf')} -", - Some(&locals), - Some(&locals), - ) - .unwrap(); +", Some(&locals), Some(&locals)).unwrap(); let document = locals.get_item("document").unwrap().unwrap(); let input: FileDocumentInput = document.extract().unwrap(); - assert_eq!( - locals - .get_item("reader") - .unwrap() - .unwrap() - .getattr("reads") - .unwrap() - .extract::() - .unwrap(), - 0 - ); + assert_eq!(locals.get_item("reader").unwrap().unwrap().getattr("reads").unwrap().extract::().unwrap(), 0); assert!(input.reader.is_some()); - let path: FileDocumentInput = locals - .get_item("path_document") - .unwrap() - .unwrap() - .extract() - .unwrap(); + let path: FileDocumentInput = locals.get_item("path_document").unwrap().unwrap().extract().unwrap(); assert!(matches!(path.input, OcrDocumentInput::Path { .. })); }); } diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/errors.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/errors.rs index c7f239054de..655cd9ea1ee 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/errors.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/errors.rs @@ -103,13 +103,9 @@ pub(super) fn public_exception( std::io::ErrorKind::PermissionDenied => 13, _ => 5, }); - return Ok(PyErr::from_value( - py.import("builtins")?.getattr("OSError")?.call1(( - errno, - source.to_string(), - path.into_pyobject(py)?.call_method0("__fspath__")?, - ))?, - )); + return Ok(PyErr::from_value(py.import("builtins")?.getattr("OSError")?.call1(( + errno, source.to_string(), path, + ))?)); } raise_public(py, classify(error), model, provider, None) } diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/host.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/host.rs index 5ccb65780df..91645aa04c1 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/host.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/host.rs @@ -1,11 +1,12 @@ +use std::sync::Arc; + use pyo3::gc::{PyTraverseError, PyVisit}; use pyo3::prelude::*; use pyo3::types::PyDict; use litellm_auth::ResolvedCredential; -use litellm_core::call_lifecycle::host::HostPhase; use litellm_core::ocr::hooks::{OcrDuringCallRequest, OcrPostCallRequest}; -use litellm_core::ocr::{OcrCall, OcrHostOperation, OcrHostResult, OcrProjectedRequest}; +use litellm_core::ocr::{OcrCall, OcrHostOperation, OcrHostResult}; use litellm_python_interop::{ from_py_preserving_errors as from_py, to_py_preserving_errors as to_py, }; @@ -13,50 +14,32 @@ use litellm_python_interop::{ use super::document::PythonFileReader; use super::errors::to_pyerr as ocr_error_to_pyerr; use super::project::project; -use super::{callbacks, errors}; +use super::{ASYNC_SIGNATURE, SIGNATURE, callbacks, errors}; use crate::auth::PythonTokenProvider; -use crate::lifecycle::{PythonCallState, PythonRoute, Signature, missing_state}; -use crate::marshal::Projection; +use crate::lifecycle::{OperationClass, PythonCallState, PythonRoute, missing_state, now}; pub(super) struct PythonOcrHost { state: PythonCallState, - signature: &'static Signature, - retained: Option, + projected: Option, } -pub(super) struct OcrRetained { - pub model: String, - pub provider: &'static str, - pub secret_fields: Vec<&'static str>, - pub azure_ad_token_provider: Option, - pub payload: Option, - pub reader: Option, +struct ProjectedOcrHost { + model: String, + provider: &'static str, + secret_fields: Vec<&'static str>, + azure_ad_token_provider: Option, + pre_call: Option, + payload: Option, + reader: Option, + reader_failed: bool, } -pub(super) struct PythonPayload { - pub body: Py, - pub headers: Py, +struct CapturedOcrPayload { + body: Py, + headers: Py, } -impl PythonPayload { - fn from_request(py: Python<'_>, request: &OcrDuringCallRequest) -> PyResult { - let body = to_py(py, &request.body)?.into_bound(py).cast_into::()?; - let headers = PyDict::new(py); - for (name, value) in &request.headers { - headers.set_item(name, value)?; - } - Ok(Self { body: body.unbind(), headers: headers.unbind() }) - } - - fn write_back(&self, py: Python<'_>, mut request: OcrDuringCallRequest) -> PyResult { - request.body = from_py(self.body.bind(py))?; - request.headers = self.headers.bind(py) - .iter() - .map(|(name, value)| Ok((name.extract::()?, value.extract::()?))) - .collect::>>()?; - Ok(request) - } - +impl CapturedOcrPayload { fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> { visit.call(&self.body)?; visit.call(&self.headers) @@ -64,36 +47,52 @@ impl PythonPayload { } impl PythonOcrHost { - pub(super) fn new(state: PythonCallState, signature: &'static Signature) -> Self { + pub(super) fn new(state: PythonCallState) -> Self { Self { state, - signature, - retained: None, + projected: None, } } - fn retained(&self) -> PyResult<&OcrRetained> { - self.retained.as_ref().ok_or_else(missing_state) + fn projected(&self) -> PyResult<&ProjectedOcrHost> { + self.projected.as_ref().ok_or_else(missing_state) } - fn retained_mut(&mut self) -> PyResult<&mut OcrRetained> { - self.retained.as_mut().ok_or_else(missing_state) + fn projected_mut(&mut self) -> PyResult<&mut ProjectedOcrHost> { + self.projected.as_mut().ok_or_else(missing_state) } fn project(&mut self, py: Python<'_>) -> PyResult { - let arguments = self.signature.bind(self.state.args.bind(py), self.state.kwargs.bind(py))?; - let Projection { native, retained } = project(py, &arguments)?; - let host_token_provider = retained.azure_ad_token_provider.is_some(); - self.retained = Some(retained); - Ok(OcrHostResult::Request(Ok(OcrProjectedRequest { - request: Box::new(native), - intercepts_requests: true, - host_token_provider, - }))) + let signature = if self.state.asynchronous { + &ASYNC_SIGNATURE + } else { + &SIGNATURE + }; + let arguments = signature.bind(self.state.args.bind(py), self.state.kwargs.bind(py))?; + let projected = project(py, &arguments)?; + let has_token_provider = projected.azure_ad_token_provider.is_some(); + self.projected = Some(ProjectedOcrHost { + model: projected.request.model.clone(), + provider: projected.request.provider_name(), + secret_fields: projected.secret_fields, + azure_ad_token_provider: projected.azure_ad_token_provider, + pre_call: None, + payload: None, + reader: projected.reader, + reader_failed: false, + }); + Ok(OcrHostResult::Request(Ok(( + Box::new( + projected + .request + .with_host_hooks(Arc::new(BridgeOcrHooks), None), + ), + has_token_provider, + )))) } fn acquire_azure_ad_token(&self, py: Python<'_>) -> PyResult { - self.retained()? + self.projected()? .azure_ad_token_provider .as_ref() .ok_or_else(missing_state)? @@ -103,21 +102,42 @@ impl PythonOcrHost { fn during_call( &mut self, py: Python<'_>, - request: OcrDuringCallRequest, + mut request: OcrDuringCallRequest, ) -> PyResult { - let retained = self.retained()?; + let projected = self.projected()?; + let pre_call = projected.pre_call.as_ref().ok_or_else(missing_state)?; let logger = self.state.logger()?; - callbacks::update_logging( + logger.update_ocr( py, - logger, &self.state.kwargs, - &request, - &retained.secret_fields, + pre_call, + &projected.secret_fields, + &request.url, )?; - let payload = PythonPayload::from_request(py, &request)?; - callbacks::pre_call(py, logger, &request, &payload)?; - let request = payload.write_back(py, request)?; - self.retained_mut()?.payload = Some(payload); + let body = to_py(py, &request.body)? + .into_bound(py) + .cast_into::()?; + let headers = PyDict::new(py); + for (name, value) in &request.headers { + headers.set_item(name, value)?; + } + logger.pre_ocr( + py, + request.api_key.as_deref(), + &body, + &headers, + &request.url, + )?; + request.body = from_py(&body)?; + request.headers = headers + .iter() + .map(|(name, value)| Ok((name.extract::()?, value.extract::()?))) + .collect::>>()?; + let projected = self.projected_mut()?; + projected.payload = Some(CapturedOcrPayload { + body: body.unbind(), + headers: headers.unbind(), + }); Ok(request) } @@ -126,24 +146,27 @@ impl PythonOcrHost { py: Python<'_>, request: OcrPostCallRequest, ) -> PyResult { - let payload = self.retained()?.payload.as_ref().ok_or_else(missing_state)?; - callbacks::post_call( + let projected = self.projected()?; + let payload = projected.payload.as_ref(); + self.state.logger()?.post_ocr( py, - self.state.logger()?, &request.original_response, - payload, + payload.map(|payload| &payload.body), + payload.map(|payload| &payload.headers), )?; Ok(request) } fn map_failure(&mut self, py: Python<'_>, error: litellm_core::ocr::Error) -> PyResult<()> { - self.state.finish(py)?; - let (model, provider) = match &self.retained { - Some(retained) => (retained.model.as_str(), retained.provider), + if self.state.end.is_none() { + self.state.end = Some(now(py)?); + } + let (model, provider) = match &self.projected { + Some(projected) => (projected.model.as_str(), projected.provider), None => ("", ""), }; let mapped = match self.state.error.take() { - Some(host_error) if matches!(error, litellm_core::ocr::Error::HostDocumentRead) => { + Some(host_error) if self.projected.as_ref().is_some_and(|host| host.reader_failed) => { PyErr::from_value(host_error.into_bound(py).into_any()) } Some(host_error) => errors::public_host_exception(py, &host_error, model, provider)?, @@ -165,8 +188,10 @@ impl PythonRoute for PythonOcrHost { &mut self.state } - fn phase(operation: &OcrHostOperation) -> Option { - operation.phase() + fn classify(operation: &OcrHostOperation) -> OperationClass { + operation + .phase() + .map_or(OperationClass::Route, OperationClass::Phase) } fn lifecycle_result() -> OcrHostResult { @@ -185,24 +210,16 @@ impl PythonRoute for PythonOcrHost { Ok(match operation { OcrHostOperation::ProjectRequest => self.project(py)?, OcrHostOperation::ReadDocument => { - let reader = self - .retained_mut()? - .reader - .take() - .ok_or_else(missing_state)?; - match reader.read(py) { - Ok(content) => OcrHostResult::Document(Ok(content)), - Err(error) if error.is_instance_of::(py) => { - self.state.retain_first_error(py, error, false); - OcrHostResult::Document(Err(litellm_core::ocr::Error::HostDocumentRead)) - } - Err(error) => return Err(error), - } + let reader = self.projected_mut()?.reader.take().ok_or_else(missing_state)?; + let result = reader.read(py); + self.projected_mut()?.reader_failed = result.is_err(); + OcrHostResult::Document(Ok(result?)) } OcrHostOperation::AcquireAzureAdToken => { OcrHostResult::AzureAdToken(Ok(self.acquire_azure_ad_token(py)?)) } OcrHostOperation::PreCall(request) => { + self.projected_mut()?.pre_call = Some((&request).into()); OcrHostResult::PreCall(Ok(request)) } OcrHostOperation::DuringCall(request) => { @@ -212,7 +229,7 @@ impl PythonRoute for PythonOcrHost { OcrHostResult::PostCall(Ok(self.post_call(py, request)?)) } OcrHostOperation::ConstructResponse(response) => { - self.state.finish(py)?; + self.state.end = Some(now(py)?); self.state.response = Some(callbacks::response(py, response.as_ref())?); OcrHostResult::Lifecycle(Ok(())) } @@ -227,22 +244,30 @@ impl PythonRoute for PythonOcrHost { } fn cleanup(&mut self) { - self.retained = None; + self.projected = None; } fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> { - let Some(retained) = &self.retained else { + let Some(projected) = &self.projected else { return Ok(()); }; - if let Some(provider) = &retained.azure_ad_token_provider { + if let Some(provider) = &projected.azure_ad_token_provider { provider.traverse(visit)?; } - if let Some(reader) = &retained.reader { + if let Some(reader) = &projected.reader { reader.traverse(visit)?; } - if let Some(payload) = &retained.payload { + if let Some(payload) = &projected.payload { payload.traverse(visit)?; } Ok(()) } } + +struct BridgeOcrHooks; + +impl litellm_core::ocr::hooks::OcrHooks for BridgeOcrHooks { + fn intercepts_requests(&self) -> bool { + true + } +} diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs index cbd32792843..7973dc45c9b 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs @@ -68,7 +68,7 @@ fn call( kwargs.copy()?.unbind(), asynchronous, signature.name, - )?, signature); + )?); run_call(py, call, host) } diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/project.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/project.rs index 57cddac48cb..a78aa5b80b7 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/project.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/project.rs @@ -10,34 +10,54 @@ use litellm_core::ocr::{ }; use litellm_python_interop::from_py_preserving_errors as from_py; -use super::document::FileDocumentInput; use super::errors::to_pyerr as ocr_error_to_pyerr; -use super::host::OcrRetained; +use super::document::{FileDocumentInput, PythonFileReader}; +use crate::auth::PythonTokenProvider; use crate::lifecycle::BoundArguments; -use crate::marshal::{BoundRouteInputs, Projection}; +use crate::marshal::BoundRouteInputs; /// Positional parameters of `ocr()` that are never projected into /// `optional_params`. const BOUND_FIELDS: &[&str] = &["model", "document", "timeout", "input_sources"]; -fn project_document(document: &Bound<'_, PyAny>) -> PyResult> { +pub(super) struct ProjectedOcrCall { + pub request: LiteLLMOcrRequest, + pub azure_ad_token_provider: Option, + pub secret_fields: Vec<&'static str>, + pub reader: Option, +} + +enum ProjectedDocument { + File(FileDocumentInput), + Url(serde_json::Value), +} + +impl ProjectedDocument { + fn into_native(self) -> Result { + match self { + Self::File(file) => Ok(file), + Self::Url(value) => Ok(FileDocumentInput { + input: OcrDocument::try_from(value)?.into(), + reader: None, + }), + } + } +} + +fn project_document(document: &Bound<'_, PyAny>) -> PyResult { let kind: String = document.get_item("type")?.extract()?; if kind != "file" { - let value: serde_json::Value = from_py(document)?; - return Ok(OcrDocument::try_from(value).map(|document| FileDocumentInput { - input: document.into(), - reader: None, - })); + return Ok(ProjectedDocument::Url(from_py(document)?)); } - document.extract().map(Ok) + document.extract().map(ProjectedDocument::File) } /// Pure core assembly; every failure here is a typed `ocr::Error`. fn build_request( inputs: BoundRouteInputs, - document: Result, -) -> Result, litellm_core::ocr::Error> { - let document = document?; + document: ProjectedDocument, +) -> Result { + let document = document.into_native()?; let BoundRouteInputs { model, custom_llm_provider, @@ -63,23 +83,18 @@ fn build_request( input_sources, }, )?; - Ok(Projection { - retained: OcrRetained { - model: request.model.clone(), - provider: request.provider_name(), - azure_ad_token_provider, - secret_fields, - reader: document.reader, - payload: None, - }, - native: request, + Ok(ProjectedOcrCall { + request, + azure_ad_token_provider, + secret_fields, + reader: document.reader, }) } pub(super) fn project( py: Python<'_>, arguments: &BoundArguments<'_>, -) -> PyResult> { +) -> PyResult { let model: String = arguments.extract("model")?; let custom_llm_provider: Option = arguments.optional("custom_llm_provider")?; let document = project_document(&arguments.required("document")?)?; @@ -127,7 +142,7 @@ mod tests { ) .unwrap(); assert!(matches!( - project_document(&file).unwrap().unwrap().input, + project_document(&file).unwrap().into_native().unwrap().input, litellm_core::ocr::OcrDocumentInput::Bytes { bytes, mime_type, .. } if bytes == b"%PDF-1.4"[..] && mime_type.as_deref() == Some("application/pdf") )); @@ -140,7 +155,7 @@ mod tests { ) .unwrap(); assert!(matches!( - project_document(&original).unwrap().unwrap().input, + project_document(&original).unwrap().into_native().unwrap().input, litellm_core::ocr::OcrDocumentInput::Document(OcrDocument::DocumentUrl { document_url, .. }) if document_url == "https://example.com/a.pdf" )); @@ -154,10 +169,7 @@ mod tests { let document = py .eval(c"{'type': 'mystery', 'mystery': 'x'}", None, None) .unwrap(); - let error = project_document(&document) - .unwrap() - .err() - .unwrap(); + let error = project_document(&document).unwrap().into_native().err().unwrap(); assert!(error.to_string().contains("document")); }); } @@ -169,16 +181,14 @@ mod tests { let missing = py.eval(c"{}", None, None).unwrap(); assert!( project_document(&missing) - .err() - .unwrap() + .err().unwrap() .is_instance_of::(py) ); let non_string = py.eval(c"{'type': 1}", None, None).unwrap(); assert!( project_document(&non_string) - .err() - .unwrap() + .err().unwrap() .is_instance_of::(py) ); @@ -192,9 +202,8 @@ class Document: document = Document() ", ); - let error = project_document(&locals.get_item("document").unwrap().unwrap()) - .err() - .unwrap(); + let error = + project_document(&locals.get_item("document").unwrap().unwrap()).err().unwrap(); assert!( error .value(py) @@ -223,11 +232,8 @@ document = Document() ", ); let document = locals.get_item("document").unwrap().unwrap(); - let projected = project_document(&document).unwrap().unwrap(); - assert!(matches!( - projected.input, - litellm_core::ocr::OcrDocumentInput::Bytes { .. } - )); + let projected = project_document(&document).unwrap().into_native().unwrap(); + assert!(matches!(projected.input, litellm_core::ocr::OcrDocumentInput::Bytes { .. })); let reads: Vec = document.getattr("reads").unwrap().extract().unwrap(); assert_eq!(reads, ["type", "mime_type", "file"]); }); diff --git a/litellm/proxy/ocr_endpoints/endpoints.py b/litellm/proxy/ocr_endpoints/endpoints.py index 3c858f0a69d..283f704bc56 100644 --- a/litellm/proxy/ocr_endpoints/endpoints.py +++ b/litellm/proxy/ocr_endpoints/endpoints.py @@ -29,12 +29,10 @@ def _build_document_from_upload( filename: str | None, content_type: str | None, ) -> dict[str, str]: - supplied_mime: Final = content_type.split(";")[0].strip() if content_type else None - mime_type: Final = ( - get_mime_type(filename) - if filename and (not supplied_mime or supplied_mime == "application/octet-stream") - else supplied_mime - ) + mime_type: Final = content_type.split(";")[0].strip() if content_type else None + if not mime_type or mime_type == "application/octet-stream": + if filename: + mime_type = get_mime_type(filename) return convert_file_document_to_url_document( { # mutable-ok: OCR file document TypedDict handed to the converter diff --git a/litellm/rust_bridge/bindings.py b/litellm/rust_bridge/bindings.py index 32f6e9ec46b..d16f150a2aa 100644 --- a/litellm/rust_bridge/bindings.py +++ b/litellm/rust_bridge/bindings.py @@ -47,8 +47,3 @@ def native_exception_types() -> tuple[type[BaseException], type[BaseException]] if not isinstance(declined, type) or not isinstance(upstream, type): return None return declined, upstream - - -def native_decline_types() -> tuple[type[BaseException], ...]: - exceptions: Final = native_exception_types() - return (exceptions[0],) if exceptions is not None else () diff --git a/litellm/rust_bridge/ocr.py b/litellm/rust_bridge/ocr.py index 35c399de73b..7ede521e6f3 100644 --- a/litellm/rust_bridge/ocr.py +++ b/litellm/rust_bridge/ocr.py @@ -1,6 +1,6 @@ from __future__ import annotations -from collections.abc import Callable, Coroutine, Mapping, Sequence +from collections.abc import Callable, Coroutine, Mapping from types import MappingProxyType from typing import Final, Protocol, cast # noqa: TID251 # validates dynamically loaded native callables @@ -8,86 +8,6 @@ from pydantic import TypeAdapter from litellm.llms.base_llm.ocr.transformation import PROVIDER_NATIVE_RESPONSE_KEY, OCRResponse from litellm.rust_bridge.bindings import NativeBinding -from litellm.types.utils import CustomPricingLiteLLMParams - - -class OcrLoggingProtocol(Protocol): - def update_from_kwargs( - self, - *, - kwargs: dict[str, object], - model: str, - optional_params: dict[str, object], - litellm_params: dict[str, object], - custom_llm_provider: str, - ) -> object: ... - - def pre_call(self, *, input: str, api_key: str | None, additional_args: dict[str, object]) -> object: ... - - def post_call(self, *, original_response: object, additional_args: dict[str, object]) -> object: ... - - -def _redact(params: Mapping[str, object], secret_fields: Sequence[str]) -> dict[str, object]: - return { # mutable-ok: update_from_kwargs takes dict - name: "****" if name in secret_fields else value - for name, value in params.items() - if name != "proxy_server_request" - } - - -def update_logging( - logger: OcrLoggingProtocol, - kwargs: Mapping[str, object], - model: str, - custom_llm_provider: str, - optional_params: Mapping[str, object], - secret_fields: Sequence[str], - url: str, -) -> None: - logger.update_from_kwargs( - kwargs=_redact(kwargs, secret_fields), - model=model, - optional_params=_redact(optional_params, secret_fields), - litellm_params={ # mutable-ok: update_from_kwargs takes dict - "litellm_call_id": kwargs.get("litellm_call_id"), - "api_base": url, - **{ - name: kwargs[name] for name in ("logger_fn", "litellm_request_debug") if name in kwargs - }, # mutable-ok: splat into the dict above - **{ # mutable-ok: splat into the dict above - name: kwargs[name] - for name in CustomPricingLiteLLMParams.model_fields - if name in kwargs and kwargs[name] is not None - }, - }, - custom_llm_provider=custom_llm_provider, - ) - - -def pre_call( - logger: OcrLoggingProtocol, - api_key: str | None, - body: dict[str, object], - headers: dict[str, str], - url: str, -) -> None: - logger.pre_call( - input="OCR document processing", - api_key=api_key, - additional_args={"complete_input_dict": body, "headers": headers, "api_base": url}, - ) - - -def post_call( - logger: OcrLoggingProtocol, - original_response: object, - body: dict[str, object], - headers: dict[str, str], -) -> None: - logger.post_call( - original_response=original_response, - additional_args={"complete_input_dict": body, "headers": headers}, - ) class RustOcr(Protocol): diff --git a/tests/test_litellm/rust_bridge/test_bindings.py b/tests/test_litellm/rust_bridge/test_bindings.py index 7e6fea7adfb..88036a5a556 100644 --- a/tests/test_litellm/rust_bridge/test_bindings.py +++ b/tests/test_litellm/rust_bridge/test_bindings.py @@ -33,21 +33,3 @@ def test_binding_validates_native_attribute( binding: Final = bindings.NativeBinding("route", validate=lambda item: item if isinstance(item, int) else None) assert binding.load() == expected - - -@pytest.mark.parametrize("available", [False, True]) -def test_decline_accessor_only_catches_admission_declines(monkeypatch: pytest.MonkeyPatch, available: bool) -> None: - class Declined(Exception): - pass - - class Upstream(Exception): - pass - - native: Final = SimpleNamespace(RustBridgeDeclined=Declined, RustUpstreamError=Upstream) if available else None - monkeypatch.setattr(bindings, "get_native_bridge", lambda: native) - assert bindings.native_decline_types() == ((Declined,) if available else ()) - with pytest.raises(Upstream): - try: - raise Upstream("already dispatched") - except bindings.native_decline_types(): - pytest.fail("upstream failure was allowed to replay") diff --git a/tests/test_litellm/rust_bridge/test_ocr_lifecycle.py b/tests/test_litellm/rust_bridge/test_ocr_lifecycle.py index 96983968bee..baf5295d571 100644 --- a/tests/test_litellm/rust_bridge/test_ocr_lifecycle.py +++ b/tests/test_litellm/rust_bridge/test_ocr_lifecycle.py @@ -8,7 +8,7 @@ import litellm from litellm.llms.base_llm.ocr.transformation import OCRResponse from litellm.ocr import legacy from litellm.rust_bridge import bindings, configuration -from litellm.rust_bridge.ocr import NATIVE_AOCR, NATIVE_OCR, post_call, pre_call, update_logging +from litellm.rust_bridge.ocr import NATIVE_AOCR, NATIVE_OCR @pytest.fixture(autouse=True) @@ -21,72 +21,6 @@ def isolated_ocr_configuration(monkeypatch: pytest.MonkeyPatch) -> Generator[Non configuration.reset_rust_configuration() -def test_logging_redacts_views_and_preserves_opaque_arguments_and_pricing() -> None: - logger: Final = Mock() - opaque: Final = object() - logger_fn: Final = object() - kwargs: Final = { - "vertex_credentials": "secret", - "proxy_server_request": opaque, - "metadata": opaque, - "logger_fn": logger_fn, - "litellm_request_debug": False, - "litellm_call_id": "call-id", - "input_cost_per_token": 0, - "output_cost_per_token": None, - } - optional: Final = {"vertex_credentials": "secret", "pages": [1], "proxy_server_request": opaque} - update_logging(logger, kwargs, "model", "vertex_ai", optional, ("vertex_credentials",), "https://provider") - logger.update_from_kwargs.assert_called_once_with( - kwargs={ - "vertex_credentials": "****", - "metadata": opaque, - "logger_fn": logger_fn, - "litellm_request_debug": False, - "litellm_call_id": "call-id", - "input_cost_per_token": 0, - "output_cost_per_token": None, - }, - model="model", - custom_llm_provider="vertex_ai", - optional_params={"vertex_credentials": "****", "pages": [1]}, - litellm_params={ - "litellm_call_id": "call-id", - "api_base": "https://provider", - "logger_fn": logger_fn, - "litellm_request_debug": False, - "input_cost_per_token": 0, - }, - ) - assert kwargs["vertex_credentials"] == optional["vertex_credentials"] == "secret" - assert kwargs["proxy_server_request"] is optional["proxy_server_request"] is opaque - assert logger.update_from_kwargs.call_args.kwargs["kwargs"]["metadata"] is opaque - assert logger.update_from_kwargs.call_args.kwargs["optional_params"]["pages"] is optional["pages"] - - -def test_logging_callbacks_receive_captured_payload_roots_and_propagate_errors() -> None: - logger: Final = Mock() - body: Final[dict[str, object]] = {"document": "original"} - headers: Final = {"authorization": "key"} - response: Final = object() - pre_call(logger, "key", body, headers, "https://provider") - post_call(logger, response, body, headers) - logger.pre_call.assert_called_once_with( - input="OCR document processing", - api_key="key", - additional_args={"complete_input_dict": body, "headers": headers, "api_base": "https://provider"}, - ) - for callback in (logger.pre_call, logger.post_call): - assert callback.call_args.kwargs["additional_args"]["complete_input_dict"] is body - assert callback.call_args.kwargs["additional_args"]["headers"] is headers - assert logger.post_call.call_args.kwargs["original_response"] is response - failure: Final = RuntimeError("callback failed") - failing_logger: Final = Mock(pre_call=Mock(side_effect=failure)) - with pytest.raises(RuntimeError) as caught: - pre_call(failing_logger, None, body, headers, "https://provider") - assert caught.value is failure - - @pytest.mark.asyncio @pytest.mark.parametrize("asynchronous", [False, True]) async def test_unavailable_native_uses_legacy(monkeypatch: pytest.MonkeyPatch, asynchronous: bool) -> None: diff --git a/tests/test_litellm_rust/test_ocr.py b/tests/test_litellm_rust/test_ocr.py index 86913f388db..a83748c0dd6 100644 --- a/tests/test_litellm_rust/test_ocr.py +++ b/tests/test_litellm_rust/test_ocr.py @@ -1,15 +1,14 @@ +import json import asyncio import base64 -import gc -import json -import os import threading import weakref -from collections.abc import Coroutine, Generator +import gc +from pathlib import Path +from collections.abc import Generator from datetime import datetime, timezone from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer from io import BytesIO -from pathlib import Path from typing import Final, Protocol import httpx @@ -341,12 +340,7 @@ def test_native_lifecycle_core_encodes_python_file_input( @pytest.mark.parametrize("asynchronous", [False, True]) @pytest.mark.parametrize("source", ["path", "pathlike", "reader", "text", "bytes"]) @pytest.mark.asyncio -async def test_native_file_sources_preserve_sdk_behavior( - ocr_server: tuple[ThreadingHTTPServer, list[dict[str, object]]], - tmp_path: Path, - asynchronous: bool, - source: str, -) -> None: +async def test_native_file_sources_preserve_sdk_behavior(ocr_server, tmp_path, asynchronous, source): server, requests = ocr_server path: Final = tmp_path / "scan.png" content: Final = b"document bytes" @@ -394,12 +388,9 @@ async def test_native_file_sources_preserve_sdk_behavior( "api_base": f"http://127.0.0.1:{server.server_port}", } response: Final = await litellm.aocr(**kwargs) if asynchronous else litellm.ocr(**kwargs) - assert isinstance(response, OCRResponse) assert response.pages[0].markdown == "native OCR response" assert len(requests) == 1 - body: Final = requests[0]["body"] - assert isinstance(body, dict) - assert body["document"] == { + assert requests[0]["body"]["document"] == { "type": "image_url", "image_url": "data:image/png;base64," + base64.b64encode(content).decode(), } @@ -409,11 +400,7 @@ async def test_native_file_sources_preserve_sdk_behavior( @pytest.mark.parametrize("asynchronous", [False, True]) @pytest.mark.asyncio -async def test_native_file_failures_preserve_identity_and_do_not_send( - ocr_server: tuple[ThreadingHTTPServer, list[dict[str, object]]], - tmp_path: Path, - asynchronous: bool, -) -> None: +async def test_native_file_failures_preserve_identity_and_do_not_send(ocr_server, tmp_path, asynchronous): from litellm.rust_bridge import _native server, requests = ocr_server @@ -429,47 +416,37 @@ async def test_native_file_failures_preserve_identity_and_do_not_send( reader: Final = Reader() path: Final = tmp_path / "missing.pdf" - - async def invoke(source: Reader | Path) -> None: - if asynchronous: - await _native.aocr( - model="mistral/mistral-ocr-latest", - document={"type": "file", "file": source}, - api_key="test-key", - api_base=f"http://127.0.0.1:{server.server_port}", - ) - else: - _native.ocr( - model="mistral/mistral-ocr-latest", - document={"type": "file", "file": source}, - api_key="test-key", - api_base=f"http://127.0.0.1:{server.server_port}", - ) - + kwargs: Final = { + "model": "mistral/mistral-ocr-latest", + "api_key": "test-key", + "api_base": f"http://127.0.0.1:{server.server_port}", + } for source, error in ((reader, KeyError), (path, FileNotFoundError)): with pytest.raises(error) as caught: - await invoke(source) + if asynchronous: + await _native.aocr(document={"type": "file", "file": source}, **kwargs) + else: + _native.ocr(document={"type": "file", "file": source}, **kwargs) if source is reader: assert caught.value is failure else: - assert isinstance(caught.value, FileNotFoundError) - assert getattr(caught.value, "filename") == str(path) + assert caught.value.filename == str(path) assert reader.calls == 1 assert requests == [] -def test_unstarted_native_file_call_does_not_read_and_releases_reader() -> None: +def test_unstarted_native_file_call_does_not_read_and_releases_reader(): from litellm.rust_bridge import _native class Reader: - pending: Coroutine[object, object, OCRResponse] | None = None - def read(self) -> bytes: raise AssertionError("unstarted call consumed its document") - def create_call() -> tuple[Coroutine[object, object, OCRResponse], weakref.ReferenceType[Reader]]: + def create_call(): reader: Final = Reader() - pending: Final = _native.aocr(model="mistral/mistral-ocr-latest", document={"type": "file", "file": reader}) + pending: Final = _native.aocr( + model="mistral/mistral-ocr-latest", document={"type": "file", "file": reader} + ) reader.pending = pending return pending, weakref.ref(reader) @@ -480,49 +457,6 @@ def test_unstarted_native_file_call_does_not_read_and_releases_reader() -> None: assert reference() is None -@pytest.mark.skipif(not hasattr(os, "mkfifo"), reason="requires Unix named pipes") -@pytest.mark.asyncio -async def test_native_path_cancellation_waits_for_read_completion( - ocr_server: tuple[ThreadingHTTPServer, list[dict[str, object]]], tmp_path: Path -) -> None: - from litellm.rust_bridge import _native - - server, requests = ocr_server - path: Final = tmp_path / "document.pdf" - os.mkfifo(path) - entered: Final = threading.Event() - release: Final = threading.Event() - - def supply_document() -> None: - with path.open("wb") as stream: - stream.write(b"document") - stream.flush() - entered.set() - release.wait(5) - - writer: Final = threading.Thread(target=supply_document, daemon=True) - writer.start() - task: Final = asyncio.create_task( - _native.aocr( - model="mistral/mistral-ocr-latest", - document={"type": "file", "file": path}, - api_key="test-key", - api_base=f"http://127.0.0.1:{server.server_port}", - ) - ) - try: - assert await asyncio.to_thread(entered.wait, 3) - task.cancel() - await asyncio.sleep(0.05) - assert not task.done() - finally: - release.set() - await asyncio.to_thread(writer.join, 3) - with pytest.raises(asyncio.CancelledError): - await task - assert requests == [] - - @pytest.mark.parametrize("asynchronous", [False, True]) @pytest.mark.parametrize("model", ["mistral/mistral-ocr-latest", "azure_ai/doc-intelligence/prebuilt-read"]) @pytest.mark.asyncio