refactor callback

This commit is contained in:
Yujong Lee 2026-09-16 07:07:41 -07:00
parent 75548d30f5
commit 1771e5bb68
29 changed files with 799 additions and 472 deletions

View file

@ -223,9 +223,11 @@ 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

@ -816,7 +816,8 @@ 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,10 +361,12 @@ 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())
@ -503,10 +505,12 @@ 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,7 +738,8 @@ 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();
@ -857,7 +858,9 @@ 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());
}
@ -913,7 +916,9 @@ 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();
@ -984,6 +989,7 @@ 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"
@ -1000,7 +1006,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!({}));
let mut request = wire_request("reducto/parse-v3", &base, json!({"use_cache":true}));
request.hooks = Arc::new(RewriteDocument);
perform_ocr(request).await.unwrap();

View file

@ -338,8 +338,12 @@ 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,9 +39,10 @@ impl OcrClient {
) -> Result<LiteLLMOcrResponse, crate::ocr::Error> {
use super::{
NativeOutcome, OcrAdmission, OcrCall, OcrCallStep, OcrHookHost, OcrHost,
OcrHostOperation, OcrHostResult,
OcrHostOperation, OcrHostResult, OcrProjectedRequest,
};
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())
@ -54,14 +55,15 @@ impl OcrClient {
loop {
match call.resume(result.take()).await? {
OcrCallStep::Host(OcrHostOperation::ProjectRequest) => {
result = Some(OcrHostResult::Request(Ok((
Box::new(request.take().ok_or_else(|| {
result = Some(OcrHostResult::Request(Ok(OcrProjectedRequest {
request: Box::new(request.take().ok_or_else(|| {
crate::ocr::Error::InvalidRequest(
"OCR request was already projected".into(),
)
})?),
false,
))))
intercepts_requests,
host_token_provider: false,
})))
}
OcrCallStep::Host(operation) => result = Some(host.invoke(operation).await),
OcrCallStep::Complete(response) => return Ok(response),

View file

@ -15,7 +15,10 @@ 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),
@ -201,9 +204,14 @@ 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);
@ -215,9 +223,16 @@ 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,6 +8,8 @@ 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::{ResolvedOcrRequest, LiteLLMOcrResponse, PreparedOcrRequest};
use super::types::{LiteLLMOcrResponse, PreparedOcrRequest, ResolvedOcrRequest};
use crate::call_lifecycle::{CallLifecycle, CallLifecycleContext};
use crate::llms::base_llm::ocr::transformation::OcrResponseContext;

View file

@ -22,6 +22,7 @@ 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,6 +1,7 @@
use std::future::Future;
use std::pin::Pin;
use std::sync::Arc;
use std::time::{SystemTime, UNIX_EPOCH};
use tokio::sync::{mpsc, oneshot};
@ -9,8 +10,8 @@ use super::hooks::{
OcrDuringCallRequest, OcrHookFuture, OcrHooks, OcrLogFuture, OcrPostCallRequest,
OcrPreCallRequest,
};
use super::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrClient, OcrDocumentInput, OcrFileContent};
use super::types::ResolvedOcrRequest;
use super::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrClient, OcrDocumentInput, OcrFileContent};
use crate::call_lifecycle::host::{
HostCall, HostCallFuture, HostCallStep, HostFailure, HostLifecycle, HostPhase,
};
@ -82,8 +83,14 @@ impl OcrHostOperation {
}
}
pub struct OcrProjectedRequest {
pub request: Box<LiteLLMOcrRequest>,
pub intercepts_requests: bool,
pub host_token_provider: bool,
}
pub enum OcrHostResult {
Request(Result<(Box<LiteLLMOcrRequest>, bool), super::Error>),
Request(Result<OcrProjectedRequest, super::Error>),
Document(Result<OcrFileContent, super::Error>),
Lifecycle(Result<(), HostFailure<super::Error>>),
AzureAdToken(Result<ResolvedCredential, litellm_auth::Error>),
@ -160,9 +167,10 @@ impl OcrCall {
Some(OcrHostResult::Request(result)) if self.projecting => {
self.projecting = false;
match result {
Ok((request, azure_ad_token_provider)) => {
self.execution.request = Some(*request);
self.execution.azure_ad_token_provider = azure_ad_token_provider;
Ok(projected) => {
self.execution.set_request(*projected.request);
self.execution.intercepts_requests = projected.intercepts_requests;
self.execution.host_token_provider = projected.host_token_provider;
}
Err(error) => self.accept(Err(HostFailure::Error(error))),
}
@ -268,6 +276,15 @@ 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 {
@ -333,11 +350,33 @@ struct OcrExecution {
pending_result: Option<oneshot::Sender<OcrHostResult>>,
execution: Option<tokio::task::JoinHandle<Result<LiteLLMOcrResponse, super::Error>>>,
completed: bool,
azure_ad_token_provider: bool,
intercepts_requests: bool,
host_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 {
@ -350,7 +389,8 @@ impl OcrExecution {
pending_result: None,
execution: None,
completed: false,
azure_ad_token_provider: false,
intercepts_requests: false,
host_token_provider: false,
terminal: Arc::default(),
}
}
@ -366,11 +406,16 @@ 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");
};
@ -428,7 +473,10 @@ 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;
@ -444,15 +492,25 @@ 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);
@ -461,8 +519,7 @@ impl OcrExecution {
fn start(&mut self, mut request: ResolvedOcrRequest) {
let client = self.client.take().expect("admitted OCR call has a client");
let intercepts_requests = request.hooks.intercepts_requests();
if self.azure_ad_token_provider {
if self.host_token_provider {
request.azure_ad_token_provider = Some(TokenProviderHandle::new(Arc::new(
OcrAzureAdTokenProvider {
operations: self.operations_tx.clone(),
@ -471,7 +528,7 @@ impl OcrExecution {
}
request.hooks = Arc::new(ProtocolHooks {
operations: self.operations_tx.clone(),
intercepts_requests,
intercepts_requests: self.intercepts_requests,
terminal: self.terminal.clone(),
});
self.execution = Some(tokio::spawn(async move {
@ -517,6 +574,13 @@ 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>,
@ -743,36 +807,60 @@ 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,
OcrDecline, OcrDocument, OcrHost, OcrHostOperation, OcrHostResult, OcrProjectedRequest,
};
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);
}
@ -781,42 +869,80 @@ 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;
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,
});
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())
}
OcrCallStep::Complete(_) => break,
}
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,
});
}
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};
@ -836,7 +962,8 @@ 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;
@ -1034,6 +1161,8 @@ 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");
@ -1287,10 +1416,7 @@ mod tests {
result = Some(OcrHostResult::Lifecycle(Ok(())))
}
OcrHostOperation::ProjectRequest => {
result = Some(OcrHostResult::Request(Ok((
Box::new(request.take().unwrap()),
false,
))))
result = Some(projected_request(request.take().unwrap()))
}
OcrHostOperation::AcquireAzureAdToken => {
panic!("test request has no token provider")
@ -1312,7 +1438,9 @@ 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"),
@ -1345,10 +1473,7 @@ mod tests {
let error = loop {
match call.resume(result.take()).await {
Ok(OcrCallStep::Host(OcrHostOperation::ProjectRequest)) => {
result = Some(OcrHostResult::Request(Ok((
Box::new(request.take().unwrap()),
false,
))));
result = Some(projected_request(request.take().unwrap()));
}
Ok(OcrCallStep::Host(operation)) => {
if let OcrHostOperation::PostCall(request) = &operation {
@ -1412,7 +1537,7 @@ mod tests {
});
result = Some(match operation {
OcrHostOperation::ProjectRequest => {
OcrHostResult::Request(Ok((Box::new(request.take().unwrap()), false)))
projected_request(request.take().unwrap())
}
operation => host.invoke(operation).await,
});
@ -1472,7 +1597,9 @@ 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()),
@ -1488,7 +1615,7 @@ mod tests {
panic!("finalization failure used provider/success dispatch")
}
OcrHostOperation::ProjectRequest => {
OcrHostResult::Request(Ok((Box::new(request.take().unwrap()), false)))
projected_request(request.take().unwrap())
}
operation => host.invoke(operation).await,
});
@ -1498,7 +1625,9 @@ 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);
}
@ -1525,10 +1654,7 @@ mod tests {
match call.resume(result.take()).await.unwrap() {
OcrCallStep::Host(OcrHostOperation::PreCall(_)) => break,
OcrCallStep::Host(OcrHostOperation::ProjectRequest) => {
result = Some(OcrHostResult::Request(Ok((
Box::new(request.take().unwrap()),
false,
))))
result = Some(projected_request(request.take().unwrap()))
}
OcrCallStep::Host(operation) => result = Some(host.invoke(operation).await),
OcrCallStep::Complete(_) => panic!("provider executed before pre-call result"),
@ -1751,7 +1877,7 @@ mod tests {
_ = entered.notified() => break,
step = call.resume(result.take()) => {
result = Some(match step.unwrap() {
OcrCallStep::Host(OcrHostOperation::ProjectRequest) => OcrHostResult::Request(Ok((Box::new(request.take().unwrap()), false))),
OcrCallStep::Host(OcrHostOperation::ProjectRequest) => projected_request(request.take().unwrap()),
OcrCallStep::Host(operation) => NoopOcrHost.invoke(operation).await,
OcrCallStep::Complete(_) => panic!("pending provider completed"),
});
@ -1778,7 +1904,9 @@ 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,12 +18,13 @@ 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,
OcrHookHost, OcrHost, OcrHostOperation, OcrHostResult, OcrProjectedRequest,
};
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::{ResolvedOcrRequest, OcrConnection, OcrDocument, PreparedOcrRequest};
use super::types::{OcrConnection, OcrDocument, PreparedOcrRequest, ResolvedOcrRequest};
pub(crate) async fn transform_request_body<B>(
client: &OcrClient,
@ -28,6 +28,7 @@ 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(),
@ -78,6 +79,7 @@ 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 serde::{Deserialize, Serialize};
use bytes::Bytes;
use serde::{Deserialize, Serialize};
use serde_json::{Map, Value};
use serde_with::serde_as;
@ -93,7 +93,10 @@ 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,
}
}
}
@ -324,7 +327,6 @@ impl LiteLLMOcrRequest {
config,
})
}
}
impl<D> LiteLLMOcrRequest<D> {
@ -383,7 +385,6 @@ impl<D> LiteLLMOcrRequest<D> {
..self
}
}
}
impl LiteLLMOcrRequest {

View file

@ -25,17 +25,12 @@ 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 classify(operation: &<Self::Call as NativeCall>::Operation) -> OperationClass;
fn phase(operation: &<Self::Call as NativeCall>::Operation) -> Option<HostPhase>;
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;
@ -167,12 +162,8 @@ impl<R: PythonRoute> PythonLifecycle<R> {
HostFailure::Cancelled(native)
};
let state = self.route.state_mut();
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();
}
state.retain_first_error(py, error, cancelled && phase != Some(HostPhase::DeploymentFailure));
let _ = state.finish(py);
failure
}
@ -215,10 +206,7 @@ impl<R: PythonRoute> PythonLifecycle<R> {
}
HostStep::Ready(NativeCallStep::Host(operation)) => operation,
};
let phase = match R::classify(&operation) {
OperationClass::Phase(phase) => Some(phase),
OperationClass::Route => None,
};
let phase = R::phase(&operation);
let result = match phase {
Some(phase) => match self.route.state_mut().invoke(py, phase) {
Ok(HostStep::Suspend(awaitable)) => {
@ -499,6 +487,19 @@ 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)?;
@ -694,8 +695,8 @@ mod tests {
&mut self.0
}
fn classify(_: &()) -> OperationClass {
OperationClass::Route
fn phase(_: &()) -> Option<HostPhase> {
None
}
fn lifecycle_result() {}

View file

@ -13,6 +13,11 @@ 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,142 +3,55 @@ use pyo3::types::PyDict;
use serde_json::Value;
use litellm_core::ocr::LiteLLMOcrResponse;
use litellm_core::ocr::hooks::OcrPreCallRequest;
use litellm_core::ocr::hooks::OcrDuringCallRequest;
use litellm_python_interop::to_py_preserving_errors as to_py;
use super::host::PythonPayload;
use crate::lifecycle::PythonLogger;
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(
pub(super) fn update_logging(
py: Python<'_>,
params: &Bound<'_, PyDict>,
logger: &PythonLogger,
kwargs: &Py<PyDict>,
request: &OcrDuringCallRequest,
secret_fields: &[&str],
) -> 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())
) -> 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(())
}
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,9 +56,11 @@ 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()
@ -93,7 +95,9 @@ 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.",
@ -107,7 +111,10 @@ 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,
}),
})
}
}
@ -127,12 +134,21 @@ 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):
@ -143,12 +159,31 @@ 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,9 +103,13 @@ 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,
))?));
return Ok(PyErr::from_value(
py.import("builtins")?.getattr("OSError")?.call1((
errno,
source.to_string(),
path.into_pyobject(py)?.call_method0("__fspath__")?,
))?,
));
}
raise_public(py, classify(error), model, provider, None)
}

View file

@ -1,12 +1,11 @@
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};
use litellm_core::ocr::{OcrCall, OcrHostOperation, OcrHostResult, OcrProjectedRequest};
use litellm_python_interop::{
from_py_preserving_errors as from_py, to_py_preserving_errors as to_py,
};
@ -14,32 +13,50 @@ use litellm_python_interop::{
use super::document::PythonFileReader;
use super::errors::to_pyerr as ocr_error_to_pyerr;
use super::project::project;
use super::{ASYNC_SIGNATURE, SIGNATURE, callbacks, errors};
use super::{callbacks, errors};
use crate::auth::PythonTokenProvider;
use crate::lifecycle::{OperationClass, PythonCallState, PythonRoute, missing_state, now};
use crate::lifecycle::{PythonCallState, PythonRoute, Signature, missing_state};
use crate::marshal::Projection;
pub(super) struct PythonOcrHost {
state: PythonCallState,
projected: Option<ProjectedOcrHost>,
signature: &'static Signature,
retained: Option<OcrRetained>,
}
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 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 CapturedOcrPayload {
body: Py<PyDict>,
headers: Py<PyDict>,
pub(super) struct PythonPayload {
pub body: Py<PyDict>,
pub headers: Py<PyDict>,
}
impl CapturedOcrPayload {
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)
}
fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> {
visit.call(&self.body)?;
visit.call(&self.headers)
@ -47,52 +64,36 @@ impl CapturedOcrPayload {
}
impl PythonOcrHost {
pub(super) fn new(state: PythonCallState) -> Self {
pub(super) fn new(state: PythonCallState, signature: &'static Signature) -> Self {
Self {
state,
projected: None,
signature,
retained: None,
}
}
fn projected(&self) -> PyResult<&ProjectedOcrHost> {
self.projected.as_ref().ok_or_else(missing_state)
fn retained(&self) -> PyResult<&OcrRetained> {
self.retained.as_ref().ok_or_else(missing_state)
}
fn projected_mut(&mut self) -> PyResult<&mut ProjectedOcrHost> {
self.projected.as_mut().ok_or_else(missing_state)
fn retained_mut(&mut self) -> PyResult<&mut OcrRetained> {
self.retained.as_mut().ok_or_else(missing_state)
}
fn project(&mut self, py: Python<'_>) -> PyResult<OcrHostResult> {
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,
))))
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,
})))
}
fn acquire_azure_ad_token(&self, py: Python<'_>) -> PyResult<ResolvedCredential> {
self.projected()?
self.retained()?
.azure_ad_token_provider
.as_ref()
.ok_or_else(missing_state)?
@ -102,42 +103,21 @@ impl PythonOcrHost {
fn during_call(
&mut self,
py: Python<'_>,
mut request: OcrDuringCallRequest,
request: OcrDuringCallRequest,
) -> PyResult<OcrDuringCallRequest> {
let projected = self.projected()?;
let pre_call = projected.pre_call.as_ref().ok_or_else(missing_state)?;
let retained = self.retained()?;
let logger = self.state.logger()?;
logger.update_ocr(
callbacks::update_logging(
py,
logger,
&self.state.kwargs,
pre_call,
&projected.secret_fields,
&request.url,
&request,
&retained.secret_fields,
)?;
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(),
});
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);
Ok(request)
}
@ -146,27 +126,24 @@ impl PythonOcrHost {
py: Python<'_>,
request: OcrPostCallRequest,
) -> PyResult<OcrPostCallRequest> {
let projected = self.projected()?;
let payload = projected.payload.as_ref();
self.state.logger()?.post_ocr(
let payload = self.retained()?.payload.as_ref().ok_or_else(missing_state)?;
callbacks::post_call(
py,
self.state.logger()?,
&request.original_response,
payload.map(|payload| &payload.body),
payload.map(|payload| &payload.headers),
payload,
)?;
Ok(request)
}
fn map_failure(&mut self, py: Python<'_>, error: litellm_core::ocr::Error) -> PyResult<()> {
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),
self.state.finish(py)?;
let (model, provider) = match &self.retained {
Some(retained) => (retained.model.as_str(), retained.provider),
None => ("", ""),
};
let mapped = match self.state.error.take() {
Some(host_error) if self.projected.as_ref().is_some_and(|host| host.reader_failed) => {
Some(host_error) if matches!(error, litellm_core::ocr::Error::HostDocumentRead) => {
PyErr::from_value(host_error.into_bound(py).into_any())
}
Some(host_error) => errors::public_host_exception(py, &host_error, model, provider)?,
@ -188,10 +165,8 @@ impl PythonRoute for PythonOcrHost {
&mut self.state
}
fn classify(operation: &OcrHostOperation) -> OperationClass {
operation
.phase()
.map_or(OperationClass::Route, OperationClass::Phase)
fn phase(operation: &OcrHostOperation) -> Option<HostPhase> {
operation.phase()
}
fn lifecycle_result() -> OcrHostResult {
@ -210,16 +185,24 @@ impl PythonRoute for PythonOcrHost {
Ok(match operation {
OcrHostOperation::ProjectRequest => self.project(py)?,
OcrHostOperation::ReadDocument => {
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?))
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),
}
}
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) => {
@ -229,7 +212,7 @@ impl PythonRoute for PythonOcrHost {
OcrHostResult::PostCall(Ok(self.post_call(py, request)?))
}
OcrHostOperation::ConstructResponse(response) => {
self.state.end = Some(now(py)?);
self.state.finish(py)?;
self.state.response = Some(callbacks::response(py, response.as_ref())?);
OcrHostResult::Lifecycle(Ok(()))
}
@ -244,30 +227,22 @@ impl PythonRoute for PythonOcrHost {
}
fn cleanup(&mut self) {
self.projected = None;
self.retained = None;
}
fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> {
let Some(projected) = &self.projected else {
let Some(retained) = &self.retained else {
return Ok(());
};
if let Some(provider) = &projected.azure_ad_token_provider {
if let Some(provider) = &retained.azure_ad_token_provider {
provider.traverse(visit)?;
}
if let Some(reader) = &projected.reader {
if let Some(reader) = &retained.reader {
reader.traverse(visit)?;
}
if let Some(payload) = &projected.payload {
if let Some(payload) = &retained.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,54 +10,34 @@ 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::document::{FileDocumentInput, PythonFileReader};
use crate::auth::PythonTokenProvider;
use super::host::OcrRetained;
use crate::lifecycle::BoundArguments;
use crate::marshal::BoundRouteInputs;
use crate::marshal::{BoundRouteInputs, Projection};
/// Positional parameters of `ocr()` that are never projected into
/// `optional_params`.
const BOUND_FIELDS: &[&str] = &["model", "document", "timeout", "input_sources"];
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> {
fn project_document(document: &Bound<'_, PyAny>) -> PyResult<Result<FileDocumentInput, litellm_core::ocr::Error>> {
let kind: String = document.get_item("type")?.extract()?;
if kind != "file" {
return Ok(ProjectedDocument::Url(from_py(document)?));
let value: serde_json::Value = from_py(document)?;
return Ok(OcrDocument::try_from(value).map(|document| FileDocumentInput {
input: document.into(),
reader: None,
}));
}
document.extract().map(ProjectedDocument::File)
document.extract().map(Ok)
}
/// Pure core assembly; every failure here is a typed `ocr::Error`.
fn build_request(
inputs: BoundRouteInputs,
document: ProjectedDocument,
) -> Result<ProjectedOcrCall, litellm_core::ocr::Error> {
let document = document.into_native()?;
document: Result<FileDocumentInput, litellm_core::ocr::Error>,
) -> Result<Projection<LiteLLMOcrRequest, OcrRetained>, litellm_core::ocr::Error> {
let document = document?;
let BoundRouteInputs {
model,
custom_llm_provider,
@ -83,18 +63,23 @@ fn build_request(
input_sources,
},
)?;
Ok(ProjectedOcrCall {
request,
azure_ad_token_provider,
secret_fields,
reader: document.reader,
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,
})
}
pub(super) fn project(
py: Python<'_>,
arguments: &BoundArguments<'_>,
) -> PyResult<ProjectedOcrCall> {
) -> PyResult<Projection<LiteLLMOcrRequest, OcrRetained>> {
let model: String = arguments.extract("model")?;
let custom_llm_provider: Option<String> = arguments.optional("custom_llm_provider")?;
let document = project_document(&arguments.required("document")?)?;
@ -142,7 +127,7 @@ mod tests {
)
.unwrap();
assert!(matches!(
project_document(&file).unwrap().into_native().unwrap().input,
project_document(&file).unwrap().unwrap().input,
litellm_core::ocr::OcrDocumentInput::Bytes { bytes, mime_type, .. }
if bytes == b"%PDF-1.4"[..] && mime_type.as_deref() == Some("application/pdf")
));
@ -155,7 +140,7 @@ mod tests {
)
.unwrap();
assert!(matches!(
project_document(&original).unwrap().into_native().unwrap().input,
project_document(&original).unwrap().unwrap().input,
litellm_core::ocr::OcrDocumentInput::Document(OcrDocument::DocumentUrl { document_url, .. })
if document_url == "https://example.com/a.pdf"
));
@ -169,7 +154,10 @@ mod tests {
let document = py
.eval(c"{'type': 'mystery', 'mystery': 'x'}", None, None)
.unwrap();
let error = project_document(&document).unwrap().into_native().err().unwrap();
let error = project_document(&document)
.unwrap()
.err()
.unwrap();
assert!(error.to_string().contains("document"));
});
}
@ -181,14 +169,16 @@ 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)
);
@ -202,8 +192,9 @@ 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)
@ -232,8 +223,11 @@ document = Document()
",
);
let document = locals.get_item("document").unwrap().unwrap();
let projected = project_document(&document).unwrap().into_native().unwrap();
assert!(matches!(projected.input, litellm_core::ocr::OcrDocumentInput::Bytes { .. }));
let projected = project_document(&document).unwrap().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

@ -4,7 +4,7 @@ from typing import Final, cast # noqa: TID251 # native binding selects a sync
from litellm.llms.base_llm.ocr.transformation import OCRResponse
from litellm.ocr import legacy
from litellm.ocr.legacy import convert_file_document_to_url_document, get_mime_type
from litellm.rust_bridge.bindings import native_exception_types
from litellm.rust_bridge.bindings import native_decline_types
from litellm.rust_bridge.configuration import rust_ocr_enabled
from litellm.rust_bridge.ocr import NATIVE_AOCR, NATIVE_OCR
@ -19,7 +19,7 @@ def ocr(
if native is not None:
try:
return native(*args, **kwargs)
except _decline_types():
except native_decline_types():
pass
fallback: Final = cast( # cast-ok: forward the original call shape through the legacy @client decorator
Callable[..., OCRResponse | Coroutine[object, object, OCRResponse]], legacy.ocr
@ -32,14 +32,10 @@ async def aocr(*args: object, **kwargs: object) -> OCRResponse: # kwargs-ok: pr
if native is not None:
try:
return await native(*args, **kwargs)
except _decline_types():
except native_decline_types():
pass
fallback: Final = cast( # cast-ok: forward the original call shape through the legacy @client decorator
Callable[..., Awaitable[OCRResponse]], legacy.aocr
)
return await fallback(*args, **kwargs)
def _decline_types() -> tuple[type[BaseException], ...]:
exception_types: Final = native_exception_types()
return (exception_types[0],) if exception_types is not None else ()

View file

@ -29,10 +29,12 @@ def _build_document_from_upload(
filename: str | None,
content_type: str | None,
) -> dict[str, str]:
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)
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
)
return convert_file_document_to_url_document(
{

View file

@ -47,3 +47,8 @@ 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
from collections.abc import Callable, Coroutine, Mapping, Sequence
from types import MappingProxyType
from typing import Final, Protocol, cast # noqa: TID251 # validates dynamically loaded native callables
@ -8,6 +8,84 @@ 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 {
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={
"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},
**{
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,3 +33,21 @@ 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
from litellm.rust_bridge.ocr import NATIVE_AOCR, NATIVE_OCR, post_call, pre_call, update_logging
@pytest.fixture(autouse=True)
@ -21,6 +21,72 @@ 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,14 +1,15 @@
import json
import asyncio
import base64
import gc
import json
import os
import threading
import weakref
import gc
from pathlib import Path
from collections.abc import Generator
from collections.abc import Coroutine, 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
@ -340,7 +341,12 @@ 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, tmp_path, asynchronous, source):
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:
server, requests = ocr_server
path: Final = tmp_path / "scan.png"
content: Final = b"document bytes"
@ -388,9 +394,12 @@ async def test_native_file_sources_preserve_sdk_behavior(ocr_server, tmp_path, a
"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
assert requests[0]["body"]["document"] == {
body: Final = requests[0]["body"]
assert isinstance(body, dict)
assert body["document"] == {
"type": "image_url",
"image_url": "data:image/png;base64," + base64.b64encode(content).decode(),
}
@ -400,7 +409,11 @@ async def test_native_file_sources_preserve_sdk_behavior(ocr_server, tmp_path, a
@pytest.mark.parametrize("asynchronous", [False, True])
@pytest.mark.asyncio
async def test_native_file_failures_preserve_identity_and_do_not_send(ocr_server, tmp_path, asynchronous):
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:
from litellm.rust_bridge import _native
server, requests = ocr_server
@ -416,37 +429,47 @@ async def test_native_file_failures_preserve_identity_and_do_not_send(ocr_server
reader: Final = Reader()
path: Final = tmp_path / "missing.pdf"
kwargs: Final = {
"model": "mistral/mistral-ocr-latest",
"api_key": "test-key",
"api_base": f"http://127.0.0.1:{server.server_port}",
}
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}",
)
for source, error in ((reader, KeyError), (path, FileNotFoundError)):
with pytest.raises(error) as caught:
if asynchronous:
await _native.aocr(document={"type": "file", "file": source}, **kwargs)
else:
_native.ocr(document={"type": "file", "file": source}, **kwargs)
await invoke(source)
if source is reader:
assert caught.value is failure
else:
assert caught.value.filename == str(path)
assert isinstance(caught.value, FileNotFoundError)
assert getattr(caught.value, "filename") == str(path)
assert reader.calls == 1
assert requests == []
def test_unstarted_native_file_call_does_not_read_and_releases_reader():
def test_unstarted_native_file_call_does_not_read_and_releases_reader() -> None:
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():
def create_call() -> tuple[Coroutine[object, object, OCRResponse], weakref.ReferenceType[Reader]]:
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)
@ -457,6 +480,49 @@ def test_unstarted_native_file_call_does_not_read_and_releases_reader():
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