move io to core

This commit is contained in:
Yujong Lee 2026-09-15 21:52:27 -07:00
parent 5a7ea7ef71
commit 75548d30f5
30 changed files with 1319 additions and 716 deletions

View file

@ -2006,6 +2006,7 @@ dependencies = [
name = "litellm-python-bridge"
version = "0.1.0"
dependencies = [
"bytes",
"criterion",
"futures-util",
"litellm-auth",

View file

@ -211,10 +211,10 @@ mod tests {
lifecycle.accept::<crate::ocr::Error>(Ok(()));
}
let selected = crate::ocr::Error::InvalidRequest("provider".into());
assert_eq!(
assert!(matches!(
lifecycle.accept(Err(HostFailure::Error(selected.clone()))),
Some(selected)
);
Some(crate::ocr::Error::InvalidRequest(message)) if message == "provider"
));
lifecycle.accept::<crate::ocr::Error>(Ok(()));
for phase in [
HostPhase::DeploymentFailure,
@ -222,11 +222,10 @@ mod tests {
HostPhase::AsyncFailure,
] {
assert_eq!(lifecycle.phase(), phase);
assert_eq!(
assert!(
lifecycle.accept(Err(HostFailure::Error(crate::ocr::Error::InvalidRequest(
"callback".into()
)))),
None
)))).is_none()
);
}
assert_eq!(lifecycle.phase(), HostPhase::Complete);
@ -236,10 +235,10 @@ mod tests {
fn cancellation_skips_terminal_dispatch() {
let mut lifecycle = HostLifecycle::new(true);
let error = crate::ocr::Error::InvalidRequest("cancelled".into());
assert_eq!(
assert!(matches!(
lifecycle.accept(Err(HostFailure::Cancelled(error.clone()))),
Some(error)
);
Some(crate::ocr::Error::InvalidRequest(message)) if message == "cancelled"
));
assert_eq!(lifecycle.phase(), HostPhase::Complete);
}
}

View file

@ -99,7 +99,6 @@ impl BaseOcrConfig for AzureAICohereParseConfig {
CohereParseConfig.transform_ocr_response(model, raw_response, request_format)
}
fn validate_request_body(&self, body: &Value) -> Result<(), crate::ocr::Error> {
let document = crate::ocr::prepare::body_document(body)?;
validate_document(&document)?;

View file

@ -541,7 +541,6 @@ impl BaseOcrConfig for AzureDocumentIntelligenceOCRConfig {
) -> Result<DocumentIntelligenceRequest, crate::ocr::Error> {
build_request(document)
}
}
impl AzureDocumentIntelligenceOCRConfig {
@ -813,11 +812,11 @@ mod tests {
&base,
json!({"pages":[2,0,0,1],"features":["keyValuePairs","languages"], "future_option": {"nested":null}, "extra_body":{"provider_option":false}}),
);
request.document = serde_json::from_value(json!({
request.document = serde_json::from_value::<OcrDocument>(json!({
"type":"document_url",
"document_url":"https://example.com/document.pdf"
}))
.unwrap();
.unwrap().into();
perform_ocr(request).await.unwrap();
server.await.unwrap();

View file

@ -100,7 +100,6 @@ impl BaseOcrConfig for AzureAIOCRConfig {
MistralOCRConfig.transform_ocr_response(model, raw_response, request_format)
}
fn validate_request_body(&self, body: &Value) -> Result<(), crate::ocr::Error> {
validate_inline_document(&crate::ocr::prepare::body_document(body)?)
}

View file

@ -349,7 +349,7 @@ mod tests {
#[tokio::test]
async fn composed_body_preserves_native_document_fields_and_untyped_overrides() {
let mut request = crate::ocr::test_support::wire_request(
let request = crate::ocr::test_support::wire_request(
"cohere/parse",
"https://example.com",
json!({
@ -361,10 +361,10 @@ mod tests {
}
}),
);
request.document = serde_json::from_value(json!({
let request = request.with_document(serde_json::from_value(json!({
"type":"image_url","image_url":"https://example.com/original.png"
}))
.unwrap();
.unwrap());
let request = crate::ocr::prepare::prepare_request(request);
let http = CohereParseConfig
.prepare_request(&request, &crate::ocr::test_support::ocr_client())
@ -461,12 +461,11 @@ mod tests {
"pages":[{"markdown":{"images":[{"image_base64":42}]}}]
}))
.unwrap();
assert_eq!(
assert!(matches!(
normalize_response("parse", response).unwrap_err(),
crate::ocr::Error::ResponseField {
path: "pages[0].markdown.images[0].image_base64".into()
}
);
crate::ocr::Error::ResponseField { path }
if path == "pages[0].markdown.images[0].image_base64"
));
}
#[test]
@ -504,13 +503,10 @@ mod tests {
"https://example.com",
json!({"output_format":null,"req_format":null}),
);
let request = crate::ocr::types::LiteLLMOcrRequest {
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(),
..request
};
.unwrap());
assert_eq!(
request.response_format().unwrap(),
crate::ocr::types::OcrResponseFormat::Litellm
@ -677,10 +673,10 @@ mod tests {
json!({"type":"image_url","image_url":""}),
json!({"type":"image_url","image_url":"data:application/pdf;base64,YQ=="}),
] {
assert_eq!(
assert!(matches!(
validate_document(&serde_json::from_value(value).unwrap()),
Err(crate::ocr::Error::CohereImageOnly)
);
));
}
assert!(serde_json::from_value::<CohereOptions>(json!({"output_format":"html"})).is_err());
for format in ["markdown", "blocks"] {

View file

@ -206,12 +206,10 @@ mod tests {
#[test]
fn explicit_null_model_does_not_use_the_missing_model_default() {
let response = serde_json::from_value(json!({"model":null})).unwrap();
assert_eq!(
assert!(matches!(
normalize_response("fallback", response).unwrap_err(),
crate::ocr::Error::ResponseField {
path: "model".into()
}
);
crate::ocr::Error::ResponseField { path } if path == "model"
));
}
#[test]
@ -241,10 +239,10 @@ mod tests {
false,
)
.unwrap_err();
assert_eq!(
assert!(matches!(
error,
crate::ocr::Error::ResponseField { path: path.into() }
);
crate::ocr::Error::ResponseField { path: actual } if actual == path
));
}
}

View file

@ -738,8 +738,7 @@ mod tests {
"result":{"chunks":[]}
}))])
.await;
let mut request = wire_request(model, &base, options);
request.document = request.document.with_source(source.into());
let request = crate::ocr::test_support::with_source(wire_request(model, &base, options), source);
perform_ocr(request).await.unwrap();
server.await.unwrap();
@ -857,8 +856,8 @@ mod tests {
#[case("data:application/pdf;base64,INVALID!")]
#[tokio::test]
async fn rejects_invalid_document_sources_before_network(#[case] source: &str) {
let mut request = wire_request("reducto/parse-v3", "http://127.0.0.1:1", json!({}));
request.document = request.document.with_source(source.into());
let request = crate::ocr::test_support::with_source(
wire_request("reducto/parse-v3", "http://127.0.0.1:1", json!({})), source);
assert!(perform_ocr(request).await.is_err());
}
@ -913,8 +912,8 @@ mod tests {
async fn facade_omits_native_response_by_default_and_preserves_auth_priority() {
let raw = json!({"job_id":"job-1","result":{"chunks":[]}});
let (base, seen, server) = mock_server(vec![MockResponse::json(raw)]).await;
let mut request = wire_request("reducto/parse-v3", &base, json!({}));
request.document = request.document.with_source("reducto://ready.pdf".into());
let mut request = crate::ocr::test_support::with_source(
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();

View file

@ -192,7 +192,6 @@ impl BaseOcrConfig for VertexAIDeepSeekOCRConfig {
.collect(),
})
}
}
pub(crate) fn normalize_response(
@ -612,7 +611,7 @@ mod tests {
"usage":{"prompt_tokens":1}
}))])
.await;
let mut request = wire_request(
let request = wire_request(
"vertex_ai/deepseek-ocr-maas",
&base,
json!({
@ -623,9 +622,7 @@ mod tests {
"extra_body":{"provider_option":"value"}
}),
);
request.document = request
.document
.with_source("gs://bucket/document.pdf".into());
let request = crate::ocr::test_support::with_source(request, "gs://bucket/document.pdf");
let response = perform_ocr(request).await.unwrap();
server.await.unwrap();

View file

@ -109,7 +109,6 @@ impl BaseOcrConfig for VertexAIOCRConfig {
MistralOCRConfig.transform_ocr_response(model, raw_response, request_format)
}
fn validate_request_body(&self, body: &Value) -> Result<(), crate::ocr::Error> {
validate_inline_document(&crate::ocr::prepare::body_document(body)?)
}
@ -339,8 +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(direct);
let vertex = crate::ocr::prepare::prepare_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

@ -2,11 +2,28 @@ use base64::{Engine, engine::general_purpose::STANDARD};
use data_url::mime::Mime;
use data_url::{DataUrl, DataUrlError, forgiving_base64::DecodeError};
use reqwest::Url;
use std::io::Read;
use std::path::Path;
use super::types::{OcrConnection, OcrDocument};
use crate::constants::{OCR_INLINE_MAX_BYTES, OCR_MAX_FETCH_REDIRECTS};
use crate::media::{DownloadPolicy, MediaFetcher};
pub(crate) fn read_path_document(
path: &Path,
mime_type: Option<&str>,
) -> 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))
.map_err(|source| super::Error::FileRead {
path: path.to_owned(),
source: std::sync::Arc::new(source),
})?;
let name = path.file_name().map(|name| name.to_string_lossy());
encode_file_document(&bytes, name.as_deref(), mime_type)
}
pub fn encode_file_document(
bytes: &[u8],
file_name: Option<&str>,
@ -182,6 +199,27 @@ fn map_media_error(error: crate::media::Error) -> crate::ocr::Error {
mod tests {
use super::*;
#[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 error = read_path_document(&path, None).unwrap_err();
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);
assert_eq!(source.kind(), std::io::ErrorKind::NotFound);
let file = std::fs::File::create(&path).unwrap();
file.set_len(OCR_INLINE_MAX_BYTES as u64 + 1).unwrap();
let oversized = read_path_document(&path, None);
std::fs::write(&path, b"image bytes").unwrap();
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="));
}
fn document(source: &str) -> OcrDocument {
OcrDocument::DocumentUrl {
document_url: source.into(),
@ -248,10 +286,10 @@ mod tests {
#[test]
fn file_encoding_enforces_decoded_size_limit() {
let bytes = vec![b'a'; OCR_INLINE_MAX_BYTES + 1];
assert_eq!(
assert!(matches!(
encode_file_document(&bytes, None, None),
Err(crate::ocr::Error::InlineDocumentTooLarge)
);
));
let document = encode_file_document(&bytes[..OCR_INLINE_MAX_BYTES], None, None).unwrap();
let inline = InlineDocument::parse(document.source()).unwrap().unwrap();
assert_eq!(
@ -282,10 +320,10 @@ mod tests {
] {
let inline = InlineDocument::parse(source).unwrap().unwrap();
assert_eq!(inline.decode(expected.len()).unwrap(), expected);
assert_eq!(
assert!(matches!(
inline.decode(expected.len() - 1),
Err(crate::ocr::Error::InlineDocumentTooLarge)
);
));
}
}

View file

@ -1,4 +1,4 @@
#[derive(Clone, Debug, PartialEq, Eq, thiserror::Error)]
#[derive(Clone, Debug, thiserror::Error)]
pub enum Error {
#[error("upstream OCR error ({status}): {body}")]
Provider {
@ -8,6 +8,14 @@ pub enum Error {
},
#[error("File is empty or could not be read")]
EmptyFile,
#[error("Failed to read OCR file {}: {source}", path.display())]
FileRead {
path: std::path::PathBuf,
#[source]
source: std::sync::Arc<std::io::Error>,
},
#[error("OCR document preparation task failed: {0}")]
DocumentTask(#[source] std::sync::Arc<tokio::task::JoinError>),
#[error("Invalid MIME type: {0}")]
InvalidMimeType(String),
#[error(

View file

@ -2,13 +2,13 @@ use std::sync::Arc;
use super::OcrClient;
use super::hooks::{OcrHooks, OcrLifecycleHooks, OcrPostCallRequest};
use super::types::{LiteLLMOcrRequest, LiteLLMOcrResponse, PreparedOcrRequest};
use super::types::{ResolvedOcrRequest, LiteLLMOcrResponse, PreparedOcrRequest};
use crate::call_lifecycle::{CallLifecycle, CallLifecycleContext};
use crate::llms::base_llm::ocr::transformation::OcrResponseContext;
pub(crate) async fn perform_ocr_request(
client: &OcrClient,
request: LiteLLMOcrRequest,
request: ResolvedOcrRequest,
) -> Result<LiteLLMOcrResponse, super::Error> {
request.response_format()?;
let context = CallLifecycleContext::new(
@ -43,7 +43,7 @@ pub(crate) struct PreparedOcrCall {
impl PreparedOcrCall {
pub(crate) async fn prepare(
client: OcrClient,
request: LiteLLMOcrRequest,
request: ResolvedOcrRequest,
) -> Result<Self, super::Error> {
let request = super::prepare::prepare_request(request);
let http = request.config.prepare_request(&request, &client).await?;

View file

@ -2,7 +2,7 @@ use std::future::Future;
use std::pin::Pin;
use std::sync::Arc;
use super::types::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrDocument};
use super::types::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrDocument, ResolvedOcrRequest};
use crate::call_lifecycle::{CallLifecycleContext, CallLifecycleHooks, CallLifecycleTiming};
use serde::Serialize;
use serde_json::Value;
@ -75,19 +75,19 @@ pub(crate) struct OcrLifecycleHooks {
pub provider_name: String,
}
impl CallLifecycleHooks<LiteLLMOcrRequest, LiteLLMOcrRequest, LiteLLMOcrResponse>
impl CallLifecycleHooks<ResolvedOcrRequest, ResolvedOcrRequest, LiteLLMOcrResponse>
for OcrLifecycleHooks
{
type Error = super::Error;
type PreCallFuture<'a> = OcrHookFuture<'a, LiteLLMOcrRequest>;
type DuringCallFuture<'a> = OcrHookFuture<'a, LiteLLMOcrRequest>;
type PreCallFuture<'a> = OcrHookFuture<'a, ResolvedOcrRequest>;
type DuringCallFuture<'a> = OcrHookFuture<'a, ResolvedOcrRequest>;
type SuccessFuture<'a> = OcrLogFuture<'a>;
type FailureFuture<'a> = OcrLogFuture<'a>;
fn async_pre_call_hook<'a>(
&'a self,
_context: &'a CallLifecycleContext,
request: LiteLLMOcrRequest,
request: ResolvedOcrRequest,
) -> Self::PreCallFuture<'a> {
Box::pin(async move {
if !self.hooks.intercepts_requests() {
@ -118,7 +118,7 @@ impl CallLifecycleHooks<LiteLLMOcrRequest, LiteLLMOcrRequest, LiteLLMOcrResponse
fn async_during_call_hook<'a>(
&'a self,
_context: &'a CallLifecycleContext,
request: LiteLLMOcrRequest,
request: ResolvedOcrRequest,
) -> Self::DuringCallFuture<'a> {
Box::pin(async move { Ok(request) })
}

View file

@ -9,7 +9,8 @@ use super::hooks::{
OcrDuringCallRequest, OcrHookFuture, OcrHooks, OcrLogFuture, OcrPostCallRequest,
OcrPreCallRequest,
};
use super::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrClient};
use super::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrClient, OcrDocumentInput, OcrFileContent};
use super::types::ResolvedOcrRequest;
use crate::call_lifecycle::host::{
HostCall, HostCallFuture, HostCallStep, HostFailure, HostLifecycle, HostPhase,
};
@ -50,6 +51,7 @@ impl OcrAdmission {
#[derive(Clone, Debug)]
pub enum OcrHostOperation {
ProjectRequest,
ReadDocument,
Lifecycle(HostPhase),
ConstructResponse(Arc<LiteLLMOcrResponse>),
MapFailure(super::Error),
@ -82,6 +84,7 @@ impl OcrHostOperation {
pub enum OcrHostResult {
Request(Result<(Box<LiteLLMOcrRequest>, bool), super::Error>),
Document(Result<OcrFileContent, super::Error>),
Lifecycle(Result<(), HostFailure<super::Error>>),
AzureAdToken(Result<ResolvedCredential, litellm_auth::Error>),
PreCall(Result<OcrPreCallRequest, super::Error>),
@ -179,6 +182,7 @@ impl OcrCall {
if self.lifecycle.phase() == HostPhase::Execute {
if self.execution.request.is_none()
&& self.execution.execution.is_none()
&& self.execution.preparation.is_none()
&& !self.execution.completed
{
self.projecting = true;
@ -322,6 +326,8 @@ struct PendingOperation {
struct OcrExecution {
client: Option<OcrClient>,
request: Option<LiteLLMOcrRequest>,
preparation: Option<tokio::task::JoinHandle<Result<ResolvedOcrRequest, super::Error>>>,
reading: bool,
operations_tx: mpsc::UnboundedSender<PendingOperation>,
operations_rx: mpsc::UnboundedReceiver<PendingOperation>,
pending_result: Option<oneshot::Sender<OcrHostResult>>,
@ -337,6 +343,8 @@ impl OcrExecution {
Self {
client: Some(client),
request: None,
preparation: None,
reading: false,
operations_tx,
operations_rx,
pending_result: None,
@ -356,11 +364,35 @@ impl OcrExecution {
"OCR call cannot be resumed after completion".into(),
));
}
let result = if self.reading {
let Some(OcrHostResult::Document(content)) = result else {
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 OcrDocumentInput::HostReader { mime_type } = &request.document else {
unreachable!("only host readers request document reads");
};
let document = OcrDocumentInput::Bytes {
bytes: content.bytes,
file_name: content.file_name,
mime_type: mime_type.clone(),
};
self.request = Some(request.with_document(document));
None
} else {
result
};
match (self.pending_result.take(), result) {
(Some(sender), Some(result)) => sender.send(result).map_err(|_| {
super::Error::InvalidRequest("OCR host operation was abandoned".into())
})?,
(None, None) if self.execution.is_none() => self.start(),
(None, None) if self.execution.is_none() => {
if let Some(operation) = self.prepare().await? {
return Ok(OcrCallStep::Host(operation));
}
}
(Some(sender), None) => {
self.pending_result = Some(sender);
return Err(super::Error::InvalidRequest(
@ -394,12 +426,41 @@ impl OcrExecution {
}
}
fn start(&mut self) {
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");
if matches!(request.document, OcrDocumentInput::HostReader { .. }) {
self.request = Some(request);
self.reading = true;
return Ok(Some(OcrHostOperation::ReadDocument));
}
if let OcrDocumentInput::Document(document) = &request.document {
let document = document.clone();
self.start(request.with_document(document));
return Ok(None);
}
self.preparation = Some(tokio::task::spawn_blocking(move || {
let document = match &request.document {
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())?
}
_ => unreachable!("only native file inputs require preparation"),
};
Ok(request.with_document(document))
}));
}
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);
Ok(None)
}
fn start(&mut self, mut request: ResolvedOcrRequest) {
let client = self.client.take().expect("admitted OCR call has a client");
let mut request = self
.request
.take()
.expect("admitted OCR call has a request");
let intercepts_requests = request.hooks.intercepts_requests();
if self.azure_ad_token_provider {
request.azure_ad_token_provider = Some(TokenProviderHandle::new(Arc::new(
@ -427,6 +488,11 @@ impl OcrExecution {
async fn stop(&mut self) {
self.cancel();
if let Some(preparation) = self.preparation.as_mut() {
let _ = preparation.await;
}
self.preparation = None;
self.request = None;
if let Some(execution) = self.execution.as_mut() {
let _ = execution.await;
}
@ -439,6 +505,9 @@ impl Drop for OcrExecution {
if let Some(execution) = &self.execution {
execution.abort();
}
if let Some(preparation) = &self.preparation {
preparation.abort();
}
}
}
@ -580,6 +649,9 @@ impl OcrHost for NoopOcrHost {
OcrHostOperation::ProjectRequest => OcrHostResult::Request(Err(
super::Error::InvalidRequest("OCR host has no request projection".into()),
)),
OcrHostOperation::ReadDocument => OcrHostResult::Document(Err(
super::Error::InvalidRequest("OCR host has no document reader".into()),
)),
OcrHostOperation::Lifecycle(_)
| OcrHostOperation::ConstructResponse(_)
| OcrHostOperation::MapFailure(_)
@ -615,6 +687,9 @@ impl OcrHost for OcrHookHost {
OcrHostOperation::ProjectRequest => OcrHostResult::Request(Err(
super::Error::InvalidRequest("OCR hook host has no request projection".into()),
)),
OcrHostOperation::ReadDocument => OcrHostResult::Document(Err(
super::Error::InvalidRequest("OCR hook host has no document reader".into()),
)),
OcrHostOperation::Success {
context,
response,
@ -671,6 +746,104 @@ mod tests {
OcrDecline, OcrDocument, OcrHost, OcrHostOperation, OcrHostResult,
};
#[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 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 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);
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);
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();
assert_eq!(body["document"], expected);
}
#[tokio::test]
async fn core_requests_a_host_read_once_before_pre_call() {
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);
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 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,
});
}
OcrCallStep::Complete(_) => break,
}
}
server.await.unwrap();
assert_eq!((reads, pre_calls, seen.lock().unwrap().len()), (1, 1, 1));
}
#[tokio::test]
async fn cancellation_acknowledges_blocking_preparation_completion() {
use std::sync::atomic::{AtomicBool, Ordering};
let mut execution = super::OcrExecution::new(crate::ocr::test_support::ocr_client());
let finished = Arc::new(AtomicBool::new(false));
let completed = finished.clone();
let (entered_tx, entered_rx) = tokio::sync::oneshot::channel();
let (release_tx, release_rx) = std::sync::mpsc::channel();
execution.preparation = Some(tokio::task::spawn_blocking(move || {
entered_tx.send(()).unwrap();
release_rx.recv().unwrap();
completed.store(true, Ordering::SeqCst);
Err(crate::ocr::Error::EmptyFile)
}));
entered_rx.await.unwrap();
let mut stop = Box::pin(execution.stop());
std::future::poll_fn(|cx| {
assert!(stop.as_mut().poll(cx).is_pending());
std::task::Poll::Ready(())
}).await;
assert!(!finished.load(Ordering::SeqCst));
release_tx.send(()).unwrap();
stop.await;
assert!(finished.load(Ordering::SeqCst));
assert!(execution.preparation.is_none());
}
#[test]
fn request_boundary_selects_mistral_and_rejects_unknown_providers() {
let document = OcrDocument::try_from(
@ -1139,7 +1312,7 @@ mod tests {
Ok(request)
}));
}
OcrHostOperation::PostCall(_) => 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"),
@ -1299,7 +1472,7 @@ mod tests {
OcrHostResult::Lifecycle(Err(HostFailure::Error(selected.clone())))
}
OcrHostOperation::Failure { error, .. } => {
assert_eq!(error, selected);
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()),
@ -1325,7 +1498,7 @@ mod tests {
}
};
server.await.unwrap();
assert_eq!(error, selected);
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);
}
@ -1364,7 +1537,7 @@ mod tests {
let selected = crate::ocr::Error::InvalidRequest("cancelled".into());
assert!(matches!(
call.interrupt(HostFailure::Cancelled(selected.clone())).await,
Err(error) if error == selected
Err(crate::ocr::Error::InvalidRequest(message)) if message == "cancelled"
));
assert!(
call.resume(Some(OcrHostResult::Lifecycle(Ok(()))))
@ -1605,7 +1778,7 @@ mod tests {
)
.await
.unwrap();
assert!(matches!(result, Err(error) if error == selected));
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

@ -22,8 +22,8 @@ pub use lifecycle::{
};
pub use provider_config::{get_api_key_env_var, get_health_check_document};
pub use types::{
LiteLLMOcrRequest, LiteLLMOcrResponse, OcrConnection, OcrCredentialInputs, OcrDocument,
OcrPage, OcrPageDimensions, OcrPageImage, OcrTransportConfig, OcrUsageInfo,
LiteLLMOcrRequest, LiteLLMOcrResponse, OcrConnection, OcrConnectionInputs, OcrCredentialInputs,
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::{LiteLLMOcrRequest, OcrConnection, OcrDocument, PreparedOcrRequest};
use super::types::{ResolvedOcrRequest, OcrConnection, OcrDocument, PreparedOcrRequest};
pub(crate) async fn transform_request_body<B>(
client: &OcrClient,
@ -111,7 +111,7 @@ pub(crate) fn credential_env(name: &str) -> Option<String> {
std::env::var(name).ok()
}
pub(crate) fn prepare_request(request: LiteLLMOcrRequest) -> PreparedOcrRequest {
pub(crate) fn prepare_request(request: ResolvedOcrRequest) -> PreparedOcrRequest {
use litellm_auth::{InputSource, Sourced};
let credentials = request.credentials.clone();

View file

@ -205,10 +205,10 @@ mod tests {
#[case("Mistral")]
#[case("unknown")]
fn invalid_provider_names_are_rejected(#[case] provider: &str) {
assert_eq!(
assert!(matches!(
resolve_provider_config("model", Some(provider)),
Err(crate::ocr::Error::InvalidProvider(provider.into()))
);
Err(crate::ocr::Error::InvalidProvider(value)) if value == provider
));
}
#[rstest]

View file

@ -5,7 +5,7 @@ use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::TcpListener;
use crate::ocr::{
LiteLLMOcrRequest, LiteLLMOcrResponse, OcrClient, OcrCredentialInputs, OcrDocument,
LiteLLMOcrRequest, LiteLLMOcrResponse, OcrClient, OcrConnectionInputs, OcrDocument,
};
pub(crate) fn ocr_client() -> OcrClient {
@ -23,7 +23,7 @@ pub(crate) async fn perform_ocr(
}
pub(crate) fn wire_request(model: &str, base: &str, options: Value) -> LiteLLMOcrRequest {
let request = LiteLLMOcrRequest::new(
LiteLLMOcrRequest::from_inputs(
model.into(),
OcrDocument::try_from(
json!({"type":"document_url","document_url":"data:application/pdf;base64,YWJj"}),
@ -31,23 +31,28 @@ pub(crate) fn wire_request(model: &str, base: &str, options: Value) -> LiteLLMOc
.unwrap(),
None,
options.as_object().unwrap().clone().into(),
OcrConnectionInputs {
api_key: Some("test-key".into()),
api_base: Some(base.into()),
timeout: Some(std::time::Duration::from_secs(2)),
..Default::default()
},
)
.unwrap();
let transport = request.transport.clone().with_overrides(
Vec::new(),
Default::default(),
Some(std::time::Duration::from_secs(2)),
);
request.with_connection_inputs(
OcrCredentialInputs::new(
Some("test-key".into()),
Default::default(),
Some(base.into()),
Default::default(),
),
transport,
Default::default(),
)
.unwrap()
}
pub(crate) fn resolved_request(request: LiteLLMOcrRequest) -> super::types::ResolvedOcrRequest {
let super::OcrDocumentInput::Document(document) = &request.document else {
panic!("wire fixture must contain a URL document");
};
let document = document.clone();
request.with_document(document)
}
pub(crate) fn with_source(request: LiteLLMOcrRequest, source: &str) -> LiteLLMOcrRequest {
let request = resolved_request(request);
let document = request.document.clone().with_source(source.into());
request.with_document(document.into())
}
pub(crate) struct MockResponse {

View file

@ -1,8 +1,10 @@
use std::collections::BTreeMap;
use std::path::PathBuf;
use std::sync::Arc;
use std::time::Duration;
use serde::{Deserialize, Serialize};
use bytes::Bytes;
use serde_json::{Map, Value};
use serde_with::serde_as;
@ -66,6 +68,40 @@ impl TryFrom<Value> for OcrDocument {
}
}
#[derive(Clone, Debug)]
pub enum OcrDocumentInput {
Document(OcrDocument),
Path {
path: PathBuf,
mime_type: Option<String>,
},
Bytes {
bytes: Bytes,
file_name: Option<String>,
mime_type: Option<String>,
},
HostReader {
mime_type: Option<String>,
},
}
impl From<OcrDocument> for OcrDocumentInput {
fn from(document: OcrDocument) -> Self {
Self::Document(document)
}
}
impl From<PathBuf> for OcrDocumentInput {
fn from(path: PathBuf) -> Self {
Self::Path { path, mime_type: None }
}
}
pub struct OcrFileContent {
pub bytes: Bytes,
pub file_name: Option<String>,
}
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "lowercase")]
pub enum OcrResponseFormat {
@ -143,6 +179,38 @@ fn nonblank(value: Option<String>) -> Option<String> {
.filter(|value| !value.is_empty())
}
/// Caller-supplied connection overrides for a [`LiteLLMOcrRequest`], in the
/// shape hosts receive them: JSON-ish headers, optional timeout, optional
/// credentials, and per-field provenance in `input_sources`.
#[derive(Clone, Debug, Default)]
pub struct OcrConnectionInputs {
pub api_key: Option<String>,
pub api_base: Option<String>,
pub extra_headers: Map<String, Value>,
pub timeout: Option<Duration>,
pub input_sources: BTreeMap<String, InputSource>,
}
impl OcrConnectionInputs {
fn source(&self, name: &str) -> InputSource {
self.input_sources.get(name).copied().unwrap_or_default()
}
fn header_pairs(&self) -> Result<Vec<(String, String)>, super::Error> {
self.extra_headers
.iter()
.map(|(name, value)| {
value
.as_str()
.map(|value| (name.clone(), value.to_string()))
.ok_or_else(|| super::Error::RequestField {
path: format!("extra_headers.{name}"),
})
})
.collect()
}
}
#[derive(Clone)]
pub struct OcrConnection {
pub api_key: Option<String>,
@ -199,9 +267,9 @@ pub(crate) struct ResolvedOcrCredentials {
pub api_base: Option<Sourced<String>>,
}
pub struct LiteLLMOcrRequest {
pub struct LiteLLMOcrRequest<D = OcrDocumentInput> {
pub model: String,
pub document: OcrDocument,
pub document: D,
pub credentials: OcrCredentialInputs,
pub transport: OcrTransportConfig,
pub hooks: Arc<dyn OcrHooks>,
@ -215,7 +283,7 @@ pub struct LiteLLMOcrRequest {
impl LiteLLMOcrRequest {
pub fn new(
model: String,
document: OcrDocument,
document: impl Into<OcrDocumentInput>,
custom_llm_provider: Option<&str>,
optional_params: CallArguments,
) -> Result<Self, super::Error> {
@ -245,7 +313,7 @@ impl LiteLLMOcrRequest {
Ok(Self {
model,
document,
document: document.into(),
credentials: OcrCredentialInputs::default(),
transport,
hooks: Arc::new(NoopOcrHooks),
@ -257,6 +325,24 @@ impl LiteLLMOcrRequest {
})
}
}
impl<D> LiteLLMOcrRequest<D> {
pub(crate) fn with_document<T>(self, document: T) -> LiteLLMOcrRequest<T> {
LiteLLMOcrRequest {
model: self.model,
document,
credentials: self.credentials,
transport: self.transport,
hooks: self.hooks,
litellm_call_id: self.litellm_call_id,
optional_params: self.optional_params,
input_sources: self.input_sources,
azure_ad_token_provider: self.azure_ad_token_provider,
config: self.config,
}
}
pub(crate) fn response_format(&self) -> Result<OcrResponseFormat, super::Error> {
self.optional_params
.get("req_format")
@ -297,8 +383,42 @@ impl LiteLLMOcrRequest {
..self
}
}
}
impl LiteLLMOcrRequest {
/// Builds a request from host-shaped inputs in one step: provider
/// resolution, optional-param validation, header/timeout overrides and
/// sourced credentials. Hosts should prefer this over sequencing
/// [`Self::new`], [`OcrTransportConfig::with_overrides`] and
/// [`Self::with_connection_inputs`] by hand.
pub fn from_inputs(
model: String,
document: impl Into<OcrDocumentInput>,
custom_llm_provider: Option<&str>,
optional_params: CallArguments,
connection: OcrConnectionInputs,
) -> Result<Self, super::Error> {
let request = Self::new(model, document, custom_llm_provider, optional_params)?;
let transport = request.transport.clone().with_overrides(
connection.header_pairs()?,
connection.source("extra_headers"),
connection.timeout,
);
let (api_key_source, api_base_source) =
(connection.source("api_key"), connection.source("api_base"));
let credentials = OcrCredentialInputs::new(
connection.api_key,
api_key_source,
connection.api_base,
api_base_source,
);
Ok(request.with_connection_inputs(credentials, transport, connection.input_sources))
}
}
pub(crate) type ResolvedOcrRequest = LiteLLMOcrRequest<OcrDocument>;
pub(crate) struct PreparedOcrRequest {
pub model: String,
pub document: OcrDocument,
@ -311,7 +431,7 @@ pub(crate) struct PreparedOcrRequest {
}
impl PreparedOcrRequest {
pub(crate) fn new(request: LiteLLMOcrRequest, connection: OcrConnection) -> Self {
pub(crate) fn new(request: ResolvedOcrRequest, connection: OcrConnection) -> Self {
let LiteLLMOcrRequest {
model,
document,
@ -446,6 +566,84 @@ mod tests {
use super::*;
use serde_json::json;
fn document() -> OcrDocument {
OcrDocument::try_from(
json!({"type":"document_url","document_url":"data:application/pdf;base64,YWJj"}),
)
.unwrap()
}
#[test]
fn from_inputs_applies_connection_overrides_with_field_sources() {
let request = LiteLLMOcrRequest::from_inputs(
"mistral/model".into(),
document(),
None,
Default::default(),
OcrConnectionInputs {
api_key: Some(" key ".into()),
api_base: Some("".into()),
extra_headers: json!({"x-a": "1"}).as_object().unwrap().clone(),
timeout: Some(Duration::from_secs(7)),
input_sources: [
("api_key".to_string(), InputSource::Request),
("extra_headers".to_string(), InputSource::Request),
]
.into(),
},
)
.unwrap();
let api_key = request.credentials.api_key.as_ref().unwrap();
assert_eq!(api_key.clone().into_value(), "key");
assert_eq!(api_key.source(), InputSource::Request);
assert!(request.credentials.api_base.is_none());
assert_eq!(
request.transport.extra_headers,
vec![("x-a".to_string(), "1".to_string())]
);
assert_eq!(request.transport.extra_headers_source, InputSource::Request);
assert_eq!(request.transport.timeout, Duration::from_secs(7));
assert_eq!(request.input_sources.len(), 2);
let defaulted = LiteLLMOcrRequest::from_inputs(
"mistral/model".into(),
document(),
None,
Default::default(),
OcrConnectionInputs::default(),
)
.unwrap();
assert_eq!(
defaulted.transport.timeout,
OcrTransportConfig::default().timeout
);
assert_eq!(
defaulted.transport.extra_headers_source,
InputSource::Deployment
);
}
#[test]
fn from_inputs_rejects_non_string_header_values_by_path() {
let Err(error) = LiteLLMOcrRequest::from_inputs(
"mistral/model".into(),
document(),
None,
Default::default(),
OcrConnectionInputs {
extra_headers: json!({"x-a": 1}).as_object().unwrap().clone(),
..Default::default()
},
) else {
panic!("non-string header value accepted");
};
assert!(matches!(
error,
super::super::Error::RequestField { ref path } if path == "extra_headers.x-a"
));
}
#[test]
fn normalized_response_rejects_invalid_shared_fields() {
for fields in [

View file

@ -16,6 +16,7 @@ extension-module = ["pyo3/extension-module"]
panic-test = []
[dependencies]
bytes.workspace = true
futures-util.workspace = true
litellm-core.workspace = true
litellm-auth.workspace = true

View file

@ -7,8 +7,97 @@ use pyo3::types::PyDict;
use serde_json::{Map, Value};
use litellm_auth::InputSource;
use litellm_core::call_arguments::ArgumentSpec;
use litellm_python_interop::from_py_preserving_errors as from_py;
use crate::auth::{AZURE_AD_TOKEN_PROVIDER, PythonTokenProvider};
use crate::lifecycle::BoundArguments;
/// 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.
pub(crate) struct BoundRouteInputs {
pub(crate) model: String,
pub(crate) custom_llm_provider: Option<String>,
pub(crate) api_key: Option<String>,
pub(crate) api_base: Option<String>,
pub(crate) extra_headers: Map<String, Value>,
pub(crate) timeout: Option<Duration>,
/// Only the optional params core reports as consumed plus opaque unknown
/// values; bound/control fields are never serialized.
pub(crate) optional_params: Map<String, Value>,
/// Which projected fields were supplied via the proxy request body.
pub(crate) input_sources: BTreeMap<String, InputSource>,
pub(crate) azure_ad_token_provider: Option<PythonTokenProvider>,
/// Consumed params whose values must be redacted from logging.
pub(crate) secret_fields: Vec<&'static str>,
}
const CONNECTION_FIELDS: [&str; 3] = ["api_key", "api_base", "extra_headers"];
impl BoundRouteInputs {
/// `consumed` is the route's provider-selected optional param list for the
/// already-extracted `model`/`custom_llm_provider`; `bound_fields` are the
/// route's positional parameter names that must not be projected.
pub(crate) fn extract(
py: Python<'_>,
arguments: &BoundArguments<'_>,
consumed: Vec<ArgumentSpec>,
bound_fields: &[&str],
) -> PyResult<Self> {
let kwargs = arguments.kwargs();
let optional_params = project_optional_fields(kwargs, &consumed, bound_fields)?;
let input_sources = request_input_sources(
kwargs,
optional_params
.keys()
.map(String::as_str)
.chain(CONNECTION_FIELDS),
)?;
Ok(Self {
model: arguments.extract("model")?,
custom_llm_provider: arguments.optional("custom_llm_provider")?,
api_key: arguments.optional("api_key")?,
api_base: arguments.optional("api_base")?,
extra_headers: arguments
.optional::<Bound<'_, PyAny>>("extra_headers")?
.map(|value| {
from_py(&value).and_then(|value| required_object("extra_headers", value))
})
.transpose()?
.unwrap_or_default(),
timeout: bound_timeout(py, arguments)?,
optional_params,
input_sources,
azure_ad_token_provider: kwargs.get_item("azure_ad_token_provider")?.and_then(
|provider| PythonTokenProvider::select(provider, AZURE_AD_TOKEN_PROVIDER),
),
secret_fields: consumed
.into_iter()
.filter(|spec| spec.secret)
.map(|spec| spec.name)
.collect(),
})
}
}
/// Reads `timeout` through `litellm.rust_bridge.timeouts.timeout_to_seconds`
/// so Python `httpx.Timeout` objects keep their existing semantics. Unlike the
/// value-route [`optional_timeout`], non-finite or negative values are
/// rejected rather than silently dropped.
fn bound_timeout(py: Python<'_>, arguments: &BoundArguments<'_>) -> PyResult<Option<Duration>> {
arguments
.optional::<Bound<'_, PyAny>>("timeout")?
.map(|value| python_timeout_seconds(py, value.unbind()))
.transpose()?
.flatten()
.map(|seconds| {
Duration::try_from_secs_f64(seconds)
.map_err(|_| PyValueError::new_err("timeout must be a non-negative finite number"))
})
.transpose()
}
pub(crate) struct RouteOptions {
pub(crate) model: String,
pub(crate) api_key: Option<String>,

View file

@ -1,4 +1,3 @@
use pyo3::exceptions::PyBaseException;
use pyo3::prelude::*;
use pyo3::types::PyDict;
use serde_json::Value;
@ -148,17 +147,3 @@ pub(super) fn response(py: Python<'_>, response: &LiteLLMOcrResponse) -> PyResul
.call1((to_py(py, response)?,))
.map(Bound::unbind)
}
pub(super) fn map_failure(
py: Python<'_>,
error: &Py<PyBaseException>,
model: &str,
provider: &str,
kwargs: &Py<PyDict>,
) -> PyResult<Py<PyBaseException>> {
Ok(py
.import("litellm.rust_bridge.ocr")?
.getattr("map_failure")?
.call1((error, model, provider, kwargs))?
.extract()?)
}

View file

@ -1,96 +1,49 @@
use std::io::Read;
use std::path::PathBuf;
use pyo3::exceptions::{PyFileNotFoundError, PyTypeError, PyValueError};
use bytes::Bytes;
use pyo3::exceptions::PyValueError;
use pyo3::gc::{PyTraverseError, PyVisit};
use pyo3::prelude::*;
use pyo3::pybacked::PyBackedBytes;
#[cfg(test)]
use pyo3::types::PyDict;
use pyo3::types::{PyBytes, PyString};
use litellm_core::constants::OCR_INLINE_MAX_BYTES;
use litellm_core::ocr::{OcrDocument, encode_file_document};
use litellm_core::ocr::{OcrDocumentInput, OcrFileContent};
enum FileBytes {
Python(PyBackedBytes),
Native(Vec<u8>),
pub(super) struct PythonFileReader {
reader: Py<PyAny>,
name: Option<String>,
}
impl AsRef<[u8]> for FileBytes {
fn as_ref(&self) -> &[u8] {
match self {
Self::Python(bytes) => bytes,
Self::Native(bytes) => bytes,
}
impl PythonFileReader {
pub(super) fn read(&self, py: Python<'_>) -> PyResult<OcrFileContent> {
let value = py
.import("litellm.rust_bridge.ocr")?
.getattr("read_document")?
.call1((self.reader.bind(py),))?;
Ok(OcrFileContent {
bytes: extract_bytes(&value)?,
file_name: self.name.clone(),
})
}
pub(super) fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> {
visit.call(&self.reader)
}
}
fn read_file_input(
py: Python<'_>,
file: &Bound<'_, PyAny>,
) -> PyResult<(FileBytes, Option<String>)> {
if file.is_instance_of::<PyString>() {
return Err(PyValueError::new_err(
"OCR file input does not accept bare str values. Pass bytes, a pathlib.Path, or a file-like object.",
));
fn extract_bytes(value: &Bound<'_, PyAny>) -> PyResult<Bytes> {
if value.is_exact_instance_of::<PyBytes>() {
return Ok(Bytes::from_owner(value.extract::<PyBackedBytes>()?));
}
if file.is_instance(&py.import("os")?.getattr("PathLike")?)? {
let path: PathBuf = file.extract()?;
let name = path
.file_name()
.map(|value| value.to_string_lossy().into_owned());
let bytes = py
.detach(|| {
let mut bytes = Vec::new();
std::fs::File::open(&path)?
.take(OCR_INLINE_MAX_BYTES as u64 + 1)
.read_to_end(&mut bytes)?;
Ok::<_, std::io::Error>(bytes)
})
.map_err(|error| {
if error.kind() == std::io::ErrorKind::NotFound {
PyFileNotFoundError::new_err(format!("File not found: {}", path.display()))
} else {
error.into()
}
})?;
return Ok((FileBytes::Native(bytes), name));
}
if file.is_instance_of::<PyBytes>() {
return Ok((FileBytes::Python(file.extract()?), None));
}
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.",
file.get_type(),
)));
};
let name = file
.getattr_opt("name")?
.filter(|value| !value.is_none())
.map(|value| value.extract::<String>())
.transpose()?;
let value = reader.call0()?;
let bytes = if value.is_instance_of::<PyString>() {
FileBytes::Native(value.extract::<String>()?.into_bytes())
} else if value.is_instance_of::<PyBytes>() {
FileBytes::Python(value.extract()?)
} else {
return Err(PyTypeError::new_err(format!(
"OCR file read must return bytes or str, got {}",
value.get_type(),
)));
};
Ok((bytes, name))
// Bytes subclasses may retain GC edges that a native Bytes owner cannot traverse.
Ok(Bytes::copy_from_slice(value.extract::<PyBackedBytes>()?.as_ref()))
}
pub(super) struct FileDocumentInput {
bytes: FileBytes,
name: Option<String>,
mime_type: Option<String>,
pub input: OcrDocumentInput,
pub reader: Option<PythonFileReader>,
}
impl FromPyObject<'_, '_> for FileDocumentInput {
@ -103,123 +56,112 @@ 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 file = document.get_item("file").map_err(|error| {
if error.is_instance_of::<pyo3::exceptions::PyKeyError>(py) {
PyValueError::new_err("document with type='file' must include a 'file' field containing a pathlib.Path, file-like object, or bytes")
missing()
} else {
error
}
})?;
if file.is_none() {
return Err(missing());
}
if file.is_instance_of::<PyString>() {
return Err(PyValueError::new_err(
"document with type='file' must include a 'file' field containing a pathlib.Path, file-like object, or bytes",
"OCR file input does not accept bare str values. Pass bytes, a pathlib.Path, or a file-like object.",
));
}
let (bytes, name) = read_file_input(py, &file)?;
if file.is_instance(&py.import("os")?.getattr("PathLike")?)? {
return Ok(Self {
input: OcrDocumentInput::Path {
path: file.extract::<PathBuf>()?,
mime_type,
},
reader: None,
});
}
if file.is_instance_of::<PyBytes>() {
return Ok(Self {
input: OcrDocumentInput::Bytes {
bytes: extract_bytes(&file)?,
file_name: None,
mime_type,
},
reader: None,
});
}
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.",
file.get_type(),
)));
};
let name = file
.getattr_opt("name")?
.filter(|value| !value.is_none())
.map(|value| value.extract::<String>())
.transpose()?;
Ok(Self {
bytes,
name,
mime_type,
input: OcrDocumentInput::HostReader { mime_type },
reader: Some(PythonFileReader { reader: reader.unbind(), name }),
})
}
}
pub(super) fn file_document(py: Python<'_>, document: FileDocumentInput) -> PyResult<OcrDocument> {
py.detach(|| {
encode_file_document(
document.bytes.as_ref(),
document.name.as_deref(),
document.mime_type.as_deref(),
)
})
.map_err(|error| PyValueError::new_err(error.to_string()))
}
#[cfg(test)]
mod tests {
use super::*;
use pyo3::exceptions::PyTypeError;
use pyo3::types::PyDict;
#[test]
fn extraction_validates_required_file_and_optional_mime_type() {
fn projection_validates_fields_without_consuming_readers_or_opening_paths() {
Python::initialize();
Python::attach(|py| {
for expression in [c"{}", c"{'file': None}"] {
let document = py.eval(expression, None, None).unwrap();
let error = document.extract::<FileDocumentInput>().err().unwrap();
assert!(error.is_instance_of::<PyValueError>(py));
assert!(error.to_string().contains("must include a 'file' field"));
}
for expression in [
c"{'file': b'abc', 'mime_type': None}",
c"{'file': b'abc', 'mime_type': 7}",
] {
let document = py.eval(expression, None, None).unwrap();
let error = document.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 document = py.eval(c"{'file': b'abc'}", None, None).unwrap();
let input: FileDocumentInput = document.extract().unwrap();
assert_eq!(input.bytes.as_ref(), b"abc");
assert_eq!(input.name, None);
assert_eq!(input.mime_type, None);
});
}
#[test]
fn extraction_validates_mime_type_before_consuming_file() {
Python::initialize();
Python::attach(|py| {
let locals = PyDict::new(py);
py.run(
c"class Reader:
py.run(c"from pathlib import Path
failure = KeyError('reader failed')
class Reader:
def __init__(self):
self.reads = 0
def read(self):
self.reads += 1
return b'abc'
raise failure
reader = Reader()
document = {'file': reader, 'mime_type': 7}",
Some(&locals),
Some(&locals),
)
.unwrap();
document = {'file': reader}
path_document = {'file': Path('/nonexistent/ocr-projection-test.pdf')}
", Some(&locals), Some(&locals)).unwrap();
let document = locals.get_item("document").unwrap().unwrap();
let error = document.extract::<FileDocumentInput>().err().unwrap();
assert!(error.is_instance_of::<PyTypeError>(py));
let reads: usize = locals
.get_item("reader")
.unwrap()
.unwrap()
.getattr("reads")
.unwrap()
.extract()
.unwrap();
assert_eq!(reads, 0);
let input: FileDocumentInput = document.extract().unwrap();
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();
assert!(matches!(path.input, OcrDocumentInput::Path { .. }));
});
}
#[test]
fn extraction_preserves_reader_key_error_identity() {
fn exact_python_bytes_transfer_without_copying_and_outlive_the_input() {
Python::initialize();
Python::attach(|py| {
let locals = PyDict::new(py);
py.run(
c"failure = KeyError('reader failed')
class Reader:
def read(self):
raise failure
document = {'file': Reader()}",
Some(&locals),
Some(&locals),
)
.unwrap();
let document = locals.get_item("document").unwrap().unwrap();
let error = document.extract::<FileDocumentInput>().err().unwrap();
assert!(
error
.value(py)
.is(locals.get_item("failure").unwrap().unwrap())
);
let (bytes, pointer) = Python::attach(|py| {
let value = PyBytes::new(py, b"document bytes");
let pointer = value.as_bytes().as_ptr() as usize;
(extract_bytes(value.as_any()).unwrap(), pointer)
});
assert_eq!(bytes.as_ptr() as usize, pointer);
assert_eq!(bytes.as_ref(), b"document bytes");
}
}

View file

@ -1,170 +1,252 @@
use litellm_core::ocr::Error;
use litellm_core::transport::Error as TransportError;
use pyo3::exceptions::{PyRuntimeError, PyValueError};
use pyo3::exceptions::{PyBaseException, PyException};
use pyo3::prelude::*;
use pyo3::types::PyDict;
use crate::errors::{RustUpstreamError, core_error_to_pyerr};
pub(super) fn to_pyerr(error: Error) -> PyErr {
if let Error::Provider {
status,
body,
headers,
} = error
{
let mapped = attach_status(RustUpstreamError::new_err((status, body)), Some(status));
return Python::attach(|py| -> PyResult<PyErr> {
let headers =
pyo3::types::PyDict::from_sequence(&headers.into_pyobject(py)?.into_any())?;
mapped.value(py).setattr("headers", headers)?;
Ok(mapped)
})
.unwrap_or_else(|error| error);
}
let (mapped, status) = match error {
Error::MissingDocumentUrl => (
PyValueError::new_err(Error::MissingDocumentUrl.to_string()),
Some(500),
),
error @ Error::MissingField(_) => (PyValueError::new_err(error.to_string()), None),
error if error.is_request() => (
PyValueError::new_err(format!("invalid request: {error}")),
Some(400),
),
error if error.is_response() => (
PyRuntimeError::new_err(format!("invalid response: {error}")),
None,
),
error @ (Error::InvalidRequest(_) | Error::Params(_) | Error::Headers(_)) => {
(PyValueError::new_err(error.to_string()), Some(400))
}
Error::Transport(TransportError::Http { status, body }) => {
(RustUpstreamError::new_err((status, body)), Some(status))
}
other => (core_error_to_pyerr(other), None),
};
attach_status(mapped, status)
enum Public {
BadRequest,
Authentication,
NotFound,
Timeout,
RateLimit,
InternalServer,
BadGateway,
ServiceUnavailable,
Api(u16),
ApiConnection,
}
fn attach_status(error: PyErr, status: Option<u16>) -> PyErr {
if let Some(status) = status {
Python::attach(|py| {
let value = error.value(py);
value.setattr("status_code", status).ok();
value.setattr("message", value.to_string()).ok();
});
impl Public {
fn class_name(&self) -> &'static str {
match self {
Self::BadRequest => "BadRequestError",
Self::Authentication => "AuthenticationError",
Self::NotFound => "NotFoundError",
Self::Timeout => "Timeout",
Self::RateLimit => "RateLimitError",
Self::InternalServer => "InternalServerError",
Self::BadGateway => "BadGatewayError",
Self::ServiceUnavailable => "ServiceUnavailableError",
Self::Api(_) => "APIError",
Self::ApiConnection => "APIConnectionError",
}
}
error
fn from_status(status: u16) -> Self {
match status {
400 | 422 => Self::BadRequest,
401 => Self::Authentication,
404 => Self::NotFound,
408 | 504 => Self::Timeout,
429 => Self::RateLimit,
500 => Self::InternalServer,
502 => Self::BadGateway,
503 => Self::ServiceUnavailable,
other => Self::Api(other),
}
}
}
struct Failure {
public: Public,
detail: String,
headers: Vec<(String, String)>,
}
fn classify(error: Error) -> Failure {
match error {
Error::Provider {
status,
body,
headers,
} => Failure {
public: Public::from_status(status),
detail: body,
headers,
},
Error::Transport(TransportError::Http { status, body }) => Failure {
public: Public::from_status(status),
detail: body,
headers: Vec::new(),
},
error if error.is_request() => Failure {
public: Public::BadRequest,
detail: error.to_string(),
headers: Vec::new(),
},
error => Failure {
public: Public::ApiConnection,
detail: error.to_string(),
headers: Vec::new(),
},
}
}
fn exception_provider(provider: &str) -> String {
let mut chars = provider.chars();
match chars.next() {
Some(first) => format!("{}{}Exception", first.to_uppercase(), chars.as_str()),
None => "Exception".into(),
}
}
pub(super) fn public_exception(
py: Python<'_>,
error: Error,
model: &str,
provider: &str,
) -> PyResult<PyErr> {
if let Error::FileRead { path, source } = error {
let errno = source.raw_os_error().unwrap_or(match source.kind() {
std::io::ErrorKind::NotFound => 2,
std::io::ErrorKind::PermissionDenied => 13,
_ => 5,
});
return Ok(PyErr::from_value(py.import("builtins")?.getattr("OSError")?.call1((
errno, source.to_string(), path,
))?));
}
raise_public(py, classify(error), model, provider, None)
}
pub(super) fn public_host_exception(
py: Python<'_>,
error: &Py<PyBaseException>,
model: &str,
provider: &str,
) -> PyResult<PyErr> {
let value = error.bind(py);
if !value.is_instance_of::<PyException>() || is_public(value)? {
return Ok(PyErr::from_value(value.clone().into_any()));
}
let failure = Failure {
public: Public::ApiConnection,
detail: value.str()?.to_string(),
headers: Vec::new(),
};
raise_public(py, failure, model, provider, Some(value))
}
fn is_public(value: &Bound<'_, PyBaseException>) -> PyResult<bool> {
let types = value
.py()
.import("litellm")?
.getattr("LITELLM_EXCEPTION_TYPES")?;
for public in types.try_iter()? {
if value.is_instance(&public?)? {
return Ok(true);
}
}
Ok(false)
}
fn raise_public(
py: Python<'_>,
failure: Failure,
model: &str,
provider: &str,
context: Option<&Bound<'_, PyBaseException>>,
) -> PyResult<PyErr> {
let class = failure.public.class_name();
let message = format!(
"{}{} - {}",
match failure.public {
Public::RateLimit | Public::Api(_) | Public::ApiConnection => format!("{class}: "),
_ => String::new(),
},
exception_provider(provider),
failure.detail
);
let kwargs = PyDict::new(py);
kwargs.set_item("message", message)?;
kwargs.set_item("llm_provider", provider)?;
kwargs.set_item("model", model)?;
if let Public::Api(status) = failure.public {
kwargs.set_item("status_code", status)?;
}
let value = py
.import("litellm")?
.getattr(class)?
.call((), Some(&kwargs))?;
if !failure.headers.is_empty() {
let headers = PyDict::new(py);
for (name, header) in &failure.headers {
headers.set_item(name, header)?;
}
value.setattr("litellm_response_headers", headers)?;
}
if let Some(context) = context {
value.setattr("__context__", context)?;
}
Ok(PyErr::from_value(value))
}
pub(super) fn to_pyerr(error: Error) -> PyErr {
Python::attach(|py| public_exception(py, error, "", "").unwrap_or_else(|error| error))
}
#[cfg(test)]
mod tests {
use super::*;
use pyo3::exceptions::PyValueError;
#[test]
fn provider_error_retains_headers_at_the_python_boundary() {
Python::initialize();
Python::attach(|py| {
let error = to_pyerr(Error::Provider {
status: 429,
body: "rate limited".into(),
headers: vec![("retry-after".into(), "17".into())],
});
assert!(error.is_instance_of::<RustUpstreamError>(py));
assert_eq!(
error
.value(py)
.getattr("args")
.unwrap()
.extract::<(u16, String)>()
.unwrap(),
(429, "rate limited".into())
);
assert_eq!(
error
.value(py)
.getattr("headers")
.unwrap()
.get_item("retry-after")
.unwrap()
.extract::<String>()
.unwrap(),
"17"
);
});
fn classified(error: Error) -> (&'static str, String, Vec<(String, String)>) {
let failure = classify(error);
(failure.public.class_name(), failure.detail, failure.headers)
}
#[test]
fn preserves_python_validation_and_provider_details() {
Python::initialize();
Python::attach(|py| {
let mapped = to_pyerr(Error::MissingDocumentUrl);
assert!(mapped.is_instance_of::<pyo3::exceptions::PyValueError>(py));
assert_eq!(mapped.value(py).to_string(), "Document URL is required");
assert_eq!(
mapped
.value(py)
.getattr("status_code")
.unwrap()
.extract::<u16>()
.unwrap(),
500
);
let mapped = to_pyerr(Error::Transport(litellm_core::transport::Error::Http {
status: 429,
body: r#"{"message":"rate limited"}"#.to_string(),
fn provider_status_selects_public_class_and_keeps_details() {
let (class, detail, headers) = classified(Error::Provider {
status: 429,
body: "rate limited".into(),
headers: vec![("retry-after".into(), "17".into())],
});
assert_eq!(class, "RateLimitError");
assert_eq!(detail, "rate limited");
assert_eq!(headers, [("retry-after".to_string(), "17".to_string())]);
}
#[test]
fn status_table_matches_legacy_openai_mapping() {
for (status, class) in [
(400, "BadRequestError"),
(401, "AuthenticationError"),
(404, "NotFoundError"),
(408, "Timeout"),
(422, "BadRequestError"),
(429, "RateLimitError"),
(500, "InternalServerError"),
(502, "BadGatewayError"),
(503, "ServiceUnavailableError"),
(504, "Timeout"),
(418, "APIError"),
] {
let (mapped, _, _) = classified(Error::Transport(TransportError::Http {
status,
body: "x".into(),
}));
assert!(mapped.is_instance_of::<RustUpstreamError>(py));
let args: (u16, String) = mapped
.value(py)
.getattr("args")
.and_then(|args| args.extract())
.expect("OCR failures retain status and unprefixed provider message");
assert_eq!(args, (429, r#"{"message":"rate limited"}"#.to_string()));
let mapped = to_pyerr(Error::InvalidRequest("invalid format".into()));
assert!(mapped.is_instance_of::<PyValueError>(py));
assert_eq!(
mapped
.value(py)
.getattr("status_code")
.unwrap()
.extract::<u16>()
.unwrap(),
400
);
});
assert_eq!(mapped, class, "{status}");
}
}
#[test]
fn typed_request_and_response_failures_keep_python_contracts() {
Python::initialize();
Python::attach(|py| {
let request = to_pyerr(Error::RequestField {
path: "document.type".into(),
});
assert!(request.is_instance_of::<PyValueError>(py));
assert_eq!(
request
.value(py)
.getattr("status_code")
.unwrap()
.extract::<u16>()
.unwrap(),
400
);
assert_eq!(
request.value(py).to_string(),
"invalid request: invalid OCR request field: document.type"
);
let response = to_pyerr(Error::EmptyContent);
assert!(response.is_instance_of::<PyRuntimeError>(py));
assert_eq!(
response.value(py).to_string(),
"invalid response: OCR response is missing non-empty content"
);
assert!(!response.value(py).hasattr("status_code").unwrap());
});
fn request_shape_and_statusless_failures_use_their_public_families() {
assert_eq!(classified(Error::RequestFormat).0, "BadRequestError");
let (class, detail, _) = classified(Error::MissingAzureAiCredentials);
assert_eq!(class, "APIConnectionError");
assert!(detail.contains("Missing Azure AI credentials"));
let (class, detail, _) = classified(Error::Transport(TransportError::Network(
"connection reset".into(),
)));
assert_eq!(class, "APIConnectionError");
assert!(detail.contains("connection reset"));
}
#[test]
fn exception_provider_matches_legacy_capitalisation() {
assert_eq!(exception_provider("mistral"), "MistralException");
assert_eq!(exception_provider("azure_ai"), "Azure_aiException");
assert_eq!(exception_provider(""), "Exception");
}
}

View file

@ -1,42 +1,24 @@
use std::sync::Arc;
use pyo3::gc::{PyTraverseError, PyVisit};
use pyo3::prelude::*;
use pyo3::types::{PyDict, PyTuple};
use pyo3::types::PyDict;
use litellm_auth::ResolvedCredential;
use litellm_core::ocr::hooks::{OcrDuringCallRequest, OcrPostCallRequest};
use litellm_core::ocr::{OcrAdmission, OcrCall, OcrClient, OcrHostOperation, OcrHostResult};
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,
};
use super::callbacks;
use super::document::PythonFileReader;
use super::errors::to_pyerr as ocr_error_to_pyerr;
use super::project::{admitted_call, project};
use super::project::project;
use super::{ASYNC_SIGNATURE, SIGNATURE, callbacks, errors};
use crate::auth::PythonTokenProvider;
use crate::lifecycle::{
OperationClass, PythonCallState, PythonRoute, Signature, missing_state, now, run_call,
};
use crate::lifecycle::{OperationClass, PythonCallState, PythonRoute, missing_state, now};
const SIGNATURE: Signature = Signature {
name: "ocr",
parameters: &[
"model",
"document",
"api_key",
"api_base",
"timeout",
"custom_llm_provider",
"extra_headers",
],
required: 2,
};
const ASYNC_SIGNATURE: Signature = Signature {
name: "aocr",
..SIGNATURE
};
struct PythonOcrHost {
pub(super) struct PythonOcrHost {
state: PythonCallState,
projected: Option<ProjectedOcrHost>,
}
@ -48,6 +30,8 @@ struct ProjectedOcrHost {
azure_ad_token_provider: Option<PythonTokenProvider>,
pre_call: Option<callbacks::OcrLoggingFields>,
payload: Option<CapturedOcrPayload>,
reader: Option<PythonFileReader>,
reader_failed: bool,
}
struct CapturedOcrPayload {
@ -63,6 +47,13 @@ impl CapturedOcrPayload {
}
impl PythonOcrHost {
pub(super) fn new(state: PythonCallState) -> Self {
Self {
state,
projected: None,
}
}
fn projected(&self) -> PyResult<&ProjectedOcrHost> {
self.projected.as_ref().ok_or_else(missing_state)
}
@ -87,9 +78,15 @@ impl PythonOcrHost {
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),
Box::new(
projected
.request
.with_host_hooks(Arc::new(BridgeOcrHooks), None),
),
has_token_provider,
))))
}
@ -161,20 +158,21 @@ impl PythonOcrHost {
}
fn map_failure(&mut self, py: Python<'_>, error: litellm_core::ocr::Error) -> PyResult<()> {
if self.state.error.is_none() {
self.state.retain_error(py, ocr_error_to_pyerr(error));
}
if self.state.end.is_none() {
self.state.end = Some(now(py)?);
}
let error = self.state.error.as_ref().ok_or_else(missing_state)?;
let (model, provider) = match &self.projected {
Some(projected) => (projected.model.as_str(), projected.provider),
None => ("", ""),
};
let mapped = callbacks::map_failure(py, error, model, provider, &self.state.kwargs)?;
self.state
.retain_error(py, PyErr::from_value(mapped.into_bound(py).into_any()));
let mapped = match self.state.error.take() {
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)?,
None => errors::public_exception(py, error, model, provider)?,
};
self.state.retain_error(py, mapped);
Ok(())
}
}
@ -211,6 +209,12 @@ impl PythonRoute for PythonOcrHost {
fn invoke(&mut self, py: Python<'_>, operation: OcrHostOperation) -> PyResult<OcrHostResult> {
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?))
}
OcrHostOperation::AcquireAzureAdToken => {
OcrHostResult::AzureAdToken(Ok(self.acquire_azure_ad_token(py)?))
}
@ -250,6 +254,9 @@ impl PythonRoute for PythonOcrHost {
if let Some(provider) = &projected.azure_ad_token_provider {
provider.traverse(visit)?;
}
if let Some(reader) = &projected.reader {
reader.traverse(visit)?;
}
if let Some(payload) = &projected.payload {
payload.traverse(visit)?;
}
@ -257,69 +264,10 @@ impl PythonRoute for PythonOcrHost {
}
}
pub(super) struct BridgeOcrHooks;
struct BridgeOcrHooks;
impl litellm_core::ocr::hooks::OcrHooks for BridgeOcrHooks {
fn intercepts_requests(&self) -> bool {
true
}
}
fn call(
py: Python<'_>,
args: Bound<'_, PyTuple>,
kwargs: Option<Bound<'_, PyDict>>,
asynchronous: bool,
) -> PyResult<Py<PyAny>> {
let kwargs = kwargs.unwrap_or_else(|| PyDict::new(py));
let signature = if asynchronous {
&ASYNC_SIGNATURE
} else {
&SIGNATURE
};
signature.bind(&args, &kwargs)?;
let client = OcrClient::shared().map_err(ocr_error_to_pyerr)?;
let call = admitted_call(OcrCall::admit(
client,
OcrAdmission {
asynchronous,
..OcrAdmission::all()
},
))?;
let host = PythonOcrHost {
state: PythonCallState::new(
py,
args.unbind(),
kwargs.copy()?.unbind(),
asynchronous,
signature.name,
)?,
projected: None,
};
run_call(py, call, host)
}
#[pyfunction]
#[pyo3(signature = (*args, **kwargs))]
fn ocr(
py: Python<'_>,
args: Bound<'_, PyTuple>,
kwargs: Option<Bound<'_, PyDict>>,
) -> PyResult<Py<PyAny>> {
call(py, args, kwargs, false)
}
#[pyfunction]
#[pyo3(signature = (*args, **kwargs))]
fn aocr(
py: Python<'_>,
args: Bound<'_, PyTuple>,
kwargs: Option<Bound<'_, PyDict>>,
) -> PyResult<Py<PyAny>> {
call(py, args, kwargs, true)
}
pub(super) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> {
super::super::add_function(module, wrap_pyfunction!(ocr, module)?)?;
super::super::add_function(module, wrap_pyfunction!(aocr, module)?)
}

View file

@ -1,11 +1,126 @@
mod callbacks;
mod document;
mod errors;
mod lifecycle;
mod host;
mod project;
use litellm_core::ocr::{NativeOutcome, OcrAdmission, OcrCall, OcrClient};
use pyo3::prelude::*;
use pyo3::types::{PyDict, PyTuple};
use self::errors::to_pyerr as ocr_error_to_pyerr;
use self::host::PythonOcrHost;
use crate::errors::RustBridgeDeclined;
use crate::lifecycle::{PythonCallState, Signature, run_call};
const SIGNATURE: Signature = Signature {
name: "ocr",
parameters: &[
"model",
"document",
"api_key",
"api_base",
"timeout",
"custom_llm_provider",
"extra_headers",
],
required: 2,
};
const ASYNC_SIGNATURE: Signature = Signature {
name: "aocr",
..SIGNATURE
};
fn admitted_call(outcome: NativeOutcome<OcrCall>) -> PyResult<OcrCall> {
match outcome {
NativeOutcome::Completed(call) => Ok(call),
NativeOutcome::Declined(reason) => Err(RustBridgeDeclined::new_err(format!(
"native OCR admission declined: {reason:?}"
))),
}
}
fn call(
py: Python<'_>,
args: Bound<'_, PyTuple>,
kwargs: Option<Bound<'_, PyDict>>,
asynchronous: bool,
) -> PyResult<Py<PyAny>> {
let kwargs = kwargs.unwrap_or_else(|| PyDict::new(py));
let signature = if asynchronous {
&ASYNC_SIGNATURE
} else {
&SIGNATURE
};
signature.bind(&args, &kwargs)?;
let client = OcrClient::shared().map_err(ocr_error_to_pyerr)?;
let call = admitted_call(OcrCall::admit(
client,
OcrAdmission {
asynchronous,
..OcrAdmission::all()
},
))?;
let host = PythonOcrHost::new(PythonCallState::new(
py,
args.unbind(),
kwargs.copy()?.unbind(),
asynchronous,
signature.name,
)?);
run_call(py, call, host)
}
#[pyfunction]
#[pyo3(signature = (*args, **kwargs))]
fn ocr(
py: Python<'_>,
args: Bound<'_, PyTuple>,
kwargs: Option<Bound<'_, PyDict>>,
) -> PyResult<Py<PyAny>> {
call(py, args, kwargs, false)
}
#[pyfunction]
#[pyo3(signature = (*args, **kwargs))]
fn aocr(
py: Python<'_>,
args: Bound<'_, PyTuple>,
kwargs: Option<Bound<'_, PyDict>>,
) -> PyResult<Py<PyAny>> {
call(py, args, kwargs, true)
}
pub(super) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> {
lifecycle::register(module)
super::add_function(module, wrap_pyfunction!(ocr, module)?)?;
super::add_function(module, wrap_pyfunction!(aocr, module)?)
}
#[cfg(test)]
mod tests {
use litellm_core::ocr::{Error, OcrDecline};
use super::*;
#[test]
fn typed_initial_decline_uses_bridge_decline_contract() {
Python::initialize();
Python::attach(|py| {
let Err(error) = admitted_call(NativeOutcome::Declined(OcrDecline::HostOperations))
else {
panic!("unsupported host operations should decline admission");
};
assert!(error.is_instance_of::<RustBridgeDeclined>(py));
});
}
#[test]
fn post_admission_error_does_not_use_bridge_decline_contract() {
Python::initialize();
Python::attach(|py| {
let error = ocr_error_to_pyerr(Error::InvalidRequest("callback result".into()));
assert!(!error.is_instance_of::<RustBridgeDeclined>(py));
});
}
}

View file

@ -1,153 +1,113 @@
use std::collections::BTreeMap;
use std::sync::Arc;
use std::time::Duration;
//! `ProjectRequest` host step: read the bound `ocr()` arguments and hand core a
//! typed [`LiteLLMOcrRequest`]. Generic argument handling lives in
//! [`crate::marshal::BoundRouteInputs`]; this module only owns what is specific
//! to OCR: the `document` argument and the core constructor call.
use pyo3::prelude::*;
use pyo3::types::PyDict;
use serde_json::{Map, Value};
use litellm_auth::InputSource;
use litellm_core::ocr::{
LiteLLMOcrRequest, NativeOutcome, OcrCall, OcrCredentialInputs, OcrDocument,
consumed_optional_params,
LiteLLMOcrRequest, OcrConnectionInputs, OcrDocument, consumed_optional_params,
};
use litellm_python_interop::from_py_preserving_errors as from_py;
use super::errors::to_pyerr as ocr_error_to_pyerr;
use super::lifecycle::BridgeOcrHooks;
use crate::auth::{AZURE_AD_TOKEN_PROVIDER, PythonTokenProvider};
use crate::errors::RustBridgeDeclined;
use super::document::{FileDocumentInput, PythonFileReader};
use crate::auth::PythonTokenProvider;
use crate::lifecycle::BoundArguments;
use crate::marshal::{project_optional_fields, python_timeout_seconds, request_input_sources};
use crate::marshal::BoundRouteInputs;
/// 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>,
}
fn project_document(py: Python<'_>, document: &Bound<'_, PyAny>) -> PyResult<Value> {
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" {
return from_py(document);
return Ok(ProjectedDocument::Url(from_py(document)?));
}
let encoded = super::document::file_document(py, document.extract()?)?;
serde_json::to_value(encoded)
.map_err(|error| pyo3::exceptions::PyValueError::new_err(error.to_string()))
document.extract().map(ProjectedDocument::File)
}
fn header_pairs(headers: Option<Map<String, Value>>) -> Result<Vec<(String, String)>, litellm_core::ocr::Error> {
headers
.unwrap_or_default()
.into_iter()
.map(|(name, value)| {
value
.as_str()
.map(|value| (name.clone(), value.to_string()))
.ok_or_else(|| litellm_core::ocr::Error::RequestField {
path: format!("extra_headers.{name}"),
})
})
.collect()
}
fn source_for(sources: &BTreeMap<String, InputSource>, name: &str) -> InputSource {
sources.get(name).copied().unwrap_or_default()
}
pub(super) fn project(py: Python<'_>, arguments: &BoundArguments<'_>) -> PyResult<ProjectedOcrCall> {
let kwargs: &Bound<'_, PyDict> = arguments.kwargs();
let model: String = arguments.extract("model")?;
let custom_llm_provider: Option<String> = arguments.optional("custom_llm_provider")?;
let document = project_document(py, &arguments.required("document")?)?;
let api_key: Option<String> = arguments.optional("api_key")?;
let specs = consumed_optional_params(&model, custom_llm_provider.as_deref())
.map_err(ocr_error_to_pyerr)?;
let optional_params = project_optional_fields(kwargs, &specs, BOUND_FIELDS)?;
let input_sources = request_input_sources(
kwargs,
optional_params
.keys()
.map(String::as_str)
.chain(["api_key", "api_base", "extra_headers"]),
)?;
let azure_ad_token_provider = kwargs
.get_item("azure_ad_token_provider")?
.and_then(|provider| PythonTokenProvider::select(provider, AZURE_AD_TOKEN_PROVIDER));
let api_base: Option<String> = arguments.optional("api_base")?;
let extra_headers: Option<Map<String, Value>> = arguments
.get("extra_headers")?
.filter(|value| !value.is_none())
.map(|value| from_py(&value))
.transpose()?;
let timeout = arguments
.get("timeout")?
.filter(|value| !value.is_none())
.map(|value| python_timeout_seconds(py, value.unbind()))
.transpose()?
.flatten()
.map(|seconds| {
Duration::try_from_secs_f64(seconds).map_err(|_| {
litellm_core::ocr::Error::RequestField {
path: "timeout_seconds".into(),
}
})
})
.transpose()
.map_err(ocr_error_to_pyerr)?;
let request = (|| {
let core = LiteLLMOcrRequest::new(
model,
OcrDocument::try_from(document)?,
custom_llm_provider.as_deref(),
optional_params.into(),
)?;
let transport = core.transport.clone().with_overrides(
header_pairs(extra_headers)?,
source_for(&input_sources, "extra_headers"),
timeout,
);
let credentials = OcrCredentialInputs::new(
/// 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()?;
let BoundRouteInputs {
model,
custom_llm_provider,
api_key,
api_base,
extra_headers,
timeout,
optional_params,
input_sources,
azure_ad_token_provider,
secret_fields,
} = inputs;
let request = LiteLLMOcrRequest::from_inputs(
model,
document.input,
custom_llm_provider.as_deref(),
optional_params.into(),
OcrConnectionInputs {
api_key,
source_for(&input_sources, "api_key"),
api_base,
source_for(&input_sources, "api_base"),
);
Ok::<_, litellm_core::ocr::Error>(
core.with_connection_inputs(credentials, transport, input_sources)
.with_host_hooks(Arc::new(BridgeOcrHooks), None),
)
})()
.map_err(ocr_error_to_pyerr)?;
extra_headers,
timeout,
input_sources,
},
)?;
Ok(ProjectedOcrCall {
request,
azure_ad_token_provider,
secret_fields: specs
.into_iter()
.filter(|spec| spec.secret)
.map(|spec| spec.name)
.collect(),
secret_fields,
reader: document.reader,
})
}
pub(super) fn admitted_call(outcome: NativeOutcome<OcrCall>) -> PyResult<OcrCall> {
match outcome {
NativeOutcome::Completed(call) => Ok(call),
NativeOutcome::Declined(reason) => Err(RustBridgeDeclined::new_err(format!(
"native OCR admission declined: {reason:?}"
))),
}
pub(super) fn project(
py: Python<'_>,
arguments: &BoundArguments<'_>,
) -> 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")?)?;
let consumed = consumed_optional_params(&model, custom_llm_provider.as_deref())
.map_err(ocr_error_to_pyerr)?;
let inputs = BoundRouteInputs::extract(py, arguments, consumed, BOUND_FIELDS)?;
build_request(inputs, document).map_err(ocr_error_to_pyerr)
}
#[cfg(test)]
mod tests {
use litellm_core::ocr::Error;
use litellm_core::ocr::OcrDecline;
use pyo3::exceptions::{PyKeyError, PyTypeError, PyValueError};
use pyo3::exceptions::{PyKeyError, PyTypeError};
use pyo3::types::PyDict;
use super::*;
@ -171,29 +131,7 @@ mod tests {
}
#[test]
fn typed_initial_decline_uses_bridge_decline_contract() {
Python::initialize();
Python::attach(|py| {
let Err(error) = admitted_call(NativeOutcome::Declined(OcrDecline::HostOperations))
else {
panic!("unsupported host operations should decline admission");
};
assert!(error.is_instance_of::<RustBridgeDeclined>(py));
});
}
#[test]
fn post_admission_error_does_not_use_bridge_decline_contract() {
Python::initialize();
Python::attach(|py| {
let error = ocr_error_to_pyerr(Error::InvalidRequest("callback result".into()));
assert!(error.is_instance_of::<PyValueError>(py));
assert!(!error.is_instance_of::<RustBridgeDeclined>(py));
});
}
#[test]
fn file_documents_are_encoded_and_other_documents_pass_through() {
fn file_documents_remain_raw_and_other_documents_pass_through() {
Python::initialize();
Python::attach(|py| {
let file = py
@ -203,13 +141,11 @@ mod tests {
None,
)
.unwrap();
assert_eq!(
project_document(py, &file).unwrap(),
serde_json::json!({
"type": "document_url",
"document_url": "data:application/pdf;base64,JVBERi0xLjQ=",
})
);
assert!(matches!(
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")
));
let original = py
.eval(
@ -218,13 +154,11 @@ mod tests {
None,
)
.unwrap();
assert_eq!(
project_document(py, &original).unwrap(),
serde_json::json!({
"type": "document_url",
"document_url": "https://example.com/a.pdf",
})
);
assert!(matches!(
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"
));
});
}
@ -235,12 +169,7 @@ mod tests {
let document = py
.eval(c"{'type': 'mystery', 'mystery': 'x'}", None, None)
.unwrap();
let wire_document = project_document(py, &document).unwrap();
assert_eq!(
wire_document,
serde_json::json!({"type": "mystery", "mystery": "x"})
);
let error = OcrDocument::try_from(wire_document).unwrap_err();
let error = project_document(&document).unwrap().into_native().err().unwrap();
assert!(error.to_string().contains("document"));
});
}
@ -251,15 +180,15 @@ mod tests {
Python::attach(|py| {
let missing = py.eval(c"{}", None, None).unwrap();
assert!(
project_document(py, &missing)
.unwrap_err()
project_document(&missing)
.err().unwrap()
.is_instance_of::<PyKeyError>(py)
);
let non_string = py.eval(c"{'type': 1}", None, None).unwrap();
assert!(
project_document(py, &non_string)
.unwrap_err()
project_document(&non_string)
.err().unwrap()
.is_instance_of::<PyTypeError>(py)
);
@ -274,7 +203,7 @@ document = Document()
",
);
let error =
project_document(py, &locals.get_item("document").unwrap().unwrap()).unwrap_err();
project_document(&locals.get_item("document").unwrap().unwrap()).err().unwrap();
assert!(
error
.value(py)
@ -303,8 +232,8 @@ document = Document()
",
);
let document = locals.get_item("document").unwrap().unwrap();
let wire = project_document(py, &document).unwrap();
assert_eq!(wire["type"], "document_url");
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

@ -1,12 +1,11 @@
from __future__ import annotations
from collections.abc import Coroutine, Mapping
from collections.abc import Callable, Coroutine, Mapping
from types import MappingProxyType
from typing import Final, Protocol, cast # noqa: TID251 # validates dynamically loaded native callables
from pydantic import TypeAdapter
import litellm
from litellm.llms.base_llm.ocr.transformation import PROVIDER_NATIVE_RESPONSE_KEY, OCRResponse
from litellm.rust_bridge.bindings import NativeBinding
@ -32,6 +31,15 @@ NATIVE_AOCR: Final = NativeBinding("aocr", validate=_as_aocr)
_NATIVE_RESPONSE: Final = TypeAdapter(Mapping[str, object])
def read_document(reader: Callable[[], object]) -> bytes:
value: Final = reader()
if isinstance(value, str):
return value.encode("utf-8")
if isinstance(value, bytes):
return value
raise TypeError(f"OCR file read must return bytes or str, got {type(value)}")
def build_response(response: Mapping[str, object]) -> OCRResponse:
provider_native_response: Final = response.get(PROVIDER_NATIVE_RESPONSE_KEY)
normalized: Final = OCRResponse.model_validate(
@ -40,32 +48,3 @@ def build_response(response: Mapping[str, object]) -> OCRResponse:
if isinstance(provider_native_response, Mapping):
normalized.set_provider_native_response(_NATIVE_RESPONSE.validate_python(provider_native_response))
return normalized
class ExceptionMapper(Protocol):
def __call__(
self,
*,
model: str,
custom_llm_provider: str | None,
original_exception: Exception,
completion_kwargs: dict[str, object],
extra_kwargs: dict[str, object],
) -> Exception: ...
def map_failure(error: Exception, model: str, provider: str, kwargs: Mapping[str, object]) -> Exception:
mapper: Final = cast(
ExceptionMapper, litellm.exception_type
) # cast-ok: bounded adapter for the public exception mapper
try:
return mapper(
model=model,
custom_llm_provider=provider,
original_exception=error,
completion_kwargs=dict(kwargs), # mutable-ok: exception mapper requires owned kwargs
extra_kwargs=dict(kwargs), # mutable-ok: exception mapper requires owned kwargs
)
except Exception as public_error:
public_error.__context__ = error
return public_error

View file

@ -1,5 +1,10 @@
import json
import asyncio
import base64
import threading
import weakref
import gc
from pathlib import Path
from collections.abc import Generator
from datetime import datetime, timezone
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
@ -332,6 +337,126 @@ def test_native_lifecycle_core_encodes_python_file_input(
assert requests[0]["body"]["opaque_extension"] == {"nested": [None, False, 0]}
@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):
server, requests = ocr_server
path: Final = tmp_path / "scan.png"
content: Final = b"document bytes"
path.write_bytes(content)
class CustomPath:
def __fspath__(self) -> str:
return str(path)
def __str__(self) -> str:
return "/not/the/document.pdf"
class Reader:
name = "scan.png"
def __init__(self) -> None:
self.stream = BytesIO(b"prefix" + content)
self.stream.seek(6)
self.calls = 0
def read(self) -> bytes | str:
assert threading.get_ident() == caller_thread
if asynchronous:
assert asyncio.current_task() is caller_task
self.calls += 1
value: Final = self.stream.read()
return value.decode() if source == "text" else value
caller_thread: Final = threading.get_ident()
caller_task: Final = asyncio.current_task()
reader: Final = Reader()
file_input: Final = {
"path": path,
"pathlike": CustomPath(),
"reader": reader,
"text": reader,
"bytes": content,
}[source]
document: Final = {"type": "file", "file": file_input, **({"mime_type": "image/png"} if source == "bytes" else {})}
litellm.rust(True)
kwargs: Final = {
"model": "mistral/mistral-ocr-latest",
"document": document,
"api_key": "test-key",
"api_base": f"http://127.0.0.1:{server.server_port}",
}
response: Final = await litellm.aocr(**kwargs) if asynchronous else litellm.ocr(**kwargs)
assert response.pages[0].markdown == "native OCR response"
assert len(requests) == 1
assert requests[0]["body"]["document"] == {
"type": "image_url",
"image_url": "data:image/png;base64," + base64.b64encode(content).decode(),
}
assert reader.calls == (1 if source in {"reader", "text"} else 0)
assert not reader.stream.closed
@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):
from litellm.rust_bridge import _native
server, requests = ocr_server
failure: Final = KeyError("reader failed")
class Reader:
def __init__(self) -> None:
self.calls = 0
def read(self) -> bytes:
self.calls += 1
raise failure
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}",
}
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)
if source is reader:
assert caught.value is failure
else:
assert caught.value.filename == str(path)
assert reader.calls == 1
assert requests == []
def test_unstarted_native_file_call_does_not_read_and_releases_reader():
from litellm.rust_bridge import _native
class Reader:
def read(self) -> bytes:
raise AssertionError("unstarted call consumed its document")
def create_call():
reader: Final = Reader()
pending: Final = _native.aocr(
model="mistral/mistral-ocr-latest", document={"type": "file", "file": reader}
)
reader.pending = pending
return pending, weakref.ref(reader)
pending, reference = create_call()
pending.close()
del pending
gc.collect()
assert reference() is None
@pytest.mark.parametrize("asynchronous", [False, True])
@pytest.mark.parametrize("model", ["mistral/mistral-ocr-latest", "azure_ai/doc-intelligence/prebuilt-read"])
@pytest.mark.asyncio