This commit is contained in:
Yujong Lee 2026-09-15 18:31:31 -07:00
parent 5de196af63
commit 056487a283
33 changed files with 3367 additions and 2594 deletions

View file

@ -2027,6 +2027,7 @@ dependencies = [
"serde_path_to_error",
"serde_with",
"sha2 0.10.9",
"strum",
"subtle",
"thiserror 2.0.19",
"tokio",

View file

@ -4,7 +4,6 @@ version = "0.1.0"
edition.workspace = true
license.workspace = true
repository.workspace = true
autotests = false
[dependencies]
bytes.workspace = true
@ -25,6 +24,7 @@ serde.workspace = true
serde_json = { workspace = true, features = ["preserve_order"] }
serde_with.workspace = true
serde_path_to_error = "0.1"
strum.workspace = true
subtle.workspace = true
tokio = { workspace = true, features = ["sync"] }
tokio-tungstenite.workspace = true
@ -32,6 +32,7 @@ thiserror.workspace = true
sha2.workspace = true
url.workspace = true
veil.workspace = true
[features]
default = []
observability = ["dep:tracing-subscriber"]

View file

@ -120,3 +120,126 @@ impl HostLifecycle {
};
}
}
#[cfg(test)]
mod tests {
use super::{HostFailure, HostLifecycle, HostPhase};
fn run(
fail_at: Option<HostPhase>,
asynchronous: bool,
) -> (Vec<HostPhase>, Vec<crate::ocr::Error>) {
let mut lifecycle = HostLifecycle::new(asynchronous);
let mut events = Vec::new();
let mut failures = Vec::new();
while lifecycle.phase() != HostPhase::Complete {
let phase = lifecycle.phase();
events.push(phase);
let result = if Some(phase) == fail_at {
Err(HostFailure::Error(crate::ocr::Error::InvalidRequest(
"selected failure".into(),
)))
} else {
Ok(())
};
if let Some(error) = lifecycle.accept(result) {
failures.push(error);
}
}
(events, failures)
}
#[test]
fn public_outcome_is_finalized_before_a_single_terminal_dispatch() {
for asynchronous in [false, true] {
let (events, failures) = run(None, asynchronous);
assert!(failures.is_empty());
assert_eq!(
&events[events.len() - 2..],
&[HostPhase::Finalize, HostPhase::Success]
);
assert_eq!(
events
.iter()
.filter(|phase| **phase == HostPhase::Execute)
.count(),
1
);
assert_eq!(
events.contains(&HostPhase::DeploymentPostCall),
asynchronous
);
}
}
#[test]
fn only_provider_and_response_construction_failures_use_provider_mapping() {
for phase in [
HostPhase::Setup,
HostPhase::DeploymentPreCall,
HostPhase::Prepare,
HostPhase::Execute,
HostPhase::ConstructResponse,
HostPhase::DeploymentPostCall,
HostPhase::Finalize,
] {
let (events, failures) = run(Some(phase), true);
assert_eq!(failures.len(), 1);
assert!(!events.contains(&HostPhase::Success));
let mapped = matches!(phase, HostPhase::Execute | HostPhase::ConstructResponse);
assert_eq!(events.contains(&HostPhase::MapFailure), mapped);
assert_eq!(events.contains(&HostPhase::DeploymentFailure), mapped);
assert_eq!(
&events[events.len() - 2..],
&[HostPhase::Failure, HostPhase::AsyncFailure]
);
assert!(
events
.iter()
.filter(|phase| **phase == HostPhase::Execute)
.count()
<= 1
);
}
}
#[test]
fn failure_handler_errors_do_not_replace_selected_failure_or_suppress_async_dispatch() {
let mut lifecycle = HostLifecycle::new(true);
while lifecycle.phase() != HostPhase::Execute {
lifecycle.accept::<crate::ocr::Error>(Ok(()));
}
let selected = crate::ocr::Error::InvalidRequest("provider".into());
assert_eq!(
lifecycle.accept(Err(HostFailure::Error(selected.clone()))),
Some(selected)
);
lifecycle.accept::<crate::ocr::Error>(Ok(()));
for phase in [
HostPhase::DeploymentFailure,
HostPhase::Failure,
HostPhase::AsyncFailure,
] {
assert_eq!(lifecycle.phase(), phase);
assert_eq!(
lifecycle.accept(Err(HostFailure::Error(crate::ocr::Error::InvalidRequest(
"callback".into()
)))),
None
);
}
assert_eq!(lifecycle.phase(), HostPhase::Complete);
}
#[test]
fn cancellation_skips_terminal_dispatch() {
let mut lifecycle = HostLifecycle::new(true);
let error = crate::ocr::Error::InvalidRequest("cancelled".into());
assert_eq!(
lifecycle.accept(Err(HostFailure::Cancelled(error.clone()))),
Some(error)
);
assert_eq!(lifecycle.phase(), HostPhase::Complete);
}
}

View file

@ -2,9 +2,6 @@ use std::future::Future;
use std::time::{Instant, SystemTime, UNIX_EPOCH};
pub mod host;
#[cfg(test)]
#[path = "../../tests/host_lifecycle.rs"]
mod host_tests;
pub mod types;
pub use types::{

View file

@ -1,12 +1,11 @@
use crate::llms::base_llm::ocr::transformation::{BaseOcrConfig, OcrRequestContext};
use crate::llms::cohere::ocr::transformation::{CohereParseConfig, CohereRequest};
use crate::llms::cohere::ocr::{CohereOptions, CohereResponse, validate_document};
use crate::llms::cohere::ocr::{CohereOptions, validate_document};
use crate::ocr::OcrClient;
use crate::ocr::document::{inline_remote_document, validate_inline_document};
use crate::ocr::prepare::{credential_env, transform_request_body};
use crate::ocr::types::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrDocument};
use crate::url_utils::ApiUrl;
use litellm_auth_azure::AzureAuthInputs;
const AZURE_AI_API_BASE_ENV: &str = "AZURE_AI_API_BASE";
@ -16,7 +15,58 @@ pub(crate) struct AzureAICohereParseConfig;
impl BaseOcrConfig for AzureAICohereParseConfig {
type OcrParams = CohereOptions;
type ProviderRequest = CohereRequest;
type ProviderResponse = CohereResponse;
type Environment = Vec<(String, String)>;
fn get_api_key_env_var(&self) -> Option<&'static str> {
super::transformation::AzureAIOCRConfig.get_api_key_env_var()
}
fn get_health_check_document(&self) -> OcrDocument {
CohereParseConfig.get_health_check_document()
}
async fn validate_environment(
&self,
request: &LiteLLMOcrRequest,
client: &OcrClient,
) -> Result<Self::Environment, crate::ocr::Error> {
BaseOcrConfig::validate_environment(
&super::transformation::AzureAIOCRConfig,
request,
client,
)
.await
}
fn get_complete_url(
&self,
request: &LiteLLMOcrRequest,
_params: &Self::OcrParams,
_environment: &Self::Environment,
) -> Result<String, crate::ocr::Error> {
let base = request
.connection
.api_base
.clone()
.or_else(|| credential_env(AZURE_AI_API_BASE_ENV))
.filter(|base| !base.trim().is_empty())
.ok_or_else(|| {
crate::ocr::Error::Auth(litellm_auth::Error::ProviderAuthentication(
"Missing Azure AI API Base - Set AZURE_AI_API_BASE or pass api_base".into(),
))
})?;
self.get_complete_url(&base)
}
fn transform_ocr_request(
&self,
model: &str,
document: OcrDocument,
params: &CohereOptions,
headers: &[(String, String)],
) -> Result<CohereRequest, crate::ocr::Error> {
CohereParseConfig.transform_ocr_request(model, document, params, headers)
}
fn get_supported_ocr_params(&self, model: &str) -> &'static [&'static str] {
CohereParseConfig.get_supported_ocr_params(model)
@ -37,15 +87,16 @@ impl BaseOcrConfig for AzureAICohereParseConfig {
context.connection,
)
.await?;
CohereParseConfig.transform_ocr_request(model, document, optional_params, headers)
self.transform_ocr_request(model, document, optional_params, headers)
}
fn normalize_response(
fn transform_ocr_response(
&self,
model: &str,
response: CohereResponse,
raw_response: &[u8],
request_format: crate::ocr::types::OcrResponseFormat,
) -> Result<LiteLLMOcrResponse, crate::ocr::Error> {
CohereParseConfig.normalize_response(model, response)
CohereParseConfig.transform_ocr_response(model, raw_response, request_format)
}
}
@ -56,28 +107,8 @@ impl AzureAICohereParseConfig {
client: &OcrClient,
) -> Result<reqwest::Request, crate::ocr::Error> {
let params = self.map_ocr_params(&request.optional_params, &request.model)?;
let config = AzureAuthInputs {
azure_ad_token_provider: request.azure_ad_token_provider.clone(),
..AzureAuthInputs::from_sourced_optional_params(
&request.optional_params,
&request.input_sources,
)
.map_err(crate::ocr::Error::from)?
};
let base = request
.connection
.api_base
.clone()
.or_else(|| credential_env(AZURE_AI_API_BASE_ENV))
.filter(|base| !base.trim().is_empty())
.ok_or_else(|| {
crate::ocr::Error::Auth(litellm_auth::Error::ProviderAuthentication(
"Missing Azure AI API Base - Set AZURE_AI_API_BASE or pass api_base".into(),
))
})?;
let headers = super::transformation::AzureAIOCRConfig
.validate_environment(&request.connection, &config, &credential_env)
.await?;
let url = BaseOcrConfig::get_complete_url(self, request, &params, &Vec::new())?;
let headers = self.validate_environment(request, client).await?;
let remote = request.document.source().starts_with("http://")
|| request.document.source().starts_with("https://");
let body = self
@ -92,19 +123,11 @@ impl AzureAICohereParseConfig {
},
)
.await?;
transform_request_body(
client,
request,
&self.get_complete_url(&base)?,
&headers,
!remote,
body,
|body| {
let document = crate::ocr::prepare::body_document(body)?;
validate_document(&document)?;
validate_inline_document(&document)
},
)
transform_request_body(client, request, &url, &headers, !remote, body, |body| {
let document = crate::ocr::prepare::body_document(body)?;
validate_document(&document)?;
validate_inline_document(&document)
})
.await
}
}

View file

@ -470,7 +470,63 @@ pub(crate) struct AzureDocumentIntelligenceOCRConfig;
impl BaseOcrConfig for AzureDocumentIntelligenceOCRConfig {
type OcrParams = DocumentIntelligenceParams;
type ProviderRequest = DocumentIntelligenceRequest;
type ProviderResponse = AzureDocumentIntelligenceOperation;
type Environment = Vec<(String, String)>;
fn get_api_key_env_var(&self) -> Option<&'static str> {
Some(AZURE_DI_API_KEY_ENV)
}
fn resolve_connection_params(
&self,
api_key: Option<litellm_auth::Sourced<String>>,
api_base: Option<litellm_auth::Sourced<String>>,
dynamic_api_key: Option<litellm_auth::Sourced<String>>,
dynamic_api_base: Option<litellm_auth::Sourced<String>>,
) -> (
Option<litellm_auth::Sourced<String>>,
Option<litellm_auth::Sourced<String>>,
) {
(
api_key.and_then(|key| {
dynamic_api_key
.filter(|value| !value.value().is_empty())
.or(Some(key))
}),
api_base.and_then(|base| {
dynamic_api_base
.filter(|value| !value.value().is_empty())
.or(Some(base))
}),
)
}
async fn validate_environment(
&self,
request: &LiteLLMOcrRequest,
_client: &OcrClient,
) -> Result<Self::Environment, crate::ocr::Error> {
let config = AzureAuthInputs {
azure_ad_token_provider: request.azure_ad_token_provider.clone(),
..AzureAuthInputs::from_sourced_optional_params(
&request.optional_params,
&request.input_sources,
)?
};
self.validate_environment(&request.connection, &config, &credential_env)
.await
}
fn get_complete_url(
&self,
request: &LiteLLMOcrRequest,
params: &Self::OcrParams,
_environment: &Self::Environment,
) -> Result<String, crate::ocr::Error> {
let endpoint = nonblank(request.connection.api_base.clone())
.or_else(|| nonblank(credential_env(AZURE_DI_ENDPOINT_ENV)))
.ok_or_else(|| crate::ocr::Error::Auth(litellm_auth::Error::ProviderAuthentication("Missing Azure Document Intelligence API Base - Set AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT or pass api_base".into())))?;
self.get_complete_url(&endpoint, &request.model, params)
}
fn get_supported_ocr_params(&self, _model: &str) -> &'static [&'static str] {
&["pages", "features", "req_format"]
@ -498,12 +554,18 @@ impl BaseOcrConfig for AzureDocumentIntelligenceOCRConfig {
self.transform_ocr_request(model, document, optional_params, headers)
}
fn normalize_response(
fn transform_ocr_response(
&self,
model: &str,
response: AzureDocumentIntelligenceOperation,
raw_response: &[u8],
request_format: OcrResponseFormat,
) -> Result<LiteLLMOcrResponse, crate::ocr::Error> {
transform_completed_response(model, response)
crate::llms::base_llm::ocr::transformation::decode_and_normalize_response(
model,
raw_response,
request_format,
transform_completed_response,
)
}
async fn async_transform_ocr_response(
@ -527,9 +589,6 @@ impl BaseOcrConfig for AzureDocumentIntelligenceOCRConfig {
..transform_completed_response(model, decoded.data)?
})
}
}
impl AzureDocumentIntelligenceOCRConfig {
fn transform_ocr_request(
&self,
_model: &str,
@ -539,28 +598,17 @@ impl AzureDocumentIntelligenceOCRConfig {
) -> Result<DocumentIntelligenceRequest, crate::ocr::Error> {
build_request(document)
}
}
impl AzureDocumentIntelligenceOCRConfig {
pub(crate) async fn prepare_request(
&self,
request: &LiteLLMOcrRequest,
client: &OcrClient,
) -> Result<reqwest::Request, crate::ocr::Error> {
let params = self.map_ocr_params(&request.optional_params, &request.model)?;
let config = AzureAuthInputs {
azure_ad_token_provider: request.azure_ad_token_provider.clone(),
..AzureAuthInputs::from_sourced_optional_params(
&request.optional_params,
&request.input_sources,
)
.map_err(crate::ocr::Error::from)?
};
let headers = self
.validate_environment(&request.connection, &config, &credential_env)
.await?;
let endpoint = nonblank(request.connection.api_base.clone())
.or_else(|| nonblank(credential_env(AZURE_DI_ENDPOINT_ENV)))
.ok_or_else(|| crate::ocr::Error::Auth(litellm_auth::Error::ProviderAuthentication("Missing Azure Document Intelligence API Base - Set AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT or pass api_base".into())))?;
let url = self.get_complete_url(&endpoint, &request.model, &params)?;
let headers = BaseOcrConfig::validate_environment(self, request, client).await?;
let url = BaseOcrConfig::get_complete_url(self, request, &params, &headers)?;
let body = self
.async_transform_ocr_request(
&request.model,
@ -627,7 +675,7 @@ impl AzureDocumentIntelligenceOCRConfig {
let key = nonblank(connection.api_key.clone())
.map(|value| Sourced::new(value, connection.api_key_source))
.or_else(|| {
nonblank(env_lookup(AZURE_DI_API_KEY_ENV))
nonblank(self.get_api_key_env_var().and_then(env_lookup))
.map(|value| Sourced::new(value, InputSource::Environment))
});
if let Some(key) = key {
@ -706,7 +754,7 @@ mod tests {
#[test]
fn response_numbers_follow_python_validation_before_dimension_conversion() {
let response = AzureDocumentIntelligenceOCRConfig.decode_and_normalize_response(
let response = AzureDocumentIntelligenceOCRConfig.transform_ocr_response(
"model",
br#"{"status":"succeeded","analyzeResult":{"pages":[{"pageNumber":2.0,"width":" 8.5 ","height":true}]}}"#,
OcrResponseFormat::Litellm,
@ -814,4 +862,464 @@ mod tests {
(AZURE_DI_SUBSCRIPTION_HEADER.into(), "request-key".into())
);
}
use std::sync::{Arc, Mutex};
use crate::ocr::test_support::{MockResponse, mock_server, perform_ocr, wire_request};
use crate::ocr::wire::{OcrWireRequest, decode_request};
fn query_value(url: &str, key: &str) -> Option<String> {
url::Url::parse(url)
.unwrap()
.query_pairs()
.find_map(|(name, value)| (name == key).then(|| value.into_owned()))
}
#[tokio::test]
async fn facade_maps_pages_features_and_url_document() {
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({
"status":"succeeded",
"analyzeResult":{"pages":[]}
}))])
.await;
let mut request = wire_request(
"azure_ai/doc-intelligence/prebuilt-read",
&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!({
"type":"document_url",
"document_url":"https://example.com/document.pdf"
}))
.unwrap();
perform_ocr(request).await.unwrap();
server.await.unwrap();
let request = &seen.lock().unwrap()[0];
let target = request.split_whitespace().nth(1).unwrap();
let url = format!("{base}{target}");
assert_eq!(query_value(&url, "pages").as_deref(), Some("1,2,3"));
assert_eq!(
query_value(&url, "features").as_deref(),
Some("keyValuePairs,languages")
);
let body: Value = serde_json::from_str(request.split_once("\r\n\r\n").unwrap().1).unwrap();
assert_eq!(
body,
json!({"urlSource":"https://example.com/document.pdf", "future_option":{"nested":null}, "provider_option":false})
);
}
#[tokio::test]
async fn rejects_invalid_pages_features_and_format() {
for options in [
json!({"pages":[true]}),
json!({"pages":[1,"2"]}),
json!({"pages":[-1]}),
json!({"pages":"1&&features=bad"}),
json!({"features":"languages&pages=1"}),
json!({"req_format":"azure"}),
] {
let result = decode_request(OcrWireRequest {
model: "azure_ai/doc-intelligence/prebuilt-read".into(),
document: json!({"type":"document_url","document_url":"https://example.com/a.pdf"}),
api_key: Some("key".into()),
api_base: Some("http://127.0.0.1:1".into()),
custom_llm_provider: None,
extra_headers: None,
optional_params: options.as_object().unwrap().clone().into(),
input_sources: Default::default(),
timeout_seconds: None,
});
let rejected = match result {
Ok(request) => perform_ocr(request).await.is_err(),
Err(_) => true,
};
assert!(rejected, "accepted {options}");
}
}
#[tokio::test]
async fn inline_document_decodes_to_base64_source() {
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({
"status":"succeeded"
}))])
.await;
let request = wire_request("azure_ai/doc-intelligence/prebuilt-read", &base, json!({}));
perform_ocr(request).await.unwrap();
server.await.unwrap();
let request = &seen.lock().unwrap()[0];
let body: Value = serde_json::from_str(request.split_once("\r\n\r\n").unwrap().1).unwrap();
assert_eq!(body, json!({"base64Source":"YWJj"}));
}
#[tokio::test]
async fn immediate_response_normalizes_pages_and_preserves_native() {
let operation = json!({
"status":"succeeded",
"operationExtension":42,
"analyzeResult":{
"content":"A\n\nB",
"tables":[{"cells":[]}],
"keyValuePairs":[{"key":{"content":"A"}}],
"pages":[{
"pageNumber":"2",
"width":"8.5",
"height":11,
"unit":"inch",
"lines":[{"content":"A"},{"content":null},{"content":"B"}]
}]
}
});
let (base, _, server) = mock_server(vec![MockResponse::json(operation.clone())]).await;
let result = perform_ocr(wire_request(
"azure_ai/doc-intelligence/prebuilt-read",
&base,
json!({"req_format":"native"}),
))
.await
.unwrap();
server.await.unwrap();
assert_eq!(result.pages[0].index, 1);
assert_eq!(result.pages[0].markdown, "A\n\nB");
assert_eq!(
serde_json::to_value(&result.pages[0].dimensions).unwrap(),
json!({"width":816,"height":1056,"dpi":96})
);
assert_eq!(result.usage_info.as_ref().unwrap().pages_processed, Some(1));
let serialized = result.clone().into_json();
assert_eq!(serialized["content"], "A\n\nB");
assert_eq!(serialized["tables"], json!([{"cells":[]}]));
assert_eq!(
serialized["keyValuePairs"],
json!([{"key":{"content":"A"}}])
);
assert!(serialized.get("key_value_pairs").is_none());
assert_eq!(
result.provider_native_response.as_ref(),
operation.as_object()
);
}
#[tokio::test]
async fn accepted_response_polls_to_success_with_only_credentials() {
let operation = json!({"status":"succeeded","analyzeResult":{"pages":[]}});
let (base, seen, server) = mock_server(vec![
MockResponse {
status: 202,
headers: vec![("Operation-Location", "{base}/operation".into())],
body: json!({}),
},
MockResponse {
status: 200,
headers: vec![("Retry-After", "0".into())],
body: json!({"status":"running"}),
},
MockResponse::json(operation.clone()),
])
.await;
let mut request = wire_request(
"azure_ai/doc-intelligence/prebuilt-read",
&base,
json!({"req_format":"native"}),
);
request
.connection
.extra_headers
.push(("X-Trace".into(), "initial-only".into()));
let result = perform_ocr(request).await.unwrap();
server.await.unwrap();
assert_eq!(
result.provider_native_response.as_ref(),
operation.as_object()
);
let requests = seen.lock().unwrap();
assert_eq!(requests.len(), 3);
assert!(requests[0].to_ascii_lowercase().contains("x-trace:"));
for poll in &requests[1..] {
assert!(!poll.to_ascii_lowercase().contains("x-trace:"));
assert!(
poll.to_ascii_lowercase()
.contains("ocp-apim-subscription-key: test-key")
);
}
}
struct SubmissionBoundary {
request_count: Arc<Mutex<Vec<String>>>,
post_calls: Arc<Mutex<Vec<Value>>>,
}
impl crate::ocr::hooks::OcrHooks for SubmissionBoundary {
fn post_call(
&self,
request: crate::ocr::hooks::OcrPostCallRequest,
) -> crate::ocr::hooks::OcrHookFuture<'_, crate::ocr::hooks::OcrPostCallRequest> {
Box::pin(async move {
assert_eq!(self.request_count.lock().unwrap().len(), 1);
self.post_calls
.lock()
.unwrap()
.push(request.original_response.clone());
Ok(request)
})
}
}
#[tokio::test]
async fn accepted_response_runs_post_call_once_before_polling() {
let (base, seen, server) = mock_server(vec![
MockResponse {
status: 202,
headers: vec![("Operation-Location", "{base}/operation".into())],
body: json!({"submitted": true}),
},
MockResponse::json(json!({"status":"succeeded"})),
])
.await;
let post_calls = Arc::new(Mutex::new(Vec::new()));
let request = crate::ocr::LiteLLMOcrRequest {
hooks: Arc::new(SubmissionBoundary {
request_count: seen.clone(),
post_calls: post_calls.clone(),
}),
..wire_request("azure_ai/doc-intelligence/prebuilt-read", &base, json!({}))
};
perform_ocr(request).await.unwrap();
server.await.unwrap();
assert_eq!(seen.lock().unwrap().len(), 2);
assert_eq!(
*post_calls.lock().unwrap(),
[json!(r#"{"submitted":true}"#)]
);
}
#[tokio::test]
async fn polling_forwards_bearer_credentials() {
let (base, seen, server) = mock_server(vec![
MockResponse {
status: 202,
headers: vec![("Operation-Location", "{base}/operation".into())],
body: json!({}),
},
MockResponse::json(json!({"status":"succeeded"})),
])
.await;
let mut request = wire_request("azure_ai/doc-intelligence/prebuilt-read", &base, json!({}));
request.connection.api_key = None;
request.connection.extra_headers = vec![("Authorization".into(), "Bearer token".into())];
perform_ocr(request).await.unwrap();
server.await.unwrap();
let requests = seen.lock().unwrap();
assert!(
requests[1]
.to_ascii_lowercase()
.contains("authorization: bearer token")
);
}
#[tokio::test]
async fn polling_does_not_follow_redirects() {
let (base, seen, server) = mock_server(vec![
MockResponse {
status: 202,
headers: vec![("Operation-Location", "{base}/operation".into())],
body: json!({}),
},
MockResponse {
status: 302,
headers: vec![("Location", "{base}/redirected".into())],
body: json!({}),
},
MockResponse::json(json!({"status":"succeeded"})),
])
.await;
let error = perform_ocr(wire_request(
"azure_ai/doc-intelligence/prebuilt-read",
&base,
json!({}),
))
.await
.unwrap_err();
assert!(error.to_string().contains("status 302"), "{error}");
assert_eq!(seen.lock().unwrap().len(), 2);
server.abort();
}
#[tokio::test]
async fn polling_rejects_terminal_failure() {
let (base, _, server) = mock_server(vec![
MockResponse {
status: 202,
headers: vec![("Operation-Location", "{base}/operation".into())],
body: json!({}),
},
MockResponse::json(json!({"status":"failed"})),
])
.await;
let error = perform_ocr(wire_request(
"azure_ai/doc-intelligence/prebuilt-read",
&base,
json!({}),
))
.await
.unwrap_err();
server.await.unwrap();
assert!(error.to_string().contains("status failed"));
}
#[tokio::test]
async fn malformed_provider_pages_report_response_paths() {
for (analysis, path) in [
(json!({"pages":null}), "pages"),
(json!({"pages":[null]}), "pages[0]"),
(json!({"pages":[{"lines":null}]}), "lines"),
(json!({"pages":[{"width":"bad"}]}), "width"),
] {
let (base, _, server) = mock_server(vec![MockResponse::json(json!({
"status":"succeeded",
"analyzeResult":analysis
}))])
.await;
let error = perform_ocr(wire_request(
"azure_ai/doc-intelligence/prebuilt-read",
&base,
json!({}),
))
.await
.unwrap_err();
server.await.unwrap();
assert!(error.to_string().contains(path), "{error}");
}
}
#[tokio::test]
async fn rejects_missing_invalid_and_cross_origin_operation_locations() {
for headers in [
Vec::new(),
vec![("Operation-Location", "/relative".into())],
vec![("Operation-Location", "http://example.com/operation".into())],
vec![(
"Operation-Location",
"http://user:password@127.0.0.1/operation".into(),
)],
] {
let (base, _, server) = mock_server(vec![MockResponse {
status: 202,
headers,
body: json!({}),
}])
.await;
let error = perform_ocr(wire_request(
"azure_ai/doc-intelligence/prebuilt-read",
&base,
json!({}),
))
.await
.unwrap_err();
server.await.unwrap();
assert!(error.to_string().contains("operation-location"));
}
}
#[tokio::test]
async fn polling_deadline_bounds_retry_delay() {
let (base, _, server) = mock_server(vec![
MockResponse {
status: 202,
headers: vec![("Operation-Location", "{base}/operation".into())],
body: json!({}),
},
MockResponse {
status: 200,
headers: vec![("Retry-After", "9999".into())],
body: json!({"status":"notStarted"}),
},
])
.await;
let mut request = wire_request("azure_ai/doc-intelligence/prebuilt-read", &base, json!({}));
request.connection.poll_timeout = std::time::Duration::from_millis(100);
let error = tokio::time::timeout(std::time::Duration::from_secs(1), perform_ocr(request))
.await
.unwrap()
.unwrap_err();
server.await.unwrap();
assert!(error.to_string().contains("timed out"));
}
#[tokio::test]
async fn model_id_is_encoded_and_dot_segments_are_rejected() {
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({
"status":"succeeded"
}))])
.await;
perform_ocr(wire_request(
"azure_ai/doc-intelligence/a ?#é",
&base,
json!({}),
))
.await
.unwrap();
server.await.unwrap();
assert!(seen.lock().unwrap()[0].contains("a%20%3F%23%C3%A9:analyze"));
for model in [
"azure_ai/doc-intelligence/.",
"azure_ai/doc-intelligence/..",
] {
let error = perform_ocr(wire_request(model, "http://127.0.0.1:1", json!({})))
.await
.unwrap_err();
assert!(error.to_string().contains("dot segment"));
}
}
#[tokio::test]
async fn pre_call_guardrail_receives_caller_pages_before_mapping() {
use crate::ocr::hooks::{OcrHookFuture, OcrHooks, OcrPreCallRequest};
use std::sync::Arc;
struct RewritePages;
impl OcrHooks for RewritePages {
fn intercepts_requests(&self) -> bool {
true
}
fn pre_call(&self, request: OcrPreCallRequest) -> OcrHookFuture<'_, OcrPreCallRequest> {
Box::pin(async move {
assert_eq!(request.optional_params["pages"], json!([0, 2]));
Ok(OcrPreCallRequest {
optional_params: json!({"pages": [1]}),
..request
})
})
}
}
let (base, seen, server) =
mock_server(vec![MockResponse::json(json!({"status": "succeeded"}))]).await;
let request = wire_request(
"azure_ai/doc-intelligence/prebuilt-read",
&base,
json!({"pages": [0, 2]}),
)
.with_host_hooks(Arc::new(RewritePages), None);
perform_ocr(request).await.unwrap();
server.await.unwrap();
let requests = seen.lock().unwrap();
let target = requests[0].split_whitespace().nth(1).unwrap();
assert_eq!(
query_value(&format!("{base}{target}"), "pages").as_deref(),
Some("2")
);
assert_eq!(requests.len(), 1);
}
}

View file

@ -1,6 +1,5 @@
use crate::constants::AZURE_AI_OCR_PATH;
use crate::llms::base_llm::ocr::transformation::{BaseOcrConfig, OcrRequestContext};
use crate::llms::mistral::ocr::MistralOcrResponse;
use crate::llms::mistral::ocr::transformation::{MistralOCRConfig, MistralOcrRequest};
use crate::ocr::OcrClient;
use crate::ocr::document::{inline_remote_document, validate_inline_document};
@ -20,7 +19,46 @@ pub(crate) struct AzureAIOCRConfig;
impl BaseOcrConfig for AzureAIOCRConfig {
type OcrParams = OpaqueParams;
type ProviderRequest = MistralOcrRequest;
type ProviderResponse = MistralOcrResponse;
type Environment = Vec<(String, String)>;
fn get_api_key_env_var(&self) -> Option<&'static str> {
Some(AZURE_AI_API_KEY_ENV)
}
async fn validate_environment(
&self,
request: &LiteLLMOcrRequest,
_client: &OcrClient,
) -> Result<Self::Environment, crate::ocr::Error> {
let config = AzureAuthInputs {
azure_ad_token_provider: request.azure_ad_token_provider.clone(),
..AzureAuthInputs::from_sourced_optional_params(
&request.optional_params,
&request.input_sources,
)?
};
self.validate_environment(&request.connection, &config, &credential_env)
.await
}
fn get_complete_url(
&self,
request: &LiteLLMOcrRequest,
_params: &Self::OcrParams,
_environment: &Self::Environment,
) -> Result<String, crate::ocr::Error> {
self.get_complete_url(request.connection.api_base.as_deref(), &credential_env)
}
fn transform_ocr_request(
&self,
model: &str,
document: OcrDocument,
params: &OpaqueParams,
headers: &[(String, String)],
) -> Result<MistralOcrRequest, crate::ocr::Error> {
MistralOCRConfig.transform_ocr_request(model, document, params, headers)
}
fn get_supported_ocr_params(&self, model: &str) -> &'static [&'static str] {
MistralOCRConfig.get_supported_ocr_params(model)
@ -40,15 +78,16 @@ impl BaseOcrConfig for AzureAIOCRConfig {
context.connection,
)
.await?;
MistralOCRConfig.transform_ocr_request(model, document, optional_params, headers)
self.transform_ocr_request(model, document, optional_params, headers)
}
fn normalize_response(
fn transform_ocr_response(
&self,
model: &str,
response: MistralOcrResponse,
raw_response: &[u8],
request_format: crate::ocr::types::OcrResponseFormat,
) -> Result<LiteLLMOcrResponse, crate::ocr::Error> {
MistralOCRConfig.normalize_response(model, response)
MistralOCRConfig.transform_ocr_response(model, raw_response, request_format)
}
}
@ -59,18 +98,8 @@ impl AzureAIOCRConfig {
client: &OcrClient,
) -> Result<reqwest::Request, crate::ocr::Error> {
let params = self.map_ocr_params(&request.optional_params, &request.model)?;
let config = AzureAuthInputs {
azure_ad_token_provider: request.azure_ad_token_provider.clone(),
..AzureAuthInputs::from_sourced_optional_params(
&request.optional_params,
&request.input_sources,
)
.map_err(crate::ocr::Error::from)?
};
let url = self.get_complete_url(request.connection.api_base.as_deref(), &credential_env)?;
let headers = self
.validate_environment(&request.connection, &config, &credential_env)
.await?;
let url = BaseOcrConfig::get_complete_url(self, request, &params, &Vec::new())?;
let headers = BaseOcrConfig::validate_environment(self, request, client).await?;
let retains_document = !request.document.source().starts_with("http://")
&& !request.document.source().starts_with("https://");
let body = self
@ -132,7 +161,7 @@ impl AzureAIOCRConfig {
let key = nonblank(connection.api_key.clone())
.map(|value| Sourced::new(value, connection.api_key_source))
.or_else(|| {
nonblank(env_lookup(AZURE_AI_API_KEY_ENV))
nonblank(self.get_api_key_env_var().and_then(env_lookup))
.map(|value| Sourced::new(value, InputSource::Environment))
});
if let Some(key) = key {
@ -259,4 +288,103 @@ mod tests {
("Authorization".into(), "Bearer request-key".into())
);
}
use std::sync::Arc;
use serde_json::{Value, json};
use crate::ocr::hooks::{OcrDuringCallRequest, OcrHookFuture, OcrHooks};
use crate::ocr::test_support::{MockResponse, mock_server, perform_ocr, wire_request};
#[tokio::test]
async fn facade_executes_azure_mistral_with_prepared_auth() {
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({
"pages":[{"index":0,"markdown":"hello"}],
"usage_info":{"pages_processed":1}
}))])
.await;
let mut request = wire_request(
"azure_ai/model",
&base,
json!({"include_image_base64":true}),
);
request.connection.api_key = None;
request.connection.extra_headers = vec![(
"Authorization".into(),
"Bearer python-prepared-token".into(),
)];
let result = perform_ocr(request).await.unwrap();
server.await.unwrap();
assert_eq!(result.pages[0].markdown, "hello");
let requests = seen.lock().unwrap();
assert_eq!(requests.len(), 1);
assert!(requests[0].starts_with("POST /providers/mistral/azure/ocr "));
assert!(
requests[0]
.to_ascii_lowercase()
.contains("authorization: bearer python-prepared-token\r\n")
);
let body: Value =
serde_json::from_str(requests[0].split_once("\r\n\r\n").unwrap().1).unwrap();
assert_eq!(
body,
json!({
"model":"model",
"document":{"type":"document_url","document_url":"data:application/pdf;base64,YWJj"},
"include_image_base64":true
})
);
}
#[tokio::test]
async fn facade_acquires_supplied_entra_token_for_final_request() {
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({"pages":[]}))]).await;
let mut request = wire_request(
"azure_ai/model",
&base,
json!({"azure_ad_token":"rust-owned-token"}),
);
request.connection.api_key = None;
perform_ocr(request).await.unwrap();
server.await.unwrap();
let requests = seen.lock().unwrap();
assert_eq!(requests.len(), 1);
assert!(
requests[0]
.to_ascii_lowercase()
.contains("authorization: bearer rust-owned-token\r\n")
);
}
struct ReplaceBodyDocument;
impl OcrHooks for ReplaceBodyDocument {
fn intercepts_requests(&self) -> bool {
true
}
fn during_call(
&self,
mut request: OcrDuringCallRequest,
) -> OcrHookFuture<'_, OcrDuringCallRequest> {
Box::pin(async move {
request.body["document"] = json!({
"type":"document_url",
"document_url":"https://example.com/not-inline.pdf"
});
Ok(request)
})
}
}
#[tokio::test]
async fn rejects_non_inline_body_after_guardrails() {
let mut request = wire_request("azure_ai/model", "http://127.0.0.1:1", json!({}));
request.hooks = Arc::new(ReplaceBodyDocument);
let error = perform_ocr(request).await.unwrap_err();
assert!(error.to_string().contains("data URI"));
}
}

View file

@ -1,18 +1,64 @@
use std::future::Future;
use std::sync::Arc;
use litellm_auth::Sourced;
use serde::Serialize;
use serde::de::DeserializeOwned;
use crate::call_arguments::{CallArguments, parse_options};
use crate::ocr::OcrClient;
use crate::ocr::hooks::OcrHooks;
use crate::ocr::types::{LiteLLMOcrResponse, OcrConnection, OcrDocument, OcrResponseFormat};
use crate::ocr::types::{
LiteLLMOcrRequest, LiteLLMOcrResponse, OcrConnection, OcrDocument, OcrResponseFormat,
};
const HEALTH_CHECK_PDF_DATA_URI: &str = "data:application/pdf;base64,JVBERi0xLjQKJeLjz9MKMyAwIG9iago8PC9UeXBlIC9QYWdlCi9QYXJlbnQgMSAwIFIKL01lZGlhQm94IFswIDAgNjEyIDc5Ml0KL0NvbnRlbnRzIDQgMCBSCi9SZXNvdXJjZXMgPDwvRm9udCA8PC9GMSAyIDAgUj4+Pj4+PgplbmRvYmoKNCAwIG9iago8PC9MZW5ndGggNDQ+PgpzdHJlYW0KQlQKL0YxIDI0IFRmCjEwMCA3MDAgVGQKKHRlc3QpIFRqCkVUCmVuZHN0cmVhbQplbmRvYmoKMiAwIG9iago8PC9UeXBlIC9Gb250Ci9TdWJ0eXBlIC9UeXBlMQovQmFzZUZvbnQgL0hlbHZldGljYT4+CmVuZG9iagoxIDAgb2JqCjw8L1R5cGUgL1BhZ2VzCi9LaWRzIFszIDAgUl0KL0NvdW50IDE+PgplbmRvYmoKNSAwIG9iago8PC9UeXBlIC9DYXRhbG9nCi9QYWdlcyAxIDAgUj4+CmVuZG9iagp0cmFpbGVyCjw8L1NpemUgNgovUm9vdCA1IDAgUj4+CnN0YXJ0eHJlZgozMjQKJSVFT0Y=";
pub(crate) trait BaseOcrConfig: Send + Sync + Sized + 'static {
type OcrParams: DeserializeOwned + Send + Sync;
type ProviderRequest: Serialize + Send;
type ProviderResponse: DeserializeOwned + Send;
type Environment: Send + Sync;
fn get_api_key_env_var(&self) -> Option<&'static str> {
None
}
fn resolve_connection_params(
&self,
api_key: Option<Sourced<String>>,
api_base: Option<Sourced<String>>,
dynamic_api_key: Option<Sourced<String>>,
dynamic_api_base: Option<Sourced<String>>,
) -> (Option<Sourced<String>>, Option<Sourced<String>>) {
(
dynamic_api_key
.filter(|value| !value.value().is_empty())
.or(api_key),
dynamic_api_base
.filter(|value| !value.value().is_empty())
.or(api_base),
)
}
fn get_health_check_document(&self) -> OcrDocument {
OcrDocument::DocumentUrl {
document_url: HEALTH_CHECK_PDF_DATA_URI.into(),
extra_fields: Default::default(),
}
}
fn validate_environment(
&self,
request: &LiteLLMOcrRequest,
client: &OcrClient,
) -> impl Future<Output = Result<Self::Environment, crate::ocr::Error>> + Send;
fn get_complete_url(
&self,
request: &LiteLLMOcrRequest,
optional_params: &Self::OcrParams,
environment: &Self::Environment,
) -> Result<String, crate::ocr::Error>;
fn get_supported_ocr_params(&self, _model: &str) -> &'static [&'static str] {
&[]
@ -30,36 +76,31 @@ pub(crate) trait BaseOcrConfig: Send + Sync + Sized + 'static {
)?)
}
fn transform_ocr_request(
&self,
model: &str,
document: OcrDocument,
optional_params: &Self::OcrParams,
headers: &[(String, String)],
) -> Result<Self::ProviderRequest, crate::ocr::Error>;
fn async_transform_ocr_request(
&self,
model: &str,
document: OcrDocument,
optional_params: &Self::OcrParams,
headers: &[(String, String)],
context: OcrRequestContext<'_>,
) -> impl Future<Output = Result<Self::ProviderRequest, crate::ocr::Error>> + Send;
_context: OcrRequestContext<'_>,
) -> impl Future<Output = Result<Self::ProviderRequest, crate::ocr::Error>> + Send {
async move { self.transform_ocr_request(model, document, optional_params, headers) }
}
fn normalize_response(
&self,
model: &str,
response: Self::ProviderResponse,
) -> Result<LiteLLMOcrResponse, crate::ocr::Error>;
fn decode_and_normalize_response(
fn transform_ocr_response(
&self,
model: &str,
raw_response: &[u8],
request_format: OcrResponseFormat,
) -> Result<LiteLLMOcrResponse, crate::ocr::Error> {
let decoded = crate::ocr::wire::decode_response::<Self::ProviderResponse>(
raw_response,
request_format == OcrResponseFormat::Native,
)?;
Ok(LiteLLMOcrResponse {
provider_native_response: decoded.native,
..self.normalize_response(model, decoded.data)?
})
}
) -> Result<LiteLLMOcrResponse, crate::ocr::Error>;
fn async_transform_ocr_response(
&self,
@ -74,9 +115,38 @@ pub(crate) trait BaseOcrConfig: Send + Sync + Sized + 'static {
)
.await?;
crate::ocr::handler::post_call(context.hooks, &bytes).await?;
self.decode_and_normalize_response(model, &bytes, context.request_format)
self.transform_ocr_response(model, &bytes, context.request_format)
}
}
fn get_error_class(
&self,
error_message: String,
status_code: u16,
headers: Vec<(String, String)>,
) -> crate::ocr::Error {
crate::ocr::Error::Provider {
status: status_code,
body: error_message,
headers,
}
}
}
pub(crate) fn decode_and_normalize_response<T: DeserializeOwned>(
model: &str,
raw_response: &[u8],
request_format: OcrResponseFormat,
normalize: impl FnOnce(&str, T) -> Result<LiteLLMOcrResponse, crate::ocr::Error>,
) -> Result<LiteLLMOcrResponse, crate::ocr::Error> {
let decoded = crate::ocr::wire::decode_response(
raw_response,
request_format == OcrResponseFormat::Native,
)?;
Ok(LiteLLMOcrResponse {
provider_native_response: decoded.native,
..normalize(model, decoded.data)?
})
}
#[derive(Clone, Copy)]

View file

@ -1,3 +1,3 @@
pub(crate) mod transformation;
pub(crate) use transformation::{CohereOptions, CohereResponse, validate_document};
pub(crate) use transformation::{CohereOptions, validate_document};

View file

@ -1,8 +1,8 @@
use crate::serde_compat::LaxI64;
use serde::{Deserialize, Serialize};
use serde_json::{Map, Value};
use serde_with::serde_as;
use crate::serde_compat::LaxI64;
use crate::constants::{COHERE_API_KEY_ENV, COHERE_PARSE_API_BASE};
use crate::llms::base_llm::ocr::transformation::{BaseOcrConfig, OcrRequestContext};
use crate::ocr::OcrClient;
@ -14,6 +14,8 @@ use crate::ocr::types::{
};
use crate::url_utils::ApiUrl;
const COHERE_PARSE_HEALTH_CHECK_IMAGE_DATA_URI: &str = "data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAIAAACQd1PeAAAADElEQVR4nGP4//8/AAX+Av4N70a4AAAAAElFTkSuQmCC";
#[derive(Clone, Copy, Debug, Default, Deserialize, Serialize)]
#[serde(rename_all = "lowercase")]
pub(crate) enum OutputFormat {
@ -80,8 +82,46 @@ struct CohereBilledUnits {
#[derive(Default)]
pub(crate) struct CohereParseConfig;
impl CohereParseConfig {
pub(crate) fn transform_ocr_request(
impl BaseOcrConfig for CohereParseConfig {
type OcrParams = CohereOptions;
type ProviderRequest = CohereRequest;
type Environment = Vec<(String, String)>;
fn get_api_key_env_var(&self) -> Option<&'static str> {
Some(COHERE_API_KEY_ENV)
}
fn get_health_check_document(&self) -> OcrDocument {
OcrDocument::ImageUrl {
image_url: COHERE_PARSE_HEALTH_CHECK_IMAGE_DATA_URI.into(),
extra_fields: Default::default(),
}
}
async fn validate_environment(
&self,
request: &LiteLLMOcrRequest,
_client: &OcrClient,
) -> Result<Self::Environment, crate::ocr::Error> {
self.validate_environment(&request.connection, &credential_env)
}
fn get_complete_url(
&self,
request: &LiteLLMOcrRequest,
_params: &Self::OcrParams,
_environment: &Self::Environment,
) -> Result<String, crate::ocr::Error> {
self.get_complete_url(
request
.connection
.api_base
.as_deref()
.unwrap_or(COHERE_PARSE_API_BASE),
)
}
fn transform_ocr_request(
&self,
model: &str,
document: OcrDocument,
@ -91,12 +131,6 @@ impl CohereParseConfig {
let image_url = image_url(document)?;
Ok(build_request(model, image_url, optional_params))
}
}
impl BaseOcrConfig for CohereParseConfig {
type OcrParams = CohereOptions;
type ProviderRequest = CohereRequest;
type ProviderResponse = CohereResponse;
fn get_supported_ocr_params(&self, _model: &str) -> &'static [&'static str] {
&["output_format", "req_format"]
@ -113,12 +147,18 @@ impl BaseOcrConfig for CohereParseConfig {
self.transform_ocr_request(model, document, optional_params, headers)
}
fn normalize_response(
fn transform_ocr_response(
&self,
model: &str,
response: CohereResponse,
raw_response: &[u8],
request_format: crate::ocr::types::OcrResponseFormat,
) -> Result<LiteLLMOcrResponse, crate::ocr::Error> {
normalize_response(model, response)
crate::llms::base_llm::ocr::transformation::decode_and_normalize_response(
model,
raw_response,
request_format,
normalize_response,
)
}
}
@ -129,14 +169,8 @@ impl CohereParseConfig {
client: &OcrClient,
) -> Result<reqwest::Request, crate::ocr::Error> {
let params = self.map_ocr_params(&request.optional_params, &request.model)?;
let headers = self.validate_environment(&request.connection, &credential_env)?;
let url = self.get_complete_url(
request
.connection
.api_base
.as_deref()
.unwrap_or(COHERE_PARSE_API_BASE),
)?;
let headers = BaseOcrConfig::validate_environment(self, request, client).await?;
let url = BaseOcrConfig::get_complete_url(self, request, &params, &headers)?;
let body = self
.async_transform_ocr_request(
&request.model,
@ -299,7 +333,11 @@ impl CohereParseConfig {
.map(str::trim)
.filter(|key| !key.is_empty())
.map(str::to_string)
.or_else(|| env_lookup(COHERE_API_KEY_ENV).filter(|key| !key.trim().is_empty()))
.or_else(|| {
self.get_api_key_env_var()
.and_then(env_lookup)
.filter(|key| !key.trim().is_empty())
})
.ok_or_else(|| {
crate::ocr::Error::Auth(litellm_auth::Error::ProviderAuthentication(
"Missing COHERE_API_KEY - set it in the environment or pass api_key".into(),

View file

@ -1,3 +1 @@
pub(crate) mod transformation;
pub(crate) use transformation::MistralOcrResponse;

View file

@ -37,9 +37,34 @@ pub(crate) struct MistralOcrResponse {
#[derive(Clone, Debug, Default)]
pub(crate) struct MistralOCRConfig;
impl MistralOCRConfig {
impl BaseOcrConfig for MistralOCRConfig {
type OcrParams = OpaqueParams;
type ProviderRequest = MistralOcrRequest;
type Environment = Vec<(String, String)>;
fn get_api_key_env_var(&self) -> Option<&'static str> {
Some(MISTRAL_API_KEY_ENV)
}
async fn validate_environment(
&self,
request: &LiteLLMOcrRequest,
_client: &OcrClient,
) -> Result<Self::Environment, crate::ocr::Error> {
self.validate_environment(&request.connection, &credential_env)
}
fn get_complete_url(
&self,
request: &LiteLLMOcrRequest,
_params: &Self::OcrParams,
_environment: &Self::Environment,
) -> Result<String, crate::ocr::Error> {
self.get_complete_url(request.connection.api_base.as_deref())
}
#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)]
pub(crate) fn transform_ocr_request(
fn transform_ocr_request(
&self,
model: &str,
document: OcrDocument,
@ -52,12 +77,6 @@ impl MistralOCRConfig {
params: optional_params.clone(),
})
}
}
impl BaseOcrConfig for MistralOCRConfig {
type OcrParams = OpaqueParams;
type ProviderRequest = MistralOcrRequest;
type ProviderResponse = MistralOcrResponse;
fn get_supported_ocr_params(&self, _model: &str) -> &'static [&'static str] {
&[
@ -88,12 +107,18 @@ impl BaseOcrConfig for MistralOCRConfig {
self.transform_ocr_request(model, document, optional_params, headers)
}
fn normalize_response(
fn transform_ocr_response(
&self,
model: &str,
response: MistralOcrResponse,
raw_response: &[u8],
request_format: crate::ocr::types::OcrResponseFormat,
) -> Result<LiteLLMOcrResponse, crate::ocr::Error> {
normalize_response(model, response)
crate::llms::base_llm::ocr::transformation::decode_and_normalize_response(
model,
raw_response,
request_format,
normalize_response,
)
}
}
@ -104,8 +129,8 @@ impl MistralOCRConfig {
client: &OcrClient,
) -> Result<reqwest::Request, crate::ocr::Error> {
let params = self.map_ocr_params(&request.optional_params, &request.model)?;
let headers = self.validate_environment(&request.connection, &credential_env)?;
let url = self.get_complete_url(request.connection.api_base.as_deref())?;
let headers = BaseOcrConfig::validate_environment(self, request, client).await?;
let url = BaseOcrConfig::get_complete_url(self, request, &params, &headers)?;
let body = self
.async_transform_ocr_request(
&request.model,
@ -170,7 +195,11 @@ impl MistralOCRConfig {
.map(str::trim)
.filter(|key| !key.is_empty())
.map(str::to_string)
.or_else(|| env_lookup(MISTRAL_API_KEY_ENV).filter(|key| !key.trim().is_empty()))
.or_else(|| {
self.get_api_key_env_var()
.and_then(env_lookup)
.filter(|key| !key.trim().is_empty())
})
.ok_or(litellm_auth::Error::MissingApiKey {
provider: "Mistral",
environment_variable: MISTRAL_API_KEY_ENV,
@ -292,11 +321,7 @@ mod tests {
fn raw_response_transform_keeps_native_payload_separate_from_typed_normalization() {
let raw = br#"{"pages":[{"index":"2","markdown":"text"}],"provider_extension":false}"#;
let response = MistralOCRConfig
.decode_and_normalize_response(
"model",
raw,
crate::ocr::types::OcrResponseFormat::Native,
)
.transform_ocr_response("model", raw, crate::ocr::types::OcrResponseFormat::Native)
.unwrap();
assert_eq!(response.pages[0].index, 2);
let native = response.provider_native_response.unwrap();
@ -305,7 +330,7 @@ mod tests {
assert!(response.extra_fields.is_empty());
assert!(
MistralOCRConfig
.decode_and_normalize_response(
.transform_ocr_response(
"model",
br#"{"pages":[{"index":0}]}"#,
crate::ocr::types::OcrResponseFormat::Litellm

View file

@ -81,7 +81,37 @@ pub(crate) struct ReductoParseV3Config;
impl BaseOcrConfig for ReductoParseV3Config {
type OcrParams = ReductoV3Params;
type ProviderRequest = ReductoV3Request;
type ProviderResponse = ReductoResponse;
type Environment = Vec<(String, String)>;
async fn validate_environment(
&self,
request: &LiteLLMOcrRequest,
_client: &OcrClient,
) -> Result<Self::Environment, crate::ocr::Error> {
validate_environment(&request.connection, &credential_env)
}
fn get_complete_url(
&self,
request: &LiteLLMOcrRequest,
_params: &Self::OcrParams,
_environment: &Self::Environment,
) -> Result<String, crate::ocr::Error> {
get_complete_url(request.connection.api_base.as_deref())
}
fn transform_ocr_request(
&self,
_model: &str,
document: OcrDocument,
params: &Self::OcrParams,
_headers: &[(String, String)],
) -> Result<Self::ProviderRequest, crate::ocr::Error> {
Ok(ReductoV3Request {
input: uploaded_file_id(document)?,
params: params.clone(),
})
}
fn get_supported_ocr_params(&self, _model: &str) -> &'static [&'static str] {
&["formatting", "retrieval", "settings"]
@ -108,12 +138,18 @@ impl BaseOcrConfig for ReductoParseV3Config {
})
}
fn normalize_response(
fn transform_ocr_response(
&self,
model: &str,
response: ReductoResponse,
raw_response: &[u8],
request_format: crate::ocr::types::OcrResponseFormat,
) -> Result<LiteLLMOcrResponse, crate::ocr::Error> {
normalize_response(model, response)
crate::llms::base_llm::ocr::transformation::decode_and_normalize_response(
model,
raw_response,
request_format,
normalize_response,
)
}
}
@ -124,8 +160,8 @@ impl ReductoParseV3Config {
client: &OcrClient,
) -> Result<reqwest::Request, crate::ocr::Error> {
let params = self.map_ocr_params(&request.optional_params, &request.model)?;
let headers = validate_environment(&request.connection, &credential_env)?;
let url = get_complete_url(request.connection.api_base.as_deref())?;
let headers = self.validate_environment(request, client).await?;
let url = self.get_complete_url(request, &params, &headers)?;
let (document, headers) = guardrail_document(request, &url, &headers).await?;
let body = self
.async_transform_ocr_request(
@ -154,7 +190,36 @@ pub(crate) struct ReductoParseLegacyConfig;
impl BaseOcrConfig for ReductoParseLegacyConfig {
type OcrParams = ReductoLegacyParams;
type ProviderRequest = ReductoLegacyRequest;
type ProviderResponse = ReductoResponse;
type Environment = Vec<(String, String)>;
async fn validate_environment(
&self,
request: &LiteLLMOcrRequest,
client: &OcrClient,
) -> Result<Self::Environment, crate::ocr::Error> {
ReductoParseV3Config
.validate_environment(request, client)
.await
}
fn get_complete_url(
&self,
request: &LiteLLMOcrRequest,
params: &Self::OcrParams,
environment: &Self::Environment,
) -> Result<String, crate::ocr::Error> {
ReductoParseV3Config.get_complete_url(request, params, environment)
}
fn transform_ocr_request(
&self,
_model: &str,
document: OcrDocument,
params: &Self::OcrParams,
_headers: &[(String, String)],
) -> Result<Self::ProviderRequest, crate::ocr::Error> {
Ok(build_legacy_body(uploaded_file_id(document)?, params))
}
fn get_supported_ocr_params(&self, _model: &str) -> &'static [&'static str] {
&["enhance"]
@ -178,12 +243,13 @@ impl BaseOcrConfig for ReductoParseLegacyConfig {
Ok(build_legacy_body(file_id, optional_params))
}
fn normalize_response(
fn transform_ocr_response(
&self,
model: &str,
response: ReductoResponse,
raw_response: &[u8],
request_format: crate::ocr::types::OcrResponseFormat,
) -> Result<LiteLLMOcrResponse, crate::ocr::Error> {
normalize_response(model, response)
ReductoParseV3Config.transform_ocr_response(model, raw_response, request_format)
}
}
@ -194,8 +260,8 @@ impl ReductoParseLegacyConfig {
client: &OcrClient,
) -> Result<reqwest::Request, crate::ocr::Error> {
let params = self.map_ocr_params(&request.optional_params, &request.model)?;
let headers = validate_environment(&request.connection, &credential_env)?;
let url = get_complete_url(request.connection.api_base.as_deref())?;
let headers = self.validate_environment(request, client).await?;
let url = self.get_complete_url(request, &params, &headers)?;
let (document, headers) = guardrail_document(request, &url, &headers).await?;
let body = self
.async_transform_ocr_request(
@ -218,6 +284,21 @@ impl ReductoParseLegacyConfig {
}
}
fn uploaded_file_id(document: OcrDocument) -> Result<ReductoFileId, crate::ocr::Error> {
if !document.source().starts_with(REDUCTO_ID_PREFIX) {
return Err(crate::ocr::Error::ReductoSource);
}
if document.source()[REDUCTO_ID_PREFIX.len()..]
.trim()
.is_empty()
{
return Err(crate::ocr::Error::RequestField {
path: "document file id".into(),
});
}
Ok(ReductoFileId(document.source().into()))
}
fn present_nullable<'de, D: Deserializer<'de>, T: Deserialize<'de>>(
deserializer: D,
) -> Result<Option<Option<T>>, D::Error> {
@ -605,4 +686,335 @@ mod tests {
connection.extra_headers
);
}
use std::sync::Arc;
use rstest::rstest;
use crate::ocr::hooks::{OcrDuringCallRequest, OcrHookFuture, OcrHooks, OcrPostCallRequest};
use crate::ocr::test_support::{MockResponse, mock_server, perform_ocr, wire_request};
fn request_body(request: &str) -> Value {
serde_json::from_str(request.split_once("\r\n\r\n").unwrap().1).unwrap()
}
#[rstest]
#[case(
"reducto/parse-v3",
json!({
"formatting":{"table_output_format":"html"},
"retrieval":{"chunk_mode":"section"},
"settings":{"ocr_system":"standard"},
"future_ocr_option":true,
"extra_body":{"provider_option":"value"}
}),
"reducto://already.pdf",
json!({
"input":"reducto://already.pdf",
"formatting":{"table_output_format":"html"},
"retrieval":{"chunk_mode":"section"},
"settings":{"ocr_system":"standard"},
"future_ocr_option":true,
"provider_option":"value"
})
)]
#[case(
"reducto/parse-legacy",
json!({
"enhance":{"agentic":[{"type":"table"}]},
"future_ocr_option":true,
"extra_body":{"provider_option":"value"}
}),
"reducto://legacy.pdf",
json!({
"document_url":"reducto://legacy.pdf",
"options":{"enhance":{"agentic":[{"type":"table"}]}},
"future_ocr_option":true,
"provider_option":"value"
})
)]
#[tokio::test]
async fn request_mapping_matches_python(
#[case] model: &str,
#[case] options: Value,
#[case] source: &str,
#[case] expected: Value,
) {
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({
"result":{"chunks":[]}
}))])
.await;
let mut request = wire_request(model, &base, options);
request.document = request.document.with_source(source.into());
perform_ocr(request).await.unwrap();
server.await.unwrap();
let requests = seen.lock().unwrap();
assert_eq!(requests.len(), 1);
assert!(requests[0].starts_with("POST /parse "));
assert_eq!(request_body(&requests[0]), expected);
}
#[rstest]
#[case("parse-v3")]
#[case("parse-legacy")]
#[tokio::test]
async fn data_uri_upload_preserves_multipart_headers(#[case] model: &str) {
let (base, seen, server) = mock_server(vec![
MockResponse::json(json!({"file_id":"reducto://uploaded.pdf"})),
MockResponse::json(json!({"result":{"chunks":[{"content":"hello"}]}})),
])
.await;
let mut request = wire_request(&format!("reducto/{model}"), &base, json!({}));
request.connection.extra_headers = vec![
("Content-Type".into(), "application/json".into()),
("X-Trace".into(), "upload-test".into()),
];
let response = perform_ocr(request).await.unwrap();
server.await.unwrap();
assert_eq!(response.pages[0].markdown, "hello");
let requests = seen.lock().unwrap();
assert_eq!(requests.len(), 2);
assert!(requests[0].starts_with("POST /upload "));
assert!(
requests[0]
.to_ascii_lowercase()
.contains("content-type: multipart/form-data; boundary=")
);
assert!(requests[0].contains("x-trace: upload-test"));
assert!(requests[0].contains("application/pdf"));
assert!(requests[0].contains("abc"));
assert!(requests[1].starts_with("POST /parse "));
}
struct ParseBoundary {
request_count: Arc<std::sync::Mutex<Vec<String>>>,
}
impl OcrHooks for ParseBoundary {
fn post_call(&self, request: OcrPostCallRequest) -> OcrHookFuture<'_, OcrPostCallRequest> {
Box::pin(async move {
assert_eq!(self.request_count.lock().unwrap().len(), 2);
assert_eq!(
request.original_response,
json!(r#"{"result":{"chunks":[]}}"#)
);
Ok(request)
})
}
}
#[tokio::test]
async fn post_call_stays_after_reducto_upload_and_parse() {
let (base, seen, server) = mock_server(vec![
MockResponse::json(json!({"file_id":"reducto://uploaded.pdf"})),
MockResponse::json(json!({"result":{"chunks":[]}})),
])
.await;
let request = crate::ocr::LiteLLMOcrRequest {
hooks: Arc::new(ParseBoundary {
request_count: seen.clone(),
}),
..wire_request("reducto/parse-v3", &base, json!({}))
};
perform_ocr(request).await.unwrap();
server.await.unwrap();
assert_eq!(seen.lock().unwrap().len(), 2);
}
#[rstest]
#[case(json!({"file_id":""}))]
#[case(json!({}))]
#[case(json!({"file_id":null}))]
#[tokio::test]
async fn invalid_upload_ids_stop_before_parse(#[case] response: Value) {
let (base, seen, server) = mock_server(vec![MockResponse::json(response)]).await;
let error = perform_ocr(wire_request("reducto/parse-v3", &base, json!({})))
.await
.unwrap_err();
server.await.unwrap();
assert!(error.to_string().contains("file_id"));
assert_eq!(seen.lock().unwrap().len(), 1);
}
#[tokio::test]
async fn upload_failure_stops_before_parse() {
let (base, seen, server) = mock_server(vec![MockResponse {
status: 503,
headers: vec![],
body: json!({"error":"unavailable"}),
}])
.await;
assert!(
perform_ocr(wire_request("reducto/parse-v3", &base, json!({})))
.await
.is_err()
);
server.await.unwrap();
assert_eq!(seen.lock().unwrap().len(), 1);
}
#[rstest]
#[case("https://example.com/a.pdf")]
#[case("reducto://")]
#[case("data:application/pdf;base64")]
#[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());
assert!(perform_ocr(request).await.is_err());
}
#[test]
fn response_normalization_groups_blocks_and_distinguishes_null_result() {
use crate::llms::reducto::ocr::transformation::{ReductoResponse, normalize_response};
let raw = json!({"usage":{"num_pages":"2","credits":"3"},"result":{"type":"full","chunks":[
{"blocks":[{
"type":"Table",
"content":"B",
"bbox":{"left":0.1,"top":0.2,"width":0.8,"height":0.3,"page":2,"original_page":4},
"confidence":"high",
"granular_confidence":{"parse_confidence":0.95,"extract_confidence":null},
"image_url":null
}]},
{"blocks":[{"content":"A","bbox":{"page":1},"type":"Text"},{"content":"C","bbox":{"page":1}}]}
]}});
let response: ReductoResponse = serde_json::from_value(raw).unwrap();
let normalized = normalize_response("parse-v3", response)
.unwrap()
.into_json();
assert_eq!(normalized["pages"][0]["markdown"], "A\n\nC");
assert_eq!(normalized["pages"][1]["markdown"], "B");
assert_eq!(normalized["pages"][1]["blocks"][0]["type"], "Table");
assert_eq!(
normalized["pages"][1]["blocks"][0]["bbox"],
json!({"left":0.1,"top":0.2,"width":0.8,"height":0.3,"page":2,"original_page":4})
);
assert_eq!(normalized["pages"][1]["blocks"][0]["confidence"], "high");
assert_eq!(
normalized["pages"][1]["blocks"][0]["granular_confidence"]["parse_confidence"],
0.95
);
assert!(normalized["pages"][1]["blocks"][0]["image_url"].is_null());
assert_eq!(normalized["usage_info"]["pages_processed"], 2);
assert_eq!(normalized["usage_info"]["credits"], 3.0);
let missing: ReductoResponse =
serde_json::from_value(json!({"chunks":[{"content":"text"}]})).unwrap();
let missing = normalize_response("parse-v3", missing).unwrap();
assert_eq!(missing.pages[0].markdown, "text");
let null: ReductoResponse = serde_json::from_value(
json!({"result":null,"chunks":[{"content":"ignored"}],"usage":null}),
)
.unwrap();
let null = normalize_response("parse-v3", null).unwrap();
assert!(null.pages.is_empty());
}
#[tokio::test]
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());
request.connection.extra_headers = vec![("authorization".into(), "Bearer existing".into())];
let response = perform_ocr(request).await.unwrap();
server.await.unwrap();
assert_eq!(response.provider_native_response, None);
assert!(
seen.lock().unwrap()[0]
.to_ascii_lowercase()
.contains("authorization: bearer existing")
);
}
struct RewriteDocument;
struct RewriteHeaders;
impl OcrHooks for RewriteHeaders {
fn intercepts_requests(&self) -> bool {
true
}
fn during_call(
&self,
request: OcrDuringCallRequest,
) -> OcrHookFuture<'_, OcrDuringCallRequest> {
Box::pin(async move {
Ok(OcrDuringCallRequest {
headers: vec![("authorization".into(), "Bearer guarded".into())],
..request
})
})
}
}
#[rstest]
#[case("reducto/parse-v3")]
#[case("reducto/parse-legacy")]
#[tokio::test]
async fn guardrail_headers_reach_upload_and_parse(#[case] model: &str) {
let (base, seen, server) = mock_server(vec![
MockResponse::json(json!({"file_id":"reducto://uploaded.pdf"})),
MockResponse::json(json!({"result":{"chunks":[]}})),
])
.await;
let mut request = wire_request(model, &base, json!({}));
request.connection.extra_headers = vec![("authorization".into(), "Bearer original".into())];
request.hooks = Arc::new(RewriteHeaders);
perform_ocr(request).await.unwrap();
server.await.unwrap();
let requests = seen.lock().unwrap();
assert_eq!(requests.len(), 2);
assert!(requests[0].starts_with("POST /upload "));
assert!(requests[1].starts_with("POST /parse "));
for request in requests.iter() {
assert!(request.contains("authorization: Bearer guarded"));
assert!(!request.contains("Bearer original"));
}
}
impl OcrHooks for RewriteDocument {
fn intercepts_requests(&self) -> bool {
true
}
fn during_call(
&self,
request: OcrDuringCallRequest,
) -> OcrHookFuture<'_, OcrDuringCallRequest> {
Box::pin(async move {
assert_eq!(
request.body["document_url"],
"data:application/pdf;base64,YWJj"
);
Ok(OcrDuringCallRequest {
body: json!({"type":"document_url","document_url":"reducto://guarded.pdf"}),
..request
})
})
}
}
#[tokio::test]
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!({}));
request.hooks = Arc::new(RewriteDocument);
perform_ocr(request).await.unwrap();
server.await.unwrap();
let requests = seen.lock().unwrap();
assert_eq!(requests.len(), 1);
assert!(requests[0].starts_with("POST /parse "));
assert!(requests[0].contains("reducto://guarded.pdf"));
}
}

View file

@ -100,7 +100,38 @@ pub(crate) struct VertexAIDeepSeekOCRConfig;
impl BaseOcrConfig for VertexAIDeepSeekOCRConfig {
type OcrParams = DeepSeekOcrParams;
type ProviderRequest = DeepSeekOcrRequest;
type ProviderResponse = DeepSeekOcrResponse;
type Environment = vertex::VertexEnvironment;
fn get_api_key_env_var(&self) -> Option<&'static str> {
VertexAIOCRConfig.get_api_key_env_var()
}
async fn validate_environment(
&self,
request: &LiteLLMOcrRequest,
client: &OcrClient,
) -> Result<Self::Environment, crate::ocr::Error> {
BaseOcrConfig::validate_environment(&VertexAIOCRConfig, request, client).await
}
fn get_complete_url(
&self,
request: &LiteLLMOcrRequest,
_params: &Self::OcrParams,
environment: &Self::Environment,
) -> Result<String, crate::ocr::Error> {
let config = VertexConfig::from_sourced_optional_params(
&request.optional_params,
&request.input_sources,
)?;
let location = vertex::get_vertex_ai_location(&config, &credential_env)
.unwrap_or_else(|| DEFAULT_LOCATION.to_string());
self.get_complete_url(
request.connection.api_base.as_deref(),
&environment.project_id,
&location,
)
}
async fn async_transform_ocr_request(
&self,
@ -113,17 +144,21 @@ impl BaseOcrConfig for VertexAIDeepSeekOCRConfig {
self.transform_ocr_request(model, document, optional_params, headers)
}
fn normalize_response(
fn transform_ocr_response(
&self,
model: &str,
response: DeepSeekOcrResponse,
raw_response: &[u8],
request_format: crate::ocr::types::OcrResponseFormat,
) -> Result<LiteLLMOcrResponse, crate::ocr::Error> {
normalize_response(model, response)
crate::llms::base_llm::ocr::transformation::decode_and_normalize_response(
model,
raw_response,
request_format,
normalize_response,
)
}
}
impl VertexAIDeepSeekOCRConfig {
pub(crate) fn transform_ocr_request(
fn transform_ocr_request(
&self,
model: &str,
document: OcrDocument,
@ -148,28 +183,18 @@ impl VertexAIDeepSeekOCRConfig {
.collect(),
})
}
}
impl VertexAIDeepSeekOCRConfig {
pub(crate) async fn prepare_request(
&self,
request: &LiteLLMOcrRequest,
client: &OcrClient,
) -> Result<reqwest::Request, crate::ocr::Error> {
let params = self.map_ocr_params(&request.optional_params, &request.model)?;
let config = VertexConfig::from_sourced_optional_params(
&request.optional_params,
&request.input_sources,
)
.map_err(crate::ocr::Error::from)?;
let authentication = VertexAIOCRConfig
.validate_environment(&request.connection, &config, client)
.await?;
let location = vertex::get_vertex_ai_location(&config, &credential_env)
.unwrap_or_else(|| DEFAULT_LOCATION.to_string());
let url = self.get_complete_url(
request.connection.api_base.as_deref(),
&authentication.project_id,
&location,
)?;
let authentication = self.validate_environment(request, client).await?;
let url = BaseOcrConfig::get_complete_url(self, request, &params, &authentication)?;
let body = self
.async_transform_ocr_request(
&request.model,
@ -393,7 +418,11 @@ impl VertexAIDeepSeekOCRConfig {
#[cfg(test)]
mod tests {
use super::{VertexAIDeepSeekOCRConfig, provider_model};
use super::{
DeepSeekOcrParams, DeepSeekOcrResponse, VertexAIDeepSeekOCRConfig, normalize_response,
provider_model,
};
use serde_json::{Value, json};
#[test]
fn unconsumed_options_remain_available_for_body_composition() {
@ -437,4 +466,247 @@ mod tests {
"https://aiplatform.googleapis.com/v1/projects/proj-1/locations/europe-west4/endpoints/openapi/chat/completions"
);
}
use rstest::rstest;
use crate::llms::base_llm::ocr::transformation::BaseOcrConfig;
use crate::ocr::types::OcrDocument;
fn document() -> OcrDocument {
serde_json::from_value(json!({"type":"image_url","image_url":"gs://bucket/a.png"})).unwrap()
}
#[rstest]
#[case("stream", json!(true))]
#[case("temperature", json!(0.1))]
#[case("max_tokens", json!(1024))]
#[case("top_p", json!(0.9))]
#[case("n", json!(2))]
#[case("stop", json!("done"))]
#[case("stop", json!(["done", "stop"]))]
#[case("temperature", json!(null))]
fn request_mapping_matches_python(#[case] name: &str, #[case] value: Value) {
let params: DeepSeekOcrParams =
serde_json::from_value(json!({name: value.clone(), "ignored": true})).unwrap();
let result = serde_json::to_value(
VertexAIDeepSeekOCRConfig
.transform_ocr_request("deepseek-ai/deepseek-ocr-maas", document(), &params, &[])
.unwrap(),
)
.unwrap();
assert_eq!(result["model"], "deepseek-ai/deepseek-ocr-maas");
assert_eq!(
result["messages"][0]["content"][0],
json!({"type":"image_url","image_url":"gs://bucket/a.png"})
);
assert_eq!(result[name], value);
assert!(result.get("ignored").is_none());
}
#[rstest]
#[case(json!({"type":"image_url","image_url":"data:image/png;base64,AA=="}))]
#[case(json!({"type":"document_url","document_url":"data:application/pdf;base64,AA=="}))]
fn request_maps_both_document_types_to_image_content(#[case] document: Value) {
let source = document
.get("image_url")
.or_else(|| document.get("document_url"))
.unwrap()
.clone();
let request = VertexAIDeepSeekOCRConfig
.transform_ocr_request(
"deepseek-ai/deepseek-ocr-maas",
serde_json::from_value(document).unwrap(),
&DeepSeekOcrParams::default(),
&[],
)
.unwrap();
let result = serde_json::to_value(request).unwrap();
assert_eq!(
result["messages"][0]["content"][0],
json!({"type":"image_url","image_url":source})
);
}
#[rstest]
#[case(json!("# hello"), "# hello")]
#[case(json!("{broken"), "{broken")]
#[case(json!(" {\"pages\":[]} "), " {\"pages\":[]} ")]
#[case(json!({"pages":[]}), "")]
#[case(json!("[]"), "[]")]
#[case(json!("{\"pages\":[{\"markdown\":\"json text\"}]}"), "json text")]
#[case(json!({"pages":[{"markdown":"object"}]}), "object")]
fn response_transform_handles_text_json_and_objects(
#[case] content: Value,
#[case] expected: &str,
) {
let has_pages = content
.as_object()
.is_some_and(|data| data.contains_key("pages"))
|| content
.as_str()
.is_some_and(|text| text.contains("\"pages\""));
let response: DeepSeekOcrResponse = serde_json::from_value(
json!({"choices":[{"message":{"content":content}}],"usage":{"prompt_tokens":1}}),
)
.unwrap();
let result = normalize_response("model", response).unwrap().into_json();
assert_eq!(result["pages"][0]["markdown"], expected);
assert_eq!(result["pages"][0]["index"], 0);
if has_pages {
assert!(result["usage_info"].is_null());
} else {
assert_eq!(result["usage_info"]["prompt_tokens"], 1);
}
}
#[test]
fn structured_result_maps_pages_usage_model_and_annotation() {
let response: DeepSeekOcrResponse = serde_json::from_value(json!({
"choices":[{"message":{"content":{
"pages":[{"index":2,"markdown":"page","images":[{"id":"one"}],"dimensions":{"width":10}}],
"model":"provider-model",
"usage_info":{"pages_processed":1},
"document_annotation":{"language":"en"},
"future":"kept"
}}}]
}))
.unwrap();
let result = normalize_response("requested", response)
.unwrap()
.into_json();
assert_eq!(result["pages"][0]["index"], 2);
assert_eq!(result["pages"][0]["images"][0]["id"], "one");
assert_eq!(result["model"], "provider-model");
assert_eq!(result["usage_info"]["pages_processed"], 1);
assert_eq!(result["document_annotation"]["language"], "en");
assert!(result.get("future").is_none());
}
#[test]
fn response_transform_rejects_missing_empty_and_malformed_content() {
for value in [
json!({"choices":[]}),
json!({"choices":[{"message":{"content":{}}}]}),
json!({"choices":[{"message":{"content":""}}]}),
json!({"choices":[{"message":{"content":"{\"pages\":[{\"markdown\":42}]}"}}]}),
json!({"choices":[{"message":{"content":{"pages":[{"markdown":42}]}}}]}),
] {
let result = serde_json::from_value::<DeepSeekOcrResponse>(value)
.map_err(|_| ())
.and_then(|response| normalize_response("model", response).map_err(|_| ()));
assert!(result.is_err());
}
}
#[test]
fn structured_content_preserves_usage_presence_and_shared_page_defaults() {
for (usage, expected) in [(json!(null), None), (json!({"pages_processed":2}), Some(2))] {
let response = serde_json::from_value(json!({
"choices":[{"message":{"content":{
"pages":[42, {"index":"2", "images":[{"id":"kept"}], "ignored":true}],
"usage_info":usage
}}}],
"usage":{"pages_processed":99}
}))
.unwrap();
let normalized = normalize_response("model", response).unwrap();
assert_eq!(normalized.pages.len(), 1);
assert_eq!(normalized.pages[0].index, 2);
assert_eq!(normalized.pages[0].markdown, "");
assert!(normalized.pages[0].extra_fields.is_empty());
assert_eq!(
normalized
.usage_info
.and_then(|usage| usage.pages_processed),
expected
);
}
}
use crate::ocr::test_support::{MockResponse, mock_server, perform_ocr, wire_request};
use litellm_auth::InputSource;
fn request_body(request: &str) -> Value {
serde_json::from_str(request.split_once("\r\n\r\n").unwrap().1).unwrap()
}
#[tokio::test]
async fn facade_executes_vertex_deepseek_at_the_openai_endpoint() {
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({
"choices":[{"message":{"content":"recognized"}}],
"usage":{"prompt_tokens":1}
}))])
.await;
let mut request = wire_request(
"vertex_ai/deepseek-ocr-maas",
&base,
json!({
"vertex_project":"project-1",
"vertex_location":"europe-west4",
"temperature":0.1,
"future_ocr_option":true,
"extra_body":{"provider_option":"value"}
}),
);
request.document = request
.document
.with_source("gs://bucket/document.pdf".into());
let response = perform_ocr(request).await.unwrap();
server.await.unwrap();
assert_eq!(response.pages[0].markdown, "recognized");
assert_eq!(
response.usage_info.unwrap().extra_fields["prompt_tokens"],
1
);
let requests = seen.lock().unwrap();
assert!(requests[0].starts_with(
"POST /v1/projects/project-1/locations/europe-west4/endpoints/openapi/chat/completions "
));
assert!(
requests[0]
.to_ascii_lowercase()
.contains("authorization: bearer test-key")
);
let body = request_body(&requests[0]);
assert_eq!(body["model"], "deepseek-ai/deepseek-ocr-maas");
assert_eq!(body["temperature"], 0.1);
assert_eq!(body["future_ocr_option"], true);
assert_eq!(body["provider_option"], "value");
assert!(body.get("vertex_project").is_none());
assert!(body.get("extra_body").is_none());
assert_eq!(
body["messages"][0]["content"][0],
json!({"type":"image_url","image_url":"gs://bucket/document.pdf"})
);
}
#[test]
fn host_registration_selects_deepseek_without_affecting_mistral() {
assert!(crate::ocr::wire::is_supported_request(
"deepseek-ocr-maas",
Some("vertex_ai")
));
assert!(crate::ocr::wire::is_supported_request(
"mistral-ocr-maas",
Some("vertex_ai")
));
}
#[tokio::test]
async fn request_controlled_api_base_is_rejected_before_vertex_auth() {
let mut request = wire_request(
"vertex_ai/deepseek-ocr-maas",
"https://caller.example",
json!({"vertex_project":"project-1"}),
);
request.connection.api_base_source = InputSource::Request;
let error = perform_ocr(request).await.unwrap_err();
assert!(
error
.to_string()
.contains("request-controlled Vertex AI endpoint")
);
}
}

View file

@ -1,6 +1,7 @@
use litellm_auth_gcp::{self as vertex, VertexConfig};
use super::common_utils::validate_destination;
use crate::llms::base_llm::ocr::transformation::{BaseOcrConfig, OcrRequestContext};
use crate::llms::mistral::ocr::MistralOcrResponse;
use crate::llms::mistral::ocr::transformation::{MistralOCRConfig, MistralOcrRequest};
use crate::ocr::OcrClient;
use crate::ocr::document::{inline_remote_document, validate_inline_document};
@ -8,7 +9,7 @@ use crate::ocr::prepare::{credential_env, transform_request_body};
use crate::ocr::types::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrConnection, OcrDocument};
use crate::params::OpaqueParams;
use crate::url_utils::ApiUrl;
use litellm_auth_gcp::{self as vertex, VertexConfig};
const DEFAULT_LOCATION: &str = "us-central1";
#[derive(Clone, Debug, Default)]
@ -17,7 +18,54 @@ pub(crate) struct VertexAIOCRConfig;
impl BaseOcrConfig for VertexAIOCRConfig {
type OcrParams = OpaqueParams;
type ProviderRequest = MistralOcrRequest;
type ProviderResponse = MistralOcrResponse;
type Environment = vertex::VertexEnvironment;
fn get_api_key_env_var(&self) -> Option<&'static str> {
Some("VERTEX_AI_API_KEY")
}
async fn validate_environment(
&self,
request: &LiteLLMOcrRequest,
client: &OcrClient,
) -> Result<Self::Environment, crate::ocr::Error> {
let config = VertexConfig::from_sourced_optional_params(
&request.optional_params,
&request.input_sources,
)?;
self.validate_environment(&request.connection, &config, client)
.await
}
fn get_complete_url(
&self,
request: &LiteLLMOcrRequest,
_params: &Self::OcrParams,
environment: &Self::Environment,
) -> Result<String, crate::ocr::Error> {
let config = VertexConfig::from_sourced_optional_params(
&request.optional_params,
&request.input_sources,
)?;
let location = vertex::get_vertex_ai_location(&config, &credential_env)
.unwrap_or_else(|| DEFAULT_LOCATION.to_string());
self.get_complete_url(
request.connection.api_base.as_deref(),
&environment.project_id,
&location,
&request.model,
)
}
fn transform_ocr_request(
&self,
model: &str,
document: OcrDocument,
params: &OpaqueParams,
headers: &[(String, String)],
) -> Result<MistralOcrRequest, crate::ocr::Error> {
MistralOCRConfig.transform_ocr_request(model, document, params, headers)
}
fn get_supported_ocr_params(&self, model: &str) -> &'static [&'static str] {
MistralOCRConfig.get_supported_ocr_params(model)
@ -37,15 +85,16 @@ impl BaseOcrConfig for VertexAIOCRConfig {
context.connection,
)
.await?;
MistralOCRConfig.transform_ocr_request(model, document, optional_params, headers)
self.transform_ocr_request(model, document, optional_params, headers)
}
fn normalize_response(
fn transform_ocr_response(
&self,
model: &str,
response: MistralOcrResponse,
raw_response: &[u8],
request_format: crate::ocr::types::OcrResponseFormat,
) -> Result<LiteLLMOcrResponse, crate::ocr::Error> {
MistralOCRConfig.normalize_response(model, response)
MistralOCRConfig.transform_ocr_response(model, raw_response, request_format)
}
}
@ -56,22 +105,8 @@ impl VertexAIOCRConfig {
client: &OcrClient,
) -> Result<reqwest::Request, crate::ocr::Error> {
let params = self.map_ocr_params(&request.optional_params, &request.model)?;
let config = VertexConfig::from_sourced_optional_params(
&request.optional_params,
&request.input_sources,
)
.map_err(crate::ocr::Error::from)?;
let authentication = self
.validate_environment(&request.connection, &config, client)
.await?;
let location = vertex::get_vertex_ai_location(&config, &credential_env)
.unwrap_or_else(|| DEFAULT_LOCATION.to_string());
let url = self.get_complete_url(
request.connection.api_base.as_deref(),
&authentication.project_id,
&location,
&request.model,
)?;
let authentication = BaseOcrConfig::validate_environment(self, request, client).await?;
let url = BaseOcrConfig::get_complete_url(self, request, &params, &authentication)?;
let retains_document = !request.document.source().starts_with("http://")
&& !request.document.source().starts_with("https://");
let body = self
@ -193,4 +228,173 @@ mod tests {
.is_err()
);
}
use serde_json::{Value, json};
use crate::ocr::test_support::{MockResponse, mock_server, perform_ocr, wire_request};
use litellm_auth::InputSource;
fn request_body(request: &str) -> Value {
serde_json::from_str(request.split_once("\r\n\r\n").unwrap().1).unwrap()
}
#[tokio::test]
async fn facade_executes_vertex_mistral_with_resolved_project_and_location() {
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({
"pages":[{"index":0,"markdown":"hello"}],
"usage_info":{"pages_processed":1}
}))])
.await;
let request = wire_request(
"vertex_ai/mistral-ocr-maas",
&base,
json!({
"vertex_project":"project-1",
"vertex_location":"europe-west4",
"extract_footer":true
}),
);
let response = perform_ocr(request).await.unwrap();
server.await.unwrap();
assert_eq!(response.pages[0].markdown, "hello");
let requests = seen.lock().unwrap();
assert_eq!(requests.len(), 1);
assert!(requests[0].starts_with(
"POST /v1/projects/project-1/locations/europe-west4/publishers/mistralai/models/mistral-ocr-maas:rawPredict "
));
assert!(
requests[0]
.to_ascii_lowercase()
.contains("authorization: bearer test-key")
);
assert_eq!(
request_body(&requests[0]),
json!({
"model":"mistral-ocr-maas",
"document":{"type":"document_url","document_url":"data:application/pdf;base64,YWJj"},
"extract_footer":true
})
);
}
#[tokio::test]
async fn supplied_authorization_is_forwarded_without_a_static_token() {
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({"pages":[]}))]).await;
let mut request = wire_request(
"vertex_ai/model",
&base,
json!({"vertex_project":"project-1"}),
);
request.connection.api_key = None;
request.connection.extra_headers = vec![("authorization".into(), "Bearer supplied".into())];
perform_ocr(request).await.unwrap();
server.await.unwrap();
assert!(
seen.lock().unwrap()[0]
.to_ascii_lowercase()
.contains("authorization: bearer supplied")
);
}
#[tokio::test]
async fn invalid_credentials_fail_before_provider_http() {
let request = wire_request(
"vertex_ai/model",
"http://127.0.0.1:1",
json!({"vertex_credentials": true}),
);
let error = perform_ocr(request).await.unwrap_err();
assert!(error.to_string().contains("vertex_credentials"));
}
#[tokio::test]
async fn request_controlled_api_base_is_rejected_before_vertex_auth() {
let mut request = wire_request(
"vertex_ai/mistral-ocr-maas",
"https://caller.example",
json!({"vertex_project":"project-1"}),
);
request.connection.api_base_source = InputSource::Request;
let error = perform_ocr(request).await.unwrap_err();
assert!(
error
.to_string()
.contains("request-controlled Vertex AI endpoint")
);
}
#[tokio::test]
async fn configs_build_complete_requests_and_share_mistral_normalization() {
use std::time::Duration;
use crate::llms::base_llm::ocr::transformation::BaseOcrConfig;
use crate::llms::mistral::ocr::transformation::MistralOCRConfig;
use crate::llms::vertex_ai::ocr::transformation::VertexAIOCRConfig;
use crate::ocr::test_support::ocr_client;
let client = ocr_client();
let options = json!({
"pages": [0, 2],
"include_image_base64": true,
"vertex_project": "project-1",
"vertex_location": "us-central1",
"unknown": "preserved"
});
let direct = wire_request(
"mistral/mistral-ocr-maas",
"https://mistral.test",
options.clone(),
);
let vertex = wire_request("vertex_ai/mistral-ocr-maas", "https://vertex.test", options);
let direct_http = MistralOCRConfig
.prepare_request(&direct, &client)
.await
.unwrap();
let vertex_http = VertexAIOCRConfig
.prepare_request(&vertex, &client)
.await
.unwrap();
assert_eq!(direct_http.url().as_str(), "https://mistral.test/v1/ocr");
assert_eq!(
vertex_http.url().as_str(),
"https://vertex.test/v1/projects/project-1/locations/us-central1/publishers/mistralai/models/mistral-ocr-maas:rawPredict"
);
for http in [&direct_http, &vertex_http] {
assert_eq!(http.method(), reqwest::Method::POST);
assert_eq!(http.headers()["authorization"], "Bearer test-key");
assert_eq!(http.headers()["content-type"], "application/json");
assert_eq!(http.timeout(), Some(&Duration::from_secs(2)));
let body: Value =
serde_json::from_slice(http.body().unwrap().as_bytes().unwrap()).unwrap();
assert_eq!(
body,
json!({
"model": "mistral-ocr-maas",
"document": {"type": "document_url", "document_url": "data:application/pdf;base64,YWJj"},
"pages": [0, 2],
"include_image_base64": true,
"unknown": "preserved"
})
);
}
let payload = serde_json::to_vec(
&json!({"pages": [{"index": 0, "markdown": "hello"}], "extra": "preserved"}),
)
.unwrap();
let direct_response = MistralOCRConfig
.transform_ocr_response(&direct.model, &payload, Default::default())
.unwrap()
.into_json();
let vertex_response = VertexAIOCRConfig
.transform_ocr_response(&vertex.model, &payload, Default::default())
.unwrap()
.into_json();
assert_eq!(direct_response, vertex_response);
assert_eq!(direct_response["model"], "mistral-ocr-maas");
assert_eq!(direct_response["object"], "ocr");
assert!(direct_response.get("extra").is_none());
}
}

View file

@ -1,5 +1,11 @@
#[derive(Clone, Debug, PartialEq, Eq, thiserror::Error)]
pub enum Error {
#[error("upstream OCR error ({status}): {body}")]
Provider {
status: u16,
body: String,
headers: Vec<(String, String)>,
},
#[error("File is empty or could not be read")]
EmptyFile,
#[error("Invalid MIME type: {0}")]

View file

@ -23,7 +23,7 @@ pub(crate) async fn perform_ocr_request(
let context = CallLifecycleContext::new(
"ocr",
request.model.clone(),
request.config.provider().as_str(),
request.provider_name(),
request
.litellm_call_id
.clone()
@ -54,6 +54,7 @@ impl PreparedOcrCall {
client: OcrClient,
request: LiteLLMOcrRequest,
) -> Result<Self, super::Error> {
let request = super::prepare::resolve_connection_params(request);
let http = match request.config {
OcrConfigKind::Cohere => CohereParseConfig.prepare_request(&request, &client).await?,
OcrConfigKind::Mistral => MistralOCRConfig.prepare_request(&request, &client).await?,
@ -99,6 +100,30 @@ impl PreparedOcrCall {
crate::http_utils::execute_http_request(self.client.provider_http(), self.http)
.await
.map_err(super::client::transport_error)?;
if !response.status().is_success() {
let headers = response
.headers()
.iter()
.filter_map(|(name, value)| {
value
.to_str()
.ok()
.map(|value| (name.to_string(), value.to_string()))
})
.collect();
return match super::client::read_response_bytes(
response,
self.request.connection.max_response_bytes,
)
.await
{
Err(super::Error::Transport(crate::transport::Error::Http { status, body })) => {
Err(self.request.config.get_error_class(body, status, headers))
}
Err(error) => Err(error),
Ok(_) => unreachable!("non-success response produces an HTTP error"),
};
}
let model = &self.request.model;
let context = OcrResponseContext {
client: &self.client,

View file

@ -652,3 +652,966 @@ impl OcrHost for OcrHookHost {
})
}
}
#[cfg(test)]
mod tests {
use std::sync::{Arc, Mutex};
use serde_json::{Value, json};
use crate::call_lifecycle::{CallLifecycleContext, CallLifecycleTiming};
use crate::ocr::OcrClient;
use crate::ocr::hooks::{
OcrDuringCallRequest, OcrHookFuture, OcrHooks, OcrLogFuture, OcrPostCallRequest,
OcrPreCallRequest,
};
use crate::ocr::test_support::{MockResponse, mock_server, perform_ocr, wire_request};
use crate::ocr::wire::{OcrWireRequest, decode_request};
use crate::ocr::{
NativeOutcome, NoopOcrHost, OcrAdmission, OcrCall, OcrCallStep, OcrDecline, OcrHost,
OcrHostOperation, OcrHostResult,
};
#[test]
fn request_boundary_selects_mistral_and_rejects_unknown_providers() {
let request = OcrWireRequest {
model: "mistral/model".into(),
document: json!({"type":"document_url","document_url":"https://example.com/doc.pdf"}),
api_key: Some("key".into()),
api_base: None,
custom_llm_provider: None,
extra_headers: None,
optional_params: json!({"extract_header":true,"unknown":42})
.as_object()
.unwrap()
.clone()
.into(),
input_sources: Default::default(),
timeout_seconds: None,
};
assert!(decode_request(request).is_ok());
assert!(
decode_request(OcrWireRequest {
model: "model".into(),
document: json!({"type":"document_url","document_url":"https://example.com/doc.pdf"}),
api_key: Some("key".into()),
api_base: None,
custom_llm_provider: Some("unknown".into()),
extra_headers: None,
optional_params: Default::default(),
input_sources: Default::default(),
timeout_seconds: None,
})
.is_err()
);
}
#[tokio::test]
async fn facade_executes_direct_mistral_once() {
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({
"pages":[{"index":0,"markdown":"hello","custom":"preserved"}],
"usage_info":{"pages_processed":1}
}))])
.await;
let result = perform_ocr(wire_request(
"mistral/model",
&base,
json!({"pages":"0,2-4","extract_header":true,"unknown":{"nested":[null,false,0]}}),
))
.await
.unwrap();
server.await.unwrap();
assert_eq!(result.pages[0].markdown, "hello");
assert_eq!(result.pages[0].extra_fields["custom"], "preserved");
let requests = seen.lock().unwrap();
assert_eq!(requests.len(), 1);
assert!(requests[0].starts_with("POST /v1/ocr "));
assert!(
requests[0]
.to_ascii_lowercase()
.contains("authorization: bearer test-key\r\n")
);
let body: Value =
serde_json::from_str(requests[0].split_once("\r\n\r\n").unwrap().1).unwrap();
assert_eq!(
body,
json!({
"model":"model",
"document":{"type":"document_url","document_url":"data:application/pdf;base64,YWJj"},
"pages":"0,2-4",
"extract_header":true,
"unknown":{"nested":[null,false,0]}
})
);
}
#[tokio::test]
async fn facade_resolves_dynamic_connection_before_auth_and_url_preparation() {
use litellm_auth::{InputSource, Sourced};
let (base, seen, server) =
mock_server(vec![MockResponse::json(json!({"pages": []}))]).await;
let request = wire_request("mistral/model", "https://unused.invalid", json!({}));
let request = crate::ocr::LiteLLMOcrRequest {
connection: crate::ocr::OcrConnection {
dynamic_api_key: Some(Sourced::new("dynamic-key".into(), InputSource::Deployment)),
dynamic_api_base: Some(Sourced::new(base, InputSource::Deployment)),
..request.connection
},
..request
};
perform_ocr(request).await.unwrap();
server.await.unwrap();
let requests = seen.lock().unwrap();
assert_eq!(requests.len(), 1);
assert!(requests[0].starts_with("POST /v1/ocr "));
assert!(
requests[0]
.to_ascii_lowercase()
.contains("authorization: bearer dynamic-key\r\n")
);
}
#[tokio::test]
async fn provider_error_factory_preserves_status_body_and_response_headers() {
let (base, seen, server) = mock_server(vec![MockResponse {
status: 429,
headers: vec![
("retry-after", "17".into()),
("x-request-id", "ocr-request".into()),
],
body: json!({"message": "rate limited"}),
}])
.await;
let error = perform_ocr(wire_request("mistral/model", &base, json!({})))
.await
.unwrap_err();
server.await.unwrap();
let crate::ocr::Error::Provider {
status,
body,
headers,
} = error
else {
panic!("expected provider error")
};
assert_eq!(status, 429);
assert_eq!(
serde_json::from_str::<Value>(&body).unwrap(),
json!({"message": "rate limited"})
);
assert!(headers.contains(&("retry-after".into(), "17".into())));
assert!(headers.contains(&("x-request-id".into(), "ocr-request".into())));
assert_eq!(seen.lock().unwrap().len(), 1);
}
#[tokio::test]
async fn facade_retains_native_response_when_requested() {
let provider_response = json!({
"pages":[{"index":0,"markdown":"hello"}],
"usage_info":{"pages_processed":1},
"provider_only":"preserved"
});
let (base, _, server) =
mock_server(vec![MockResponse::json(provider_response.clone())]).await;
let response = perform_ocr(wire_request(
"mistral/model",
&base,
json!({"req_format":"native"}),
))
.await
.unwrap();
server.await.unwrap();
assert_eq!(
response.provider_native_response.as_ref(),
provider_response.as_object()
);
}
#[tokio::test]
async fn facade_uses_the_injected_http_client() {
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({"pages":[]}))]).await;
let mut default_headers = reqwest::header::HeaderMap::new();
default_headers.insert(
"x-transport-owner",
reqwest::header::HeaderValue::from_static("host"),
);
let provider_http = reqwest::Client::builder()
.default_headers(default_headers)
.build()
.unwrap();
OcrClient::new(provider_http)
.unwrap()
.perform(wire_request("mistral/model", &base, json!({})))
.await
.unwrap();
server.await.unwrap();
assert!(seen.lock().unwrap()[0].contains("x-transport-owner: host"));
}
struct RecordingHooks {
events: Arc<Mutex<Vec<&'static str>>>,
block: bool,
}
struct ExtensionHooks;
impl OcrHooks for ExtensionHooks {
fn intercepts_requests(&self) -> bool {
true
}
fn during_call(
&self,
mut request: OcrDuringCallRequest,
) -> OcrHookFuture<'_, OcrDuringCallRequest> {
Box::pin(async move {
assert_eq!(request.body["pages"], json!([2]));
assert_eq!(request.body.get("future"), Some(&Value::Null));
assert!(
!request
.retained_fields
.iter()
.any(|field| field == "pages" || field == "document")
);
request.body.as_object_mut().unwrap().remove("future");
request.body["hook_option"] = json!({"nested":[null,false,0]});
Ok(request)
})
}
}
#[tokio::test]
async fn composed_extensions_reach_hooks_and_removed_fields_stay_removed() {
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({"pages":[]}))]).await;
let request = crate::ocr::LiteLLMOcrRequest {
hooks: Arc::new(ExtensionHooks),
..wire_request(
"mistral/model",
&base,
json!({
"pages":[0], "future":null, "extra_body":{"pages":[2],
"document":{"type":"document_url","document_url":"data:application/pdf;base64,eHl6"}}
}),
)
};
perform_ocr(request).await.unwrap();
server.await.unwrap();
let requests = seen.lock().unwrap();
let body: Value =
serde_json::from_str(requests[0].split_once("\r\n\r\n").unwrap().1).unwrap();
assert_eq!(body["pages"], json!([2]));
assert_eq!(
body["document"]["document_url"],
"data:application/pdf;base64,eHl6"
);
assert_eq!(body["hook_option"], json!({"nested":[null,false,0]}));
assert!(body.get("future").is_none());
}
impl OcrHooks for RecordingHooks {
fn intercepts_requests(&self) -> bool {
true
}
fn pre_call(&self, request: OcrPreCallRequest) -> OcrHookFuture<'_, OcrPreCallRequest> {
Box::pin(async move {
self.events.lock().unwrap().push("pre");
if self.block {
return Err(crate::ocr::Error::InvalidRequest("blocked".into()));
}
Ok(request)
})
}
fn during_call(
&self,
request: crate::ocr::hooks::OcrDuringCallRequest,
) -> OcrHookFuture<'_, crate::ocr::hooks::OcrDuringCallRequest> {
Box::pin(async move {
self.events.lock().unwrap().push("during");
Ok(request)
})
}
fn post_call(&self, request: OcrPostCallRequest) -> OcrHookFuture<'_, OcrPostCallRequest> {
Box::pin(async move {
self.events.lock().unwrap().push("post");
Ok(request)
})
}
fn success<'a>(
&'a self,
_context: &'a CallLifecycleContext,
_response: &'a crate::ocr::LiteLLMOcrResponse,
_timing: &'a CallLifecycleTiming,
) -> OcrLogFuture<'a> {
Box::pin(async move {
self.events.lock().unwrap().push("success");
})
}
fn failure<'a>(
&'a self,
_context: &'a CallLifecycleContext,
_error: &'a crate::ocr::Error,
_timing: &'a CallLifecycleTiming,
) -> OcrLogFuture<'a> {
Box::pin(async move {
self.events.lock().unwrap().push("failure");
})
}
}
struct HeaderEditHooks;
impl OcrHooks for HeaderEditHooks {
fn intercepts_requests(&self) -> bool {
true
}
fn during_call(
&self,
mut request: OcrDuringCallRequest,
) -> OcrHookFuture<'_, OcrDuringCallRequest> {
request
.headers
.push(("x-core-callback".into(), "edited".into()));
Box::pin(async move { Ok(request) })
}
}
#[tokio::test]
async fn lifecycle_sends_headers_returned_by_the_typed_during_call_operation() {
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({"pages":[]}))]).await;
let request = crate::ocr::LiteLLMOcrRequest {
hooks: Arc::new(HeaderEditHooks),
..wire_request("mistral/model", &base, json!({}))
};
perform_ocr(request).await.unwrap();
server.await.unwrap();
assert!(seen.lock().unwrap()[0].contains("x-core-callback: edited"));
}
#[tokio::test]
async fn lifecycle_orders_hooks_and_emits_one_success() {
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({"pages":[]}))]).await;
let events = Arc::new(Mutex::new(Vec::new()));
let request = wire_request("mistral/model", &base, json!({}));
let request = crate::ocr::LiteLLMOcrRequest {
hooks: Arc::new(RecordingHooks {
events: events.clone(),
block: false,
}),
..request
};
perform_ocr(request).await.unwrap();
server.await.unwrap();
assert_eq!(
*events.lock().unwrap(),
["pre", "during", "post", "success"]
);
assert_eq!(seen.lock().unwrap().len(), 1);
}
#[tokio::test]
async fn lifecycle_blocking_prevents_execution_and_emits_one_failure() {
let events = Arc::new(Mutex::new(Vec::new()));
let request = wire_request("mistral/model", "http://127.0.0.1:1", json!({}));
let request = crate::ocr::LiteLLMOcrRequest {
hooks: Arc::new(RecordingHooks {
events: events.clone(),
block: true,
}),
..request
};
let error = perform_ocr(request).await.unwrap_err();
assert!(matches!(error, crate::ocr::Error::InvalidRequest(_)));
assert_eq!(*events.lock().unwrap(), ["pre", "failure"]);
}
#[tokio::test]
async fn upstream_failure_emits_one_terminal_failure() {
let (base, seen, server) = mock_server(vec![MockResponse {
status: 500,
headers: vec![],
body: json!({"error":"failed"}),
}])
.await;
let events = Arc::new(Mutex::new(Vec::new()));
let request = wire_request("mistral/model", &base, json!({}));
let request = crate::ocr::LiteLLMOcrRequest {
hooks: Arc::new(RecordingHooks {
events: events.clone(),
block: false,
}),
..request
};
assert!(perform_ocr(request).await.is_err());
server.await.unwrap();
assert_eq!(*events.lock().unwrap(), ["pre", "during", "failure"]);
assert_eq!(seen.lock().unwrap().len(), 1);
}
struct AdmissionSpy {
effects: Arc<Mutex<usize>>,
}
impl OcrHooks for AdmissionSpy {
fn intercepts_requests(&self) -> bool {
*self.effects.lock().unwrap() += 1;
true
}
fn pre_call(&self, request: OcrPreCallRequest) -> OcrHookFuture<'_, OcrPreCallRequest> {
*self.effects.lock().unwrap() += 1;
Box::pin(async move { Ok(request) })
}
}
#[test]
fn admission_declines_without_invoking_hooks_or_transport() {
for (admission, expected) in [
(
OcrAdmission {
provider_workflow: false,
host_operations: true,
asynchronous: false,
},
OcrDecline::ProviderWorkflow,
),
(
OcrAdmission {
provider_workflow: true,
host_operations: false,
asynchronous: false,
},
OcrDecline::HostOperations,
),
] {
let outcome = OcrCall::admit(crate::ocr::test_support::ocr_client(), admission);
assert!(matches!(outcome, NativeOutcome::Declined(reason) if reason == expected));
}
}
#[tokio::test]
async fn fallible_host_phases_do_not_replay_or_reach_transport() {
for failure_phase in ["pre", "during"] {
let request = crate::ocr::LiteLLMOcrRequest {
hooks: Arc::new(AdmissionSpy {
effects: Arc::new(Mutex::new(0)),
}),
..wire_request("mistral/model", "http://127.0.0.1:1", json!({}))
};
let NativeOutcome::Completed(mut call) =
OcrCall::admit(crate::ocr::test_support::ocr_client(), OcrAdmission::all())
else {
panic!("supported call declined")
};
let mut request = Some(request);
let mut result = None;
let mut phases = Vec::new();
let error = loop {
match call.resume(result.take()).await {
Ok(OcrCallStep::Host(operation)) => match operation {
OcrHostOperation::Lifecycle(_)
| OcrHostOperation::ConstructResponse(_)
| OcrHostOperation::MapFailure(_)
| OcrHostOperation::Success { .. }
| OcrHostOperation::Failure { .. } => {
result = Some(OcrHostResult::Lifecycle(Ok(())))
}
OcrHostOperation::ProjectRequest => {
result = Some(OcrHostResult::Request(Ok((
Box::new(request.take().unwrap()),
false,
))))
}
OcrHostOperation::AcquireAzureAdToken => {
panic!("test request has no token provider")
}
OcrHostOperation::PreCall(request) => {
phases.push("pre");
result = Some(OcrHostResult::PreCall(if failure_phase == "pre" {
Err(crate::ocr::Error::InvalidRequest("pre failed".into()))
} else {
Ok(request)
}));
}
OcrHostOperation::DuringCall(request) => {
phases.push("during");
result =
Some(OcrHostResult::DuringCall(if failure_phase == "during" {
Err(crate::ocr::Error::InvalidRequest("during failed".into()))
} else {
Ok(request)
}));
}
OcrHostOperation::PostCall(_) => panic!("transport should not be reached"),
},
Err(error) => break error,
Ok(OcrCallStep::Complete(_)) => panic!("failed call completed"),
}
};
assert!(matches!(error, crate::ocr::Error::InvalidRequest(_)));
assert_eq!(
phases
.iter()
.filter(|phase| **phase == failure_phase)
.count(),
1
);
}
}
#[tokio::test]
async fn invalid_provider_response_runs_post_call_before_normalization_failure() {
let (base, seen, server) =
mock_server(vec![MockResponse::json(json!({"pages":"invalid"}))]).await;
let mut request = Some(wire_request("mistral/model", &base, json!({})));
let NativeOutcome::Completed(mut call) =
OcrCall::admit(crate::ocr::test_support::ocr_client(), OcrAdmission::all())
else {
panic!("supported call declined")
};
let host = NoopOcrHost;
let mut result = None;
let mut post_calls = Vec::new();
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,
))));
}
Ok(OcrCallStep::Host(operation)) => {
if let OcrHostOperation::PostCall(request) = &operation {
post_calls.push(request.original_response.clone());
}
result = Some(host.invoke(operation).await);
}
Err(error) => break error,
Ok(OcrCallStep::Complete(_)) => panic!("invalid provider response completed"),
}
};
server.await.unwrap();
assert!(matches!(
error,
crate::ocr::Error::ResponseField { ref path } if path == "pages"
));
assert_eq!(seen.lock().unwrap().len(), 1);
assert_eq!(post_calls, [json!(r#"{"pages":"invalid"}"#)]);
}
#[tokio::test]
async fn direct_native_host_drives_the_same_state_machine() {
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({
"pages":[{"index":0,"markdown":"native"}]
}))])
.await;
let request = crate::ocr::LiteLLMOcrRequest {
hooks: Arc::new(AdmissionSpy {
effects: Arc::new(Mutex::new(0)),
}),
..wire_request("mistral/model", &base, json!({}))
};
let NativeOutcome::Completed(mut call) = OcrCall::admit(
crate::ocr::test_support::ocr_client(),
OcrAdmission {
asynchronous: true,
..OcrAdmission::all()
},
) else {
panic!("supported call declined")
};
let mut request = Some(request);
let host = NoopOcrHost;
let mut result = None;
let mut operations = Vec::new();
let response = loop {
match call.resume(result.take()).await.unwrap() {
OcrCallStep::Host(operation) => {
operations.push(match &operation {
OcrHostOperation::ProjectRequest => "ProjectRequest".into(),
OcrHostOperation::Lifecycle(phase) => format!("{phase:?}"),
OcrHostOperation::PreCall(_) => "PreCall".into(),
OcrHostOperation::DuringCall(_) => "DuringCall".into(),
OcrHostOperation::PostCall(_) => "PostCall".into(),
OcrHostOperation::ConstructResponse(_) => "ConstructResponse".into(),
OcrHostOperation::Success { response, .. } => {
assert_eq!(response.pages[0].markdown, "native");
"Success".into()
}
_ => panic!("unexpected OCR operation"),
});
result = Some(match operation {
OcrHostOperation::ProjectRequest => {
OcrHostResult::Request(Ok((Box::new(request.take().unwrap()), false)))
}
operation => host.invoke(operation).await,
});
}
OcrCallStep::Complete(response) => break response,
}
};
server.await.unwrap();
assert_eq!(response.pages[0].markdown, "native");
assert_eq!(seen.lock().unwrap().len(), 1);
assert_eq!(
operations,
[
"Setup",
"DeploymentPreCall",
"Prepare",
"ProjectRequest",
"PreCall",
"DuringCall",
"PostCall",
"ConstructResponse",
"DeploymentPostCall",
"Finalize",
"Success",
]
);
assert!(matches!(
call.resume(None).await,
Err(crate::ocr::Error::InvalidRequest(_))
));
}
#[tokio::test]
async fn public_finalization_failure_never_dispatches_success_or_replays_provider() {
use crate::call_lifecycle::host::{HostFailure, HostPhase};
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({"pages":[]}))]).await;
let mut request = Some(wire_request("mistral/model", &base, json!({})));
let NativeOutcome::Completed(mut call) = OcrCall::admit(
crate::ocr::test_support::ocr_client(),
OcrAdmission {
asynchronous: true,
..OcrAdmission::all()
},
) else {
panic!("supported call declined")
};
let selected = crate::ocr::Error::InvalidRequest("public metadata failed".into());
let host = NoopOcrHost;
let mut result = None;
let mut failures = Vec::new();
let error = loop {
match call.resume(result.take()).await {
Ok(OcrCallStep::Host(operation)) => {
result = Some(match operation {
OcrHostOperation::Lifecycle(HostPhase::Finalize) => {
OcrHostResult::Lifecycle(Err(HostFailure::Error(selected.clone())))
}
OcrHostOperation::Failure { error, .. } => {
assert_eq!(error, selected);
failures.push("sync");
OcrHostResult::Lifecycle(Err(HostFailure::Error(
crate::ocr::Error::InvalidRequest("failure callback failed".into()),
)))
}
OcrHostOperation::Lifecycle(HostPhase::AsyncFailure) => {
failures.push("async");
OcrHostResult::Lifecycle(Ok(()))
}
OcrHostOperation::Success { .. }
| OcrHostOperation::MapFailure(_)
| OcrHostOperation::Lifecycle(HostPhase::DeploymentFailure) => {
panic!("finalization failure used provider/success dispatch")
}
OcrHostOperation::ProjectRequest => {
OcrHostResult::Request(Ok((Box::new(request.take().unwrap()), false)))
}
operation => host.invoke(operation).await,
});
}
Ok(OcrCallStep::Complete(_)) => panic!("failed call completed successfully"),
Err(error) => break error,
}
};
server.await.unwrap();
assert_eq!(error, selected);
assert_eq!(failures, ["sync", "async"]);
assert_eq!(seen.lock().unwrap().len(), 1);
}
#[tokio::test]
async fn cancellation_at_provider_hook_prevents_execution_and_further_resumption() {
use crate::call_lifecycle::host::HostFailure;
let request = crate::ocr::LiteLLMOcrRequest {
hooks: Arc::new(AdmissionSpy {
effects: Arc::new(Mutex::new(0)),
}),
..wire_request("mistral/model", "http://127.0.0.1:1", json!({}))
};
let NativeOutcome::Completed(mut call) =
OcrCall::admit(crate::ocr::test_support::ocr_client(), OcrAdmission::all())
else {
panic!("supported call declined")
};
let mut request = Some(request);
let host = NoopOcrHost;
let mut result = None;
loop {
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,
))))
}
OcrCallStep::Host(operation) => result = Some(host.invoke(operation).await),
OcrCallStep::Complete(_) => panic!("provider executed before pre-call result"),
}
}
let selected = crate::ocr::Error::InvalidRequest("cancelled".into());
assert!(matches!(
call.interrupt(HostFailure::Cancelled(selected.clone())).await,
Err(error) if error == selected
));
assert!(
call.resume(Some(OcrHostResult::Lifecycle(Ok(()))))
.await
.is_err()
);
}
#[tokio::test]
async fn missing_host_result_preserves_pending_operation() {
use crate::call_lifecycle::host::HostPhase;
let NativeOutcome::Completed(mut call) =
OcrCall::admit(crate::ocr::test_support::ocr_client(), OcrAdmission::all())
else {
panic!("supported call declined")
};
assert!(matches!(
call.resume(None).await.unwrap(),
OcrCallStep::Host(OcrHostOperation::Lifecycle(HostPhase::Setup))
));
assert!(call.resume(None).await.is_err());
assert!(matches!(
call.resume(Some(OcrHostResult::Lifecycle(Ok(()))))
.await
.unwrap(),
OcrCallStep::Host(OcrHostOperation::Lifecycle(HostPhase::Prepare))
));
}
async fn read_bounded_response(
response: Vec<u8>,
limit: usize,
) -> Result<bytes::Bytes, crate::ocr::Error> {
use tokio::io::{AsyncReadExt, AsyncWriteExt};
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let address = listener.local_addr().unwrap();
let server = tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.unwrap();
let mut request = [0; 4096];
assert!(socket.read(&mut request).await.unwrap() > 0);
socket.write_all(&response).await.unwrap();
std::future::pending::<()>().await;
});
let response = reqwest::Client::new()
.get(format!("http://{address}"))
.send()
.await
.unwrap();
let result = tokio::time::timeout(
std::time::Duration::from_secs(2),
crate::ocr::client::read_response_bytes(response, limit),
)
.await;
server.abort();
let _ = server.await;
result.expect("bounded reads must finish without waiting for the rest of an oversized body")
}
#[tokio::test]
async fn response_limit_accepts_exact_size_and_rejects_declared_and_chunked_overflow() {
for response in [
"HTTP/1.1 200 OK\r\nContent-Length: 8\r\n\r\nabcdefgh",
"HTTP/1.1 200 OK\r\nTransfer-Encoding: chunked\r\n\r\n4\r\nabcd\r\n4\r\nefgh\r\n0\r\n\r\n",
] {
assert_eq!(
read_bounded_response(response.as_bytes().to_vec(), 8)
.await
.unwrap(),
"abcdefgh"
);
}
for response in [
"HTTP/1.1 200 OK\r\nContent-Length: 9\r\n\r\n",
"HTTP/1.1 200 OK\r\nTransfer-Encoding: chunked\r\n\r\n4\r\nabcd\r\n5\r\nefghi\r\n",
] {
assert!(matches!(
read_bounded_response(response.as_bytes().to_vec(), 8).await,
Err(crate::ocr::Error::TooLarge { limit: 8 })
));
}
}
#[tokio::test]
async fn oversized_error_retains_http_status_and_bounded_diagnostics_without_draining() {
let prefix = "x".repeat(4 * (crate::constants::UPSTREAM_ERROR_BODY_MAX_CHARS + 1));
for headers in ["Content-Length: 1000000", "Transfer-Encoding: chunked"] {
let body = if headers.starts_with("Transfer") {
format!("{:x}\r\n{prefix}\r\n", prefix.len())
} else {
prefix.clone()
};
let response = format!("HTTP/1.1 429 Too Many Requests\r\n{headers}\r\n\r\n{body}");
let error = read_bounded_response(response.into_bytes(), 4096)
.await
.unwrap_err();
match error {
crate::ocr::Error::Transport(crate::transport::Error::Http { status, body }) => {
assert_eq!(status, 429);
assert_eq!(
body,
format!(
"{}... (truncated)",
"x".repeat(crate::constants::UPSTREAM_ERROR_BODY_MAX_CHARS)
)
);
}
error => panic!("unexpected error: {error}"),
}
}
}
#[test]
fn response_limit_is_validated_and_not_forwarded_to_the_provider() {
let request = wire_request(
"mistral/model",
"http://localhost",
json!({"max_response_bytes": 123}),
);
assert_eq!(request.connection.max_response_bytes, 123);
assert!(!request.optional_params.contains_key("max_response_bytes"));
for value in [
json!(0),
json!(-1),
json!(true),
json!("123"),
json!(1.5),
json!(crate::constants::OCR_RESPONSE_MAX_BYTES + 1),
Value::Null,
] {
let wire = serde_json::from_value(json!({
"model": "mistral/model", "document": {"type": "document_url", "document_url": "data:application/pdf;base64,YWJj"},
"optional_params": {"max_response_bytes": value}
})).unwrap();
let Err(error) = decode_request(wire) else {
panic!("invalid response limit accepted")
};
assert!(error.to_string().contains("max_response_bytes"));
}
}
#[derive(Debug)]
struct PendingToken {
entered: Arc<tokio::sync::Notify>,
dropped: Arc<std::sync::atomic::AtomicBool>,
}
struct TokenFutureDrop(Arc<std::sync::atomic::AtomicBool>);
impl Drop for TokenFutureDrop {
fn drop(&mut self) {
self.0.store(true, std::sync::atomic::Ordering::SeqCst);
}
}
impl litellm_auth::TokenProvider for PendingToken {
fn acquire(&self) -> litellm_auth::TokenFuture<'_> {
Box::pin(async move {
let _guard = TokenFutureDrop(self.dropped.clone());
self.entered.notify_one();
std::future::pending().await
})
}
}
#[tokio::test]
async fn cancellation_waits_for_provider_capture_drop_even_when_acknowledgement_is_cancelled() {
use crate::call_lifecycle::host::HostFailure;
use std::future::Future;
use std::sync::atomic::{AtomicBool, Ordering};
use std::task::Poll;
for interrupt_acknowledgement in [false, true] {
let entered = Arc::new(tokio::sync::Notify::new());
let dropped = Arc::new(AtomicBool::new(false));
let request =
wire_request("azure_ai/mistral-ocr", "https://example.invalid", json!({}));
let request = crate::ocr::LiteLLMOcrRequest {
connection: crate::ocr::OcrConnection {
extra_headers: vec![("authorization".into(), "Bearer test-key".into())],
..request.connection
},
azure_ad_token_provider: Some(litellm_auth::TokenProviderHandle::new(Arc::new(
PendingToken {
entered: entered.clone(),
dropped: dropped.clone(),
},
))),
..request
};
let NativeOutcome::Completed(mut call) =
OcrCall::admit(crate::ocr::test_support::ocr_client(), OcrAdmission::all())
else {
panic!("supported call declined")
};
let mut request = Some(request);
let mut result = None;
tokio::time::timeout(std::time::Duration::from_secs(2), async {
loop {
tokio::select! {
_ = 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(operation) => NoopOcrHost.invoke(operation).await,
OcrCallStep::Complete(_) => panic!("pending provider completed"),
});
}
}
}
}).await.unwrap();
assert!(!dropped.load(Ordering::SeqCst));
let selected = crate::ocr::Error::InvalidRequest("cancelled".into());
if interrupt_acknowledgement {
let mut acknowledgement =
Box::pin(call.interrupt(HostFailure::Cancelled(selected.clone())));
std::future::poll_fn(|cx| {
assert!(acknowledgement.as_mut().poll(cx).is_pending());
Poll::Ready(())
})
.await;
drop(acknowledgement);
assert!(!dropped.load(Ordering::SeqCst));
}
let result = tokio::time::timeout(
std::time::Duration::from_secs(2),
call.interrupt(HostFailure::Cancelled(selected.clone())),
)
.await
.unwrap();
assert!(matches!(result, Err(error) if error == selected));
assert!(
dropped.load(Ordering::SeqCst),
"cancellation returned while provider captures were still alive"
);
}
}
}

View file

@ -16,32 +16,11 @@ pub use lifecycle::{
NativeOutcome, NativeResult, NoopOcrHost, OcrAdmission, OcrCall, OcrCallStep, OcrDecline,
OcrHookHost, OcrHost, OcrHostOperation, OcrHostResult,
};
pub use provider_config::{get_api_key_env_var, get_health_check_document};
pub use types::{
LiteLLMOcrRequest, LiteLLMOcrResponse, OcrConnection, OcrDocument, OcrPage, OcrPageDimensions,
OcrPageImage, OcrUsageInfo,
};
#[cfg(test)]
#[path = "../../tests/azure_ai_ocr.rs"]
mod azure_ai_tests;
#[cfg(test)]
#[path = "../../tests/azure_document_intelligence_ocr.rs"]
mod azure_document_intelligence_tests;
#[cfg(test)]
#[path = "../../tests/deepseek_ocr.rs"]
mod deepseek_tests;
#[cfg(test)]
#[path = "../../tests/reducto_ocr.rs"]
mod reducto_tests;
#[cfg(test)]
#[path = "../../tests/ocr/support.rs"]
pub(crate) mod test_support;
#[cfg(test)]
#[path = "../../tests/ocr.rs"]
pub(crate) mod tests;
#[cfg(test)]
#[path = "../../tests/vertex_ai_deepseek_ocr.rs"]
mod vertex_ai_deepseek_tests;
#[cfg(test)]
#[path = "../../tests/vertex_ai_ocr.rs"]
mod vertex_ai_tests;

View file

@ -43,7 +43,7 @@ where
.hooks
.during_call(OcrDuringCallRequest {
model: request.model.clone(),
custom_llm_provider: request.config.provider().as_str().into(),
custom_llm_provider: request.provider_name().into(),
url: url.into(),
headers: headers.to_vec(),
body,
@ -93,7 +93,7 @@ pub(crate) async fn guardrail_document(
.hooks
.during_call(OcrDuringCallRequest {
model: request.model.clone(),
custom_llm_provider: request.config.provider().as_str().into(),
custom_llm_provider: request.provider_name().into(),
url: url.into(),
headers: headers.to_vec(),
body: serde_json::to_value(&request.document).map_err(|_| {
@ -126,3 +126,50 @@ pub(crate) fn body_document(body: &Value) -> Result<OcrDocument, super::Error> {
pub(crate) fn credential_env(name: &str) -> Option<String> {
std::env::var(name).ok()
}
pub(crate) fn resolve_connection_params(request: LiteLLMOcrRequest) -> LiteLLMOcrRequest {
use litellm_auth::{InputSource, Sourced};
let connection = request.connection;
let api_base_env = match request.config.provider() {
super::provider_config::OcrProvider::Mistral => Some("MISTRAL_API_BASE"),
super::provider_config::OcrProvider::AzureAi => Some("AZURE_AI_API_BASE"),
super::provider_config::OcrProvider::Cohere
| super::provider_config::OcrProvider::Reducto
| super::provider_config::OcrProvider::VertexAi => None,
};
let dynamic_api_key = connection.dynamic_api_key.or_else(|| {
connection
.api_key
.clone()
.map(|value| Sourced::new(value, connection.api_key_source))
.or_else(|| {
request
.config
.get_api_key_env_var()
.and_then(credential_env)
.map(|value| Sourced::new(value, InputSource::Environment))
})
});
let dynamic_api_base = connection.dynamic_api_base.or_else(|| {
connection
.api_base
.clone()
.map(|value| Sourced::new(value, connection.api_base_source))
.or_else(|| {
api_base_env
.and_then(credential_env)
.map(|value| Sourced::new(value, InputSource::Environment))
})
});
LiteLLMOcrRequest {
connection: request
.config
.resolve_connection_params(super::OcrConnection {
dynamic_api_key,
dynamic_api_base,
..connection
}),
..request
}
}

View file

@ -1,3 +1,4 @@
use super::types::{OcrConnection, OcrDocument};
use crate::llms::azure_ai::ocr::cohere_parse_transformation::AzureAICohereParseConfig;
use crate::llms::azure_ai::ocr::document_intelligence::transformation::AzureDocumentIntelligenceOCRConfig;
use crate::llms::azure_ai::ocr::transformation::AzureAIOCRConfig;
@ -8,6 +9,24 @@ use crate::llms::reducto::ocr::transformation::{ReductoParseLegacyConfig, Reduct
use crate::llms::vertex_ai::ocr::deepseek_transformation::VertexAIDeepSeekOCRConfig;
use crate::llms::vertex_ai::ocr::transformation::VertexAIOCRConfig;
use crate::routing_utils::provider::{CustomLlmProvider, get_custom_llm_provider};
use litellm_auth::Sourced;
use strum::{EnumString, IntoStaticStr};
macro_rules! dispatch_config {
($config:expr, $method:ident($($argument:expr),* $(,)?)) => {
match $config {
OcrConfigKind::Cohere => CohereParseConfig.$method($($argument),*),
OcrConfigKind::Mistral => MistralOCRConfig.$method($($argument),*),
OcrConfigKind::AzureAi => AzureAIOCRConfig.$method($($argument),*),
OcrConfigKind::AzureCohere => AzureAICohereParseConfig.$method($($argument),*),
OcrConfigKind::AzureDocumentIntelligence => AzureDocumentIntelligenceOCRConfig.$method($($argument),*),
OcrConfigKind::ReductoLegacy => ReductoParseLegacyConfig.$method($($argument),*),
OcrConfigKind::ReductoV3 => ReductoParseV3Config.$method($($argument),*),
OcrConfigKind::VertexAi => VertexAIOCRConfig.$method($($argument),*),
OcrConfigKind::VertexDeepSeek => VertexAIDeepSeekOCRConfig.$method($($argument),*),
}
};
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub(crate) enum OcrConfigKind {
@ -36,23 +55,80 @@ impl OcrConfigKind {
}
pub(crate) fn get_supported_ocr_params(self, model: &str) -> &'static [&'static str] {
match self {
Self::Cohere => CohereParseConfig.get_supported_ocr_params(model),
Self::Mistral => MistralOCRConfig.get_supported_ocr_params(model),
Self::AzureAi => AzureAIOCRConfig.get_supported_ocr_params(model),
Self::AzureCohere => AzureAICohereParseConfig.get_supported_ocr_params(model),
Self::AzureDocumentIntelligence => {
AzureDocumentIntelligenceOCRConfig.get_supported_ocr_params(model)
}
Self::ReductoLegacy => ReductoParseLegacyConfig.get_supported_ocr_params(model),
Self::ReductoV3 => ReductoParseV3Config.get_supported_ocr_params(model),
Self::VertexAi => VertexAIOCRConfig.get_supported_ocr_params(model),
Self::VertexDeepSeek => VertexAIDeepSeekOCRConfig.get_supported_ocr_params(model),
dispatch_config!(self, get_supported_ocr_params(model))
}
pub(crate) fn get_api_key_env_var(self) -> Option<&'static str> {
dispatch_config!(self, get_api_key_env_var())
}
pub(crate) fn get_health_check_document(self) -> OcrDocument {
dispatch_config!(self, get_health_check_document())
}
pub(crate) fn resolve_connection_params(self, connection: OcrConnection) -> OcrConnection {
let api_key = connection
.api_key
.map(|value| Sourced::new(value, connection.api_key_source));
let api_base = connection
.api_base
.map(|value| Sourced::new(value, connection.api_base_source));
let (api_key, api_base) = dispatch_config!(
self,
resolve_connection_params(
api_key,
api_base,
connection.dynamic_api_key,
connection.dynamic_api_base,
)
);
OcrConnection {
api_key_source: api_key
.as_ref()
.map(Sourced::source)
.unwrap_or(connection.api_key_source),
api_base_source: api_base
.as_ref()
.map(Sourced::source)
.unwrap_or(connection.api_base_source),
api_key: api_key.map(Sourced::into_value),
api_base: api_base.map(Sourced::into_value),
dynamic_api_key: None,
dynamic_api_base: None,
..connection
}
}
pub(crate) fn get_error_class(
self,
message: String,
status: u16,
headers: Vec<(String, String)>,
) -> super::Error {
dispatch_config!(self, get_error_class(message, status, headers))
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub fn get_api_key_env_var(
model: &str,
custom_llm_provider: Option<&str>,
) -> Result<Option<&'static str>, super::Error> {
Ok(resolve_provider_config(model, custom_llm_provider)?
.1
.get_api_key_env_var())
}
pub fn get_health_check_document(
model: &str,
custom_llm_provider: Option<&str>,
) -> Result<OcrDocument, super::Error> {
Ok(resolve_provider_config(model, custom_llm_provider)?
.1
.get_health_check_document())
}
#[derive(Clone, Copy, Debug, EnumString, IntoStaticStr, PartialEq, Eq)]
#[strum(serialize_all = "snake_case")]
pub(crate) enum OcrProvider {
Cohere,
Mistral,
@ -61,18 +137,6 @@ pub(crate) enum OcrProvider {
VertexAi,
}
impl OcrProvider {
pub(crate) const fn as_str(self) -> &'static str {
match self {
Self::Cohere => "cohere",
Self::Mistral => "mistral",
Self::AzureAi => "azure_ai",
Self::Reducto => "reducto",
Self::VertexAi => "vertex_ai",
}
}
}
pub(crate) fn resolve_provider_config(
model: &str,
custom_llm_provider: Option<&str>,
@ -80,30 +144,33 @@ pub(crate) fn resolve_provider_config(
let provider =
get_custom_llm_provider(model, custom_llm_provider).unwrap_or(CustomLlmProvider {
model,
custom_llm_provider: OcrProvider::Mistral.as_str(),
custom_llm_provider: OcrProvider::Mistral.into(),
});
let config = match provider.custom_llm_provider {
"cohere" => OcrConfigKind::Cohere,
"mistral" => OcrConfigKind::Mistral,
"azure_ai" if is_document_intelligence_model(provider.model) => {
let ocr_provider = provider
.custom_llm_provider
.parse::<OcrProvider>()
.map_err(|_| super::Error::InvalidProvider(provider.custom_llm_provider.to_string()))?;
let config = match ocr_provider {
OcrProvider::Cohere => OcrConfigKind::Cohere,
OcrProvider::Mistral => OcrConfigKind::Mistral,
OcrProvider::AzureAi if is_document_intelligence_model(provider.model) => {
OcrConfigKind::AzureDocumentIntelligence
}
"azure_ai"
OcrProvider::AzureAi
if provider.model.to_ascii_lowercase().contains("cohere")
&& provider.model.to_ascii_lowercase().contains("parse") =>
{
OcrConfigKind::AzureCohere
}
"azure_ai" => OcrConfigKind::AzureAi,
"reducto" if provider.model.eq_ignore_ascii_case("parse-legacy") => {
OcrProvider::AzureAi => OcrConfigKind::AzureAi,
OcrProvider::Reducto if provider.model.eq_ignore_ascii_case("parse-legacy") => {
OcrConfigKind::ReductoLegacy
}
"reducto" => OcrConfigKind::ReductoV3,
"vertex_ai" if provider.model.to_ascii_lowercase().contains("deepseek") => {
OcrProvider::Reducto => OcrConfigKind::ReductoV3,
OcrProvider::VertexAi if provider.model.to_ascii_lowercase().contains("deepseek") => {
OcrConfigKind::VertexDeepSeek
}
"vertex_ai" => OcrConfigKind::VertexAi,
value => return Err(super::Error::InvalidProvider(value.to_string())),
OcrProvider::VertexAi => OcrConfigKind::VertexAi,
};
Ok((provider.model.to_string(), config))
}
@ -116,6 +183,149 @@ fn is_document_intelligence_model(model: &str) -> bool {
#[cfg(test)]
mod tests {
use super::*;
use litellm_auth::InputSource;
#[test]
fn provider_names_round_trip_exactly() {
for provider in ["cohere", "mistral", "azure_ai", "reducto", "vertex_ai"] {
let (_, config) = resolve_provider_config("model", Some(provider)).unwrap();
let resolved: &'static str = config.provider().into();
assert_eq!(resolved, provider);
}
for provider in ["Mistral", "unknown"] {
assert_eq!(
resolve_provider_config("model", Some(provider)),
Err(crate::ocr::Error::InvalidProvider(provider.into()))
);
}
}
#[test]
fn health_check_documents_are_valid_for_each_provider() {
for model in [
"mistral/ocr",
"azure_ai/ocr",
"azure_ai/doc-intelligence/prebuilt-layout",
"reducto/parse-v3",
"vertex_ai/mistral-ocr",
"vertex_ai/deepseek-ocr",
] {
let document = get_health_check_document(model, None).unwrap();
assert!(matches!(document, OcrDocument::DocumentUrl { .. }));
let inline = crate::ocr::document::InlineDocument::parse(document.source())
.unwrap()
.unwrap();
assert_eq!(inline.mime_type().to_string(), "application/pdf");
assert!(inline.decode(4096).unwrap().starts_with(b"%PDF-"));
}
for model in ["cohere/parse", "azure_ai/cohere-parse"] {
let document = get_health_check_document(model, None).unwrap();
crate::llms::cohere::ocr::validate_document(&document).unwrap();
let inline = crate::ocr::document::InlineDocument::parse(document.source())
.unwrap()
.unwrap();
assert_eq!(inline.mime_type().to_string(), "image/png");
assert!(
inline
.decode(4096)
.unwrap()
.starts_with(b"\x89PNG\r\n\x1a\n")
);
}
}
#[test]
fn api_key_metadata_follows_provider_overrides_and_python_defaults() {
for (model, expected) in [
("mistral/ocr", Some("MISTRAL_API_KEY")),
("cohere/parse", Some("COHERE_API_KEY")),
("azure_ai/ocr", Some("AZURE_AI_API_KEY")),
("azure_ai/cohere-parse", Some("AZURE_AI_API_KEY")),
(
"azure_ai/doc-intelligence/prebuilt-layout",
Some("AZURE_DOCUMENT_INTELLIGENCE_API_KEY"),
),
("vertex_ai/mistral-ocr", Some("VERTEX_AI_API_KEY")),
("vertex_ai/deepseek-ocr", Some("VERTEX_AI_API_KEY")),
("reducto/parse-v3", None),
("reducto/parse-legacy", None),
] {
assert_eq!(
get_api_key_env_var(model, None).unwrap(),
expected,
"{model}"
);
}
}
#[test]
fn connection_resolution_preserves_dynamic_precedence_and_input_sources() {
let connection = OcrConfigKind::Mistral.resolve_connection_params(OcrConnection {
api_key: Some("explicit-key".into()),
api_base: Some("https://explicit.test".into()),
dynamic_api_key: Some(Sourced::new("dynamic-key".into(), InputSource::Environment)),
dynamic_api_base: Some(Sourced::new(
"https://dynamic.test".into(),
InputSource::Request,
)),
..Default::default()
});
assert_eq!(connection.api_key.as_deref(), Some("dynamic-key"));
assert_eq!(connection.api_base.as_deref(), Some("https://dynamic.test"));
assert_eq!(connection.api_key_source, InputSource::Environment);
assert_eq!(connection.api_base_source, InputSource::Request);
for dynamic in [
None,
Some(Sourced::new(String::new(), InputSource::Environment)),
] {
let connection = OcrConfigKind::Mistral.resolve_connection_params(OcrConnection {
api_key: Some("explicit-key".into()),
api_base: Some("https://explicit.test".into()),
dynamic_api_key: dynamic.clone(),
dynamic_api_base: dynamic,
..Default::default()
});
assert_eq!(connection.api_key.as_deref(), Some("explicit-key"));
assert_eq!(
connection.api_base.as_deref(),
Some("https://explicit.test")
);
}
}
#[test]
fn document_intelligence_only_accepts_dynamic_values_for_explicit_fields() {
for (explicit_key, explicit_base) in [
(None, None),
(Some("key"), None),
(None, Some("base")),
(Some("key"), Some("base")),
] {
let connection =
OcrConfigKind::AzureDocumentIntelligence.resolve_connection_params(OcrConnection {
api_key: explicit_key.map(str::to_string),
api_base: explicit_base.map(str::to_string),
dynamic_api_key: Some(Sourced::new(
"dynamic-key".into(),
InputSource::Environment,
)),
dynamic_api_base: Some(Sourced::new(
"https://dynamic.test".into(),
InputSource::Deployment,
)),
..Default::default()
});
assert_eq!(
connection.api_key.as_deref(),
explicit_key.map(|_| "dynamic-key")
);
assert_eq!(
connection.api_base.as_deref(),
explicit_base.map(|_| "https://dynamic.test")
);
}
}
#[test]
fn provider_models_are_preserved_without_a_local_allowlist() {

View file

@ -37,13 +37,13 @@ pub(crate) fn wire_request(model: &str, base: &str, options: Value) -> LiteLLMOc
}
pub(crate) struct MockResponse {
pub status: u16,
pub headers: Vec<(&'static str, String)>,
pub body: Value,
pub(crate) status: u16,
pub(crate) headers: Vec<(&'static str, String)>,
pub(crate) body: Value,
}
impl MockResponse {
pub fn json(body: Value) -> Self {
pub(crate) fn json(body: Value) -> Self {
Self {
status: 200,
headers: vec![],

View file

@ -2,16 +2,17 @@ use std::collections::BTreeMap;
use std::sync::Arc;
use std::time::Duration;
use crate::serde_compat::{FiniteF64, LaxI64};
use serde::{Deserialize, Serialize};
use serde_json::{Map, Value};
use serde_with::serde_as;
use litellm_auth::{InputSource, Sourced, TokenProviderHandle};
use super::hooks::{NoopOcrHooks, OcrHooks};
use super::provider_config::{OcrConfigKind, resolve_provider_config};
use crate::call_arguments::CallArguments;
use crate::constants::OCR_HTTP_TIMEOUT_SECS;
use litellm_auth::{InputSource, TokenProviderHandle};
use crate::serde_compat::{FiniteF64, LaxI64};
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
#[serde(tag = "type")]
@ -63,8 +64,10 @@ pub enum OcrResponseFormat {
#[derive(Clone)]
pub struct OcrConnection {
pub api_key: Option<String>,
pub dynamic_api_key: Option<Sourced<String>>,
pub api_key_source: InputSource,
pub api_base: Option<String>,
pub dynamic_api_base: Option<Sourced<String>>,
pub api_base_source: InputSource,
pub extra_headers: Vec<(String, String)>,
pub extra_headers_source: InputSource,
@ -78,8 +81,10 @@ impl Default for OcrConnection {
fn default() -> Self {
Self {
api_key: None,
dynamic_api_key: None,
api_key_source: InputSource::Deployment,
api_base: None,
dynamic_api_base: None,
api_base_source: InputSource::Deployment,
extra_headers: Vec::new(),
extra_headers_source: InputSource::Deployment,
@ -137,7 +142,7 @@ impl LiteLLMOcrRequest {
}
pub fn provider_name(&self) -> &'static str {
self.config.provider().as_str()
self.config.provider().into()
}
pub fn with_host_hooks(

View file

@ -157,8 +157,10 @@ pub fn decode_request(wire: OcrWireRequest) -> Result<LiteLLMOcrRequest, crate::
)?;
let connection = OcrConnection {
api_key: nonblank(wire.api_key),
dynamic_api_key: None,
api_key_source,
api_base: nonblank(wire.api_base),
dynamic_api_base: None,
api_base_source,
extra_headers: headers,
extra_headers_source,

View file

@ -1,97 +0,0 @@
use std::sync::Arc;
use serde_json::{Value, json};
use super::hooks::{OcrDuringCallRequest, OcrHookFuture, OcrHooks};
use super::test_support::{MockResponse, mock_server, perform_ocr, wire_request};
#[tokio::test]
async fn facade_executes_azure_mistral_with_prepared_auth() {
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({
"pages":[{"index":0,"markdown":"hello"}],
"usage_info":{"pages_processed":1}
}))])
.await;
let mut request = wire_request(
"azure_ai/model",
&base,
json!({"include_image_base64":true}),
);
request.connection.api_key = None;
request.connection.extra_headers = vec![(
"Authorization".into(),
"Bearer python-prepared-token".into(),
)];
let result = perform_ocr(request).await.unwrap();
server.await.unwrap();
assert_eq!(result.pages[0].markdown, "hello");
let requests = seen.lock().unwrap();
assert_eq!(requests.len(), 1);
assert!(requests[0].starts_with("POST /providers/mistral/azure/ocr "));
assert!(
requests[0]
.to_ascii_lowercase()
.contains("authorization: bearer python-prepared-token\r\n")
);
let body: Value = serde_json::from_str(requests[0].split_once("\r\n\r\n").unwrap().1).unwrap();
assert_eq!(
body,
json!({
"model":"model",
"document":{"type":"document_url","document_url":"data:application/pdf;base64,YWJj"},
"include_image_base64":true
})
);
}
#[tokio::test]
async fn facade_acquires_supplied_entra_token_for_final_request() {
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({"pages":[]}))]).await;
let mut request = wire_request(
"azure_ai/model",
&base,
json!({"azure_ad_token":"rust-owned-token"}),
);
request.connection.api_key = None;
perform_ocr(request).await.unwrap();
server.await.unwrap();
let requests = seen.lock().unwrap();
assert_eq!(requests.len(), 1);
assert!(
requests[0]
.to_ascii_lowercase()
.contains("authorization: bearer rust-owned-token\r\n")
);
}
struct ReplaceBodyDocument;
impl OcrHooks for ReplaceBodyDocument {
fn intercepts_requests(&self) -> bool {
true
}
fn during_call(
&self,
mut request: OcrDuringCallRequest,
) -> OcrHookFuture<'_, OcrDuringCallRequest> {
Box::pin(async move {
request.body["document"] = json!({
"type":"document_url",
"document_url":"https://example.com/not-inline.pdf"
});
Ok(request)
})
}
}
#[tokio::test]
async fn rejects_non_inline_body_after_guardrails() {
let mut request = wire_request("azure_ai/model", "http://127.0.0.1:1", json!({}));
request.hooks = Arc::new(ReplaceBodyDocument);
let error = perform_ocr(request).await.unwrap_err();
assert!(error.to_string().contains("data URI"));
}

View file

@ -1,460 +0,0 @@
use serde_json::{Value, json};
use std::sync::{Arc, Mutex};
use super::test_support::{MockResponse, mock_server, perform_ocr, wire_request};
use super::wire::{OcrWireRequest, decode_request};
fn query_value(url: &str, key: &str) -> Option<String> {
url::Url::parse(url)
.unwrap()
.query_pairs()
.find_map(|(name, value)| (name == key).then(|| value.into_owned()))
}
#[tokio::test]
async fn facade_maps_pages_features_and_url_document() {
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({
"status":"succeeded",
"analyzeResult":{"pages":[]}
}))])
.await;
let mut request = wire_request(
"azure_ai/doc-intelligence/prebuilt-read",
&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!({
"type":"document_url",
"document_url":"https://example.com/document.pdf"
}))
.unwrap();
perform_ocr(request).await.unwrap();
server.await.unwrap();
let request = &seen.lock().unwrap()[0];
let target = request.split_whitespace().nth(1).unwrap();
let url = format!("{base}{target}");
assert_eq!(query_value(&url, "pages").as_deref(), Some("1,2,3"));
assert_eq!(
query_value(&url, "features").as_deref(),
Some("keyValuePairs,languages")
);
let body: Value = serde_json::from_str(request.split_once("\r\n\r\n").unwrap().1).unwrap();
assert_eq!(
body,
json!({"urlSource":"https://example.com/document.pdf", "future_option":{"nested":null}, "provider_option":false})
);
}
#[tokio::test]
async fn rejects_invalid_pages_features_and_format() {
for options in [
json!({"pages":[true]}),
json!({"pages":[1,"2"]}),
json!({"pages":[-1]}),
json!({"pages":"1&&features=bad"}),
json!({"features":"languages&pages=1"}),
json!({"req_format":"azure"}),
] {
let result = decode_request(OcrWireRequest {
model: "azure_ai/doc-intelligence/prebuilt-read".into(),
document: json!({"type":"document_url","document_url":"https://example.com/a.pdf"}),
api_key: Some("key".into()),
api_base: Some("http://127.0.0.1:1".into()),
custom_llm_provider: None,
extra_headers: None,
optional_params: options.as_object().unwrap().clone().into(),
input_sources: Default::default(),
timeout_seconds: None,
});
let rejected = match result {
Ok(request) => perform_ocr(request).await.is_err(),
Err(_) => true,
};
assert!(rejected, "accepted {options}");
}
}
#[tokio::test]
async fn inline_document_decodes_to_base64_source() {
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({
"status":"succeeded"
}))])
.await;
let request = wire_request("azure_ai/doc-intelligence/prebuilt-read", &base, json!({}));
perform_ocr(request).await.unwrap();
server.await.unwrap();
let request = &seen.lock().unwrap()[0];
let body: Value = serde_json::from_str(request.split_once("\r\n\r\n").unwrap().1).unwrap();
assert_eq!(body, json!({"base64Source":"YWJj"}));
}
#[tokio::test]
async fn immediate_response_normalizes_pages_and_preserves_native() {
let operation = json!({
"status":"succeeded",
"operationExtension":42,
"analyzeResult":{
"content":"A\n\nB",
"tables":[{"cells":[]}],
"keyValuePairs":[{"key":{"content":"A"}}],
"pages":[{
"pageNumber":"2",
"width":"8.5",
"height":11,
"unit":"inch",
"lines":[{"content":"A"},{"content":null},{"content":"B"}]
}]
}
});
let (base, _, server) = mock_server(vec![MockResponse::json(operation.clone())]).await;
let result = perform_ocr(wire_request(
"azure_ai/doc-intelligence/prebuilt-read",
&base,
json!({"req_format":"native"}),
))
.await
.unwrap();
server.await.unwrap();
assert_eq!(result.pages[0].index, 1);
assert_eq!(result.pages[0].markdown, "A\n\nB");
assert_eq!(
serde_json::to_value(&result.pages[0].dimensions).unwrap(),
json!({"width":816,"height":1056,"dpi":96})
);
assert_eq!(result.usage_info.as_ref().unwrap().pages_processed, Some(1));
let serialized = result.clone().into_json();
assert_eq!(serialized["content"], "A\n\nB");
assert_eq!(serialized["tables"], json!([{"cells":[]}]));
assert_eq!(
serialized["keyValuePairs"],
json!([{"key":{"content":"A"}}])
);
assert!(serialized.get("key_value_pairs").is_none());
assert_eq!(
result.provider_native_response.as_ref(),
operation.as_object()
);
}
#[tokio::test]
async fn accepted_response_polls_to_success_with_only_credentials() {
let operation = json!({"status":"succeeded","analyzeResult":{"pages":[]}});
let (base, seen, server) = mock_server(vec![
MockResponse {
status: 202,
headers: vec![("Operation-Location", "{base}/operation".into())],
body: json!({}),
},
MockResponse {
status: 200,
headers: vec![("Retry-After", "0".into())],
body: json!({"status":"running"}),
},
MockResponse::json(operation.clone()),
])
.await;
let mut request = wire_request(
"azure_ai/doc-intelligence/prebuilt-read",
&base,
json!({"req_format":"native"}),
);
request
.connection
.extra_headers
.push(("X-Trace".into(), "initial-only".into()));
let result = perform_ocr(request).await.unwrap();
server.await.unwrap();
assert_eq!(
result.provider_native_response.as_ref(),
operation.as_object()
);
let requests = seen.lock().unwrap();
assert_eq!(requests.len(), 3);
assert!(requests[0].to_ascii_lowercase().contains("x-trace:"));
for poll in &requests[1..] {
assert!(!poll.to_ascii_lowercase().contains("x-trace:"));
assert!(
poll.to_ascii_lowercase()
.contains("ocp-apim-subscription-key: test-key")
);
}
}
struct SubmissionBoundary {
request_count: Arc<Mutex<Vec<String>>>,
post_calls: Arc<Mutex<Vec<Value>>>,
}
impl super::hooks::OcrHooks for SubmissionBoundary {
fn post_call(
&self,
request: super::hooks::OcrPostCallRequest,
) -> super::hooks::OcrHookFuture<'_, super::hooks::OcrPostCallRequest> {
Box::pin(async move {
assert_eq!(self.request_count.lock().unwrap().len(), 1);
self.post_calls
.lock()
.unwrap()
.push(request.original_response.clone());
Ok(request)
})
}
}
#[tokio::test]
async fn accepted_response_runs_post_call_once_before_polling() {
let (base, seen, server) = mock_server(vec![
MockResponse {
status: 202,
headers: vec![("Operation-Location", "{base}/operation".into())],
body: json!({"submitted": true}),
},
MockResponse::json(json!({"status":"succeeded"})),
])
.await;
let post_calls = Arc::new(Mutex::new(Vec::new()));
let request = super::LiteLLMOcrRequest {
hooks: Arc::new(SubmissionBoundary {
request_count: seen.clone(),
post_calls: post_calls.clone(),
}),
..wire_request("azure_ai/doc-intelligence/prebuilt-read", &base, json!({}))
};
perform_ocr(request).await.unwrap();
server.await.unwrap();
assert_eq!(seen.lock().unwrap().len(), 2);
assert_eq!(
*post_calls.lock().unwrap(),
[json!(r#"{"submitted":true}"#)]
);
}
#[tokio::test]
async fn polling_forwards_bearer_credentials() {
let (base, seen, server) = mock_server(vec![
MockResponse {
status: 202,
headers: vec![("Operation-Location", "{base}/operation".into())],
body: json!({}),
},
MockResponse::json(json!({"status":"succeeded"})),
])
.await;
let mut request = wire_request("azure_ai/doc-intelligence/prebuilt-read", &base, json!({}));
request.connection.api_key = None;
request.connection.extra_headers = vec![("Authorization".into(), "Bearer token".into())];
perform_ocr(request).await.unwrap();
server.await.unwrap();
let requests = seen.lock().unwrap();
assert!(
requests[1]
.to_ascii_lowercase()
.contains("authorization: bearer token")
);
}
#[tokio::test]
async fn polling_does_not_follow_redirects() {
let (base, seen, server) = mock_server(vec![
MockResponse {
status: 202,
headers: vec![("Operation-Location", "{base}/operation".into())],
body: json!({}),
},
MockResponse {
status: 302,
headers: vec![("Location", "{base}/redirected".into())],
body: json!({}),
},
MockResponse::json(json!({"status":"succeeded"})),
])
.await;
let error = perform_ocr(wire_request(
"azure_ai/doc-intelligence/prebuilt-read",
&base,
json!({}),
))
.await
.unwrap_err();
assert!(error.to_string().contains("status 302"), "{error}");
assert_eq!(seen.lock().unwrap().len(), 2);
server.abort();
}
#[tokio::test]
async fn polling_rejects_terminal_failure() {
let (base, _, server) = mock_server(vec![
MockResponse {
status: 202,
headers: vec![("Operation-Location", "{base}/operation".into())],
body: json!({}),
},
MockResponse::json(json!({"status":"failed"})),
])
.await;
let error = perform_ocr(wire_request(
"azure_ai/doc-intelligence/prebuilt-read",
&base,
json!({}),
))
.await
.unwrap_err();
server.await.unwrap();
assert!(error.to_string().contains("status failed"));
}
#[tokio::test]
async fn malformed_provider_pages_report_response_paths() {
for (analysis, path) in [
(json!({"pages":null}), "pages"),
(json!({"pages":[null]}), "pages[0]"),
(json!({"pages":[{"lines":null}]}), "lines"),
(json!({"pages":[{"width":"bad"}]}), "width"),
] {
let (base, _, server) = mock_server(vec![MockResponse::json(json!({
"status":"succeeded",
"analyzeResult":analysis
}))])
.await;
let error = perform_ocr(wire_request(
"azure_ai/doc-intelligence/prebuilt-read",
&base,
json!({}),
))
.await
.unwrap_err();
server.await.unwrap();
assert!(error.to_string().contains(path), "{error}");
}
}
#[tokio::test]
async fn rejects_missing_invalid_and_cross_origin_operation_locations() {
for headers in [
Vec::new(),
vec![("Operation-Location", "/relative".into())],
vec![("Operation-Location", "http://example.com/operation".into())],
vec![(
"Operation-Location",
"http://user:password@127.0.0.1/operation".into(),
)],
] {
let (base, _, server) = mock_server(vec![MockResponse {
status: 202,
headers,
body: json!({}),
}])
.await;
let error = perform_ocr(wire_request(
"azure_ai/doc-intelligence/prebuilt-read",
&base,
json!({}),
))
.await
.unwrap_err();
server.await.unwrap();
assert!(error.to_string().contains("operation-location"));
}
}
#[tokio::test]
async fn polling_deadline_bounds_retry_delay() {
let (base, _, server) = mock_server(vec![
MockResponse {
status: 202,
headers: vec![("Operation-Location", "{base}/operation".into())],
body: json!({}),
},
MockResponse {
status: 200,
headers: vec![("Retry-After", "9999".into())],
body: json!({"status":"notStarted"}),
},
])
.await;
let mut request = wire_request("azure_ai/doc-intelligence/prebuilt-read", &base, json!({}));
request.connection.poll_timeout = std::time::Duration::from_millis(100);
let error = tokio::time::timeout(std::time::Duration::from_secs(1), perform_ocr(request))
.await
.unwrap()
.unwrap_err();
server.await.unwrap();
assert!(error.to_string().contains("timed out"));
}
#[tokio::test]
async fn model_id_is_encoded_and_dot_segments_are_rejected() {
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({
"status":"succeeded"
}))])
.await;
perform_ocr(wire_request(
"azure_ai/doc-intelligence/a ?#é",
&base,
json!({}),
))
.await
.unwrap();
server.await.unwrap();
assert!(seen.lock().unwrap()[0].contains("a%20%3F%23%C3%A9:analyze"));
for model in [
"azure_ai/doc-intelligence/.",
"azure_ai/doc-intelligence/..",
] {
let error = perform_ocr(wire_request(model, "http://127.0.0.1:1", json!({})))
.await
.unwrap_err();
assert!(error.to_string().contains("dot segment"));
}
}
#[tokio::test]
async fn pre_call_guardrail_receives_caller_pages_before_mapping() {
use crate::ocr::hooks::{OcrHookFuture, OcrHooks, OcrPreCallRequest};
use std::sync::Arc;
struct RewritePages;
impl OcrHooks for RewritePages {
fn intercepts_requests(&self) -> bool {
true
}
fn pre_call(&self, request: OcrPreCallRequest) -> OcrHookFuture<'_, OcrPreCallRequest> {
Box::pin(async move {
assert_eq!(request.optional_params["pages"], json!([0, 2]));
Ok(OcrPreCallRequest {
optional_params: json!({"pages": [1]}),
..request
})
})
}
}
let (base, seen, server) =
mock_server(vec![MockResponse::json(json!({"status": "succeeded"}))]).await;
let request = wire_request(
"azure_ai/doc-intelligence/prebuilt-read",
&base,
json!({"pages": [0, 2]}),
)
.with_host_hooks(Arc::new(RewritePages), None);
perform_ocr(request).await.unwrap();
server.await.unwrap();
let requests = seen.lock().unwrap();
let target = requests[0].split_whitespace().nth(1).unwrap();
assert_eq!(
query_value(&format!("{base}{target}"), "pages").as_deref(),
Some("2")
);
assert_eq!(requests.len(), 1);
}

View file

@ -1,158 +0,0 @@
use rstest::rstest;
use serde_json::{Value, json};
use crate::llms::vertex_ai::ocr::deepseek_transformation::{
DeepSeekOcrParams, DeepSeekOcrResponse, VertexAIDeepSeekOCRConfig, normalize_response,
};
use crate::ocr::types::OcrDocument;
fn document() -> OcrDocument {
serde_json::from_value(json!({"type":"image_url","image_url":"gs://bucket/a.png"})).unwrap()
}
#[rstest]
#[case("stream", json!(true))]
#[case("temperature", json!(0.1))]
#[case("max_tokens", json!(1024))]
#[case("top_p", json!(0.9))]
#[case("n", json!(2))]
#[case("stop", json!("done"))]
#[case("stop", json!(["done", "stop"]))]
#[case("temperature", json!(null))]
fn request_mapping_matches_python(#[case] name: &str, #[case] value: Value) {
let params: DeepSeekOcrParams =
serde_json::from_value(json!({name: value.clone(), "ignored": true})).unwrap();
let result = serde_json::to_value(
VertexAIDeepSeekOCRConfig
.transform_ocr_request("deepseek-ai/deepseek-ocr-maas", document(), &params, &[])
.unwrap(),
)
.unwrap();
assert_eq!(result["model"], "deepseek-ai/deepseek-ocr-maas");
assert_eq!(
result["messages"][0]["content"][0],
json!({"type":"image_url","image_url":"gs://bucket/a.png"})
);
assert_eq!(result[name], value);
assert!(result.get("ignored").is_none());
}
#[rstest]
#[case(json!({"type":"image_url","image_url":"data:image/png;base64,AA=="}))]
#[case(json!({"type":"document_url","document_url":"data:application/pdf;base64,AA=="}))]
fn request_maps_both_document_types_to_image_content(#[case] document: Value) {
let source = document
.get("image_url")
.or_else(|| document.get("document_url"))
.unwrap()
.clone();
let request = VertexAIDeepSeekOCRConfig
.transform_ocr_request(
"deepseek-ai/deepseek-ocr-maas",
serde_json::from_value(document).unwrap(),
&DeepSeekOcrParams::default(),
&[],
)
.unwrap();
let result = serde_json::to_value(request).unwrap();
assert_eq!(
result["messages"][0]["content"][0],
json!({"type":"image_url","image_url":source})
);
}
#[rstest]
#[case(json!("# hello"), "# hello")]
#[case(json!("{broken"), "{broken")]
#[case(json!(" {\"pages\":[]} "), " {\"pages\":[]} ")]
#[case(json!({"pages":[]}), "")]
#[case(json!("[]"), "[]")]
#[case(json!("{\"pages\":[{\"markdown\":\"json text\"}]}"), "json text")]
#[case(json!({"pages":[{"markdown":"object"}]}), "object")]
fn response_transform_handles_text_json_and_objects(
#[case] content: Value,
#[case] expected: &str,
) {
let has_pages = content
.as_object()
.is_some_and(|data| data.contains_key("pages"))
|| content
.as_str()
.is_some_and(|text| text.contains("\"pages\""));
let response: DeepSeekOcrResponse = serde_json::from_value(
json!({"choices":[{"message":{"content":content}}],"usage":{"prompt_tokens":1}}),
)
.unwrap();
let result = normalize_response("model", response).unwrap().into_json();
assert_eq!(result["pages"][0]["markdown"], expected);
assert_eq!(result["pages"][0]["index"], 0);
if has_pages {
assert!(result["usage_info"].is_null());
} else {
assert_eq!(result["usage_info"]["prompt_tokens"], 1);
}
}
#[test]
fn structured_result_maps_pages_usage_model_and_annotation() {
let response: DeepSeekOcrResponse = serde_json::from_value(json!({
"choices":[{"message":{"content":{
"pages":[{"index":2,"markdown":"page","images":[{"id":"one"}],"dimensions":{"width":10}}],
"model":"provider-model",
"usage_info":{"pages_processed":1},
"document_annotation":{"language":"en"},
"future":"kept"
}}}]
}))
.unwrap();
let result = normalize_response("requested", response)
.unwrap()
.into_json();
assert_eq!(result["pages"][0]["index"], 2);
assert_eq!(result["pages"][0]["images"][0]["id"], "one");
assert_eq!(result["model"], "provider-model");
assert_eq!(result["usage_info"]["pages_processed"], 1);
assert_eq!(result["document_annotation"]["language"], "en");
assert!(result.get("future").is_none());
}
#[test]
fn response_transform_rejects_missing_empty_and_malformed_content() {
for value in [
json!({"choices":[]}),
json!({"choices":[{"message":{"content":{}}}]}),
json!({"choices":[{"message":{"content":""}}]}),
json!({"choices":[{"message":{"content":"{\"pages\":[{\"markdown\":42}]}"}}]}),
json!({"choices":[{"message":{"content":{"pages":[{"markdown":42}]}}}]}),
] {
let result = serde_json::from_value::<DeepSeekOcrResponse>(value)
.map_err(|_| ())
.and_then(|response| normalize_response("model", response).map_err(|_| ()));
assert!(result.is_err());
}
}
#[test]
fn structured_content_preserves_usage_presence_and_shared_page_defaults() {
for (usage, expected) in [(json!(null), None), (json!({"pages_processed":2}), Some(2))] {
let response = serde_json::from_value(json!({
"choices":[{"message":{"content":{
"pages":[42, {"index":"2", "images":[{"id":"kept"}], "ignored":true}],
"usage_info":usage
}}}],
"usage":{"pages_processed":99}
}))
.unwrap();
let normalized = normalize_response("model", response).unwrap();
assert_eq!(normalized.pages.len(), 1);
assert_eq!(normalized.pages[0].index, 2);
assert_eq!(normalized.pages[0].markdown, "");
assert!(normalized.pages[0].extra_fields.is_empty());
assert_eq!(
normalized
.usage_info
.and_then(|usage| usage.pages_processed),
expected
);
}
}

View file

@ -1,116 +0,0 @@
use crate::call_lifecycle::host::{HostFailure, HostLifecycle, HostPhase};
fn run(fail_at: Option<HostPhase>, asynchronous: bool) -> (Vec<HostPhase>, Vec<crate::ocr::Error>) {
let mut lifecycle = HostLifecycle::new(asynchronous);
let mut events = Vec::new();
let mut failures = Vec::new();
while lifecycle.phase() != HostPhase::Complete {
let phase = lifecycle.phase();
events.push(phase);
let result = if Some(phase) == fail_at {
Err(HostFailure::Error(crate::ocr::Error::InvalidRequest(
"selected failure".into(),
)))
} else {
Ok(())
};
if let Some(error) = lifecycle.accept(result) {
failures.push(error);
}
}
(events, failures)
}
#[test]
fn public_outcome_is_finalized_before_a_single_terminal_dispatch() {
for asynchronous in [false, true] {
let (events, failures) = run(None, asynchronous);
assert!(failures.is_empty());
assert_eq!(
&events[events.len() - 2..],
&[HostPhase::Finalize, HostPhase::Success]
);
assert_eq!(
events
.iter()
.filter(|phase| **phase == HostPhase::Execute)
.count(),
1
);
assert_eq!(
events.contains(&HostPhase::DeploymentPostCall),
asynchronous
);
}
}
#[test]
fn only_provider_and_response_construction_failures_use_provider_mapping() {
for phase in [
HostPhase::Setup,
HostPhase::DeploymentPreCall,
HostPhase::Prepare,
HostPhase::Execute,
HostPhase::ConstructResponse,
HostPhase::DeploymentPostCall,
HostPhase::Finalize,
] {
let (events, failures) = run(Some(phase), true);
assert_eq!(failures.len(), 1);
assert!(!events.contains(&HostPhase::Success));
let mapped = matches!(phase, HostPhase::Execute | HostPhase::ConstructResponse);
assert_eq!(events.contains(&HostPhase::MapFailure), mapped);
assert_eq!(events.contains(&HostPhase::DeploymentFailure), mapped);
assert_eq!(
&events[events.len() - 2..],
&[HostPhase::Failure, HostPhase::AsyncFailure]
);
assert!(
events
.iter()
.filter(|phase| **phase == HostPhase::Execute)
.count()
<= 1
);
}
}
#[test]
fn failure_handler_errors_do_not_replace_selected_failure_or_suppress_async_dispatch() {
let mut lifecycle = HostLifecycle::new(true);
while lifecycle.phase() != HostPhase::Execute {
lifecycle.accept::<crate::ocr::Error>(Ok(()));
}
let selected = crate::ocr::Error::InvalidRequest("provider".into());
assert_eq!(
lifecycle.accept(Err(HostFailure::Error(selected.clone()))),
Some(selected)
);
lifecycle.accept::<crate::ocr::Error>(Ok(()));
for phase in [
HostPhase::DeploymentFailure,
HostPhase::Failure,
HostPhase::AsyncFailure,
] {
assert_eq!(lifecycle.phase(), phase);
assert_eq!(
lifecycle.accept(Err(HostFailure::Error(crate::ocr::Error::InvalidRequest(
"callback".into()
)))),
None
);
}
assert_eq!(lifecycle.phase(), HostPhase::Complete);
}
#[test]
fn cancellation_skips_terminal_dispatch() {
let mut lifecycle = HostLifecycle::new(true);
let error = crate::ocr::Error::InvalidRequest("cancelled".into());
assert_eq!(
lifecycle.accept(Err(HostFailure::Cancelled(error.clone()))),
Some(error)
);
assert_eq!(lifecycle.phase(), HostPhase::Complete);
}

View file

@ -1,895 +0,0 @@
use std::sync::{Arc, Mutex};
use serde_json::{Value, json};
use super::OcrClient;
use super::hooks::{
OcrDuringCallRequest, OcrHookFuture, OcrHooks, OcrLogFuture, OcrPostCallRequest,
OcrPreCallRequest,
};
use super::test_support::{MockResponse, mock_server, perform_ocr, wire_request};
use super::wire::{OcrWireRequest, decode_request};
use super::{
NativeOutcome, NoopOcrHost, OcrAdmission, OcrCall, OcrCallStep, OcrDecline, OcrHost,
OcrHostOperation, OcrHostResult,
};
use crate::call_lifecycle::{CallLifecycleContext, CallLifecycleTiming};
#[test]
fn request_boundary_selects_mistral_and_rejects_unknown_providers() {
let request = OcrWireRequest {
model: "mistral/model".into(),
document: json!({"type":"document_url","document_url":"https://example.com/doc.pdf"}),
api_key: Some("key".into()),
api_base: None,
custom_llm_provider: None,
extra_headers: None,
optional_params: json!({"extract_header":true,"unknown":42})
.as_object()
.unwrap()
.clone()
.into(),
input_sources: Default::default(),
timeout_seconds: None,
};
assert!(decode_request(request).is_ok());
assert!(
decode_request(OcrWireRequest {
model: "model".into(),
document: json!({"type":"document_url","document_url":"https://example.com/doc.pdf"}),
api_key: Some("key".into()),
api_base: None,
custom_llm_provider: Some("unknown".into()),
extra_headers: None,
optional_params: Default::default(),
input_sources: Default::default(),
timeout_seconds: None,
})
.is_err()
);
}
#[tokio::test]
async fn facade_executes_direct_mistral_once() {
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({
"pages":[{"index":0,"markdown":"hello","custom":"preserved"}],
"usage_info":{"pages_processed":1}
}))])
.await;
let result = perform_ocr(wire_request(
"mistral/model",
&base,
json!({"pages":"0,2-4","extract_header":true,"unknown":{"nested":[null,false,0]}}),
))
.await
.unwrap();
server.await.unwrap();
assert_eq!(result.pages[0].markdown, "hello");
assert_eq!(result.pages[0].extra_fields["custom"], "preserved");
let requests = seen.lock().unwrap();
assert_eq!(requests.len(), 1);
assert!(requests[0].starts_with("POST /v1/ocr "));
assert!(
requests[0]
.to_ascii_lowercase()
.contains("authorization: bearer test-key\r\n")
);
let body: Value = serde_json::from_str(requests[0].split_once("\r\n\r\n").unwrap().1).unwrap();
assert_eq!(
body,
json!({
"model":"model",
"document":{"type":"document_url","document_url":"data:application/pdf;base64,YWJj"},
"pages":"0,2-4",
"extract_header":true,
"unknown":{"nested":[null,false,0]}
})
);
}
#[tokio::test]
async fn facade_retains_native_response_when_requested() {
let provider_response = json!({
"pages":[{"index":0,"markdown":"hello"}],
"usage_info":{"pages_processed":1},
"provider_only":"preserved"
});
let (base, _, server) = mock_server(vec![MockResponse::json(provider_response.clone())]).await;
let response = perform_ocr(wire_request(
"mistral/model",
&base,
json!({"req_format":"native"}),
))
.await
.unwrap();
server.await.unwrap();
assert_eq!(
response.provider_native_response.as_ref(),
provider_response.as_object()
);
}
#[tokio::test]
async fn facade_uses_the_injected_http_client() {
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({"pages":[]}))]).await;
let mut default_headers = reqwest::header::HeaderMap::new();
default_headers.insert(
"x-transport-owner",
reqwest::header::HeaderValue::from_static("host"),
);
let provider_http = reqwest::Client::builder()
.default_headers(default_headers)
.build()
.unwrap();
OcrClient::new(provider_http)
.unwrap()
.perform(wire_request("mistral/model", &base, json!({})))
.await
.unwrap();
server.await.unwrap();
assert!(seen.lock().unwrap()[0].contains("x-transport-owner: host"));
}
struct RecordingHooks {
events: Arc<Mutex<Vec<&'static str>>>,
block: bool,
}
struct ExtensionHooks;
impl OcrHooks for ExtensionHooks {
fn intercepts_requests(&self) -> bool {
true
}
fn during_call(
&self,
mut request: OcrDuringCallRequest,
) -> OcrHookFuture<'_, OcrDuringCallRequest> {
Box::pin(async move {
assert_eq!(request.body["pages"], json!([2]));
assert_eq!(request.body.get("future"), Some(&Value::Null));
assert!(
!request
.retained_fields
.iter()
.any(|field| field == "pages" || field == "document")
);
request.body.as_object_mut().unwrap().remove("future");
request.body["hook_option"] = json!({"nested":[null,false,0]});
Ok(request)
})
}
}
#[tokio::test]
async fn composed_extensions_reach_hooks_and_removed_fields_stay_removed() {
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({"pages":[]}))]).await;
let request = super::LiteLLMOcrRequest {
hooks: Arc::new(ExtensionHooks),
..wire_request(
"mistral/model",
&base,
json!({
"pages":[0], "future":null, "extra_body":{"pages":[2],
"document":{"type":"document_url","document_url":"data:application/pdf;base64,eHl6"}}
}),
)
};
perform_ocr(request).await.unwrap();
server.await.unwrap();
let requests = seen.lock().unwrap();
let body: Value = serde_json::from_str(requests[0].split_once("\r\n\r\n").unwrap().1).unwrap();
assert_eq!(body["pages"], json!([2]));
assert_eq!(
body["document"]["document_url"],
"data:application/pdf;base64,eHl6"
);
assert_eq!(body["hook_option"], json!({"nested":[null,false,0]}));
assert!(body.get("future").is_none());
}
impl OcrHooks for RecordingHooks {
fn intercepts_requests(&self) -> bool {
true
}
fn pre_call(&self, request: OcrPreCallRequest) -> OcrHookFuture<'_, OcrPreCallRequest> {
Box::pin(async move {
self.events.lock().unwrap().push("pre");
if self.block {
return Err(crate::ocr::Error::InvalidRequest("blocked".into()));
}
Ok(request)
})
}
fn during_call(
&self,
request: super::hooks::OcrDuringCallRequest,
) -> OcrHookFuture<'_, super::hooks::OcrDuringCallRequest> {
Box::pin(async move {
self.events.lock().unwrap().push("during");
Ok(request)
})
}
fn post_call(&self, request: OcrPostCallRequest) -> OcrHookFuture<'_, OcrPostCallRequest> {
Box::pin(async move {
self.events.lock().unwrap().push("post");
Ok(request)
})
}
fn success<'a>(
&'a self,
_context: &'a CallLifecycleContext,
_response: &'a super::LiteLLMOcrResponse,
_timing: &'a CallLifecycleTiming,
) -> OcrLogFuture<'a> {
Box::pin(async move {
self.events.lock().unwrap().push("success");
})
}
fn failure<'a>(
&'a self,
_context: &'a CallLifecycleContext,
_error: &'a crate::ocr::Error,
_timing: &'a CallLifecycleTiming,
) -> OcrLogFuture<'a> {
Box::pin(async move {
self.events.lock().unwrap().push("failure");
})
}
}
struct HeaderEditHooks;
impl OcrHooks for HeaderEditHooks {
fn intercepts_requests(&self) -> bool {
true
}
fn during_call(
&self,
mut request: OcrDuringCallRequest,
) -> OcrHookFuture<'_, OcrDuringCallRequest> {
request
.headers
.push(("x-core-callback".into(), "edited".into()));
Box::pin(async move { Ok(request) })
}
}
#[tokio::test]
async fn lifecycle_sends_headers_returned_by_the_typed_during_call_operation() {
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({"pages":[]}))]).await;
let request = super::LiteLLMOcrRequest {
hooks: Arc::new(HeaderEditHooks),
..wire_request("mistral/model", &base, json!({}))
};
perform_ocr(request).await.unwrap();
server.await.unwrap();
assert!(seen.lock().unwrap()[0].contains("x-core-callback: edited"));
}
#[tokio::test]
async fn lifecycle_orders_hooks_and_emits_one_success() {
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({"pages":[]}))]).await;
let events = Arc::new(Mutex::new(Vec::new()));
let request = wire_request("mistral/model", &base, json!({}));
let request = super::LiteLLMOcrRequest {
hooks: Arc::new(RecordingHooks {
events: events.clone(),
block: false,
}),
..request
};
perform_ocr(request).await.unwrap();
server.await.unwrap();
assert_eq!(
*events.lock().unwrap(),
["pre", "during", "post", "success"]
);
assert_eq!(seen.lock().unwrap().len(), 1);
}
#[tokio::test]
async fn lifecycle_blocking_prevents_execution_and_emits_one_failure() {
let events = Arc::new(Mutex::new(Vec::new()));
let request = wire_request("mistral/model", "http://127.0.0.1:1", json!({}));
let request = super::LiteLLMOcrRequest {
hooks: Arc::new(RecordingHooks {
events: events.clone(),
block: true,
}),
..request
};
let error = perform_ocr(request).await.unwrap_err();
assert!(matches!(error, crate::ocr::Error::InvalidRequest(_)));
assert_eq!(*events.lock().unwrap(), ["pre", "failure"]);
}
#[tokio::test]
async fn upstream_failure_emits_one_terminal_failure() {
let (base, seen, server) = mock_server(vec![MockResponse {
status: 500,
headers: vec![],
body: json!({"error":"failed"}),
}])
.await;
let events = Arc::new(Mutex::new(Vec::new()));
let request = wire_request("mistral/model", &base, json!({}));
let request = super::LiteLLMOcrRequest {
hooks: Arc::new(RecordingHooks {
events: events.clone(),
block: false,
}),
..request
};
assert!(perform_ocr(request).await.is_err());
server.await.unwrap();
assert_eq!(*events.lock().unwrap(), ["pre", "during", "failure"]);
assert_eq!(seen.lock().unwrap().len(), 1);
}
struct AdmissionSpy {
effects: Arc<Mutex<usize>>,
}
impl OcrHooks for AdmissionSpy {
fn intercepts_requests(&self) -> bool {
*self.effects.lock().unwrap() += 1;
true
}
fn pre_call(&self, request: OcrPreCallRequest) -> OcrHookFuture<'_, OcrPreCallRequest> {
*self.effects.lock().unwrap() += 1;
Box::pin(async move { Ok(request) })
}
}
#[test]
fn admission_declines_without_invoking_hooks_or_transport() {
for (admission, expected) in [
(
OcrAdmission {
provider_workflow: false,
host_operations: true,
asynchronous: false,
},
OcrDecline::ProviderWorkflow,
),
(
OcrAdmission {
provider_workflow: true,
host_operations: false,
asynchronous: false,
},
OcrDecline::HostOperations,
),
] {
let outcome = OcrCall::admit(super::test_support::ocr_client(), admission);
assert!(matches!(outcome, NativeOutcome::Declined(reason) if reason == expected));
}
}
#[tokio::test]
async fn fallible_host_phases_do_not_replay_or_reach_transport() {
for failure_phase in ["pre", "during"] {
let request = super::LiteLLMOcrRequest {
hooks: Arc::new(AdmissionSpy {
effects: Arc::new(Mutex::new(0)),
}),
..wire_request("mistral/model", "http://127.0.0.1:1", json!({}))
};
let NativeOutcome::Completed(mut call) =
OcrCall::admit(super::test_support::ocr_client(), OcrAdmission::all())
else {
panic!("supported call declined")
};
let mut request = Some(request);
let mut result = None;
let mut phases = Vec::new();
let error = loop {
match call.resume(result.take()).await {
Ok(OcrCallStep::Host(operation)) => match operation {
OcrHostOperation::Lifecycle(_)
| OcrHostOperation::ConstructResponse(_)
| OcrHostOperation::MapFailure(_)
| OcrHostOperation::Success { .. }
| OcrHostOperation::Failure { .. } => {
result = Some(OcrHostResult::Lifecycle(Ok(())))
}
OcrHostOperation::ProjectRequest => {
result = Some(OcrHostResult::Request(Ok((
Box::new(request.take().unwrap()),
false,
))))
}
OcrHostOperation::AcquireAzureAdToken => {
panic!("test request has no token provider")
}
OcrHostOperation::PreCall(request) => {
phases.push("pre");
result = Some(OcrHostResult::PreCall(if failure_phase == "pre" {
Err(crate::ocr::Error::InvalidRequest("pre failed".into()))
} else {
Ok(request)
}));
}
OcrHostOperation::DuringCall(request) => {
phases.push("during");
result = Some(OcrHostResult::DuringCall(if failure_phase == "during" {
Err(crate::ocr::Error::InvalidRequest("during failed".into()))
} else {
Ok(request)
}));
}
OcrHostOperation::PostCall(_) => panic!("transport should not be reached"),
},
Err(error) => break error,
Ok(OcrCallStep::Complete(_)) => panic!("failed call completed"),
}
};
assert!(matches!(error, crate::ocr::Error::InvalidRequest(_)));
assert_eq!(
phases
.iter()
.filter(|phase| **phase == failure_phase)
.count(),
1
);
}
}
#[tokio::test]
async fn invalid_provider_response_runs_post_call_before_normalization_failure() {
let (base, seen, server) =
mock_server(vec![MockResponse::json(json!({"pages":"invalid"}))]).await;
let mut request = Some(wire_request("mistral/model", &base, json!({})));
let NativeOutcome::Completed(mut call) =
OcrCall::admit(super::test_support::ocr_client(), OcrAdmission::all())
else {
panic!("supported call declined")
};
let host = NoopOcrHost;
let mut result = None;
let mut post_calls = Vec::new();
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,
))));
}
Ok(OcrCallStep::Host(operation)) => {
if let OcrHostOperation::PostCall(request) = &operation {
post_calls.push(request.original_response.clone());
}
result = Some(host.invoke(operation).await);
}
Err(error) => break error,
Ok(OcrCallStep::Complete(_)) => panic!("invalid provider response completed"),
}
};
server.await.unwrap();
assert!(matches!(
error,
crate::ocr::Error::ResponseField { ref path } if path == "pages"
));
assert_eq!(seen.lock().unwrap().len(), 1);
assert_eq!(post_calls, [json!(r#"{"pages":"invalid"}"#)]);
}
#[tokio::test]
async fn direct_native_host_drives_the_same_state_machine() {
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({
"pages":[{"index":0,"markdown":"native"}]
}))])
.await;
let request = super::LiteLLMOcrRequest {
hooks: Arc::new(AdmissionSpy {
effects: Arc::new(Mutex::new(0)),
}),
..wire_request("mistral/model", &base, json!({}))
};
let NativeOutcome::Completed(mut call) = OcrCall::admit(
super::test_support::ocr_client(),
OcrAdmission {
asynchronous: true,
..OcrAdmission::all()
},
) else {
panic!("supported call declined")
};
let mut request = Some(request);
let host = NoopOcrHost;
let mut result = None;
let mut operations = Vec::new();
let response = loop {
match call.resume(result.take()).await.unwrap() {
OcrCallStep::Host(operation) => {
operations.push(match &operation {
OcrHostOperation::ProjectRequest => "ProjectRequest".into(),
OcrHostOperation::Lifecycle(phase) => format!("{phase:?}"),
OcrHostOperation::PreCall(_) => "PreCall".into(),
OcrHostOperation::DuringCall(_) => "DuringCall".into(),
OcrHostOperation::PostCall(_) => "PostCall".into(),
OcrHostOperation::ConstructResponse(_) => "ConstructResponse".into(),
OcrHostOperation::Success { response, .. } => {
assert_eq!(response.pages[0].markdown, "native");
"Success".into()
}
_ => panic!("unexpected OCR operation"),
});
result = Some(match operation {
OcrHostOperation::ProjectRequest => {
OcrHostResult::Request(Ok((Box::new(request.take().unwrap()), false)))
}
operation => host.invoke(operation).await,
});
}
OcrCallStep::Complete(response) => break response,
}
};
server.await.unwrap();
assert_eq!(response.pages[0].markdown, "native");
assert_eq!(seen.lock().unwrap().len(), 1);
assert_eq!(
operations,
[
"Setup",
"DeploymentPreCall",
"Prepare",
"ProjectRequest",
"PreCall",
"DuringCall",
"PostCall",
"ConstructResponse",
"DeploymentPostCall",
"Finalize",
"Success",
]
);
assert!(matches!(
call.resume(None).await,
Err(crate::ocr::Error::InvalidRequest(_))
));
}
#[tokio::test]
async fn public_finalization_failure_never_dispatches_success_or_replays_provider() {
use crate::call_lifecycle::host::{HostFailure, HostPhase};
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({"pages":[]}))]).await;
let mut request = Some(wire_request("mistral/model", &base, json!({})));
let NativeOutcome::Completed(mut call) = OcrCall::admit(
super::test_support::ocr_client(),
OcrAdmission {
asynchronous: true,
..OcrAdmission::all()
},
) else {
panic!("supported call declined")
};
let selected = crate::ocr::Error::InvalidRequest("public metadata failed".into());
let host = NoopOcrHost;
let mut result = None;
let mut failures = Vec::new();
let error = loop {
match call.resume(result.take()).await {
Ok(OcrCallStep::Host(operation)) => {
result = Some(match operation {
OcrHostOperation::Lifecycle(HostPhase::Finalize) => {
OcrHostResult::Lifecycle(Err(HostFailure::Error(selected.clone())))
}
OcrHostOperation::Failure { error, .. } => {
assert_eq!(error, selected);
failures.push("sync");
OcrHostResult::Lifecycle(Err(HostFailure::Error(
crate::ocr::Error::InvalidRequest("failure callback failed".into()),
)))
}
OcrHostOperation::Lifecycle(HostPhase::AsyncFailure) => {
failures.push("async");
OcrHostResult::Lifecycle(Ok(()))
}
OcrHostOperation::Success { .. }
| OcrHostOperation::MapFailure(_)
| OcrHostOperation::Lifecycle(HostPhase::DeploymentFailure) => {
panic!("finalization failure used provider/success dispatch")
}
OcrHostOperation::ProjectRequest => {
OcrHostResult::Request(Ok((Box::new(request.take().unwrap()), false)))
}
operation => host.invoke(operation).await,
});
}
Ok(OcrCallStep::Complete(_)) => panic!("failed call completed successfully"),
Err(error) => break error,
}
};
server.await.unwrap();
assert_eq!(error, selected);
assert_eq!(failures, ["sync", "async"]);
assert_eq!(seen.lock().unwrap().len(), 1);
}
#[tokio::test]
async fn cancellation_at_provider_hook_prevents_execution_and_further_resumption() {
use crate::call_lifecycle::host::HostFailure;
let request = super::LiteLLMOcrRequest {
hooks: Arc::new(AdmissionSpy {
effects: Arc::new(Mutex::new(0)),
}),
..wire_request("mistral/model", "http://127.0.0.1:1", json!({}))
};
let NativeOutcome::Completed(mut call) =
OcrCall::admit(super::test_support::ocr_client(), OcrAdmission::all())
else {
panic!("supported call declined")
};
let mut request = Some(request);
let host = NoopOcrHost;
let mut result = None;
loop {
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,
))))
}
OcrCallStep::Host(operation) => result = Some(host.invoke(operation).await),
OcrCallStep::Complete(_) => panic!("provider executed before pre-call result"),
}
}
let selected = crate::ocr::Error::InvalidRequest("cancelled".into());
assert!(matches!(
call.interrupt(HostFailure::Cancelled(selected.clone())).await,
Err(error) if error == selected
));
assert!(
call.resume(Some(OcrHostResult::Lifecycle(Ok(()))))
.await
.is_err()
);
}
#[tokio::test]
async fn missing_host_result_preserves_pending_operation() {
use crate::call_lifecycle::host::HostPhase;
let NativeOutcome::Completed(mut call) =
OcrCall::admit(super::test_support::ocr_client(), OcrAdmission::all())
else {
panic!("supported call declined")
};
assert!(matches!(
call.resume(None).await.unwrap(),
OcrCallStep::Host(OcrHostOperation::Lifecycle(HostPhase::Setup))
));
assert!(call.resume(None).await.is_err());
assert!(matches!(
call.resume(Some(OcrHostResult::Lifecycle(Ok(()))))
.await
.unwrap(),
OcrCallStep::Host(OcrHostOperation::Lifecycle(HostPhase::Prepare))
));
}
async fn read_bounded_response(
response: Vec<u8>,
limit: usize,
) -> Result<bytes::Bytes, crate::ocr::Error> {
use tokio::io::{AsyncReadExt, AsyncWriteExt};
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let address = listener.local_addr().unwrap();
let server = tokio::spawn(async move {
let (mut socket, _) = listener.accept().await.unwrap();
let mut request = [0; 4096];
assert!(socket.read(&mut request).await.unwrap() > 0);
socket.write_all(&response).await.unwrap();
std::future::pending::<()>().await;
});
let response = reqwest::Client::new()
.get(format!("http://{address}"))
.send()
.await
.unwrap();
let result = tokio::time::timeout(
std::time::Duration::from_secs(2),
super::client::read_response_bytes(response, limit),
)
.await;
server.abort();
let _ = server.await;
result.expect("bounded reads must finish without waiting for the rest of an oversized body")
}
#[tokio::test]
async fn response_limit_accepts_exact_size_and_rejects_declared_and_chunked_overflow() {
for response in [
"HTTP/1.1 200 OK\r\nContent-Length: 8\r\n\r\nabcdefgh",
"HTTP/1.1 200 OK\r\nTransfer-Encoding: chunked\r\n\r\n4\r\nabcd\r\n4\r\nefgh\r\n0\r\n\r\n",
] {
assert_eq!(
read_bounded_response(response.as_bytes().to_vec(), 8)
.await
.unwrap(),
"abcdefgh"
);
}
for response in [
"HTTP/1.1 200 OK\r\nContent-Length: 9\r\n\r\n",
"HTTP/1.1 200 OK\r\nTransfer-Encoding: chunked\r\n\r\n4\r\nabcd\r\n5\r\nefghi\r\n",
] {
assert!(matches!(
read_bounded_response(response.as_bytes().to_vec(), 8).await,
Err(crate::ocr::Error::TooLarge { limit: 8 })
));
}
}
#[tokio::test]
async fn oversized_error_retains_http_status_and_bounded_diagnostics_without_draining() {
let prefix = "x".repeat(4 * (crate::constants::UPSTREAM_ERROR_BODY_MAX_CHARS + 1));
for headers in ["Content-Length: 1000000", "Transfer-Encoding: chunked"] {
let body = if headers.starts_with("Transfer") {
format!("{:x}\r\n{prefix}\r\n", prefix.len())
} else {
prefix.clone()
};
let response = format!("HTTP/1.1 429 Too Many Requests\r\n{headers}\r\n\r\n{body}");
let error = read_bounded_response(response.into_bytes(), 4096)
.await
.unwrap_err();
match error {
crate::ocr::Error::Transport(crate::transport::Error::Http { status, body }) => {
assert_eq!(status, 429);
assert_eq!(
body,
format!(
"{}... (truncated)",
"x".repeat(crate::constants::UPSTREAM_ERROR_BODY_MAX_CHARS)
)
);
}
error => panic!("unexpected error: {error}"),
}
}
}
#[test]
fn response_limit_is_validated_and_not_forwarded_to_the_provider() {
let request = wire_request(
"mistral/model",
"http://localhost",
json!({"max_response_bytes": 123}),
);
assert_eq!(request.connection.max_response_bytes, 123);
assert!(!request.optional_params.contains_key("max_response_bytes"));
for value in [
json!(0),
json!(-1),
json!(true),
json!("123"),
json!(1.5),
json!(crate::constants::OCR_RESPONSE_MAX_BYTES + 1),
Value::Null,
] {
let wire = serde_json::from_value(json!({
"model": "mistral/model", "document": {"type": "document_url", "document_url": "data:application/pdf;base64,YWJj"},
"optional_params": {"max_response_bytes": value}
})).unwrap();
let Err(error) = decode_request(wire) else {
panic!("invalid response limit accepted")
};
assert!(error.to_string().contains("max_response_bytes"));
}
}
#[derive(Debug)]
struct PendingToken {
entered: Arc<tokio::sync::Notify>,
dropped: Arc<std::sync::atomic::AtomicBool>,
}
struct TokenFutureDrop(Arc<std::sync::atomic::AtomicBool>);
impl Drop for TokenFutureDrop {
fn drop(&mut self) {
self.0.store(true, std::sync::atomic::Ordering::SeqCst);
}
}
impl litellm_auth::TokenProvider for PendingToken {
fn acquire(&self) -> litellm_auth::TokenFuture<'_> {
Box::pin(async move {
let _guard = TokenFutureDrop(self.dropped.clone());
self.entered.notify_one();
std::future::pending().await
})
}
}
#[tokio::test]
async fn cancellation_waits_for_provider_capture_drop_even_when_acknowledgement_is_cancelled() {
use crate::call_lifecycle::host::HostFailure;
use std::future::Future;
use std::sync::atomic::{AtomicBool, Ordering};
use std::task::Poll;
for interrupt_acknowledgement in [false, true] {
let entered = Arc::new(tokio::sync::Notify::new());
let dropped = Arc::new(AtomicBool::new(false));
let request = wire_request("azure_ai/mistral-ocr", "https://example.invalid", json!({}));
let request = super::LiteLLMOcrRequest {
connection: super::OcrConnection {
extra_headers: vec![("authorization".into(), "Bearer test-key".into())],
..request.connection
},
azure_ad_token_provider: Some(litellm_auth::TokenProviderHandle::new(Arc::new(
PendingToken {
entered: entered.clone(),
dropped: dropped.clone(),
},
))),
..request
};
let NativeOutcome::Completed(mut call) =
OcrCall::admit(super::test_support::ocr_client(), OcrAdmission::all())
else {
panic!("supported call declined")
};
let mut request = Some(request);
let mut result = None;
tokio::time::timeout(std::time::Duration::from_secs(2), async {
loop {
tokio::select! {
_ = 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(operation) => NoopOcrHost.invoke(operation).await,
OcrCallStep::Complete(_) => panic!("pending provider completed"),
});
}
}
}
}).await.unwrap();
assert!(!dropped.load(Ordering::SeqCst));
let selected = crate::ocr::Error::InvalidRequest("cancelled".into());
if interrupt_acknowledgement {
let mut acknowledgement =
Box::pin(call.interrupt(HostFailure::Cancelled(selected.clone())));
std::future::poll_fn(|cx| {
assert!(acknowledgement.as_mut().poll(cx).is_pending());
Poll::Ready(())
})
.await;
drop(acknowledgement);
assert!(!dropped.load(Ordering::SeqCst));
}
let result = tokio::time::timeout(
std::time::Duration::from_secs(2),
call.interrupt(HostFailure::Cancelled(selected.clone())),
)
.await
.unwrap();
assert!(matches!(result, Err(error) if error == selected));
assert!(
dropped.load(Ordering::SeqCst),
"cancellation returned while provider captures were still alive"
);
}
}

View file

@ -1,331 +0,0 @@
use std::sync::Arc;
use rstest::rstest;
use serde_json::{Value, json};
use super::hooks::{OcrDuringCallRequest, OcrHookFuture, OcrHooks, OcrPostCallRequest};
use super::test_support::{MockResponse, mock_server, perform_ocr, wire_request};
fn request_body(request: &str) -> Value {
serde_json::from_str(request.split_once("\r\n\r\n").unwrap().1).unwrap()
}
#[rstest]
#[case(
"reducto/parse-v3",
json!({
"formatting":{"table_output_format":"html"},
"retrieval":{"chunk_mode":"section"},
"settings":{"ocr_system":"standard"},
"future_ocr_option":true,
"extra_body":{"provider_option":"value"}
}),
"reducto://already.pdf",
json!({
"input":"reducto://already.pdf",
"formatting":{"table_output_format":"html"},
"retrieval":{"chunk_mode":"section"},
"settings":{"ocr_system":"standard"},
"future_ocr_option":true,
"provider_option":"value"
})
)]
#[case(
"reducto/parse-legacy",
json!({
"enhance":{"agentic":[{"type":"table"}]},
"future_ocr_option":true,
"extra_body":{"provider_option":"value"}
}),
"reducto://legacy.pdf",
json!({
"document_url":"reducto://legacy.pdf",
"options":{"enhance":{"agentic":[{"type":"table"}]}},
"future_ocr_option":true,
"provider_option":"value"
})
)]
#[tokio::test]
async fn request_mapping_matches_python(
#[case] model: &str,
#[case] options: Value,
#[case] source: &str,
#[case] expected: Value,
) {
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({
"result":{"chunks":[]}
}))])
.await;
let mut request = wire_request(model, &base, options);
request.document = request.document.with_source(source.into());
perform_ocr(request).await.unwrap();
server.await.unwrap();
let requests = seen.lock().unwrap();
assert_eq!(requests.len(), 1);
assert!(requests[0].starts_with("POST /parse "));
assert_eq!(request_body(&requests[0]), expected);
}
#[rstest]
#[case("parse-v3")]
#[case("parse-legacy")]
#[tokio::test]
async fn data_uri_upload_preserves_multipart_headers(#[case] model: &str) {
let (base, seen, server) = mock_server(vec![
MockResponse::json(json!({"file_id":"reducto://uploaded.pdf"})),
MockResponse::json(json!({"result":{"chunks":[{"content":"hello"}]}})),
])
.await;
let mut request = wire_request(&format!("reducto/{model}"), &base, json!({}));
request.connection.extra_headers = vec![
("Content-Type".into(), "application/json".into()),
("X-Trace".into(), "upload-test".into()),
];
let response = perform_ocr(request).await.unwrap();
server.await.unwrap();
assert_eq!(response.pages[0].markdown, "hello");
let requests = seen.lock().unwrap();
assert_eq!(requests.len(), 2);
assert!(requests[0].starts_with("POST /upload "));
assert!(
requests[0]
.to_ascii_lowercase()
.contains("content-type: multipart/form-data; boundary=")
);
assert!(requests[0].contains("x-trace: upload-test"));
assert!(requests[0].contains("application/pdf"));
assert!(requests[0].contains("abc"));
assert!(requests[1].starts_with("POST /parse "));
}
struct ParseBoundary {
request_count: Arc<std::sync::Mutex<Vec<String>>>,
}
impl OcrHooks for ParseBoundary {
fn post_call(&self, request: OcrPostCallRequest) -> OcrHookFuture<'_, OcrPostCallRequest> {
Box::pin(async move {
assert_eq!(self.request_count.lock().unwrap().len(), 2);
assert_eq!(
request.original_response,
json!(r#"{"result":{"chunks":[]}}"#)
);
Ok(request)
})
}
}
#[tokio::test]
async fn post_call_stays_after_reducto_upload_and_parse() {
let (base, seen, server) = mock_server(vec![
MockResponse::json(json!({"file_id":"reducto://uploaded.pdf"})),
MockResponse::json(json!({"result":{"chunks":[]}})),
])
.await;
let request = super::LiteLLMOcrRequest {
hooks: Arc::new(ParseBoundary {
request_count: seen.clone(),
}),
..wire_request("reducto/parse-v3", &base, json!({}))
};
perform_ocr(request).await.unwrap();
server.await.unwrap();
assert_eq!(seen.lock().unwrap().len(), 2);
}
#[rstest]
#[case(json!({"file_id":""}))]
#[case(json!({}))]
#[case(json!({"file_id":null}))]
#[tokio::test]
async fn invalid_upload_ids_stop_before_parse(#[case] response: Value) {
let (base, seen, server) = mock_server(vec![MockResponse::json(response)]).await;
let error = perform_ocr(wire_request("reducto/parse-v3", &base, json!({})))
.await
.unwrap_err();
server.await.unwrap();
assert!(error.to_string().contains("file_id"));
assert_eq!(seen.lock().unwrap().len(), 1);
}
#[tokio::test]
async fn upload_failure_stops_before_parse() {
let (base, seen, server) = mock_server(vec![MockResponse {
status: 503,
headers: vec![],
body: json!({"error":"unavailable"}),
}])
.await;
assert!(
perform_ocr(wire_request("reducto/parse-v3", &base, json!({})))
.await
.is_err()
);
server.await.unwrap();
assert_eq!(seen.lock().unwrap().len(), 1);
}
#[rstest]
#[case("https://example.com/a.pdf")]
#[case("reducto://")]
#[case("data:application/pdf;base64")]
#[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());
assert!(perform_ocr(request).await.is_err());
}
#[test]
fn response_normalization_groups_blocks_and_distinguishes_null_result() {
use crate::llms::reducto::ocr::transformation::{ReductoResponse, normalize_response};
let raw = json!({"usage":{"num_pages":"2","credits":"3"},"result":{"type":"full","chunks":[
{"blocks":[{
"type":"Table",
"content":"B",
"bbox":{"left":0.1,"top":0.2,"width":0.8,"height":0.3,"page":2,"original_page":4},
"confidence":"high",
"granular_confidence":{"parse_confidence":0.95,"extract_confidence":null},
"image_url":null
}]},
{"blocks":[{"content":"A","bbox":{"page":1},"type":"Text"},{"content":"C","bbox":{"page":1}}]}
]}});
let response: ReductoResponse = serde_json::from_value(raw).unwrap();
let normalized = normalize_response("parse-v3", response)
.unwrap()
.into_json();
assert_eq!(normalized["pages"][0]["markdown"], "A\n\nC");
assert_eq!(normalized["pages"][1]["markdown"], "B");
assert_eq!(normalized["pages"][1]["blocks"][0]["type"], "Table");
assert_eq!(
normalized["pages"][1]["blocks"][0]["bbox"],
json!({"left":0.1,"top":0.2,"width":0.8,"height":0.3,"page":2,"original_page":4})
);
assert_eq!(normalized["pages"][1]["blocks"][0]["confidence"], "high");
assert_eq!(
normalized["pages"][1]["blocks"][0]["granular_confidence"]["parse_confidence"],
0.95
);
assert!(normalized["pages"][1]["blocks"][0]["image_url"].is_null());
assert_eq!(normalized["usage_info"]["pages_processed"], 2);
assert_eq!(normalized["usage_info"]["credits"], 3.0);
let missing: ReductoResponse =
serde_json::from_value(json!({"chunks":[{"content":"text"}]})).unwrap();
let missing = normalize_response("parse-v3", missing).unwrap();
assert_eq!(missing.pages[0].markdown, "text");
let null: ReductoResponse = serde_json::from_value(
json!({"result":null,"chunks":[{"content":"ignored"}],"usage":null}),
)
.unwrap();
let null = normalize_response("parse-v3", null).unwrap();
assert!(null.pages.is_empty());
}
#[tokio::test]
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());
request.connection.extra_headers = vec![("authorization".into(), "Bearer existing".into())];
let response = perform_ocr(request).await.unwrap();
server.await.unwrap();
assert_eq!(response.provider_native_response, None);
assert!(
seen.lock().unwrap()[0]
.to_ascii_lowercase()
.contains("authorization: bearer existing")
);
}
struct RewriteDocument;
struct RewriteHeaders;
impl OcrHooks for RewriteHeaders {
fn intercepts_requests(&self) -> bool {
true
}
fn during_call(
&self,
request: OcrDuringCallRequest,
) -> OcrHookFuture<'_, OcrDuringCallRequest> {
Box::pin(async move {
Ok(OcrDuringCallRequest {
headers: vec![("authorization".into(), "Bearer guarded".into())],
..request
})
})
}
}
#[rstest]
#[case("reducto/parse-v3")]
#[case("reducto/parse-legacy")]
#[tokio::test]
async fn guardrail_headers_reach_upload_and_parse(#[case] model: &str) {
let (base, seen, server) = mock_server(vec![
MockResponse::json(json!({"file_id":"reducto://uploaded.pdf"})),
MockResponse::json(json!({"result":{"chunks":[]}})),
])
.await;
let mut request = wire_request(model, &base, json!({}));
request.connection.extra_headers = vec![("authorization".into(), "Bearer original".into())];
request.hooks = Arc::new(RewriteHeaders);
perform_ocr(request).await.unwrap();
server.await.unwrap();
let requests = seen.lock().unwrap();
assert_eq!(requests.len(), 2);
assert!(requests[0].starts_with("POST /upload "));
assert!(requests[1].starts_with("POST /parse "));
for request in requests.iter() {
assert!(request.contains("authorization: Bearer guarded"));
assert!(!request.contains("Bearer original"));
}
}
impl OcrHooks for RewriteDocument {
fn intercepts_requests(&self) -> bool {
true
}
fn during_call(
&self,
request: OcrDuringCallRequest,
) -> OcrHookFuture<'_, OcrDuringCallRequest> {
Box::pin(async move {
assert_eq!(
request.body["document_url"],
"data:application/pdf;base64,YWJj"
);
Ok(OcrDuringCallRequest {
body: json!({"type":"document_url","document_url":"reducto://guarded.pdf"}),
..request
})
})
}
}
#[tokio::test]
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!({}));
request.hooks = Arc::new(RewriteDocument);
perform_ocr(request).await.unwrap();
server.await.unwrap();
let requests = seen.lock().unwrap();
assert_eq!(requests.len(), 1);
assert!(requests[0].starts_with("POST /parse "));
assert!(requests[0].contains("reducto://guarded.pdf"));
}

