mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
parent
92e1a2ded7
commit
60d046ccc4
30 changed files with 465 additions and 1551 deletions
|
|
@ -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<CallbackMethod> {
|
||||
match self {
|
||||
Self::SyncSuccess => Some(CallbackMethod::LoggingHook),
|
||||
Self::AsyncSuccess => Some(CallbackMethod::AsyncLoggingHook),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
pub const fn marker(self) -> Option<LoggedMarker> {
|
||||
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<CallbackId> {
|
||||
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<CallbackKind>,
|
||||
}
|
||||
|
||||
pub fn plan_success(facts: &SuccessFacts) -> Vec<SuccessDispatch> {
|
||||
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<CallbackFamily> {
|
||||
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<CallbackId>,
|
||||
stream: bool,
|
||||
position: Position,
|
||||
}
|
||||
|
||||
impl DispatchCursor {
|
||||
pub fn start(
|
||||
family: CallbackFamily,
|
||||
targets: Vec<CallbackId>,
|
||||
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<F: FnMut(CallbackId, CallbackMethod) -> bool>(F);
|
||||
|
||||
impl<F: FnMut(CallbackId, CallbackMethod) -> bool> DispatchFacts for Gate<F> {
|
||||
fn eligible(&mut self, target: CallbackId, method: CallbackMethod) -> bool {
|
||||
(self.0)(target, method)
|
||||
}
|
||||
}
|
||||
|
||||
fn drain(cursor: &mut DispatchCursor, facts: &mut dyn DispatchFacts) -> Vec<DispatchStep> {
|
||||
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<CallbackId> {
|
||||
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);
|
||||
}
|
||||
}
|
||||
|
|
@ -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);
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -39,10 +39,9 @@ impl OcrClient {
|
|||
) -> Result<LiteLLMOcrResponse, crate::ocr::Error> {
|
||||
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),
|
||||
|
|
|
|||
|
|
@ -15,10 +15,7 @@ pub(crate) fn read_path_document(
|
|||
) -> Result<OcrDocument, super::Error> {
|
||||
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::<u128>()));
|
||||
let path = std::env::temp_dir().join(format!("ocr-document-{:032x}.png", rand::random::<u128>()));
|
||||
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 {
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
||||
|
|
|
|||
|
|
@ -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<String>,
|
||||
pub url: String,
|
||||
pub headers: Vec<(String, String)>,
|
||||
|
|
|
|||
|
|
@ -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<LiteLLMOcrRequest>,
|
||||
pub intercepts_requests: bool,
|
||||
pub host_token_provider: bool,
|
||||
}
|
||||
|
||||
pub enum OcrHostResult {
|
||||
Request(Result<OcrProjectedRequest, super::Error>),
|
||||
Request(Result<(Box<LiteLLMOcrRequest>, bool), super::Error>),
|
||||
Document(Result<OcrFileContent, super::Error>),
|
||||
Lifecycle(Result<(), HostFailure<super::Error>>),
|
||||
AzureAdToken(Result<ResolvedCredential, litellm_auth::Error>),
|
||||
|
|
@ -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<super::Error>>) {
|
||||
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<oneshot::Sender<OcrHostResult>>,
|
||||
execution: Option<tokio::task::JoinHandle<Result<LiteLLMOcrResponse, super::Error>>>,
|
||||
completed: bool,
|
||||
intercepts_requests: bool,
|
||||
host_token_provider: bool,
|
||||
azure_ad_token_provider: bool,
|
||||
terminal: Arc<std::sync::Mutex<Option<(CallLifecycleContext, CallLifecycleTiming)>>>,
|
||||
}
|
||||
|
||||
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::<u128>()));
|
||||
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<Option<OcrHostOperation>, 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<std::sync::Mutex<Option<(CallLifecycleContext, CallLifecycleTiming)>>>,
|
||||
}
|
||||
|
||||
fn epoch_seconds() -> f64 {
|
||||
SystemTime::now()
|
||||
.duration_since(UNIX_EPOCH)
|
||||
.unwrap_or_default()
|
||||
.as_secs_f64()
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
struct OcrAzureAdTokenProvider {
|
||||
operations: mpsc::UnboundedSender<PendingOperation>,
|
||||
|
|
@ -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<Mutex<Vec<OcrDocument>>>);
|
||||
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::<u128>()));
|
||||
let path = std::env::temp_dir().join(format!("ocr-sdk-{:032x}.pdf", rand::random::<u128>()));
|
||||
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<Mutex<Vec<crate::ocr::Error>>>);
|
||||
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::<u128>()));
|
||||
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"
|
||||
|
|
|
|||
|
|
@ -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)]
|
||||
|
|
|
|||
|
|
@ -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<B>(
|
||||
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(),
|
||||
|
|
|
|||
|
|
@ -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<OcrDocument> for OcrDocumentInput {
|
|||
|
||||
impl From<PathBuf> 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<D> LiteLLMOcrRequest<D> {
|
||||
|
|
@ -385,6 +383,7 @@ impl<D> LiteLLMOcrRequest<D> {
|
|||
..self
|
||||
}
|
||||
}
|
||||
|
||||
}
|
||||
|
||||
impl LiteLLMOcrRequest {
|
||||
|
|
|
|||
|
|
@ -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: &<Self::Call as NativeCall>::Operation) -> Option<HostPhase>;
|
||||
fn classify(operation: &<Self::Call as NativeCall>::Operation) -> OperationClass;
|
||||
fn lifecycle_result() -> <Self::Call as NativeCall>::Result;
|
||||
fn map_error(error: <Self::Call as NativeCall>::Error) -> PyErr;
|
||||
fn host_error(message: String) -> <Self::Call as NativeCall>::Error;
|
||||
|
|
@ -162,8 +167,12 @@ impl<R: PythonRoute> PythonLifecycle<R> {
|
|||
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<R: PythonRoute> PythonLifecycle<R> {
|
|||
}
|
||||
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<HostPhase> {
|
||||
None
|
||||
fn classify(_: &()) -> OperationClass {
|
||||
OperationClass::Route
|
||||
}
|
||||
|
||||
fn lifecycle_result() {}
|
||||
|
|
|
|||
|
|
@ -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<Native, Retained> {
|
||||
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.
|
||||
|
|
|
|||
|
|
@ -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<PyDict>,
|
||||
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::<PyDict>()?,
|
||||
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<PyDict>>,
|
||||
headers: Option<&Py<PyDict>>,
|
||||
) -> 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<Vec<String>> {
|
||||
py.import("litellm.types.utils")?
|
||||
.getattr("CustomPricingLiteLLMParams")?
|
||||
.getattr("model_fields")?
|
||||
.cast_into::<PyDict>()?
|
||||
.keys()
|
||||
.iter()
|
||||
.map(|name| name.extract::<String>())
|
||||
.collect()
|
||||
}
|
||||
|
||||
fn redact(
|
||||
py: Python<'_>,
|
||||
logger: &PythonLogger,
|
||||
kwargs: &Py<PyDict>,
|
||||
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<Py<PyDict>> {
|
||||
let redacted = PyDict::new(py);
|
||||
for (name, value) in params {
|
||||
let name = name.extract::<String>()?;
|
||||
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<Py<PyAny>> {
|
||||
|
|
|
|||
|
|
@ -36,11 +36,11 @@ fn extract_bytes(value: &Bound<'_, PyAny>) -> PyResult<Bytes> {
|
|||
return Ok(Bytes::from_owner(value.extract::<PyBackedBytes>()?));
|
||||
}
|
||||
// Bytes subclasses may retain GC edges that a native Bytes owner cannot traverse.
|
||||
Ok(Bytes::copy_from_slice(
|
||||
value.extract::<PyBackedBytes>()?.as_ref(),
|
||||
))
|
||||
Ok(Bytes::copy_from_slice(value.extract::<PyBackedBytes>()?.as_ref()))
|
||||
}
|
||||
|
||||
|
||||
|
||||
pub(super) struct FileDocumentInput {
|
||||
pub input: OcrDocumentInput,
|
||||
pub reader: Option<PythonFileReader>,
|
||||
|
|
@ -56,11 +56,9 @@ impl FromPyObject<'_, '_> for FileDocumentInput {
|
|||
Err(error) if error.is_instance_of::<pyo3::exceptions::PyKeyError>(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::<pyo3::exceptions::PyKeyError>(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::<FileDocumentInput>().err().unwrap();
|
||||
assert!(error.is_instance_of::<PyValueError>(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::<FileDocumentInput>()
|
||||
.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::<FileDocumentInput>().err().unwrap();
|
||||
assert!(error.is_instance_of::<PyTypeError>(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::<usize>()
|
||||
.unwrap(),
|
||||
0
|
||||
);
|
||||
assert_eq!(locals.get_item("reader").unwrap().unwrap().getattr("reads").unwrap().extract::<usize>().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 { .. }));
|
||||
});
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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<OcrRetained>,
|
||||
projected: Option<ProjectedOcrHost>,
|
||||
}
|
||||
|
||||
pub(super) struct OcrRetained {
|
||||
pub model: String,
|
||||
pub provider: &'static str,
|
||||
pub secret_fields: Vec<&'static str>,
|
||||
pub azure_ad_token_provider: Option<PythonTokenProvider>,
|
||||
pub payload: Option<PythonPayload>,
|
||||
pub reader: Option<PythonFileReader>,
|
||||
struct ProjectedOcrHost {
|
||||
model: String,
|
||||
provider: &'static str,
|
||||
secret_fields: Vec<&'static str>,
|
||||
azure_ad_token_provider: Option<PythonTokenProvider>,
|
||||
pre_call: Option<callbacks::OcrLoggingFields>,
|
||||
payload: Option<CapturedOcrPayload>,
|
||||
reader: Option<PythonFileReader>,
|
||||
reader_failed: bool,
|
||||
}
|
||||
|
||||
pub(super) struct PythonPayload {
|
||||
pub body: Py<PyDict>,
|
||||
pub headers: Py<PyDict>,
|
||||
struct CapturedOcrPayload {
|
||||
body: Py<PyDict>,
|
||||
headers: Py<PyDict>,
|
||||
}
|
||||
|
||||
impl PythonPayload {
|
||||
fn from_request(py: Python<'_>, request: &OcrDuringCallRequest) -> PyResult<Self> {
|
||||
let body = to_py(py, &request.body)?.into_bound(py).cast_into::<PyDict>()?;
|
||||
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<OcrDuringCallRequest> {
|
||||
request.body = from_py(self.body.bind(py))?;
|
||||
request.headers = self.headers.bind(py)
|
||||
.iter()
|
||||
.map(|(name, value)| Ok((name.extract::<String>()?, value.extract::<String>()?)))
|
||||
.collect::<PyResult<Vec<_>>>()?;
|
||||
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<OcrHostResult> {
|
||||
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<ResolvedCredential> {
|
||||
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<OcrDuringCallRequest> {
|
||||
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::<PyDict>()?;
|
||||
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::<String>()?, value.extract::<String>()?)))
|
||||
.collect::<PyResult<Vec<_>>>()?;
|
||||
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<OcrPostCallRequest> {
|
||||
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<HostPhase> {
|
||||
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::<pyo3::exceptions::PyException>(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
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -68,7 +68,7 @@ fn call(
|
|||
kwargs.copy()?.unbind(),
|
||||
asynchronous,
|
||||
signature.name,
|
||||
)?, signature);
|
||||
)?);
|
||||
run_call(py, call, host)
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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<Result<FileDocumentInput, litellm_core::ocr::Error>> {
|
||||
pub(super) struct ProjectedOcrCall {
|
||||
pub request: LiteLLMOcrRequest,
|
||||
pub azure_ad_token_provider: Option<PythonTokenProvider>,
|
||||
pub secret_fields: Vec<&'static str>,
|
||||
pub reader: Option<PythonFileReader>,
|
||||
}
|
||||
|
||||
enum ProjectedDocument {
|
||||
File(FileDocumentInput),
|
||||
Url(serde_json::Value),
|
||||
}
|
||||
|
||||
impl ProjectedDocument {
|
||||
fn into_native(self) -> Result<FileDocumentInput, litellm_core::ocr::Error> {
|
||||
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<ProjectedDocument> {
|
||||
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<FileDocumentInput, litellm_core::ocr::Error>,
|
||||
) -> Result<Projection<LiteLLMOcrRequest, OcrRetained>, litellm_core::ocr::Error> {
|
||||
let document = document?;
|
||||
document: ProjectedDocument,
|
||||
) -> Result<ProjectedOcrCall, litellm_core::ocr::Error> {
|
||||
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<Projection<LiteLLMOcrRequest, OcrRetained>> {
|
||||
) -> PyResult<ProjectedOcrCall> {
|
||||
let model: String = arguments.extract("model")?;
|
||||
let custom_llm_provider: Option<String> = 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::<PyKeyError>(py)
|
||||
);
|
||||
|
||||
let non_string = py.eval(c"{'type': 1}", None, None).unwrap();
|
||||
assert!(
|
||||
project_document(&non_string)
|
||||
.err()
|
||||
.unwrap()
|
||||
.err().unwrap()
|
||||
.is_instance_of::<PyTypeError>(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<String> = document.getattr("reads").unwrap().extract().unwrap();
|
||||
assert_eq!(reads, ["type", "mime_type", "file"]);
|
||||
});
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 ()
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue