diff --git a/litellm-rust/Cargo.lock b/litellm-rust/Cargo.lock index e4432e39eda..a0cdf9568de 100644 --- a/litellm-rust/Cargo.lock +++ b/litellm-rust/Cargo.lock @@ -2006,6 +2006,7 @@ dependencies = [ name = "litellm-python-bridge" version = "0.1.0" dependencies = [ + "bytes", "criterion", "futures-util", "litellm-auth", diff --git a/litellm-rust/crates/core/src/call_lifecycle/host.rs b/litellm-rust/crates/core/src/call_lifecycle/host.rs index a62b61f1b34..f9e18b9f116 100644 --- a/litellm-rust/crates/core/src/call_lifecycle/host.rs +++ b/litellm-rust/crates/core/src/call_lifecycle/host.rs @@ -211,10 +211,10 @@ mod tests { lifecycle.accept::(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::(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); } } diff --git a/litellm-rust/crates/core/src/llms/azure_ai/ocr/cohere_parse_transformation.rs b/litellm-rust/crates/core/src/llms/azure_ai/ocr/cohere_parse_transformation.rs index 0dacb02e383..add70c2596d 100644 --- a/litellm-rust/crates/core/src/llms/azure_ai/ocr/cohere_parse_transformation.rs +++ b/litellm-rust/crates/core/src/llms/azure_ai/ocr/cohere_parse_transformation.rs @@ -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)?; diff --git a/litellm-rust/crates/core/src/llms/azure_ai/ocr/document_intelligence/transformation.rs b/litellm-rust/crates/core/src/llms/azure_ai/ocr/document_intelligence/transformation.rs index f53a4d43680..50d5b703152 100644 --- a/litellm-rust/crates/core/src/llms/azure_ai/ocr/document_intelligence/transformation.rs +++ b/litellm-rust/crates/core/src/llms/azure_ai/ocr/document_intelligence/transformation.rs @@ -541,7 +541,6 @@ impl BaseOcrConfig for AzureDocumentIntelligenceOCRConfig { ) -> Result { 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::(json!({ "type":"document_url", "document_url":"https://example.com/document.pdf" })) - .unwrap(); + .unwrap().into(); perform_ocr(request).await.unwrap(); server.await.unwrap(); diff --git a/litellm-rust/crates/core/src/llms/azure_ai/ocr/transformation.rs b/litellm-rust/crates/core/src/llms/azure_ai/ocr/transformation.rs index c7e5c13f930..dffe0aa9b05 100644 --- a/litellm-rust/crates/core/src/llms/azure_ai/ocr/transformation.rs +++ b/litellm-rust/crates/core/src/llms/azure_ai/ocr/transformation.rs @@ -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)?) } diff --git a/litellm-rust/crates/core/src/llms/cohere/ocr/transformation.rs b/litellm-rust/crates/core/src/llms/cohere/ocr/transformation.rs index af593f465b4..fb16445e11d 100644 --- a/litellm-rust/crates/core/src/llms/cohere/ocr/transformation.rs +++ b/litellm-rust/crates/core/src/llms/cohere/ocr/transformation.rs @@ -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::(json!({"output_format":"html"})).is_err()); for format in ["markdown", "blocks"] { diff --git a/litellm-rust/crates/core/src/llms/mistral/ocr/transformation.rs b/litellm-rust/crates/core/src/llms/mistral/ocr/transformation.rs index 0606c37c162..1d84b16c196 100644 --- a/litellm-rust/crates/core/src/llms/mistral/ocr/transformation.rs +++ b/litellm-rust/crates/core/src/llms/mistral/ocr/transformation.rs @@ -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 + )); } } diff --git a/litellm-rust/crates/core/src/llms/reducto/ocr/transformation.rs b/litellm-rust/crates/core/src/llms/reducto/ocr/transformation.rs index bfaa8c44064..7492a2ef6ec 100644 --- a/litellm-rust/crates/core/src/llms/reducto/ocr/transformation.rs +++ b/litellm-rust/crates/core/src/llms/reducto/ocr/transformation.rs @@ -738,8 +738,7 @@ mod tests { "result":{"chunks":[]} }))]) .await; - let 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(); diff --git a/litellm-rust/crates/core/src/llms/vertex_ai/ocr/deepseek_transformation.rs b/litellm-rust/crates/core/src/llms/vertex_ai/ocr/deepseek_transformation.rs index 87af0125710..036f58dd3f6 100644 --- a/litellm-rust/crates/core/src/llms/vertex_ai/ocr/deepseek_transformation.rs +++ b/litellm-rust/crates/core/src/llms/vertex_ai/ocr/deepseek_transformation.rs @@ -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(); diff --git a/litellm-rust/crates/core/src/llms/vertex_ai/ocr/transformation.rs b/litellm-rust/crates/core/src/llms/vertex_ai/ocr/transformation.rs index 4329a4f05f0..beca4dc141d 100644 --- a/litellm-rust/crates/core/src/llms/vertex_ai/ocr/transformation.rs +++ b/litellm-rust/crates/core/src/llms/vertex_ai/ocr/transformation.rs @@ -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 diff --git a/litellm-rust/crates/core/src/ocr/document.rs b/litellm-rust/crates/core/src/ocr/document.rs index 827a1c47365..0f737e381ef 100644 --- a/litellm-rust/crates/core/src/ocr/document.rs +++ b/litellm-rust/crates/core/src/ocr/document.rs @@ -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 { + 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::())); + 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) - ); + )); } } diff --git a/litellm-rust/crates/core/src/ocr/error.rs b/litellm-rust/crates/core/src/ocr/error.rs index 9f4c2f34ae1..0c70c478305 100644 --- a/litellm-rust/crates/core/src/ocr/error.rs +++ b/litellm-rust/crates/core/src/ocr/error.rs @@ -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, + }, + #[error("OCR document preparation task failed: {0}")] + DocumentTask(#[source] std::sync::Arc), #[error("Invalid MIME type: {0}")] InvalidMimeType(String), #[error( diff --git a/litellm-rust/crates/core/src/ocr/handler.rs b/litellm-rust/crates/core/src/ocr/handler.rs index 65ffaa38477..d66da5c81d5 100644 --- a/litellm-rust/crates/core/src/ocr/handler.rs +++ b/litellm-rust/crates/core/src/ocr/handler.rs @@ -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 { 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 { let request = super::prepare::prepare_request(request); let http = request.config.prepare_request(&request, &client).await?; diff --git a/litellm-rust/crates/core/src/ocr/hooks.rs b/litellm-rust/crates/core/src/ocr/hooks.rs index 0f7ef898281..20cb3841c87 100644 --- a/litellm-rust/crates/core/src/ocr/hooks.rs +++ b/litellm-rust/crates/core/src/ocr/hooks.rs @@ -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 +impl CallLifecycleHooks 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( &'a self, _context: &'a CallLifecycleContext, - request: LiteLLMOcrRequest, + request: ResolvedOcrRequest, ) -> Self::DuringCallFuture<'a> { Box::pin(async move { Ok(request) }) } diff --git a/litellm-rust/crates/core/src/ocr/lifecycle.rs b/litellm-rust/crates/core/src/ocr/lifecycle.rs index 9c8a058d3d6..e013c9b9c62 100644 --- a/litellm-rust/crates/core/src/ocr/lifecycle.rs +++ b/litellm-rust/crates/core/src/ocr/lifecycle.rs @@ -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), MapFailure(super::Error), @@ -82,6 +84,7 @@ impl OcrHostOperation { pub enum OcrHostResult { Request(Result<(Box, bool), super::Error>), + Document(Result), Lifecycle(Result<(), HostFailure>), AzureAdToken(Result), PreCall(Result), @@ -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, request: Option, + preparation: Option>>, + reading: bool, operations_tx: mpsc::UnboundedSender, operations_rx: mpsc::UnboundedReceiver, pending_result: Option>, @@ -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, 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>>); + 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::())); + 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" diff --git a/litellm-rust/crates/core/src/ocr/mod.rs b/litellm-rust/crates/core/src/ocr/mod.rs index bf9eecd5c79..eb3162cc79e 100644 --- a/litellm-rust/crates/core/src/ocr/mod.rs +++ b/litellm-rust/crates/core/src/ocr/mod.rs @@ -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)] diff --git a/litellm-rust/crates/core/src/ocr/prepare.rs b/litellm-rust/crates/core/src/ocr/prepare.rs index 2acbffcb3be..785f4a761b6 100644 --- a/litellm-rust/crates/core/src/ocr/prepare.rs +++ b/litellm-rust/crates/core/src/ocr/prepare.rs @@ -3,7 +3,7 @@ use serde_json::Value; use super::OcrClient; use super::hooks::OcrDuringCallRequest; -use super::types::{LiteLLMOcrRequest, OcrConnection, OcrDocument, PreparedOcrRequest}; +use super::types::{ResolvedOcrRequest, OcrConnection, OcrDocument, PreparedOcrRequest}; pub(crate) async fn transform_request_body( client: &OcrClient, @@ -111,7 +111,7 @@ pub(crate) fn credential_env(name: &str) -> Option { 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(); diff --git a/litellm-rust/crates/core/src/ocr/provider_config.rs b/litellm-rust/crates/core/src/ocr/provider_config.rs index 7c8f2692894..1b41ba8c356 100644 --- a/litellm-rust/crates/core/src/ocr/provider_config.rs +++ b/litellm-rust/crates/core/src/ocr/provider_config.rs @@ -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] diff --git a/litellm-rust/crates/core/src/ocr/test_support.rs b/litellm-rust/crates/core/src/ocr/test_support.rs index 784bc051b8f..476ffb53475 100644 --- a/litellm-rust/crates/core/src/ocr/test_support.rs +++ b/litellm-rust/crates/core/src/ocr/test_support.rs @@ -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 { diff --git a/litellm-rust/crates/core/src/ocr/types.rs b/litellm-rust/crates/core/src/ocr/types.rs index 958d1b11e97..b7483224a32 100644 --- a/litellm-rust/crates/core/src/ocr/types.rs +++ b/litellm-rust/crates/core/src/ocr/types.rs @@ -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 for OcrDocument { } } +#[derive(Clone, Debug)] +pub enum OcrDocumentInput { + Document(OcrDocument), + Path { + path: PathBuf, + mime_type: Option, + }, + Bytes { + bytes: Bytes, + file_name: Option, + mime_type: Option, + }, + HostReader { + mime_type: Option, + }, +} + +impl From for OcrDocumentInput { + fn from(document: OcrDocument) -> Self { + Self::Document(document) + } +} + +impl From for OcrDocumentInput { + fn from(path: PathBuf) -> Self { + Self::Path { path, mime_type: None } + } +} + +pub struct OcrFileContent { + pub bytes: Bytes, + pub file_name: Option, +} + #[derive(Clone, Copy, Debug, Default, PartialEq, Eq, Serialize, Deserialize)] #[serde(rename_all = "lowercase")] pub enum OcrResponseFormat { @@ -143,6 +179,38 @@ fn nonblank(value: Option) -> Option { .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, + pub api_base: Option, + pub extra_headers: Map, + pub timeout: Option, + pub input_sources: BTreeMap, +} + +impl OcrConnectionInputs { + fn source(&self, name: &str) -> InputSource { + self.input_sources.get(name).copied().unwrap_or_default() + } + + fn header_pairs(&self) -> Result, 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, @@ -199,9 +267,9 @@ pub(crate) struct ResolvedOcrCredentials { pub api_base: Option>, } -pub struct LiteLLMOcrRequest { +pub struct LiteLLMOcrRequest { pub model: String, - pub document: OcrDocument, + pub document: D, pub credentials: OcrCredentialInputs, pub transport: OcrTransportConfig, pub hooks: Arc, @@ -215,7 +283,7 @@ pub struct LiteLLMOcrRequest { impl LiteLLMOcrRequest { pub fn new( model: String, - document: OcrDocument, + document: impl Into, custom_llm_provider: Option<&str>, optional_params: CallArguments, ) -> Result { @@ -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 LiteLLMOcrRequest { + pub(crate) fn with_document(self, document: T) -> LiteLLMOcrRequest { + 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 { 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, + custom_llm_provider: Option<&str>, + optional_params: CallArguments, + connection: OcrConnectionInputs, + ) -> Result { + 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; + 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 [ diff --git a/litellm-rust/crates/python-bridge/Cargo.toml b/litellm-rust/crates/python-bridge/Cargo.toml index 1562d4c1021..6dde7c71af6 100644 --- a/litellm-rust/crates/python-bridge/Cargo.toml +++ b/litellm-rust/crates/python-bridge/Cargo.toml @@ -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 diff --git a/litellm-rust/crates/python-bridge/src/marshal.rs b/litellm-rust/crates/python-bridge/src/marshal.rs index e9efd1c61ec..a11763aaebb 100644 --- a/litellm-rust/crates/python-bridge/src/marshal.rs +++ b/litellm-rust/crates/python-bridge/src/marshal.rs @@ -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, + pub(crate) api_key: Option, + pub(crate) api_base: Option, + pub(crate) extra_headers: Map, + pub(crate) timeout: Option, + /// Only the optional params core reports as consumed plus opaque unknown + /// values; bound/control fields are never serialized. + pub(crate) optional_params: Map, + /// Which projected fields were supplied via the proxy request body. + pub(crate) input_sources: BTreeMap, + pub(crate) azure_ad_token_provider: Option, + /// 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, + bound_fields: &[&str], + ) -> PyResult { + 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::>("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> { + arguments + .optional::>("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, diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/callbacks.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/callbacks.rs index d4caa008052..65896279715 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/callbacks.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/callbacks.rs @@ -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, - model: &str, - provider: &str, - kwargs: &Py, -) -> PyResult> { - Ok(py - .import("litellm.rust_bridge.ocr")? - .getattr("map_failure")? - .call1((error, model, provider, kwargs))? - .extract()?) -} diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/document.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/document.rs index 85c54823abe..6df86e0dc54 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/document.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/document.rs @@ -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), +pub(super) struct PythonFileReader { + reader: Py, + name: Option, } -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 { + 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)> { - if file.is_instance_of::() { - 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 { + if value.is_exact_instance_of::() { + return Ok(Bytes::from_owner(value.extract::()?)); } - 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::() { - 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::()) - .transpose()?; - let value = reader.call0()?; - let bytes = if value.is_instance_of::() { - FileBytes::Native(value.extract::()?.into_bytes()) - } else if value.is_instance_of::() { - 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::()?.as_ref())) } + + pub(super) struct FileDocumentInput { - bytes: FileBytes, - name: Option, - mime_type: Option, + pub input: OcrDocumentInput, + pub reader: Option, } impl FromPyObject<'_, '_> for FileDocumentInput { @@ -103,123 +56,112 @@ impl FromPyObject<'_, '_> for FileDocumentInput { Err(error) if error.is_instance_of::(py) => None, Err(error) => return Err(error), }; + let missing = || PyValueError::new_err( + "document with type='file' must include a 'file' field containing a pathlib.Path, file-like object, or bytes", + ); let file = document.get_item("file").map_err(|error| { if error.is_instance_of::(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::() { 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::()?, + mime_type, + }, + reader: None, + }); + } + if file.is_instance_of::() { + 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::()) + .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 { - 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::().err().unwrap(); assert!(error.is_instance_of::(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::().err().unwrap(); + for expression in [c"{'file': b'abc', 'mime_type': None}", c"{'file': b'abc', 'mime_type': 7}"] { + let error = py.eval(expression, None, None).unwrap().extract::().err().unwrap(); assert!(error.is_instance_of::(py)); } - let 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::().err().unwrap(); - assert!(error.is_instance_of::(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::().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::().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"); } } diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/errors.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/errors.rs index 9f25de3dd14..655cd9ea1ee 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/errors.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/errors.rs @@ -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 { - 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) -> 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 { + 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, + model: &str, + provider: &str, +) -> PyResult { + let value = error.bind(py); + if !value.is_instance_of::() || 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 { + 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 { + 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::(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::() - .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::(py)); - assert_eq!(mapped.value(py).to_string(), "Document URL is required"); - assert_eq!( - mapped - .value(py) - .getattr("status_code") - .unwrap() - .extract::() - .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::(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::(py)); - assert_eq!( - mapped - .value(py) - .getattr("status_code") - .unwrap() - .extract::() - .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::(py)); - assert_eq!( - request - .value(py) - .getattr("status_code") - .unwrap() - .extract::() - .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::(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"); } } diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/lifecycle.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/host.rs similarity index 73% rename from litellm-rust/crates/python-bridge/src/routes/ocr/lifecycle.rs rename to litellm-rust/crates/python-bridge/src/routes/ocr/host.rs index a703e9e736a..91645aa04c1 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/lifecycle.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/host.rs @@ -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, } @@ -48,6 +30,8 @@ struct ProjectedOcrHost { azure_ad_token_provider: Option, pre_call: Option, payload: Option, + reader: Option, + 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 { 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>, - asynchronous: bool, -) -> PyResult> { - 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>, -) -> PyResult> { - call(py, args, kwargs, false) -} - -#[pyfunction] -#[pyo3(signature = (*args, **kwargs))] -fn aocr( - py: Python<'_>, - args: Bound<'_, PyTuple>, - kwargs: Option>, -) -> PyResult> { - 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)?) -} diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs index b7f9613a5a0..7973dc45c9b 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs @@ -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) -> PyResult { + 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>, + asynchronous: bool, +) -> PyResult> { + 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>, +) -> PyResult> { + call(py, args, kwargs, false) +} + +#[pyfunction] +#[pyo3(signature = (*args, **kwargs))] +fn aocr( + py: Python<'_>, + args: Bound<'_, PyTuple>, + kwargs: Option>, +) -> PyResult> { + 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::(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::(py)); + }); + } } diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/project.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/project.rs index accc8000cdb..a78aa5b80b7 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/project.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/project.rs @@ -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, pub secret_fields: Vec<&'static str>, + pub reader: Option, } -fn project_document(py: Python<'_>, document: &Bound<'_, PyAny>) -> PyResult { +enum ProjectedDocument { + File(FileDocumentInput), + Url(serde_json::Value), +} + +impl ProjectedDocument { + fn into_native(self) -> Result { + match self { + Self::File(file) => Ok(file), + Self::Url(value) => Ok(FileDocumentInput { + input: OcrDocument::try_from(value)?.into(), + reader: None, + }), + } + } +} + +fn project_document(document: &Bound<'_, PyAny>) -> PyResult { let kind: String = document.get_item("type")?.extract()?; if kind != "file" { - 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>) -> Result, 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, name: &str) -> InputSource { - sources.get(name).copied().unwrap_or_default() -} - -pub(super) fn project(py: Python<'_>, arguments: &BoundArguments<'_>) -> PyResult { - let kwargs: &Bound<'_, PyDict> = arguments.kwargs(); - let model: String = arguments.extract("model")?; - let custom_llm_provider: Option = arguments.optional("custom_llm_provider")?; - let document = project_document(py, &arguments.required("document")?)?; - let api_key: Option = 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 = arguments.optional("api_base")?; - let extra_headers: Option> = 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 { + 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) -> PyResult { - 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 { + let model: String = arguments.extract("model")?; + let custom_llm_provider: Option = 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::(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::(py)); - assert!(!error.is_instance_of::(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::(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::(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 = document.getattr("reads").unwrap().extract().unwrap(); assert_eq!(reads, ["type", "mime_type", "file"]); }); diff --git a/litellm/rust_bridge/ocr.py b/litellm/rust_bridge/ocr.py index 66ec39b5280..7ede521e6f3 100644 --- a/litellm/rust_bridge/ocr.py +++ b/litellm/rust_bridge/ocr.py @@ -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 diff --git a/tests/test_litellm_rust/test_ocr.py b/tests/test_litellm_rust/test_ocr.py index 87fa8b38666..a83748c0dd6 100644 --- a/tests/test_litellm_rust/test_ocr.py +++ b/tests/test_litellm_rust/test_ocr.py @@ -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