View file

@ -1,88 +0,0 @@
use serde_json::{Value, json};
use super::test_support::{MockResponse, mock_server, perform_ocr, wire_request};
use litellm_auth::InputSource;
fn request_body(request: &str) -> Value {
serde_json::from_str(request.split_once("\r\n\r\n").unwrap().1).unwrap()
}
#[tokio::test]
async fn facade_executes_vertex_deepseek_at_the_openai_endpoint() {
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({
"choices":[{"message":{"content":"recognized"}}],
"usage":{"prompt_tokens":1}
}))])
.await;
let mut request = wire_request(
"vertex_ai/deepseek-ocr-maas",
&base,
json!({
"vertex_project":"project-1",
"vertex_location":"europe-west4",
"temperature":0.1,
"future_ocr_option":true,
"extra_body":{"provider_option":"value"}
}),
);
request.document = request
.document
.with_source("gs://bucket/document.pdf".into());
let response = perform_ocr(request).await.unwrap();
server.await.unwrap();
assert_eq!(response.pages[0].markdown, "recognized");
assert_eq!(
response.usage_info.unwrap().extra_fields["prompt_tokens"],
1
);
let requests = seen.lock().unwrap();
assert!(requests[0].starts_with(
"POST /v1/projects/project-1/locations/europe-west4/endpoints/openapi/chat/completions "
));
assert!(
requests[0]
.to_ascii_lowercase()
.contains("authorization: bearer test-key")
);
let body = request_body(&requests[0]);
assert_eq!(body["model"], "deepseek-ai/deepseek-ocr-maas");
assert_eq!(body["temperature"], 0.1);
assert_eq!(body["future_ocr_option"], true);
assert_eq!(body["provider_option"], "value");
assert!(body.get("vertex_project").is_none());
assert!(body.get("extra_body").is_none());
assert_eq!(
body["messages"][0]["content"][0],
json!({"type":"image_url","image_url":"gs://bucket/document.pdf"})
);
}
#[test]
fn host_registration_selects_deepseek_without_affecting_mistral() {
assert!(crate::ocr::wire::is_supported_request(
"deepseek-ocr-maas",
Some("vertex_ai")
));
assert!(crate::ocr::wire::is_supported_request(
"mistral-ocr-maas",
Some("vertex_ai")
));
}
#[tokio::test]
async fn request_controlled_api_base_is_rejected_before_vertex_auth() {
let mut request = wire_request(
"vertex_ai/deepseek-ocr-maas",
"https://caller.example",
json!({"vertex_project":"project-1"}),
);
request.connection.api_base_source = InputSource::Request;
let error = perform_ocr(request).await.unwrap_err();
assert!(
error
.to_string()
.contains("request-controlled Vertex AI endpoint")
);
}

