From 1771e5bb68e56ed4a9afafd66b9986eb422a4add Mon Sep 17 00:00:00 2001 From: Yujong Lee Date: Wed, 16 Sep 2026 07:07:41 -0700 Subject: [PATCH] refactor callback --- .../crates/core/src/call_lifecycle/host.rs | 8 +- .../document_intelligence/transformation.rs | 3 +- .../src/llms/cohere/ocr/transformation.rs | 16 +- .../src/llms/reducto/ocr/transformation.rs | 14 +- .../src/llms/vertex_ai/ocr/transformation.rs | 8 +- litellm-rust/crates/core/src/ocr/client.rs | 12 +- litellm-rust/crates/core/src/ocr/document.rs | 27 +- litellm-rust/crates/core/src/ocr/error.rs | 2 + litellm-rust/crates/core/src/ocr/handler.rs | 2 +- litellm-rust/crates/core/src/ocr/hooks.rs | 1 + litellm-rust/crates/core/src/ocr/lifecycle.rs | 270 +++++++++++++----- litellm-rust/crates/core/src/ocr/mod.rs | 5 +- litellm-rust/crates/core/src/ocr/prepare.rs | 4 +- litellm-rust/crates/core/src/ocr/types.rs | 9 +- .../crates/python-bridge/src/lifecycle/mod.rs | 37 +-- .../crates/python-bridge/src/marshal.rs | 5 + .../python-bridge/src/routes/ocr/callbacks.rs | 173 +++-------- .../python-bridge/src/routes/ocr/document.rs | 63 +++- .../python-bridge/src/routes/ocr/errors.rs | 10 +- .../python-bridge/src/routes/ocr/host.rs | 209 ++++++-------- .../python-bridge/src/routes/ocr/mod.rs | 2 +- .../python-bridge/src/routes/ocr/project.rs | 90 +++--- litellm/ocr/main.py | 10 +- litellm/proxy/ocr_endpoints/endpoints.py | 10 +- litellm/rust_bridge/bindings.py | 5 + litellm/rust_bridge/ocr.py | 80 +++++- .../test_litellm/rust_bridge/test_bindings.py | 18 ++ .../rust_bridge/test_ocr_lifecycle.py | 68 ++++- tests/test_litellm_rust/test_ocr.py | 110 +++++-- 29 files changed, 799 insertions(+), 472 deletions(-) diff --git a/litellm-rust/crates/core/src/call_lifecycle/host.rs b/litellm-rust/crates/core/src/call_lifecycle/host.rs index f9e18b9f116..0e93b95b947 100644 --- a/litellm-rust/crates/core/src/call_lifecycle/host.rs +++ b/litellm-rust/crates/core/src/call_lifecycle/host.rs @@ -223,9 +223,11 @@ mod tests { ] { assert_eq!(lifecycle.phase(), phase); assert!( - lifecycle.accept(Err(HostFailure::Error(crate::ocr::Error::InvalidRequest( - "callback".into() - )))).is_none() + lifecycle + .accept(Err(HostFailure::Error(crate::ocr::Error::InvalidRequest( + "callback".into() + )))) + .is_none() ); } assert_eq!(lifecycle.phase(), HostPhase::Complete); 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 50d5b703152..9529a726ac3 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 @@ -816,7 +816,8 @@ mod tests { "type":"document_url", "document_url":"https://example.com/document.pdf" })) - .unwrap().into(); + .unwrap() + .into(); perform_ocr(request).await.unwrap(); server.await.unwrap(); 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 fb16445e11d..fc11f62833c 100644 --- a/litellm-rust/crates/core/src/llms/cohere/ocr/transformation.rs +++ b/litellm-rust/crates/core/src/llms/cohere/ocr/transformation.rs @@ -361,10 +361,12 @@ mod tests { } }), ); - let request = request.with_document(serde_json::from_value(json!({ - "type":"image_url","image_url":"https://example.com/original.png" - })) - .unwrap()); + let request = request.with_document( + serde_json::from_value(json!({ + "type":"image_url","image_url":"https://example.com/original.png" + })) + .unwrap(), + ); let request = crate::ocr::prepare::prepare_request(request); let http = CohereParseConfig .prepare_request(&request, &crate::ocr::test_support::ocr_client()) @@ -503,10 +505,12 @@ mod tests { "https://example.com", json!({"output_format":null,"req_format":null}), ); - let request = request.with_document(serde_json::from_value( + let request = request.with_document( + serde_json::from_value( json!({"type":"image_url","image_url":"https://example.com/a.png"}), ) - .unwrap()); + .unwrap(), + ); assert_eq!( request.response_format().unwrap(), crate::ocr::types::OcrResponseFormat::Litellm 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 7492a2ef6ec..4b05f1ee168 100644 --- a/litellm-rust/crates/core/src/llms/reducto/ocr/transformation.rs +++ b/litellm-rust/crates/core/src/llms/reducto/ocr/transformation.rs @@ -738,7 +738,8 @@ mod tests { "result":{"chunks":[]} }))]) .await; - let request = crate::ocr::test_support::with_source(wire_request(model, &base, options), source); + let request = + crate::ocr::test_support::with_source(wire_request(model, &base, options), source); perform_ocr(request).await.unwrap(); server.await.unwrap(); @@ -857,7 +858,9 @@ mod tests { #[tokio::test] async fn rejects_invalid_document_sources_before_network(#[case] source: &str) { let request = crate::ocr::test_support::with_source( - wire_request("reducto/parse-v3", "http://127.0.0.1:1", json!({})), source); + wire_request("reducto/parse-v3", "http://127.0.0.1:1", json!({})), + source, + ); assert!(perform_ocr(request).await.is_err()); } @@ -913,7 +916,9 @@ mod tests { let raw = json!({"job_id":"job-1","result":{"chunks":[]}}); let (base, seen, server) = mock_server(vec![MockResponse::json(raw)]).await; let mut request = crate::ocr::test_support::with_source( - wire_request("reducto/parse-v3", &base, json!({})), "reducto://ready.pdf"); + wire_request("reducto/parse-v3", &base, json!({})), + "reducto://ready.pdf", + ); request.transport.extra_headers = vec![("authorization".into(), "Bearer existing".into())]; let response = perform_ocr(request).await.unwrap(); @@ -984,6 +989,7 @@ mod tests { request: OcrDuringCallRequest, ) -> OcrHookFuture<'_, OcrDuringCallRequest> { Box::pin(async move { + assert_eq!(request.optional_params["use_cache"], json!(true)); assert_eq!( request.body["document_url"], "data:application/pdf;base64,YWJj" @@ -1000,7 +1006,7 @@ mod tests { async fn guardrail_rewrites_document_before_upload() { let (base, seen, server) = mock_server(vec![MockResponse::json(json!({"result":{"chunks":[]}}))]).await; - let mut request = wire_request("reducto/parse-v3", &base, json!({})); + let mut request = wire_request("reducto/parse-v3", &base, json!({"use_cache":true})); request.hooks = Arc::new(RewriteDocument); perform_ocr(request).await.unwrap(); 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 beca4dc141d..a3eb35f16f2 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 @@ -338,8 +338,12 @@ mod tests { options.clone(), ); let vertex = wire_request("vertex_ai/mistral-ocr-maas", "https://vertex.test", options); - let direct = crate::ocr::prepare::prepare_request(crate::ocr::test_support::resolved_request(direct)); - let vertex = crate::ocr::prepare::prepare_request(crate::ocr::test_support::resolved_request(vertex)); + let direct = crate::ocr::prepare::prepare_request( + crate::ocr::test_support::resolved_request(direct), + ); + let vertex = crate::ocr::prepare::prepare_request( + crate::ocr::test_support::resolved_request(vertex), + ); let direct_http = MistralOCRConfig .prepare_request(&direct, &client) .await diff --git a/litellm-rust/crates/core/src/ocr/client.rs b/litellm-rust/crates/core/src/ocr/client.rs index 5881519855c..6ceb4eedbd2 100644 --- a/litellm-rust/crates/core/src/ocr/client.rs +++ b/litellm-rust/crates/core/src/ocr/client.rs @@ -39,9 +39,10 @@ impl OcrClient { ) -> Result { use super::{ NativeOutcome, OcrAdmission, OcrCall, OcrCallStep, OcrHookHost, OcrHost, - OcrHostOperation, OcrHostResult, + OcrHostOperation, OcrHostResult, OcrProjectedRequest, }; + let intercepts_requests = request.hooks.intercepts_requests(); let host = OcrHookHost::new(request.hooks.clone()); let mut request = Some(request); let NativeOutcome::Completed(mut call) = OcrCall::admit(self.clone(), OcrAdmission::all()) @@ -54,14 +55,15 @@ impl OcrClient { loop { match call.resume(result.take()).await? { OcrCallStep::Host(OcrHostOperation::ProjectRequest) => { - result = Some(OcrHostResult::Request(Ok(( - Box::new(request.take().ok_or_else(|| { + result = Some(OcrHostResult::Request(Ok(OcrProjectedRequest { + request: Box::new(request.take().ok_or_else(|| { crate::ocr::Error::InvalidRequest( "OCR request was already projected".into(), ) })?), - false, - )))) + intercepts_requests, + host_token_provider: false, + }))) } OcrCallStep::Host(operation) => result = Some(host.invoke(operation).await), OcrCallStep::Complete(response) => return Ok(response), diff --git a/litellm-rust/crates/core/src/ocr/document.rs b/litellm-rust/crates/core/src/ocr/document.rs index 0f737e381ef..e9a13f0fd11 100644 --- a/litellm-rust/crates/core/src/ocr/document.rs +++ b/litellm-rust/crates/core/src/ocr/document.rs @@ -15,7 +15,10 @@ pub(crate) fn read_path_document( ) -> 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)) + .and_then(|file| { + file.take(OCR_INLINE_MAX_BYTES as u64 + 1) + .read_to_end(&mut bytes) + }) .map_err(|source| super::Error::FileRead { path: path.to_owned(), source: std::sync::Arc::new(source), @@ -201,9 +204,14 @@ mod tests { #[test] fn path_preparation_preserves_io_causes_and_enforces_the_inline_limit() { - let path = std::env::temp_dir().join(format!("ocr-document-{:032x}.png", rand::random::())); + 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 { + let super::super::Error::FileRead { + path: failed_path, + source, + } = error + else { panic!("missing path must produce a typed file error"); }; assert_eq!(failed_path, path); @@ -215,9 +223,16 @@ mod tests { let image = read_path_document(&path, None).unwrap(); let overridden = read_path_document(&path, Some("application/pdf")).unwrap(); std::fs::remove_file(&path).unwrap(); - assert!(matches!(oversized, Err(super::super::Error::InlineDocumentTooLarge))); - assert!(matches!(image, OcrDocument::ImageUrl { image_url, .. } if image_url == "data:image/png;base64,aW1hZ2UgYnl0ZXM=")); - assert!(matches!(overridden, OcrDocument::DocumentUrl { document_url, .. } if document_url == "data:application/pdf;base64,aW1hZ2UgYnl0ZXM=")); + assert!(matches!( + oversized, + Err(super::super::Error::InlineDocumentTooLarge) + )); + assert!( + matches!(image, OcrDocument::ImageUrl { image_url, .. } if image_url == "data:image/png;base64,aW1hZ2UgYnl0ZXM=") + ); + assert!( + matches!(overridden, OcrDocument::DocumentUrl { document_url, .. } if document_url == "data:application/pdf;base64,aW1hZ2UgYnl0ZXM=") + ); } fn document(source: &str) -> OcrDocument { diff --git a/litellm-rust/crates/core/src/ocr/error.rs b/litellm-rust/crates/core/src/ocr/error.rs index 0c70c478305..a8ac969577f 100644 --- a/litellm-rust/crates/core/src/ocr/error.rs +++ b/litellm-rust/crates/core/src/ocr/error.rs @@ -8,6 +8,8 @@ pub enum Error { }, #[error("File is empty or could not be read")] EmptyFile, + #[error("Host OCR document read failed")] + HostDocumentRead, #[error("Failed to read OCR file {}: {source}", path.display())] FileRead { path: std::path::PathBuf, diff --git a/litellm-rust/crates/core/src/ocr/handler.rs b/litellm-rust/crates/core/src/ocr/handler.rs index d66da5c81d5..7e42111da0a 100644 --- a/litellm-rust/crates/core/src/ocr/handler.rs +++ b/litellm-rust/crates/core/src/ocr/handler.rs @@ -2,7 +2,7 @@ use std::sync::Arc; use super::OcrClient; use super::hooks::{OcrHooks, OcrLifecycleHooks, OcrPostCallRequest}; -use super::types::{ResolvedOcrRequest, LiteLLMOcrResponse, PreparedOcrRequest}; +use super::types::{LiteLLMOcrResponse, PreparedOcrRequest, ResolvedOcrRequest}; use crate::call_lifecycle::{CallLifecycle, CallLifecycleContext}; use crate::llms::base_llm::ocr::transformation::OcrResponseContext; diff --git a/litellm-rust/crates/core/src/ocr/hooks.rs b/litellm-rust/crates/core/src/ocr/hooks.rs index 20cb3841c87..ee9b0f776d9 100644 --- a/litellm-rust/crates/core/src/ocr/hooks.rs +++ b/litellm-rust/crates/core/src/ocr/hooks.rs @@ -22,6 +22,7 @@ pub struct OcrPreCallRequest { pub struct OcrDuringCallRequest { pub model: String, pub custom_llm_provider: String, + pub optional_params: Value, pub api_key: Option, pub url: String, pub headers: Vec<(String, String)>, diff --git a/litellm-rust/crates/core/src/ocr/lifecycle.rs b/litellm-rust/crates/core/src/ocr/lifecycle.rs index e013c9b9c62..dd88e11dca6 100644 --- a/litellm-rust/crates/core/src/ocr/lifecycle.rs +++ b/litellm-rust/crates/core/src/ocr/lifecycle.rs @@ -1,6 +1,7 @@ use std::future::Future; use std::pin::Pin; use std::sync::Arc; +use std::time::{SystemTime, UNIX_EPOCH}; use tokio::sync::{mpsc, oneshot}; @@ -9,8 +10,8 @@ use super::hooks::{ OcrDuringCallRequest, OcrHookFuture, OcrHooks, OcrLogFuture, OcrPostCallRequest, OcrPreCallRequest, }; -use super::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrClient, OcrDocumentInput, OcrFileContent}; use super::types::ResolvedOcrRequest; +use super::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrClient, OcrDocumentInput, OcrFileContent}; use crate::call_lifecycle::host::{ HostCall, HostCallFuture, HostCallStep, HostFailure, HostLifecycle, HostPhase, }; @@ -82,8 +83,14 @@ impl OcrHostOperation { } } +pub struct OcrProjectedRequest { + pub request: Box, + pub intercepts_requests: bool, + pub host_token_provider: bool, +} + pub enum OcrHostResult { - Request(Result<(Box, bool), super::Error>), + Request(Result), Document(Result), Lifecycle(Result<(), HostFailure>), AzureAdToken(Result), @@ -160,9 +167,10 @@ impl OcrCall { Some(OcrHostResult::Request(result)) if self.projecting => { self.projecting = false; match result { - Ok((request, azure_ad_token_provider)) => { - self.execution.request = Some(*request); - self.execution.azure_ad_token_provider = azure_ad_token_provider; + Ok(projected) => { + self.execution.set_request(*projected.request); + self.execution.intercepts_requests = projected.intercepts_requests; + self.execution.host_token_provider = projected.host_token_provider; } Err(error) => self.accept(Err(HostFailure::Error(error))), } @@ -268,6 +276,15 @@ impl OcrCall { fn accept(&mut self, result: Result<(), HostFailure>) { let cancelled = matches!(&result, Err(HostFailure::Cancelled(_))); if let Some(error) = self.lifecycle.accept(result) { + if let Some((_, timing)) = self + .execution + .terminal + .lock() + .unwrap_or_else(|error| error.into_inner()) + .as_mut() + { + timing.end_time = epoch_seconds(); + } if cancelled { self.error = Some(error); } else { @@ -333,11 +350,33 @@ struct OcrExecution { pending_result: Option>, execution: Option>>, completed: bool, - azure_ad_token_provider: bool, + intercepts_requests: bool, + host_token_provider: bool, terminal: Arc>>, } impl OcrExecution { + fn set_request(&mut self, mut request: LiteLLMOcrRequest) { + let call_id = request + .litellm_call_id + .clone() + .unwrap_or_else(|| format!("ocr-{:032x}", rand::random::())); + let context = CallLifecycleContext::new( + "ocr", + request.model.clone(), + request.provider_name(), + call_id.clone(), + ); + request.litellm_call_id = Some(call_id); + let start = epoch_seconds(); + *self + .terminal + .lock() + .unwrap_or_else(|error| error.into_inner()) = + Some((context, CallLifecycleTiming::new(start, start, Vec::new()))); + self.request = Some(request); + } + fn new(client: OcrClient) -> Self { let (operations_tx, operations_rx) = mpsc::unbounded_channel(); Self { @@ -350,7 +389,8 @@ impl OcrExecution { pending_result: None, execution: None, completed: false, - azure_ad_token_provider: false, + intercepts_requests: false, + host_token_provider: false, terminal: Arc::default(), } } @@ -366,11 +406,16 @@ impl OcrExecution { } let result = if self.reading { let Some(OcrHostResult::Document(content)) = result else { - return Err(super::Error::InvalidRequest("OCR document read result is required".into())); + return Err(super::Error::InvalidRequest( + "OCR document read result is required".into(), + )); }; self.reading = false; let content = content?; - let request = self.request.take().expect("pending document read has a request"); + let request = self + .request + .take() + .expect("pending document read has a request"); let OcrDocumentInput::HostReader { mime_type } = &request.document else { unreachable!("only host readers request document reads"); }; @@ -428,7 +473,10 @@ impl OcrExecution { async fn prepare(&mut self) -> Result, super::Error> { if self.preparation.is_none() { - let request = self.request.take().expect("admitted OCR call has a request"); + let request = self + .request + .take() + .expect("admitted OCR call has a request"); if matches!(request.document, OcrDocumentInput::HostReader { .. }) { self.request = Some(request); self.reading = true; @@ -444,15 +492,25 @@ impl OcrExecution { OcrDocumentInput::Path { path, mime_type } => { super::document::read_path_document(path, mime_type.as_deref())? } - OcrDocumentInput::Bytes { bytes, file_name, mime_type } => { - super::document::encode_file_document(bytes, file_name.as_deref(), mime_type.as_deref())? - } + OcrDocumentInput::Bytes { + bytes, + file_name, + mime_type, + } => super::document::encode_file_document( + bytes, + file_name.as_deref(), + mime_type.as_deref(), + )?, _ => unreachable!("only native file inputs require preparation"), }; Ok(request.with_document(document)) })); } - let result = self.preparation.as_mut().expect("document preparation started").await; + let result = self + .preparation + .as_mut() + .expect("document preparation started") + .await; self.preparation = None; let request = result.map_err(|error| super::Error::DocumentTask(Arc::new(error)))??; self.start(request); @@ -461,8 +519,7 @@ impl OcrExecution { fn start(&mut self, mut request: ResolvedOcrRequest) { let client = self.client.take().expect("admitted OCR call has a client"); - let intercepts_requests = request.hooks.intercepts_requests(); - if self.azure_ad_token_provider { + if self.host_token_provider { request.azure_ad_token_provider = Some(TokenProviderHandle::new(Arc::new( OcrAzureAdTokenProvider { operations: self.operations_tx.clone(), @@ -471,7 +528,7 @@ impl OcrExecution { } request.hooks = Arc::new(ProtocolHooks { operations: self.operations_tx.clone(), - intercepts_requests, + intercepts_requests: self.intercepts_requests, terminal: self.terminal.clone(), }); self.execution = Some(tokio::spawn(async move { @@ -517,6 +574,13 @@ struct ProtocolHooks { terminal: Arc>>, } +fn epoch_seconds() -> f64 { + SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap_or_default() + .as_secs_f64() +} + #[derive(Debug)] struct OcrAzureAdTokenProvider { operations: mpsc::UnboundedSender, @@ -743,36 +807,60 @@ mod tests { use crate::ocr::test_support::{MockResponse, mock_server, perform_ocr, wire_request}; use crate::ocr::{ LiteLLMOcrRequest, NativeOutcome, NoopOcrHost, OcrAdmission, OcrCall, OcrCallStep, - OcrDecline, OcrDocument, OcrHost, OcrHostOperation, OcrHostResult, + OcrDecline, OcrDocument, OcrHost, OcrHostOperation, OcrHostResult, OcrProjectedRequest, }; + fn projected_request(request: LiteLLMOcrRequest) -> OcrHostResult { + OcrHostResult::Request(Ok(OcrProjectedRequest { + intercepts_requests: request.hooks.intercepts_requests(), + request: Box::new(request), + host_token_provider: false, + })) + } + #[tokio::test] async fn sdk_paths_are_read_during_execution_and_hooks_observe_normalized_documents() { struct CaptureDocument(Arc>>); impl OcrHooks for CaptureDocument { - fn intercepts_requests(&self) -> bool { true } + fn intercepts_requests(&self) -> bool { + true + } fn pre_call(&self, request: OcrPreCallRequest) -> OcrHookFuture<'_, OcrPreCallRequest> { self.0.lock().unwrap().push(request.document.clone()); Box::pin(async { Ok(request) }) } } let (base, seen, server) = mock_server(vec![MockResponse::json(json!({"pages":[]}))]).await; - let path = std::env::temp_dir().join(format!("ocr-sdk-{:032x}.pdf", rand::random::())); + 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); + "mistral/model".into(), + path.clone(), + None, + Default::default(), + crate::ocr::OcrConnectionInputs { + api_base: Some(base), + api_key: Some("test-key".into()), + ..Default::default() + }, + ) + .unwrap() + .with_host_hooks(Arc::new(CaptureDocument(documents.clone())), None); std::fs::write(&path, b"sdk document").unwrap(); let result = perform_ocr(request).await; std::fs::remove_file(&path).unwrap(); result.unwrap(); server.await.unwrap(); let expected = json!({"type":"document_url","document_url":"data:application/pdf;base64,c2RrIGRvY3VtZW50"}); - assert_eq!(serde_json::to_value(&documents.lock().unwrap()[0]).unwrap(), expected); + assert_eq!( + serde_json::to_value(&documents.lock().unwrap()[0]).unwrap(), + expected + ); let requests = seen.lock().unwrap(); assert_eq!(requests.len(), 1); - let body: Value = serde_json::from_str(requests[0].split_once("\r\n\r\n").unwrap().1).unwrap(); + let body: Value = + serde_json::from_str(requests[0].split_once("\r\n\r\n").unwrap().1).unwrap(); assert_eq!(body["document"], expected); } @@ -781,42 +869,80 @@ mod tests { let (base, seen, server) = mock_server(vec![MockResponse::json(json!({"pages":[]}))]).await; let request = wire_request("mistral/model", &base, json!({})) .with_document(crate::ocr::OcrDocumentInput::HostReader { mime_type: None }) - .with_host_hooks(Arc::new(AdmissionSpy { effects: Arc::new(Mutex::new(0)) }), None); + .with_host_hooks( + Arc::new(AdmissionSpy { + effects: Arc::new(Mutex::new(0)), + }), + None, + ); let mut request = Some(request); - let NativeOutcome::Completed(mut call) = OcrCall::admit(crate::ocr::test_support::ocr_client(), OcrAdmission::all()) else { panic!("admission declined"); }; + let NativeOutcome::Completed(mut call) = + OcrCall::admit(crate::ocr::test_support::ocr_client(), OcrAdmission::all()) + else { + panic!("admission declined"); + }; let mut result = None; let mut reads = 0; let mut pre_calls = 0; - loop { - match call.resume(result.take()).await.unwrap() { - OcrCallStep::Host(operation) => { - result = Some(match operation { - OcrHostOperation::ProjectRequest => { - assert_eq!(reads, 0); - OcrHostResult::Request(Ok((Box::new(request.take().unwrap()), false))) - } - OcrHostOperation::ReadDocument => { - reads += 1; - OcrHostResult::Document(Ok(crate::ocr::OcrFileContent { - bytes: bytes::Bytes::from_static(b"image"), file_name: Some("scan.png".into()), - })) - } - OcrHostOperation::PreCall(request) => { - assert_eq!(reads, 1); - pre_calls += 1; - assert!(matches!(&request.document, OcrDocument::ImageUrl { image_url, .. } if image_url == "data:image/png;base64,aW1hZ2U=")); - OcrHostResult::PreCall(Ok(request)) - } - operation => NoopOcrHost.invoke(operation).await, - }); + while let OcrCallStep::Host(operation) = call.resume(result.take()).await.unwrap() { + result = Some(match operation { + OcrHostOperation::ProjectRequest => { + assert_eq!(reads, 0); + projected_request(request.take().unwrap()) } - OcrCallStep::Complete(_) => break, - } + OcrHostOperation::ReadDocument => { + reads += 1; + OcrHostResult::Document(Ok(crate::ocr::OcrFileContent { + bytes: bytes::Bytes::from_static(b"image"), + file_name: Some("scan.png".into()), + })) + } + OcrHostOperation::PreCall(request) => { + assert_eq!(reads, 1); + pre_calls += 1; + assert!( + matches!(&request.document, OcrDocument::ImageUrl { image_url, .. } if image_url == "data:image/png;base64,aW1hZ2U=") + ); + OcrHostResult::PreCall(Ok(request)) + } + operation => NoopOcrHost.invoke(operation).await, + }); } server.await.unwrap(); assert_eq!((reads, pre_calls, seen.lock().unwrap().len()), (1, 1, 1)); } + #[tokio::test] + async fn sdk_file_failure_dispatches_the_typed_error_once() { + struct CaptureFailure(Arc>>); + impl OcrHooks for CaptureFailure { + fn failure<'a>( + &'a self, + _: &'a CallLifecycleContext, + error: &'a crate::ocr::Error, + _: &'a CallLifecycleTiming, + ) -> OcrLogFuture<'a> { + self.0.lock().unwrap().push(error.clone()); + Box::pin(async {}) + } + } + let path = + std::env::temp_dir().join(format!("missing-ocr-{:032x}.pdf", rand::random::())); + let failures = Arc::new(Mutex::new(Vec::new())); + let request = wire_request("mistral/model", "http://127.0.0.1:1", json!({})) + .with_document(path.clone().into()) + .with_host_hooks(Arc::new(CaptureFailure(failures.clone())), None); + let error = perform_ocr(request).await.unwrap_err(); + assert!( + matches!(error, crate::ocr::Error::FileRead { path: actual, .. } if actual == path) + ); + let failures = failures.lock().unwrap(); + assert_eq!(failures.len(), 1); + assert!( + matches!(&failures[0], crate::ocr::Error::FileRead { source, .. } if source.kind() == std::io::ErrorKind::NotFound) + ); + } + #[tokio::test] async fn cancellation_acknowledges_blocking_preparation_completion() { use std::sync::atomic::{AtomicBool, Ordering}; @@ -836,7 +962,8 @@ mod tests { std::future::poll_fn(|cx| { assert!(stop.as_mut().poll(cx).is_pending()); std::task::Poll::Ready(()) - }).await; + }) + .await; assert!(!finished.load(Ordering::SeqCst)); release_tx.send(()).unwrap(); stop.await; @@ -1034,6 +1161,8 @@ mod tests { mut request: OcrDuringCallRequest, ) -> OcrHookFuture<'_, OcrDuringCallRequest> { Box::pin(async move { + assert_eq!(request.optional_params["pages"], json!([0])); + assert_eq!(request.optional_params["extra_body"]["pages"], json!([2])); assert_eq!(request.body["pages"], json!([2])); assert_eq!(request.body.get("future"), Some(&Value::Null)); request.body.as_object_mut().unwrap().remove("future"); @@ -1287,10 +1416,7 @@ mod tests { result = Some(OcrHostResult::Lifecycle(Ok(()))) } OcrHostOperation::ProjectRequest => { - result = Some(OcrHostResult::Request(Ok(( - Box::new(request.take().unwrap()), - false, - )))) + result = Some(projected_request(request.take().unwrap())) } OcrHostOperation::AcquireAzureAdToken => { panic!("test request has no token provider") @@ -1312,7 +1438,9 @@ mod tests { Ok(request) })); } - OcrHostOperation::PostCall(_) | OcrHostOperation::ReadDocument => panic!("transport should not be reached"), + OcrHostOperation::PostCall(_) | OcrHostOperation::ReadDocument => { + panic!("transport should not be reached") + } }, Err(error) => break error, Ok(OcrCallStep::Complete(_)) => panic!("failed call completed"), @@ -1345,10 +1473,7 @@ mod tests { let error = loop { match call.resume(result.take()).await { Ok(OcrCallStep::Host(OcrHostOperation::ProjectRequest)) => { - result = Some(OcrHostResult::Request(Ok(( - Box::new(request.take().unwrap()), - false, - )))); + result = Some(projected_request(request.take().unwrap())); } Ok(OcrCallStep::Host(operation)) => { if let OcrHostOperation::PostCall(request) = &operation { @@ -1412,7 +1537,7 @@ mod tests { }); result = Some(match operation { OcrHostOperation::ProjectRequest => { - OcrHostResult::Request(Ok((Box::new(request.take().unwrap()), false))) + projected_request(request.take().unwrap()) } operation => host.invoke(operation).await, }); @@ -1472,7 +1597,9 @@ mod tests { OcrHostResult::Lifecycle(Err(HostFailure::Error(selected.clone()))) } OcrHostOperation::Failure { error, .. } => { - assert!(matches!(error, crate::ocr::Error::InvalidRequest(message) if message == "public metadata failed")); + assert!( + matches!(error, crate::ocr::Error::InvalidRequest(message) if message == "public metadata failed") + ); failures.push("sync"); OcrHostResult::Lifecycle(Err(HostFailure::Error( crate::ocr::Error::InvalidRequest("failure callback failed".into()), @@ -1488,7 +1615,7 @@ mod tests { panic!("finalization failure used provider/success dispatch") } OcrHostOperation::ProjectRequest => { - OcrHostResult::Request(Ok((Box::new(request.take().unwrap()), false))) + projected_request(request.take().unwrap()) } operation => host.invoke(operation).await, }); @@ -1498,7 +1625,9 @@ mod tests { } }; server.await.unwrap(); - assert!(matches!(error, crate::ocr::Error::InvalidRequest(message) if message == "public metadata failed")); + assert!( + matches!(error, crate::ocr::Error::InvalidRequest(message) if message == "public metadata failed") + ); assert_eq!(failures, ["sync", "async"]); assert_eq!(seen.lock().unwrap().len(), 1); } @@ -1525,10 +1654,7 @@ mod tests { match call.resume(result.take()).await.unwrap() { OcrCallStep::Host(OcrHostOperation::PreCall(_)) => break, OcrCallStep::Host(OcrHostOperation::ProjectRequest) => { - result = Some(OcrHostResult::Request(Ok(( - Box::new(request.take().unwrap()), - false, - )))) + result = Some(projected_request(request.take().unwrap())) } OcrCallStep::Host(operation) => result = Some(host.invoke(operation).await), OcrCallStep::Complete(_) => panic!("provider executed before pre-call result"), @@ -1751,7 +1877,7 @@ mod tests { _ = entered.notified() => break, step = call.resume(result.take()) => { result = Some(match step.unwrap() { - OcrCallStep::Host(OcrHostOperation::ProjectRequest) => OcrHostResult::Request(Ok((Box::new(request.take().unwrap()), false))), + OcrCallStep::Host(OcrHostOperation::ProjectRequest) => projected_request(request.take().unwrap()), OcrCallStep::Host(operation) => NoopOcrHost.invoke(operation).await, OcrCallStep::Complete(_) => panic!("pending provider completed"), }); @@ -1778,7 +1904,9 @@ mod tests { ) .await .unwrap(); - assert!(matches!(result, Err(crate::ocr::Error::InvalidRequest(message)) if message == "cancelled")); + assert!( + matches!(result, Err(crate::ocr::Error::InvalidRequest(message)) if message == "cancelled") + ); assert!( dropped.load(Ordering::SeqCst), "cancellation returned while provider captures were still alive" diff --git a/litellm-rust/crates/core/src/ocr/mod.rs b/litellm-rust/crates/core/src/ocr/mod.rs index eb3162cc79e..95bb827c104 100644 --- a/litellm-rust/crates/core/src/ocr/mod.rs +++ b/litellm-rust/crates/core/src/ocr/mod.rs @@ -18,12 +18,13 @@ pub use client::{OcrClient, ocr}; pub use document::{encode_file_document, mime_type_for_name, upload_mime_type}; pub use lifecycle::{ NativeOutcome, NativeResult, NoopOcrHost, OcrAdmission, OcrCall, OcrCallStep, OcrDecline, - OcrHookHost, OcrHost, OcrHostOperation, OcrHostResult, + OcrHookHost, OcrHost, OcrHostOperation, OcrHostResult, OcrProjectedRequest, }; pub use provider_config::{get_api_key_env_var, get_health_check_document}; pub use types::{ LiteLLMOcrRequest, LiteLLMOcrResponse, OcrConnection, OcrConnectionInputs, OcrCredentialInputs, - OcrDocument, OcrDocumentInput, OcrFileContent, OcrPage, OcrPageDimensions, OcrPageImage, OcrTransportConfig, OcrUsageInfo, + OcrDocument, OcrDocumentInput, OcrFileContent, OcrPage, OcrPageDimensions, OcrPageImage, + OcrTransportConfig, OcrUsageInfo, }; #[cfg(test)] diff --git a/litellm-rust/crates/core/src/ocr/prepare.rs b/litellm-rust/crates/core/src/ocr/prepare.rs index 785f4a761b6..19e8912ee7d 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::{ResolvedOcrRequest, OcrConnection, OcrDocument, PreparedOcrRequest}; +use super::types::{OcrConnection, OcrDocument, PreparedOcrRequest, ResolvedOcrRequest}; pub(crate) async fn transform_request_body( client: &OcrClient, @@ -28,6 +28,7 @@ where .during_call(OcrDuringCallRequest { model: request.model.clone(), custom_llm_provider: request.provider_name().into(), + optional_params: Value::Object(request.optional_params.clone().into()), api_key: request.connection.api_key.clone(), url: url.into(), headers: headers.to_vec(), @@ -78,6 +79,7 @@ pub(crate) async fn guardrail_document( .during_call(OcrDuringCallRequest { model: request.model.clone(), custom_llm_provider: request.provider_name().into(), + optional_params: Value::Object(request.optional_params.clone().into()), api_key: request.connection.api_key.clone(), url: url.into(), headers: headers.to_vec(), diff --git a/litellm-rust/crates/core/src/ocr/types.rs b/litellm-rust/crates/core/src/ocr/types.rs index b7483224a32..d92cd546be4 100644 --- a/litellm-rust/crates/core/src/ocr/types.rs +++ b/litellm-rust/crates/core/src/ocr/types.rs @@ -3,8 +3,8 @@ use std::path::PathBuf; use std::sync::Arc; use std::time::Duration; -use serde::{Deserialize, Serialize}; use bytes::Bytes; +use serde::{Deserialize, Serialize}; use serde_json::{Map, Value}; use serde_with::serde_as; @@ -93,7 +93,10 @@ impl From for OcrDocumentInput { impl From for OcrDocumentInput { fn from(path: PathBuf) -> Self { - Self::Path { path, mime_type: None } + Self::Path { + path, + mime_type: None, + } } } @@ -324,7 +327,6 @@ impl LiteLLMOcrRequest { config, }) } - } impl LiteLLMOcrRequest { @@ -383,7 +385,6 @@ impl LiteLLMOcrRequest { ..self } } - } impl LiteLLMOcrRequest { diff --git a/litellm-rust/crates/python-bridge/src/lifecycle/mod.rs b/litellm-rust/crates/python-bridge/src/lifecycle/mod.rs index a11162f1c55..3a9e305c9f6 100644 --- a/litellm-rust/crates/python-bridge/src/lifecycle/mod.rs +++ b/litellm-rust/crates/python-bridge/src/lifecycle/mod.rs @@ -25,17 +25,12 @@ use bindings::DeploymentHooks; pub(crate) use bindings::PythonLogger; use handle::{Execution, ExecutionBody, ExecutionStep}; -pub(crate) enum OperationClass { - Phase(HostPhase), - Route, -} - pub(crate) trait PythonRoute: Send + Sync { type Call: NativeCall + 'static; fn state(&self) -> &PythonCallState; fn state_mut(&mut self) -> &mut PythonCallState; - fn classify(operation: &::Operation) -> OperationClass; + fn phase(operation: &::Operation) -> Option; fn lifecycle_result() -> ::Result; fn map_error(error: ::Error) -> PyErr; fn host_error(message: String) -> ::Error; @@ -167,12 +162,8 @@ impl PythonLifecycle { HostFailure::Cancelled(native) }; let state = self.route.state_mut(); - if state.error.is_none() || (cancelled && phase != Some(HostPhase::DeploymentFailure)) { - state.retain_error(py, error); - } - if state.end.is_none() { - state.end = now(py).ok(); - } + state.retain_first_error(py, error, cancelled && phase != Some(HostPhase::DeploymentFailure)); + let _ = state.finish(py); failure } @@ -215,10 +206,7 @@ impl PythonLifecycle { } HostStep::Ready(NativeCallStep::Host(operation)) => operation, }; - let phase = match R::classify(&operation) { - OperationClass::Phase(phase) => Some(phase), - OperationClass::Route => None, - }; + let phase = R::phase(&operation); let result = match phase { Some(phase) => match self.route.state_mut().invoke(py, phase) { Ok(HostStep::Suspend(awaitable)) => { @@ -499,6 +487,19 @@ impl PythonCallState { self.error = Some(error.into_value(py)); } + pub fn retain_first_error(&mut self, py: Python<'_>, error: PyErr, replace: bool) { + if self.error.is_none() || replace { + self.retain_error(py, error); + } + } + + pub fn finish(&mut self, py: Python<'_>) -> PyResult<()> { + if self.end.is_none() { + self.end = Some(now(py)?); + } + Ok(()) + } + pub fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> { visit.call(&self.args)?; visit.call(&self.kwargs)?; @@ -694,8 +695,8 @@ mod tests { &mut self.0 } - fn classify(_: &()) -> OperationClass { - OperationClass::Route + fn phase(_: &()) -> Option { + None } fn lifecycle_result() {} diff --git a/litellm-rust/crates/python-bridge/src/marshal.rs b/litellm-rust/crates/python-bridge/src/marshal.rs index a11763aaebb..ef3ec8392ca 100644 --- a/litellm-rust/crates/python-bridge/src/marshal.rs +++ b/litellm-rust/crates/python-bridge/src/marshal.rs @@ -13,6 +13,11 @@ use litellm_python_interop::from_py_preserving_errors as from_py; use crate::auth::{AZURE_AD_TOKEN_PROVIDER, PythonTokenProvider}; use crate::lifecycle::BoundArguments; +pub(crate) struct Projection { + pub native: Native, + pub retained: Retained, +} + /// Fields every lifecycle route reads from its bound `*args, **kwargs` before /// asking core to build the typed request. Route-specific inputs (for example /// the OCR `document`) are read separately by the route. 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 65896279715..de70f7ce6f0 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/callbacks.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/callbacks.rs @@ -3,142 +3,55 @@ use pyo3::types::PyDict; use serde_json::Value; use litellm_core::ocr::LiteLLMOcrResponse; -use litellm_core::ocr::hooks::OcrPreCallRequest; +use litellm_core::ocr::hooks::OcrDuringCallRequest; use litellm_python_interop::to_py_preserving_errors as to_py; +use super::host::PythonPayload; use crate::lifecycle::PythonLogger; -pub(super) struct OcrLoggingFields { - model: String, - custom_llm_provider: String, - optional_params: Value, -} - -impl From<&OcrPreCallRequest> for OcrLoggingFields { - fn from(request: &OcrPreCallRequest) -> Self { - Self { - model: request.model.clone(), - custom_llm_provider: request.custom_llm_provider.clone(), - optional_params: request.optional_params.clone(), - } - } -} - -impl PythonLogger { - pub(super) fn update_ocr( - &self, - py: Python<'_>, - kwargs: &Py, - pre_call: &OcrLoggingFields, - secret_fields: &[&str], - url: &str, - ) -> PyResult<()> { - let update = PyDict::new(py); - update.set_item("kwargs", redact(py, kwargs.bind(py), secret_fields)?)?; - update.set_item("model", &pre_call.model)?; - update.set_item( - "optional_params", - redact( - py, - &to_py(py, &pre_call.optional_params)? - .into_bound(py) - .cast_into::()?, - secret_fields, - )?, - )?; - let params = PyDict::new(py); - params.set_item( - "litellm_call_id", - kwargs.bind(py).get_item("litellm_call_id")?, - )?; - params.set_item("api_base", url)?; - for name in ["logger_fn", "litellm_request_debug"] { - if let Some(value) = kwargs.bind(py).get_item(name)? { - params.set_item(name, value)?; - } - } - for name in custom_pricing_fields(py)? { - if let Some(value) = kwargs.bind(py).get_item(&name)? - && !value.is_none() - { - params.set_item(name, value)?; - } - } - update.set_item("litellm_params", params)?; - update.set_item("custom_llm_provider", &pre_call.custom_llm_provider)?; - self.object(py) - .call_method("update_from_kwargs", (), Some(&update))?; - Ok(()) - } - - pub(crate) fn pre_ocr( - &self, - py: Python<'_>, - api_key: Option<&str>, - body: &Bound<'_, PyDict>, - headers: &Bound<'_, PyDict>, - url: &str, - ) -> PyResult<()> { - let additional = PyDict::new(py); - additional.set_item("complete_input_dict", body)?; - additional.set_item("headers", headers)?; - additional.set_item("api_base", url)?; - let kwargs = PyDict::new(py); - kwargs.set_item("input", "OCR document processing")?; - kwargs.set_item("api_key", api_key)?; - kwargs.set_item("additional_args", &additional)?; - self.object(py).call_method("pre_call", (), Some(&kwargs))?; - Ok(()) - } - - pub(crate) fn post_ocr( - &self, - py: Python<'_>, - original_response: &Value, - body: Option<&Py>, - headers: Option<&Py>, - ) -> PyResult<()> { - let additional = PyDict::new(py); - additional.set_item("complete_input_dict", body)?; - additional.set_item("headers", headers)?; - let kwargs = PyDict::new(py); - kwargs.set_item("original_response", to_py(py, original_response)?)?; - kwargs.set_item("additional_args", &additional)?; - self.object(py) - .call_method("post_call", (), Some(&kwargs))?; - Ok(()) - } -} - -fn custom_pricing_fields(py: Python<'_>) -> PyResult> { - py.import("litellm.types.utils")? - .getattr("CustomPricingLiteLLMParams")? - .getattr("model_fields")? - .cast_into::()? - .keys() - .iter() - .map(|name| name.extract::()) - .collect() -} - -fn redact( +pub(super) fn update_logging( py: Python<'_>, - params: &Bound<'_, PyDict>, + logger: &PythonLogger, + kwargs: &Py, + request: &OcrDuringCallRequest, secret_fields: &[&str], -) -> PyResult> { - let redacted = PyDict::new(py); - for (name, value) in params { - let name = name.extract::()?; - if name == "proxy_server_request" { - continue; - } - if secret_fields.contains(&name.as_str()) { - redacted.set_item(name, "****")?; - } else { - redacted.set_item(name, value)?; - } - } - Ok(redacted.unbind()) +) -> PyResult<()> { + py.import("litellm.rust_bridge.ocr")? + .getattr("update_logging")? + .call1(( + logger.object(py), + kwargs, + &request.model, + &request.custom_llm_provider, + to_py(py, &request.optional_params)?, + secret_fields, + &request.url, + ))?; + Ok(()) +} + +pub(super) fn pre_call( + py: Python<'_>, + logger: &PythonLogger, + request: &OcrDuringCallRequest, + payload: &PythonPayload, +) -> PyResult<()> { + py.import("litellm.rust_bridge.ocr")? + .getattr("pre_call")? + .call1((logger.object(py), request.api_key.as_deref(), &payload.body, &payload.headers, &request.url))?; + Ok(()) +} + +pub(super) fn post_call( + py: Python<'_>, + logger: &PythonLogger, + original_response: &Value, + payload: &PythonPayload, +) -> PyResult<()> { + py.import("litellm.rust_bridge.ocr")? + .getattr("post_call")? + .call1((logger.object(py), to_py(py, original_response)?, &payload.body, &payload.headers))?; + Ok(()) } pub(super) fn response(py: Python<'_>, response: &LiteLLMOcrResponse) -> PyResult> { 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 6df86e0dc54..8d414435e9d 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/document.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/document.rs @@ -36,11 +36,11 @@ fn extract_bytes(value: &Bound<'_, PyAny>) -> PyResult { return Ok(Bytes::from_owner(value.extract::()?)); } // Bytes subclasses may retain GC edges that a native Bytes owner cannot traverse. - Ok(Bytes::copy_from_slice(value.extract::()?.as_ref())) + Ok(Bytes::copy_from_slice( + value.extract::()?.as_ref(), + )) } - - pub(super) struct FileDocumentInput { pub input: OcrDocumentInput, pub reader: Option, @@ -56,9 +56,11 @@ 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 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) { missing() @@ -93,7 +95,9 @@ impl FromPyObject<'_, '_> for FileDocumentInput { reader: None, }); } - let reader = file.getattr_opt("read")?.filter(|value| value.is_callable()); + let reader = file + .getattr_opt("read")? + .filter(|value| value.is_callable()); let Some(reader) = reader else { return Err(PyValueError::new_err(format!( "Unsupported file input type: {}. Expected pathlib.Path, bytes, or a file-like object.", @@ -107,7 +111,10 @@ impl FromPyObject<'_, '_> for FileDocumentInput { .transpose()?; Ok(Self { input: OcrDocumentInput::HostReader { mime_type }, - reader: Some(PythonFileReader { reader: reader.unbind(), name }), + reader: Some(PythonFileReader { + reader: reader.unbind(), + name, + }), }) } } @@ -127,12 +134,21 @@ mod tests { let error = document.extract::().err().unwrap(); assert!(error.is_instance_of::(py)); } - for expression in [c"{'file': b'abc', 'mime_type': None}", c"{'file': b'abc', 'mime_type': 7}"] { - let error = py.eval(expression, None, None).unwrap().extract::().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 locals = PyDict::new(py); - py.run(c"from pathlib import Path + py.run( + c"from pathlib import Path failure = KeyError('reader failed') class Reader: def __init__(self): @@ -143,12 +159,31 @@ class Reader: reader = Reader() document = {'file': reader} path_document = {'file': Path('/nonexistent/ocr-projection-test.pdf')} -", Some(&locals), Some(&locals)).unwrap(); +", + Some(&locals), + Some(&locals), + ) + .unwrap(); let document = locals.get_item("document").unwrap().unwrap(); let input: FileDocumentInput = document.extract().unwrap(); - assert_eq!(locals.get_item("reader").unwrap().unwrap().getattr("reads").unwrap().extract::().unwrap(), 0); + 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(); + let path: FileDocumentInput = locals + .get_item("path_document") + .unwrap() + .unwrap() + .extract() + .unwrap(); assert!(matches!(path.input, OcrDocumentInput::Path { .. })); }); } 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 655cd9ea1ee..c7f239054de 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/errors.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/errors.rs @@ -103,9 +103,13 @@ pub(super) fn public_exception( std::io::ErrorKind::PermissionDenied => 13, _ => 5, }); - return Ok(PyErr::from_value(py.import("builtins")?.getattr("OSError")?.call1(( - errno, source.to_string(), path, - ))?)); + return Ok(PyErr::from_value( + py.import("builtins")?.getattr("OSError")?.call1(( + errno, + source.to_string(), + path.into_pyobject(py)?.call_method0("__fspath__")?, + ))?, + )); } raise_public(py, classify(error), model, provider, None) } diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/host.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/host.rs index 91645aa04c1..5ccb65780df 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/host.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/host.rs @@ -1,12 +1,11 @@ -use std::sync::Arc; - use pyo3::gc::{PyTraverseError, PyVisit}; use pyo3::prelude::*; use pyo3::types::PyDict; use litellm_auth::ResolvedCredential; +use litellm_core::call_lifecycle::host::HostPhase; use litellm_core::ocr::hooks::{OcrDuringCallRequest, OcrPostCallRequest}; -use litellm_core::ocr::{OcrCall, OcrHostOperation, OcrHostResult}; +use litellm_core::ocr::{OcrCall, OcrHostOperation, OcrHostResult, OcrProjectedRequest}; use litellm_python_interop::{ from_py_preserving_errors as from_py, to_py_preserving_errors as to_py, }; @@ -14,32 +13,50 @@ use litellm_python_interop::{ use super::document::PythonFileReader; use super::errors::to_pyerr as ocr_error_to_pyerr; use super::project::project; -use super::{ASYNC_SIGNATURE, SIGNATURE, callbacks, errors}; +use super::{callbacks, errors}; use crate::auth::PythonTokenProvider; -use crate::lifecycle::{OperationClass, PythonCallState, PythonRoute, missing_state, now}; +use crate::lifecycle::{PythonCallState, PythonRoute, Signature, missing_state}; +use crate::marshal::Projection; pub(super) struct PythonOcrHost { state: PythonCallState, - projected: Option, + signature: &'static Signature, + retained: Option, } -struct ProjectedOcrHost { - model: String, - provider: &'static str, - secret_fields: Vec<&'static str>, - azure_ad_token_provider: Option, - pre_call: Option, - payload: Option, - reader: Option, - reader_failed: bool, +pub(super) struct OcrRetained { + pub model: String, + pub provider: &'static str, + pub secret_fields: Vec<&'static str>, + pub azure_ad_token_provider: Option, + pub payload: Option, + pub reader: Option, } -struct CapturedOcrPayload { - body: Py, - headers: Py, +pub(super) struct PythonPayload { + pub body: Py, + pub headers: Py, } -impl CapturedOcrPayload { +impl PythonPayload { + fn from_request(py: Python<'_>, request: &OcrDuringCallRequest) -> PyResult { + let body = to_py(py, &request.body)?.into_bound(py).cast_into::()?; + let headers = PyDict::new(py); + for (name, value) in &request.headers { + headers.set_item(name, value)?; + } + Ok(Self { body: body.unbind(), headers: headers.unbind() }) + } + + fn write_back(&self, py: Python<'_>, mut request: OcrDuringCallRequest) -> PyResult { + request.body = from_py(self.body.bind(py))?; + request.headers = self.headers.bind(py) + .iter() + .map(|(name, value)| Ok((name.extract::()?, value.extract::()?))) + .collect::>>()?; + Ok(request) + } + fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> { visit.call(&self.body)?; visit.call(&self.headers) @@ -47,52 +64,36 @@ impl CapturedOcrPayload { } impl PythonOcrHost { - pub(super) fn new(state: PythonCallState) -> Self { + pub(super) fn new(state: PythonCallState, signature: &'static Signature) -> Self { Self { state, - projected: None, + signature, + retained: None, } } - fn projected(&self) -> PyResult<&ProjectedOcrHost> { - self.projected.as_ref().ok_or_else(missing_state) + fn retained(&self) -> PyResult<&OcrRetained> { + self.retained.as_ref().ok_or_else(missing_state) } - fn projected_mut(&mut self) -> PyResult<&mut ProjectedOcrHost> { - self.projected.as_mut().ok_or_else(missing_state) + fn retained_mut(&mut self) -> PyResult<&mut OcrRetained> { + self.retained.as_mut().ok_or_else(missing_state) } fn project(&mut self, py: Python<'_>) -> PyResult { - let signature = if self.state.asynchronous { - &ASYNC_SIGNATURE - } else { - &SIGNATURE - }; - let arguments = signature.bind(self.state.args.bind(py), self.state.kwargs.bind(py))?; - let projected = project(py, &arguments)?; - let has_token_provider = projected.azure_ad_token_provider.is_some(); - self.projected = Some(ProjectedOcrHost { - model: projected.request.model.clone(), - provider: projected.request.provider_name(), - secret_fields: projected.secret_fields, - azure_ad_token_provider: projected.azure_ad_token_provider, - pre_call: None, - payload: None, - reader: projected.reader, - reader_failed: false, - }); - Ok(OcrHostResult::Request(Ok(( - Box::new( - projected - .request - .with_host_hooks(Arc::new(BridgeOcrHooks), None), - ), - has_token_provider, - )))) + let arguments = self.signature.bind(self.state.args.bind(py), self.state.kwargs.bind(py))?; + let Projection { native, retained } = project(py, &arguments)?; + let host_token_provider = retained.azure_ad_token_provider.is_some(); + self.retained = Some(retained); + Ok(OcrHostResult::Request(Ok(OcrProjectedRequest { + request: Box::new(native), + intercepts_requests: true, + host_token_provider, + }))) } fn acquire_azure_ad_token(&self, py: Python<'_>) -> PyResult { - self.projected()? + self.retained()? .azure_ad_token_provider .as_ref() .ok_or_else(missing_state)? @@ -102,42 +103,21 @@ impl PythonOcrHost { fn during_call( &mut self, py: Python<'_>, - mut request: OcrDuringCallRequest, + request: OcrDuringCallRequest, ) -> PyResult { - let projected = self.projected()?; - let pre_call = projected.pre_call.as_ref().ok_or_else(missing_state)?; + let retained = self.retained()?; let logger = self.state.logger()?; - logger.update_ocr( + callbacks::update_logging( py, + logger, &self.state.kwargs, - pre_call, - &projected.secret_fields, - &request.url, + &request, + &retained.secret_fields, )?; - let body = to_py(py, &request.body)? - .into_bound(py) - .cast_into::()?; - let headers = PyDict::new(py); - for (name, value) in &request.headers { - headers.set_item(name, value)?; - } - logger.pre_ocr( - py, - request.api_key.as_deref(), - &body, - &headers, - &request.url, - )?; - request.body = from_py(&body)?; - request.headers = headers - .iter() - .map(|(name, value)| Ok((name.extract::()?, value.extract::()?))) - .collect::>>()?; - let projected = self.projected_mut()?; - projected.payload = Some(CapturedOcrPayload { - body: body.unbind(), - headers: headers.unbind(), - }); + let payload = PythonPayload::from_request(py, &request)?; + callbacks::pre_call(py, logger, &request, &payload)?; + let request = payload.write_back(py, request)?; + self.retained_mut()?.payload = Some(payload); Ok(request) } @@ -146,27 +126,24 @@ impl PythonOcrHost { py: Python<'_>, request: OcrPostCallRequest, ) -> PyResult { - let projected = self.projected()?; - let payload = projected.payload.as_ref(); - self.state.logger()?.post_ocr( + let payload = self.retained()?.payload.as_ref().ok_or_else(missing_state)?; + callbacks::post_call( py, + self.state.logger()?, &request.original_response, - payload.map(|payload| &payload.body), - payload.map(|payload| &payload.headers), + payload, )?; Ok(request) } fn map_failure(&mut self, py: Python<'_>, error: litellm_core::ocr::Error) -> PyResult<()> { - if self.state.end.is_none() { - self.state.end = Some(now(py)?); - } - let (model, provider) = match &self.projected { - Some(projected) => (projected.model.as_str(), projected.provider), + self.state.finish(py)?; + let (model, provider) = match &self.retained { + Some(retained) => (retained.model.as_str(), retained.provider), None => ("", ""), }; let mapped = match self.state.error.take() { - Some(host_error) if self.projected.as_ref().is_some_and(|host| host.reader_failed) => { + Some(host_error) if matches!(error, litellm_core::ocr::Error::HostDocumentRead) => { PyErr::from_value(host_error.into_bound(py).into_any()) } Some(host_error) => errors::public_host_exception(py, &host_error, model, provider)?, @@ -188,10 +165,8 @@ impl PythonRoute for PythonOcrHost { &mut self.state } - fn classify(operation: &OcrHostOperation) -> OperationClass { - operation - .phase() - .map_or(OperationClass::Route, OperationClass::Phase) + fn phase(operation: &OcrHostOperation) -> Option { + operation.phase() } fn lifecycle_result() -> OcrHostResult { @@ -210,16 +185,24 @@ impl PythonRoute for PythonOcrHost { Ok(match operation { OcrHostOperation::ProjectRequest => self.project(py)?, OcrHostOperation::ReadDocument => { - let reader = self.projected_mut()?.reader.take().ok_or_else(missing_state)?; - let result = reader.read(py); - self.projected_mut()?.reader_failed = result.is_err(); - OcrHostResult::Document(Ok(result?)) + let reader = self + .retained_mut()? + .reader + .take() + .ok_or_else(missing_state)?; + match reader.read(py) { + Ok(content) => OcrHostResult::Document(Ok(content)), + Err(error) if error.is_instance_of::(py) => { + self.state.retain_first_error(py, error, false); + OcrHostResult::Document(Err(litellm_core::ocr::Error::HostDocumentRead)) + } + Err(error) => return Err(error), + } } OcrHostOperation::AcquireAzureAdToken => { OcrHostResult::AzureAdToken(Ok(self.acquire_azure_ad_token(py)?)) } OcrHostOperation::PreCall(request) => { - self.projected_mut()?.pre_call = Some((&request).into()); OcrHostResult::PreCall(Ok(request)) } OcrHostOperation::DuringCall(request) => { @@ -229,7 +212,7 @@ impl PythonRoute for PythonOcrHost { OcrHostResult::PostCall(Ok(self.post_call(py, request)?)) } OcrHostOperation::ConstructResponse(response) => { - self.state.end = Some(now(py)?); + self.state.finish(py)?; self.state.response = Some(callbacks::response(py, response.as_ref())?); OcrHostResult::Lifecycle(Ok(())) } @@ -244,30 +227,22 @@ impl PythonRoute for PythonOcrHost { } fn cleanup(&mut self) { - self.projected = None; + self.retained = None; } fn traverse(&self, visit: &PyVisit<'_>) -> Result<(), PyTraverseError> { - let Some(projected) = &self.projected else { + let Some(retained) = &self.retained else { return Ok(()); }; - if let Some(provider) = &projected.azure_ad_token_provider { + if let Some(provider) = &retained.azure_ad_token_provider { provider.traverse(visit)?; } - if let Some(reader) = &projected.reader { + if let Some(reader) = &retained.reader { reader.traverse(visit)?; } - if let Some(payload) = &projected.payload { + if let Some(payload) = &retained.payload { payload.traverse(visit)?; } Ok(()) } } - -struct BridgeOcrHooks; - -impl litellm_core::ocr::hooks::OcrHooks for BridgeOcrHooks { - fn intercepts_requests(&self) -> bool { - true - } -} 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 7973dc45c9b..cbd32792843 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs @@ -68,7 +68,7 @@ fn call( kwargs.copy()?.unbind(), asynchronous, signature.name, - )?); + )?, signature); run_call(py, call, host) } 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 a78aa5b80b7..57cddac48cb 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/project.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/project.rs @@ -10,54 +10,34 @@ use litellm_core::ocr::{ }; use litellm_python_interop::from_py_preserving_errors as from_py; +use super::document::FileDocumentInput; use super::errors::to_pyerr as ocr_error_to_pyerr; -use super::document::{FileDocumentInput, PythonFileReader}; -use crate::auth::PythonTokenProvider; +use super::host::OcrRetained; use crate::lifecycle::BoundArguments; -use crate::marshal::BoundRouteInputs; +use crate::marshal::{BoundRouteInputs, Projection}; /// Positional parameters of `ocr()` that are never projected into /// `optional_params`. const BOUND_FIELDS: &[&str] = &["model", "document", "timeout", "input_sources"]; -pub(super) struct ProjectedOcrCall { - pub request: LiteLLMOcrRequest, - pub azure_ad_token_provider: Option, - pub secret_fields: Vec<&'static str>, - pub reader: Option, -} - -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 { +fn project_document(document: &Bound<'_, PyAny>) -> PyResult> { let kind: String = document.get_item("type")?.extract()?; if kind != "file" { - return Ok(ProjectedDocument::Url(from_py(document)?)); + let value: serde_json::Value = from_py(document)?; + return Ok(OcrDocument::try_from(value).map(|document| FileDocumentInput { + input: document.into(), + reader: None, + })); } - document.extract().map(ProjectedDocument::File) + document.extract().map(Ok) } /// Pure core assembly; every failure here is a typed `ocr::Error`. fn build_request( inputs: BoundRouteInputs, - document: ProjectedDocument, -) -> Result { - let document = document.into_native()?; + document: Result, +) -> Result, litellm_core::ocr::Error> { + let document = document?; let BoundRouteInputs { model, custom_llm_provider, @@ -83,18 +63,23 @@ fn build_request( input_sources, }, )?; - Ok(ProjectedOcrCall { - request, - azure_ad_token_provider, - secret_fields, - reader: document.reader, + Ok(Projection { + retained: OcrRetained { + model: request.model.clone(), + provider: request.provider_name(), + azure_ad_token_provider, + secret_fields, + reader: document.reader, + payload: None, + }, + native: request, }) } pub(super) fn project( py: Python<'_>, arguments: &BoundArguments<'_>, -) -> PyResult { +) -> PyResult> { let model: String = arguments.extract("model")?; let custom_llm_provider: Option = arguments.optional("custom_llm_provider")?; let document = project_document(&arguments.required("document")?)?; @@ -142,7 +127,7 @@ mod tests { ) .unwrap(); assert!(matches!( - project_document(&file).unwrap().into_native().unwrap().input, + project_document(&file).unwrap().unwrap().input, litellm_core::ocr::OcrDocumentInput::Bytes { bytes, mime_type, .. } if bytes == b"%PDF-1.4"[..] && mime_type.as_deref() == Some("application/pdf") )); @@ -155,7 +140,7 @@ mod tests { ) .unwrap(); assert!(matches!( - project_document(&original).unwrap().into_native().unwrap().input, + project_document(&original).unwrap().unwrap().input, litellm_core::ocr::OcrDocumentInput::Document(OcrDocument::DocumentUrl { document_url, .. }) if document_url == "https://example.com/a.pdf" )); @@ -169,7 +154,10 @@ mod tests { let document = py .eval(c"{'type': 'mystery', 'mystery': 'x'}", None, None) .unwrap(); - let error = project_document(&document).unwrap().into_native().err().unwrap(); + let error = project_document(&document) + .unwrap() + .err() + .unwrap(); assert!(error.to_string().contains("document")); }); } @@ -181,14 +169,16 @@ mod tests { let missing = py.eval(c"{}", None, None).unwrap(); assert!( project_document(&missing) - .err().unwrap() + .err() + .unwrap() .is_instance_of::(py) ); let non_string = py.eval(c"{'type': 1}", None, None).unwrap(); assert!( project_document(&non_string) - .err().unwrap() + .err() + .unwrap() .is_instance_of::(py) ); @@ -202,8 +192,9 @@ class Document: document = Document() ", ); - let error = - project_document(&locals.get_item("document").unwrap().unwrap()).err().unwrap(); + let error = project_document(&locals.get_item("document").unwrap().unwrap()) + .err() + .unwrap(); assert!( error .value(py) @@ -232,8 +223,11 @@ document = Document() ", ); let document = locals.get_item("document").unwrap().unwrap(); - let projected = project_document(&document).unwrap().into_native().unwrap(); - assert!(matches!(projected.input, litellm_core::ocr::OcrDocumentInput::Bytes { .. })); + let projected = project_document(&document).unwrap().unwrap(); + assert!(matches!( + projected.input, + litellm_core::ocr::OcrDocumentInput::Bytes { .. } + )); let reads: Vec = document.getattr("reads").unwrap().extract().unwrap(); assert_eq!(reads, ["type", "mime_type", "file"]); }); diff --git a/litellm/ocr/main.py b/litellm/ocr/main.py index c6fca6078e5..a4db8c5b0e5 100644 --- a/litellm/ocr/main.py +++ b/litellm/ocr/main.py @@ -4,7 +4,7 @@ from typing import Final, cast # noqa: TID251 # native binding selects a sync from litellm.llms.base_llm.ocr.transformation import OCRResponse from litellm.ocr import legacy from litellm.ocr.legacy import convert_file_document_to_url_document, get_mime_type -from litellm.rust_bridge.bindings import native_exception_types +from litellm.rust_bridge.bindings import native_decline_types from litellm.rust_bridge.configuration import rust_ocr_enabled from litellm.rust_bridge.ocr import NATIVE_AOCR, NATIVE_OCR @@ -19,7 +19,7 @@ def ocr( if native is not None: try: return native(*args, **kwargs) - except _decline_types(): + except native_decline_types(): pass fallback: Final = cast( # cast-ok: forward the original call shape through the legacy @client decorator Callable[..., OCRResponse | Coroutine[object, object, OCRResponse]], legacy.ocr @@ -32,14 +32,10 @@ async def aocr(*args: object, **kwargs: object) -> OCRResponse: # kwargs-ok: pr if native is not None: try: return await native(*args, **kwargs) - except _decline_types(): + except native_decline_types(): pass fallback: Final = cast( # cast-ok: forward the original call shape through the legacy @client decorator Callable[..., Awaitable[OCRResponse]], legacy.aocr ) return await fallback(*args, **kwargs) - -def _decline_types() -> tuple[type[BaseException], ...]: - exception_types: Final = native_exception_types() - return (exception_types[0],) if exception_types is not None else () diff --git a/litellm/proxy/ocr_endpoints/endpoints.py b/litellm/proxy/ocr_endpoints/endpoints.py index 82e0fcf5081..dc7edba28c5 100644 --- a/litellm/proxy/ocr_endpoints/endpoints.py +++ b/litellm/proxy/ocr_endpoints/endpoints.py @@ -29,10 +29,12 @@ def _build_document_from_upload( filename: str | None, content_type: str | None, ) -> dict[str, str]: - mime_type: Final = content_type.split(";")[0].strip() if content_type else None - if not mime_type or mime_type == "application/octet-stream": - if filename: - mime_type = get_mime_type(filename) + supplied_mime: Final = content_type.split(";")[0].strip() if content_type else None + mime_type: Final = ( + get_mime_type(filename) + if filename and (not supplied_mime or supplied_mime == "application/octet-stream") + else supplied_mime + ) return convert_file_document_to_url_document( { diff --git a/litellm/rust_bridge/bindings.py b/litellm/rust_bridge/bindings.py index d16f150a2aa..32f6e9ec46b 100644 --- a/litellm/rust_bridge/bindings.py +++ b/litellm/rust_bridge/bindings.py @@ -47,3 +47,8 @@ def native_exception_types() -> tuple[type[BaseException], type[BaseException]] if not isinstance(declined, type) or not isinstance(upstream, type): return None return declined, upstream + + +def native_decline_types() -> tuple[type[BaseException], ...]: + exceptions: Final = native_exception_types() + return (exceptions[0],) if exceptions is not None else () diff --git a/litellm/rust_bridge/ocr.py b/litellm/rust_bridge/ocr.py index 7ede521e6f3..fc3c5523856 100644 --- a/litellm/rust_bridge/ocr.py +++ b/litellm/rust_bridge/ocr.py @@ -1,6 +1,6 @@ from __future__ import annotations -from collections.abc import Callable, Coroutine, Mapping +from collections.abc import Callable, Coroutine, Mapping, Sequence from types import MappingProxyType from typing import Final, Protocol, cast # noqa: TID251 # validates dynamically loaded native callables @@ -8,6 +8,84 @@ from pydantic import TypeAdapter from litellm.llms.base_llm.ocr.transformation import PROVIDER_NATIVE_RESPONSE_KEY, OCRResponse from litellm.rust_bridge.bindings import NativeBinding +from litellm.types.utils import CustomPricingLiteLLMParams + + +class OcrLoggingProtocol(Protocol): + def update_from_kwargs( + self, + *, + kwargs: dict[str, object], + model: str, + optional_params: dict[str, object], + litellm_params: dict[str, object], + custom_llm_provider: str, + ) -> object: ... + + def pre_call(self, *, input: str, api_key: str | None, additional_args: dict[str, object]) -> object: ... + + def post_call(self, *, original_response: object, additional_args: dict[str, object]) -> object: ... + + +def _redact(params: Mapping[str, object], secret_fields: Sequence[str]) -> dict[str, object]: + return { + name: "****" if name in secret_fields else value + for name, value in params.items() + if name != "proxy_server_request" + } + + +def update_logging( + logger: OcrLoggingProtocol, + kwargs: Mapping[str, object], + model: str, + custom_llm_provider: str, + optional_params: Mapping[str, object], + secret_fields: Sequence[str], + url: str, +) -> None: + logger.update_from_kwargs( + kwargs=_redact(kwargs, secret_fields), + model=model, + optional_params=_redact(optional_params, secret_fields), + litellm_params={ + "litellm_call_id": kwargs.get("litellm_call_id"), + "api_base": url, + **{name: kwargs[name] for name in ("logger_fn", "litellm_request_debug") if name in kwargs}, + **{ + name: kwargs[name] + for name in CustomPricingLiteLLMParams.model_fields + if name in kwargs and kwargs[name] is not None + }, + }, + custom_llm_provider=custom_llm_provider, + ) + + +def pre_call( + logger: OcrLoggingProtocol, + api_key: str | None, + body: dict[str, object], + headers: dict[str, str], + url: str, +) -> None: + logger.pre_call( + input="OCR document processing", + api_key=api_key, + additional_args={"complete_input_dict": body, "headers": headers, "api_base": url}, + ) + + +def post_call( + logger: OcrLoggingProtocol, + original_response: object, + body: dict[str, object], + headers: dict[str, str], +) -> None: + logger.post_call( + original_response=original_response, + additional_args={"complete_input_dict": body, "headers": headers}, + ) class RustOcr(Protocol): diff --git a/tests/test_litellm/rust_bridge/test_bindings.py b/tests/test_litellm/rust_bridge/test_bindings.py index 88036a5a556..7e6fea7adfb 100644 --- a/tests/test_litellm/rust_bridge/test_bindings.py +++ b/tests/test_litellm/rust_bridge/test_bindings.py @@ -33,3 +33,21 @@ def test_binding_validates_native_attribute( binding: Final = bindings.NativeBinding("route", validate=lambda item: item if isinstance(item, int) else None) assert binding.load() == expected + + +@pytest.mark.parametrize("available", [False, True]) +def test_decline_accessor_only_catches_admission_declines(monkeypatch: pytest.MonkeyPatch, available: bool) -> None: + class Declined(Exception): + pass + + class Upstream(Exception): + pass + + native: Final = SimpleNamespace(RustBridgeDeclined=Declined, RustUpstreamError=Upstream) if available else None + monkeypatch.setattr(bindings, "get_native_bridge", lambda: native) + assert bindings.native_decline_types() == ((Declined,) if available else ()) + with pytest.raises(Upstream): + try: + raise Upstream("already dispatched") + except bindings.native_decline_types(): + pytest.fail("upstream failure was allowed to replay") diff --git a/tests/test_litellm/rust_bridge/test_ocr_lifecycle.py b/tests/test_litellm/rust_bridge/test_ocr_lifecycle.py index baf5295d571..96983968bee 100644 --- a/tests/test_litellm/rust_bridge/test_ocr_lifecycle.py +++ b/tests/test_litellm/rust_bridge/test_ocr_lifecycle.py @@ -8,7 +8,7 @@ import litellm from litellm.llms.base_llm.ocr.transformation import OCRResponse from litellm.ocr import legacy from litellm.rust_bridge import bindings, configuration -from litellm.rust_bridge.ocr import NATIVE_AOCR, NATIVE_OCR +from litellm.rust_bridge.ocr import NATIVE_AOCR, NATIVE_OCR, post_call, pre_call, update_logging @pytest.fixture(autouse=True) @@ -21,6 +21,72 @@ def isolated_ocr_configuration(monkeypatch: pytest.MonkeyPatch) -> Generator[Non configuration.reset_rust_configuration() +def test_logging_redacts_views_and_preserves_opaque_arguments_and_pricing() -> None: + logger: Final = Mock() + opaque: Final = object() + logger_fn: Final = object() + kwargs: Final = { + "vertex_credentials": "secret", + "proxy_server_request": opaque, + "metadata": opaque, + "logger_fn": logger_fn, + "litellm_request_debug": False, + "litellm_call_id": "call-id", + "input_cost_per_token": 0, + "output_cost_per_token": None, + } + optional: Final = {"vertex_credentials": "secret", "pages": [1], "proxy_server_request": opaque} + update_logging(logger, kwargs, "model", "vertex_ai", optional, ("vertex_credentials",), "https://provider") + logger.update_from_kwargs.assert_called_once_with( + kwargs={ + "vertex_credentials": "****", + "metadata": opaque, + "logger_fn": logger_fn, + "litellm_request_debug": False, + "litellm_call_id": "call-id", + "input_cost_per_token": 0, + "output_cost_per_token": None, + }, + model="model", + custom_llm_provider="vertex_ai", + optional_params={"vertex_credentials": "****", "pages": [1]}, + litellm_params={ + "litellm_call_id": "call-id", + "api_base": "https://provider", + "logger_fn": logger_fn, + "litellm_request_debug": False, + "input_cost_per_token": 0, + }, + ) + assert kwargs["vertex_credentials"] == optional["vertex_credentials"] == "secret" + assert kwargs["proxy_server_request"] is optional["proxy_server_request"] is opaque + assert logger.update_from_kwargs.call_args.kwargs["kwargs"]["metadata"] is opaque + assert logger.update_from_kwargs.call_args.kwargs["optional_params"]["pages"] is optional["pages"] + + +def test_logging_callbacks_receive_captured_payload_roots_and_propagate_errors() -> None: + logger: Final = Mock() + body: Final[dict[str, object]] = {"document": "original"} + headers: Final = {"authorization": "key"} + response: Final = object() + pre_call(logger, "key", body, headers, "https://provider") + post_call(logger, response, body, headers) + logger.pre_call.assert_called_once_with( + input="OCR document processing", + api_key="key", + additional_args={"complete_input_dict": body, "headers": headers, "api_base": "https://provider"}, + ) + for callback in (logger.pre_call, logger.post_call): + assert callback.call_args.kwargs["additional_args"]["complete_input_dict"] is body + assert callback.call_args.kwargs["additional_args"]["headers"] is headers + assert logger.post_call.call_args.kwargs["original_response"] is response + failure: Final = RuntimeError("callback failed") + failing_logger: Final = Mock(pre_call=Mock(side_effect=failure)) + with pytest.raises(RuntimeError) as caught: + pre_call(failing_logger, None, body, headers, "https://provider") + assert caught.value is failure + + @pytest.mark.asyncio @pytest.mark.parametrize("asynchronous", [False, True]) async def test_unavailable_native_uses_legacy(monkeypatch: pytest.MonkeyPatch, asynchronous: bool) -> None: diff --git a/tests/test_litellm_rust/test_ocr.py b/tests/test_litellm_rust/test_ocr.py index a83748c0dd6..86913f388db 100644 --- a/tests/test_litellm_rust/test_ocr.py +++ b/tests/test_litellm_rust/test_ocr.py @@ -1,14 +1,15 @@ -import json import asyncio import base64 +import gc +import json +import os import threading import weakref -import gc -from pathlib import Path -from collections.abc import Generator +from collections.abc import Coroutine, Generator from datetime import datetime, timezone from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer from io import BytesIO +from pathlib import Path from typing import Final, Protocol import httpx @@ -340,7 +341,12 @@ def test_native_lifecycle_core_encodes_python_file_input( @pytest.mark.parametrize("asynchronous", [False, True]) @pytest.mark.parametrize("source", ["path", "pathlike", "reader", "text", "bytes"]) @pytest.mark.asyncio -async def test_native_file_sources_preserve_sdk_behavior(ocr_server, tmp_path, asynchronous, source): +async def test_native_file_sources_preserve_sdk_behavior( + ocr_server: tuple[ThreadingHTTPServer, list[dict[str, object]]], + tmp_path: Path, + asynchronous: bool, + source: str, +) -> None: server, requests = ocr_server path: Final = tmp_path / "scan.png" content: Final = b"document bytes" @@ -388,9 +394,12 @@ async def test_native_file_sources_preserve_sdk_behavior(ocr_server, tmp_path, a "api_base": f"http://127.0.0.1:{server.server_port}", } response: Final = await litellm.aocr(**kwargs) if asynchronous else litellm.ocr(**kwargs) + assert isinstance(response, OCRResponse) assert response.pages[0].markdown == "native OCR response" assert len(requests) == 1 - assert requests[0]["body"]["document"] == { + body: Final = requests[0]["body"] + assert isinstance(body, dict) + assert body["document"] == { "type": "image_url", "image_url": "data:image/png;base64," + base64.b64encode(content).decode(), } @@ -400,7 +409,11 @@ async def test_native_file_sources_preserve_sdk_behavior(ocr_server, tmp_path, a @pytest.mark.parametrize("asynchronous", [False, True]) @pytest.mark.asyncio -async def test_native_file_failures_preserve_identity_and_do_not_send(ocr_server, tmp_path, asynchronous): +async def test_native_file_failures_preserve_identity_and_do_not_send( + ocr_server: tuple[ThreadingHTTPServer, list[dict[str, object]]], + tmp_path: Path, + asynchronous: bool, +) -> None: from litellm.rust_bridge import _native server, requests = ocr_server @@ -416,37 +429,47 @@ async def test_native_file_failures_preserve_identity_and_do_not_send(ocr_server reader: Final = Reader() path: Final = tmp_path / "missing.pdf" - kwargs: Final = { - "model": "mistral/mistral-ocr-latest", - "api_key": "test-key", - "api_base": f"http://127.0.0.1:{server.server_port}", - } + + async def invoke(source: Reader | Path) -> None: + if asynchronous: + await _native.aocr( + model="mistral/mistral-ocr-latest", + document={"type": "file", "file": source}, + api_key="test-key", + api_base=f"http://127.0.0.1:{server.server_port}", + ) + else: + _native.ocr( + model="mistral/mistral-ocr-latest", + document={"type": "file", "file": source}, + api_key="test-key", + api_base=f"http://127.0.0.1:{server.server_port}", + ) + for source, error in ((reader, KeyError), (path, FileNotFoundError)): with pytest.raises(error) as caught: - if asynchronous: - await _native.aocr(document={"type": "file", "file": source}, **kwargs) - else: - _native.ocr(document={"type": "file", "file": source}, **kwargs) + await invoke(source) if source is reader: assert caught.value is failure else: - assert caught.value.filename == str(path) + assert isinstance(caught.value, FileNotFoundError) + assert getattr(caught.value, "filename") == str(path) assert reader.calls == 1 assert requests == [] -def test_unstarted_native_file_call_does_not_read_and_releases_reader(): +def test_unstarted_native_file_call_does_not_read_and_releases_reader() -> None: from litellm.rust_bridge import _native class Reader: + pending: Coroutine[object, object, OCRResponse] | None = None + def read(self) -> bytes: raise AssertionError("unstarted call consumed its document") - def create_call(): + def create_call() -> tuple[Coroutine[object, object, OCRResponse], weakref.ReferenceType[Reader]]: reader: Final = Reader() - pending: Final = _native.aocr( - model="mistral/mistral-ocr-latest", document={"type": "file", "file": reader} - ) + pending: Final = _native.aocr(model="mistral/mistral-ocr-latest", document={"type": "file", "file": reader}) reader.pending = pending return pending, weakref.ref(reader) @@ -457,6 +480,49 @@ def test_unstarted_native_file_call_does_not_read_and_releases_reader(): assert reference() is None +@pytest.mark.skipif(not hasattr(os, "mkfifo"), reason="requires Unix named pipes") +@pytest.mark.asyncio +async def test_native_path_cancellation_waits_for_read_completion( + ocr_server: tuple[ThreadingHTTPServer, list[dict[str, object]]], tmp_path: Path +) -> None: + from litellm.rust_bridge import _native + + server, requests = ocr_server + path: Final = tmp_path / "document.pdf" + os.mkfifo(path) + entered: Final = threading.Event() + release: Final = threading.Event() + + def supply_document() -> None: + with path.open("wb") as stream: + stream.write(b"document") + stream.flush() + entered.set() + release.wait(5) + + writer: Final = threading.Thread(target=supply_document, daemon=True) + writer.start() + task: Final = asyncio.create_task( + _native.aocr( + model="mistral/mistral-ocr-latest", + document={"type": "file", "file": path}, + api_key="test-key", + api_base=f"http://127.0.0.1:{server.server_port}", + ) + ) + try: + assert await asyncio.to_thread(entered.wait, 3) + task.cancel() + await asyncio.sleep(0.05) + assert not task.done() + finally: + release.set() + await asyncio.to_thread(writer.join, 3) + with pytest.raises(asyncio.CancelledError): + await task + assert requests == [] + + @pytest.mark.parametrize("asynchronous", [False, True]) @pytest.mark.parametrize("model", ["mistral/mistral-ocr-latest", "azure_ai/doc-intelligence/prebuilt-read"]) @pytest.mark.asyncio