From 6f8f07433dda0a4b760ec80f55fe4e461dd8b269 Mon Sep 17 00:00:00 2001 From: Yujong Lee Date: Mon, 7 Sep 2026 23:07:25 -0700 Subject: [PATCH] refactor(ocr): inject upstream transport --- litellm-rust/crates/core/src/ocr/mod.rs | 102 +++++++++++++++++++--- litellm-rust/crates/core/src/ocr/types.rs | 2 +- 2 files changed, 89 insertions(+), 15 deletions(-) diff --git a/litellm-rust/crates/core/src/ocr/mod.rs b/litellm-rust/crates/core/src/ocr/mod.rs index a7e89658b00..05f2f419547 100644 --- a/litellm-rust/crates/core/src/ocr/mod.rs +++ b/litellm-rust/crates/core/src/ocr/mod.rs @@ -4,6 +4,7 @@ pub mod types; use serde_json::Value; use std::future::Future; +use std::pin::Pin; use crate::Error; use crate::error::json_type_name; @@ -20,11 +21,10 @@ use crate::lifecycle::{ }; pub trait OcrTransport { - type SendFuture<'a>: Future> - where - Self: 'a; - - fn send(&self, request: OcrTransportRequest) -> Self::SendFuture<'_>; + fn send( + &self, + request: OcrTransportRequest, + ) -> Pin> + Send + '_>>; } pub trait OcrServices: TerminalDispatcher + Clock + OcrTransport {} @@ -55,10 +55,11 @@ impl TerminalDispatcher for DefaultOcrServices { } impl OcrTransport for DefaultOcrServices { - type SendFuture<'a> = impl Future> + 'a; - - fn send(&self, request: OcrTransportRequest) -> Self::SendFuture<'_> { - async move { + fn send( + &self, + request: OcrTransportRequest, + ) -> Pin> + Send + '_>> { + Box::pin(async move { let response = buffered_post::send(buffered_post::Request { url: request.url, headers: request.headers, @@ -71,7 +72,7 @@ impl OcrTransport for DefaultOcrServices { headers: response.headers, content: response.content, }) - } + }) } } @@ -174,10 +175,10 @@ pub(crate) async fn send( .map_err(|_| Error::InvalidRequest("could not encode OCR request".into()))?; let response = transport .send(OcrTransportRequest { - url: endpoint.url, - headers, - body, - timeout_seconds: endpoint.timeout_seconds, + url: endpoint.url, + headers, + body, + timeout_seconds: endpoint.timeout_seconds, }) .await?; if !(200..300).contains(&response.status) { @@ -198,3 +199,76 @@ pub(crate) async fn send( .map_err(|_| Error::InvalidResponse("invalid OCR JSON response".into()))?; config.transform_ocr_response(&endpoint.model, response_json) } + +#[cfg(test)] +mod tests { + use std::sync::Mutex; + + use serde_json::json; + + use super::*; + + struct RecordingTransport { + request: Mutex>, + } + + impl OcrTransport for RecordingTransport { + fn send( + &self, + request: OcrTransportRequest, + ) -> Pin> + Send + '_>> + { + Box::pin(async move { + *self.request.lock().unwrap() = Some(request); + Ok(OcrTransportResponse { + status: 200, + headers: vec![], + content: serde_json::to_vec(&json!({ + "pages": [], + "model": "mistral-ocr-latest" + })) + .unwrap(), + }) + }) + } + } + + #[tokio::test] + async fn delegates_provider_io_to_transport() { + let transport = RecordingTransport { + request: Mutex::new(None), + }; + let request = SettledOcrRequest { + endpoint: OcrEndpoint { + model: "mistral-ocr-latest".into(), + custom_llm_provider: "mistral".into(), + url: "https://ocr.example/v1/ocr".into(), + timeout_seconds: 3.0, + }, + headers: vec![("Authorization".into(), "Bearer test-key".into())], + body: json!({ + "model": "mistral-ocr-latest", + "document": { + "type": "document_url", + "document_url": "data:application/pdf;base64,cGRm" + } + }), + }; + + let response = send(&transport, request).await.unwrap(); + let recorded = transport.request.into_inner().unwrap().unwrap(); + + assert_eq!(response.model, "mistral-ocr-latest"); + assert_eq!(recorded.url, "https://ocr.example/v1/ocr"); + assert_eq!(recorded.timeout_seconds, 3.0); + assert!( + recorded + .headers + .contains(&(b"Accept-Encoding".to_vec(), b"identity".to_vec())) + ); + assert_eq!( + serde_json::from_slice::(&recorded.body).unwrap()["model"], + "mistral-ocr-latest" + ); + } +} diff --git a/litellm-rust/crates/core/src/ocr/types.rs b/litellm-rust/crates/core/src/ocr/types.rs index 15fd8209c0b..e54fe2adf27 100644 --- a/litellm-rust/crates/core/src/ocr/types.rs +++ b/litellm-rust/crates/core/src/ocr/types.rs @@ -106,7 +106,7 @@ pub struct SettledOcrRequest { pub(super) body: Value, } -#[derive(Debug, PartialEq, Eq)] +#[derive(Debug, PartialEq)] pub struct OcrTransportRequest { pub url: String, pub headers: Vec<(Vec, Vec)>,