View file

@ -1,167 +0,0 @@
use serde_json::{Value, json};
use super::test_support::{MockResponse, mock_server, perform_ocr, wire_request};
use litellm_auth::InputSource;
fn request_body(request: &str) -> Value {
serde_json::from_str(request.split_once("\r\n\r\n").unwrap().1).unwrap()
}
#[tokio::test]
async fn facade_executes_vertex_mistral_with_resolved_project_and_location() {
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({
"pages":[{"index":0,"markdown":"hello"}],
"usage_info":{"pages_processed":1}
}))])
.await;
let request = wire_request(
"vertex_ai/mistral-ocr-maas",
&base,
json!({
"vertex_project":"project-1",
"vertex_location":"europe-west4",
"extract_footer":true
}),
);
let response = perform_ocr(request).await.unwrap();
server.await.unwrap();
assert_eq!(response.pages[0].markdown, "hello");
let requests = seen.lock().unwrap();
assert_eq!(requests.len(), 1);
assert!(requests[0].starts_with(
"POST /v1/projects/project-1/locations/europe-west4/publishers/mistralai/models/mistral-ocr-maas:rawPredict "
));
assert!(
requests[0]
.to_ascii_lowercase()
.contains("authorization: bearer test-key")
);
assert_eq!(
request_body(&requests[0]),
json!({
"model":"mistral-ocr-maas",
"document":{"type":"document_url","document_url":"data:application/pdf;base64,YWJj"},
"extract_footer":true
})
);
}
#[tokio::test]
async fn supplied_authorization_is_forwarded_without_a_static_token() {
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({"pages":[]}))]).await;
let mut request = wire_request(
"vertex_ai/model",
&base,
json!({"vertex_project":"project-1"}),
);
request.connection.api_key = None;
request.connection.extra_headers = vec![("authorization".into(), "Bearer supplied".into())];
perform_ocr(request).await.unwrap();
server.await.unwrap();
assert!(
seen.lock().unwrap()[0]
.to_ascii_lowercase()
.contains("authorization: bearer supplied")
);
}
#[tokio::test]
async fn invalid_credentials_fail_before_provider_http() {
let request = wire_request(
"vertex_ai/model",
"http://127.0.0.1:1",
json!({"vertex_credentials": true}),
);
let error = perform_ocr(request).await.unwrap_err();
assert!(error.to_string().contains("vertex_credentials"));
}
#[tokio::test]
async fn request_controlled_api_base_is_rejected_before_vertex_auth() {
let mut request = wire_request(
"vertex_ai/mistral-ocr-maas",
"https://caller.example",
json!({"vertex_project":"project-1"}),
);
request.connection.api_base_source = InputSource::Request;
let error = perform_ocr(request).await.unwrap_err();
assert!(
error
.to_string()
.contains("request-controlled Vertex AI endpoint")
);
}
#[tokio::test]
async fn configs_build_complete_requests_and_share_mistral_normalization() {
use std::time::Duration;
use crate::llms::base_llm::ocr::transformation::BaseOcrConfig;
use crate::llms::mistral::ocr::transformation::MistralOCRConfig;
use crate::llms::vertex_ai::ocr::transformation::VertexAIOCRConfig;
use crate::ocr::test_support::ocr_client;
let client = ocr_client();
let options = json!({
"pages": [0, 2],
"include_image_base64": true,
"vertex_project": "project-1",
"vertex_location": "us-central1",
"unknown": "preserved"
});
let direct = wire_request(
"mistral/mistral-ocr-maas",
"https://mistral.test",
options.clone(),
);
let vertex = wire_request("vertex_ai/mistral-ocr-maas", "https://vertex.test", options);
let direct_http = MistralOCRConfig
.prepare_request(&direct, &client)
.await
.unwrap();
let vertex_http = VertexAIOCRConfig
.prepare_request(&vertex, &client)
.await
.unwrap();
assert_eq!(direct_http.url().as_str(), "https://mistral.test/v1/ocr");
assert_eq!(
vertex_http.url().as_str(),
"https://vertex.test/v1/projects/project-1/locations/us-central1/publishers/mistralai/models/mistral-ocr-maas:rawPredict"
);
for http in [&direct_http, &vertex_http] {
assert_eq!(http.method(), reqwest::Method::POST);
assert_eq!(http.headers()["authorization"], "Bearer test-key");
assert_eq!(http.headers()["content-type"], "application/json");
assert_eq!(http.timeout(), Some(&Duration::from_secs(2)));
let body: Value = serde_json::from_slice(http.body().unwrap().as_bytes().unwrap()).unwrap();
assert_eq!(
body,
json!({
"model": "mistral-ocr-maas",
"document": {"type": "document_url", "document_url": "data:application/pdf;base64,YWJj"},
"pages": [0, 2],
"include_image_base64": true,
"unknown": "preserved"
})
);
}
let payload = json!({"pages": [{"index": 0, "markdown": "hello"}], "extra": "preserved"});
let direct_response = MistralOCRConfig
.normalize_response(
&direct.model,
serde_json::from_value(payload.clone()).unwrap(),
)
.unwrap()
.into_json();
let vertex_response = VertexAIOCRConfig
.normalize_response(&vertex.model, serde_json::from_value(payload).unwrap())
.unwrap()
.into_json();
assert_eq!(direct_response, vertex_response);
assert_eq!(direct_response["model"], "mistral-ocr-maas");
assert_eq!(direct_response["object"], "ocr");
assert!(direct_response.get("extra").is_none());
}

View file

@ -6,6 +6,21 @@ use pyo3::prelude::*;
use crate::errors::{RustUpstreamError, core_error_to_pyerr};
pub(super) fn to_pyerr(error: Error) -> PyErr {
if let Error::Provider {
status,
body,
headers,
} = error
{
let mapped = attach_status(RustUpstreamError::new_err((status, body)), Some(status));
return Python::attach(|py| -> PyResult<PyErr> {
let headers =
pyo3::types::PyDict::from_sequence(&headers.into_pyobject(py)?.into_any())?;
mapped.value(py).setattr("headers", headers)?;
Ok(mapped)
})
.unwrap_or_else(|error| error);
}
let (mapped, status) = match error {
Error::MissingDocumentUrl => (
PyValueError::new_err(Error::MissingDocumentUrl.to_string()),
@ -47,6 +62,39 @@ mod tests {
use super::*;
use pyo3::exceptions::PyValueError;
#[test]
fn provider_error_retains_headers_at_the_python_boundary() {
Python::initialize();
Python::attach(|py| {
let error = to_pyerr(Error::Provider {
status: 429,
body: "rate limited".into(),
headers: vec![("retry-after".into(), "17".into())],
});
assert!(error.is_instance_of::<RustUpstreamError>(py));
assert_eq!(
error
.value(py)
.getattr("args")
.unwrap()
.extract::<(u16, String)>()
.unwrap(),
(429, "rate limited".into())
);
assert_eq!(
error
.value(py)
.getattr("headers")
.unwrap()
.get_item("retry-after")
.unwrap()
.extract::<String>()
.unwrap(),
"17"
);
});
}
#[test]
fn preserves_python_validation_and_provider_details() {
Python::initialize();