Revert "refactor callback"

This reverts commit 1771e5bb68.
This commit is contained in:
Yujong Lee 2026-09-16 19:07:43 -07:00
parent 92e1a2ded7
commit 60d046ccc4
30 changed files with 465 additions and 1551 deletions

View file

@ -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);
}
}

View file

@ -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);

View file

@ -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,

View file

@ -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();

View file

@ -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

View file

@ -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();

View file

@ -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

View file

@ -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),

View file

@ -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 {

View file

@ -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,

View file

@ -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;

View file

@ -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)>,

View file

@ -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"

View file

@ -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)]

View file

@ -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(),

View file

@ -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 {

View file

@ -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() {}

View file

@ -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.

View file

@ -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>> {

View file

@ -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 { .. }));
});
}

View file

@ -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)
}

View file

@ -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
}
}

View file

@ -68,7 +68,7 @@ fn call(
kwargs.copy()?.unbind(),
asynchronous,
signature.name,
)?, signature);
)?);
run_call(py, call, host)
}

View file

@ -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"]);
});

View 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

View file

@ -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 ()

View file

@ -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):

View file

@ -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")

View file

@ -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:

View file

@ -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