remove prepare_request

This commit is contained in:
Yujong Lee 2026-09-15 19:32:11 -07:00
parent 6c893a7dba
commit 22a593d606
30 changed files with 1059 additions and 1316 deletions

View file

@ -23,11 +23,12 @@ pub struct ArgumentError {
}
pub fn parse_options<T: DeserializeOwned>(arguments: &CallArguments) -> Result<T, ArgumentError> {
use serde::de::IntoDeserializer;
serde_path_to_error::deserialize(Value::Object(arguments.0.clone()).into_deserializer())
.map_err(|error| ArgumentError {
path: error.path().to_string(),
})
let deserializer = serde::de::value::MapDeserializer::new(
arguments.iter().map(|(name, value)| (name.as_str(), value)),
);
serde_path_to_error::deserialize(deserializer).map_err(|error| ArgumentError {
path: error.path().to_string(),
})
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]

View file

@ -1,13 +1,12 @@
use crate::call_arguments::CallArguments;
use crate::llms::base_llm::ocr::transformation::{BaseOcrConfig, OcrRequestContext};
use crate::llms::cohere::ocr::transformation::{CohereParseConfig, CohereRequest};
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::{LiteLLMOcrResponse, OcrDocument, PreparedOcrRequest};
use crate::url_utils::ApiUrl;
const AZURE_AI_API_BASE_ENV: &str = "AZURE_AI_API_BASE";
use serde_json::Value;
#[derive(Default)]
pub(crate) struct AzureAICohereParseConfig;
@ -44,17 +43,10 @@ impl BaseOcrConfig for AzureAICohereParseConfig {
_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(),
))
})?;
let base = super::transformation::AzureAIOCRConfig::resolve_api_base(
request.connection.api_base.as_deref(),
&crate::ocr::prepare::credential_env,
)?;
self.get_complete_url(&base)
}
@ -72,6 +64,14 @@ impl BaseOcrConfig for AzureAICohereParseConfig {
CohereParseConfig.get_supported_ocr_params(model)
}
fn map_ocr_params(
&self,
arguments: &CallArguments,
model: &str,
) -> Result<CohereOptions, crate::ocr::Error> {
CohereParseConfig.map_ocr_params(arguments, model)
}
async fn async_transform_ocr_request(
&self,
model: &str,
@ -98,37 +98,15 @@ impl BaseOcrConfig for AzureAICohereParseConfig {
) -> Result<LiteLLMOcrResponse, crate::ocr::Error> {
CohereParseConfig.transform_ocr_response(model, raw_response, request_format)
}
}
impl AzureAICohereParseConfig {
pub(crate) async fn prepare_request(
&self,
request: &PreparedOcrRequest,
client: &OcrClient,
) -> Result<reqwest::Request, crate::ocr::Error> {
let params = self.map_ocr_params(&request.optional_params, &request.model)?;
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
.async_transform_ocr_request(
&request.model,
request.document.clone(),
&params,
&headers,
OcrRequestContext {
client,
connection: &request.connection,
},
)
.await?;
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
fn retains_document(&self, document: &OcrDocument) -> bool {
!document.is_remote()
}
fn validate_request_body(&self, body: &Value) -> Result<(), crate::ocr::Error> {
let document = crate::ocr::prepare::body_document(body)?;
validate_document(&document)?;
validate_inline_document(&document)
}
}

View file

@ -24,37 +24,16 @@ use crate::ocr::OcrClient;
use crate::ocr::client::read_json_response;
use crate::ocr::document::InlineDocument;
use crate::ocr::hooks::OcrHooks;
use crate::ocr::prepare::{credential_env, transform_request_body};
use crate::ocr::json::DecodedOcrResponse;
use crate::ocr::prepare::credential_env;
use crate::ocr::types::{
LiteLLMOcrResponse, OcrConnection, OcrCredentialInputs, OcrDocument, OcrPage,
OcrPageDimensions, OcrResponseFormat, OcrUsageInfo, PreparedOcrRequest, ResolvedOcrCredentials,
};
use crate::ocr::wire::DecodedOcrResponse;
use crate::serde_compat::{FiniteF64, LaxI64};
use crate::url_utils::ApiUrl;
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
#[serde(untagged)]
enum PagesInput {
ZeroBasedIndices(Vec<i64>),
NativeTokens(Vec<String>),
NativeRange(String),
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
#[serde(untagged)]
enum FeaturesInput {
Names(Vec<String>),
CommaSeparated(String),
}
#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)]
struct DocumentIntelligenceInputParams {
pub pages: Option<PagesInput>,
pub features: Option<FeaturesInput>,
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
#[derive(Clone, Debug, PartialEq, Serialize)]
pub(crate) struct DocumentIntelligenceParams {
#[serde(skip_serializing_if = "Option::is_none")]
pub pages: Option<String>,
@ -145,79 +124,46 @@ struct AzureDocumentIntelligenceLine {
pub content: Option<String>,
}
fn decode_input_params(
params: Map<String, Value>,
prefix: &str,
) -> Result<DocumentIntelligenceInputParams, crate::ocr::Error> {
if let Some(Value::Array(pages)) = params.get("pages") {
if pages.iter().any(Value::is_boolean) {
return Err(crate::ocr::Error::Pages("boolean page index".into()));
}
if pages
.iter()
.any(|page| page.is_number() && page.as_i64().is_none())
{
return Err(crate::ocr::Error::Pages(
"page index is out of range".into(),
));
}
if !pages.iter().all(Value::is_i64) && !pages.iter().all(Value::is_string) {
return Err(crate::ocr::Error::Pages("mixed page element types".into()));
}
}
crate::ocr::wire::decode_request_value(Value::Object(params), prefix)
}
fn normalize_ocr_params(
params: DocumentIntelligenceInputParams,
) -> Result<DocumentIntelligenceParams, crate::ocr::Error> {
Ok(DocumentIntelligenceParams {
pages: params.pages.map(normalize_pages).transpose()?.flatten(),
features: params
.features
.map(normalize_features)
.transpose()?
.flatten(),
})
}
fn normalize_pages(pages: PagesInput) -> Result<Option<String>, crate::ocr::Error> {
fn normalize_pages(pages: Option<&Value>) -> Result<Option<String>, crate::ocr::Error> {
let normalized = match pages {
PagesInput::ZeroBasedIndices(indices) => {
if indices.is_empty() {
return Ok(None);
}
indices
.into_iter()
.map(|page| {
if page < 0 {
return Err(crate::ocr::Error::Pages("negative page index".into()));
}
page.checked_add(1).ok_or_else(|| {
crate::ocr::Error::Pages("page index is out of range".into())
})
None | Some(Value::Null) => return Ok(None),
Some(Value::Array(pages)) if pages.is_empty() => return Ok(None),
Some(Value::Array(pages)) if pages.iter().all(Value::is_number) => pages
.iter()
.map(|page| {
let page = page
.as_i64()
.ok_or_else(|| crate::ocr::Error::Pages("page index is out of range".into()))?;
if page < 0 {
return Err(crate::ocr::Error::Pages("negative page index".into()));
}
page.checked_add(1)
.ok_or_else(|| crate::ocr::Error::Pages("page index is out of range".into()))
})
.collect::<Result<BTreeSet<_>, _>>()?
.into_iter()
.map(|page| page.to_string())
.collect::<Vec<_>>()
.join(","),
Some(Value::Array(tokens)) => tokens
.iter()
.map(|token| {
token.as_str().map(str::trim).ok_or_else(|| {
crate::ocr::Error::Pages("expected only integers or only strings".into())
})
.collect::<Result<BTreeSet<_>, _>>()?
.into_iter()
.map(|page| page.to_string())
.collect::<Vec<_>>()
.join(",")
}
PagesInput::NativeTokens(tokens) => {
if tokens.is_empty() {
return Ok(None);
}
tokens
.iter()
.map(|token| token.trim())
.collect::<Vec<_>>()
.join(",")
}
PagesInput::NativeRange(range) => range
})
.collect::<Result<Vec<_>, _>>()?
.join(","),
Some(Value::String(range)) => range
.split(',')
.map(str::trim)
.collect::<Vec<_>>()
.join(","),
Some(_) => {
return Err(crate::ocr::Error::Pages(
"expected an array of integers or strings, or a native page range".into(),
));
}
};
if !normalized.split(',').all(valid_page_token) {
return Err(crate::ocr::Error::Pages("invalid native page range".into()));
@ -241,10 +187,15 @@ fn valid_page_token(token: &str) -> bool {
}
}
fn normalize_features(features: FeaturesInput) -> Result<Option<String>, crate::ocr::Error> {
fn normalize_features(features: Option<&Value>) -> Result<Option<String>, crate::ocr::Error> {
let tokens = match features {
FeaturesInput::Names(names) => names,
FeaturesInput::CommaSeparated(names) => names.split(',').map(str::to_string).collect(),
None | Some(Value::Null) => return Ok(None),
Some(Value::Array(names)) => names
.iter()
.map(|name| name.as_str().ok_or(crate::ocr::Error::Features))
.collect::<Result<Vec<_>, _>>()?,
Some(Value::String(names)) => names.split(',').collect(),
Some(_) => return Err(crate::ocr::Error::Features),
};
if tokens.is_empty() {
return Ok(None);
@ -375,7 +326,7 @@ async fn read_operation_response(
crate::ocr::client::read_response_bytes(response, connection.max_response_bytes)
.await?;
crate::ocr::handler::post_call(hooks, &bytes).await?;
return crate::ocr::wire::decode_response(&bytes, native);
return crate::ocr::json::decode_response(&bytes, native);
}
let location = response
.headers()
@ -530,10 +481,10 @@ impl BaseOcrConfig for AzureDocumentIntelligenceOCRConfig {
arguments: &CallArguments,
_model: &str,
) -> Result<DocumentIntelligenceParams, crate::ocr::Error> {
normalize_ocr_params(decode_input_params(
arguments.select(&["pages", "features"]),
"optional_params",
)?)
Ok(DocumentIntelligenceParams {
pages: normalize_pages(arguments.get("pages"))?,
features: normalize_features(arguments.get("features"))?,
})
}
async fn async_transform_ocr_request(
@ -591,30 +542,10 @@ impl BaseOcrConfig for AzureDocumentIntelligenceOCRConfig {
) -> Result<DocumentIntelligenceRequest, crate::ocr::Error> {
build_request(document)
}
}
impl AzureDocumentIntelligenceOCRConfig {
pub(crate) async fn prepare_request(
&self,
request: &PreparedOcrRequest,
client: &OcrClient,
) -> Result<reqwest::Request, crate::ocr::Error> {
let params = self.map_ocr_params(&request.optional_params, &request.model)?;
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,
request.document.clone(),
&params,
&headers,
OcrRequestContext {
client,
connection: &request.connection,
},
)
.await?;
transform_request_body(client, request, &url, &headers, false, body, |_| Ok(())).await
/// The body is `urlSource`/`base64Source`, not a `document` field.
fn retains_document(&self, _document: &OcrDocument) -> bool {
false
}
}
@ -712,8 +643,8 @@ mod tests {
use serde_json::{Value, json};
fn map(value: Value) -> Result<DocumentIntelligenceParams, crate::ocr::Error> {
let fields = value.as_object().unwrap().clone();
normalize_ocr_params(decode_input_params(fields, "optional_params")?)
let arguments = serde_json::from_value(value).unwrap();
AzureDocumentIntelligenceOCRConfig.map_ocr_params(&arguments, "model")
}
#[test]
@ -763,6 +694,8 @@ mod tests {
#[case(json!([0, 1, 2]), Some("1,2,3"))]
#[case(json!([2, 0, 0, 1]), Some("1,2,3"))]
#[case(json!([]), None)]
#[case(Value::Null, None)]
#[case(json!([i64::MAX - 1]), Some("9223372036854775807"))]
#[case(json!("3-9"), Some("3-9"))]
#[case(json!("1-3, 5"), Some("1-3,5"))]
#[case(json!(["1", "3-5"]), Some("1,3-5"))]
@ -778,8 +711,14 @@ mod tests {
#[case(json!([-1]))]
#[case(json!([true, false]))]
#[case(json!([1, "2"]))]
#[case(json!(["1", 2]))]
#[case(json!([1.0]))]
#[case(json!([i64::MAX]))]
#[case(json!([u64::MAX]))]
#[case(json!([null]))]
#[case(json!([[1]]))]
#[case(json!(5))]
fn invalid_page_mapping_matches_python(#[case] input: Value) {
fn page_mapping_rejects_invalid_shapes_and_overflow(#[case] input: Value) {
assert!(map(json!({"pages": input})).is_err());
}
@ -859,7 +798,6 @@ mod tests {
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)
@ -913,21 +851,12 @@ mod tests {
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,
};
let request = wire_request(
"azure_ai/doc-intelligence/prebuilt-read",
"http://127.0.0.1:1",
options.clone(),
);
let rejected = perform_ocr(request).await.is_err();
assert!(rejected, "accepted {options}");
}
}

View file

@ -1,14 +1,16 @@
use crate::call_arguments::CallArguments;
use crate::constants::AZURE_AI_OCR_PATH;
use crate::llms::base_llm::ocr::transformation::{BaseOcrConfig, OcrRequestContext};
use crate::llms::mistral::ocr::transformation::{MistralOCRConfig, MistralOcrRequest};
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::prepare::credential_env;
use crate::ocr::types::{LiteLLMOcrResponse, OcrConnection, OcrDocument, PreparedOcrRequest};
use crate::params::OpaqueParams;
use crate::url_utils::ApiUrl;
use litellm_auth::{InputSource, Sourced};
use litellm_auth_azure::AzureAuthInputs;
use serde_json::Value;
const AZURE_AI_API_KEY_ENV: &str = "AZURE_AI_API_KEY";
const AZURE_AI_API_BASE_ENV: &str = "AZURE_AI_API_BASE";
@ -64,6 +66,14 @@ impl BaseOcrConfig for AzureAIOCRConfig {
MistralOCRConfig.get_supported_ocr_params(model)
}
fn map_ocr_params(
&self,
arguments: &CallArguments,
model: &str,
) -> Result<OpaqueParams, crate::ocr::Error> {
MistralOCRConfig.map_ocr_params(arguments, model)
}
async fn async_transform_ocr_request(
&self,
model: &str,
@ -89,53 +99,40 @@ impl BaseOcrConfig for AzureAIOCRConfig {
) -> Result<LiteLLMOcrResponse, crate::ocr::Error> {
MistralOCRConfig.transform_ocr_response(model, raw_response, request_format)
}
}
impl AzureAIOCRConfig {
pub(crate) async fn prepare_request(
&self,
request: &PreparedOcrRequest,
client: &OcrClient,
) -> Result<reqwest::Request, crate::ocr::Error> {
let params = self.map_ocr_params(&request.optional_params, &request.model)?;
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
.async_transform_ocr_request(
&request.model,
request.document.clone(),
&params,
&headers,
OcrRequestContext {
client,
connection: &request.connection,
},
)
.await?;
transform_request_body(
client,
request,
&url,
&headers,
retains_document,
body,
|body| validate_inline_document(&crate::ocr::prepare::body_document(body)?),
)
.await
fn retains_document(&self, document: &OcrDocument) -> bool {
!document.is_remote()
}
fn validate_request_body(&self, body: &Value) -> Result<(), crate::ocr::Error> {
validate_inline_document(&crate::ocr::prepare::body_document(body)?)
}
}
impl AzureAIOCRConfig {
/// Python `AzureAIOCRConfig.validate_environment` requires the endpoint
/// before it resolves credentials; keep that order so a missing base is
/// reported without invoking any token provider.
pub(super) fn resolve_api_base(
api_base: Option<&str>,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> Result<String, crate::ocr::Error> {
nonblank(api_base.map(str::to_string))
.or_else(|| nonblank(env_lookup(AZURE_AI_API_BASE_ENV)))
.ok_or(crate::ocr::Error::Auth(
litellm_auth::Error::MissingApiBase {
provider: "Azure AI",
environment_variable: AZURE_AI_API_BASE_ENV,
},
))
}
fn get_complete_url(
&self,
api_base: Option<&str>,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> Result<String, crate::ocr::Error> {
let base = nonblank(api_base.map(str::to_string))
.or_else(|| nonblank(env_lookup(AZURE_AI_API_BASE_ENV)))
.ok_or_else(|| crate::ocr::Error::Auth(litellm_auth::Error::ProviderAuthentication("Missing Azure AI API Base - Set AZURE_AI_API_BASE environment variable or pass api_base parameter".into())))?;
let base = Self::resolve_api_base(api_base, env_lookup)?;
let path: Vec<&str> = AZURE_AI_OCR_PATH.trim_matches('/').split('/').collect();
ApiUrl::parse(&base)
.and_then(|url| url.complete_path(&path))
@ -151,6 +148,7 @@ impl AzureAIOCRConfig {
config: &AzureAuthInputs,
env_lookup: &(dyn Fn(&str) -> Option<String> + Sync),
) -> Result<Vec<(String, String)>, crate::ocr::Error> {
Self::resolve_api_base(connection.api_base.as_deref(), env_lookup)?;
if crate::http_utils::has_header(&connection.extra_headers, "authorization") {
if config.azure_ad_token_provider.is_some() {
super::common_utils::resolve_entra(config, env_lookup).await?;
@ -211,10 +209,24 @@ mod tests {
);
}
#[test]
fn missing_api_base_is_structured() {
assert!(matches!(
AzureAIOCRConfig::resolve_api_base(None, &|_| None),
Err(crate::ocr::Error::Auth(
litellm_auth::Error::MissingApiBase {
provider: "Azure AI",
environment_variable: AZURE_AI_API_BASE_ENV,
}
))
));
}
#[tokio::test]
async fn supplied_authorization_precedes_keys() {
let connection = OcrConnection {
api_key: Some("request-key".into()),
api_base: Some("https://example.com".into()),
extra_headers: vec![("authorization".into(), "Bearer prepared".into())],
..Default::default()
};
@ -233,6 +245,7 @@ mod tests {
async fn request_key_precedes_environment_key() {
let connection = OcrConnection {
api_key: Some("request-key".into()),
api_base: Some("https://example.com".into()),
..Default::default()
};
assert_eq!(
@ -291,7 +304,7 @@ mod tests {
use std::sync::Arc;
use serde_json::{Value, json};
use serde_json::json;
use crate::ocr::hooks::{OcrDuringCallRequest, OcrHookFuture, OcrHooks};
use crate::ocr::test_support::{MockResponse, mock_server, perform_ocr, wire_request};

View file

@ -3,8 +3,9 @@ use std::sync::Arc;
use serde::Serialize;
use serde::de::DeserializeOwned;
use serde_json::Value;
use crate::call_arguments::{CallArguments, parse_options};
use crate::call_arguments::CallArguments;
use crate::ocr::OcrClient;
use crate::ocr::hooks::OcrHooks;
use crate::ocr::types::{
@ -12,12 +13,24 @@ use crate::ocr::types::{
PreparedOcrRequest, ResolvedOcrCredentials,
};
/// Output of `validate_environment`: whatever a provider resolves up front
/// (headers at minimum; Vertex also carries the project id).
pub(crate) trait OcrEnvironment: Send + Sync {
fn headers(&self) -> &[(String, String)];
}
impl OcrEnvironment for Vec<(String, String)> {
fn headers(&self) -> &[(String, String)] {
self
}
}
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 OcrParams: Send + Sync;
type ProviderRequest: Serialize + Send;
type Environment: Send + Sync;
type Environment: OcrEnvironment;
fn get_api_key_env_var(&self) -> Option<&'static str> {
None
@ -64,13 +77,7 @@ pub(crate) trait BaseOcrConfig: Send + Sync + Sized + 'static {
&self,
arguments: &CallArguments,
model: &str,
) -> Result<Self::OcrParams, crate::ocr::Error> {
Ok(parse_options(
&arguments
.select(self.get_supported_ocr_params(model))
.into(),
)?)
}
) -> Result<Self::OcrParams, crate::ocr::Error>;
fn transform_ocr_request(
&self,
@ -127,6 +134,57 @@ pub(crate) trait BaseOcrConfig: Send + Sync + Sized + 'static {
headers,
}
}
/// Whether the `document` field in the outgoing body is owned by the
/// provider transform and must survive guardrail body rewrites.
/// Providers that inline remote URLs return `false` for remote documents
/// so a hook may still replace the fetched payload.
fn retains_document(&self, _document: &OcrDocument) -> bool {
true
}
/// Provider-specific check applied to the composed body, both before and
/// after guardrail hooks. Defaults to accepting any body.
fn validate_request_body(&self, _body: &Value) -> Result<(), crate::ocr::Error> {
Ok(())
}
/// Rust counterpart of `BaseLLMHTTPHandler._async_prepare_ocr_request`:
/// map params, validate environment, build URL, transform, compose body.
fn prepare_request(
&self,
request: &PreparedOcrRequest,
client: &OcrClient,
) -> impl Future<Output = Result<reqwest::Request, crate::ocr::Error>> + Send {
async move {
let params = self.map_ocr_params(&request.optional_params, &request.model)?;
let environment = self.validate_environment(request, client).await?;
let url = self.get_complete_url(request, &params, &environment)?;
let headers = environment.headers();
let body = self
.async_transform_ocr_request(
&request.model,
request.document.clone(),
&params,
headers,
OcrRequestContext {
client,
connection: &request.connection,
},
)
.await?;
crate::ocr::prepare::transform_request_body(
client,
request,
&url,
headers,
self.retains_document(&request.document),
body,
|body| self.validate_request_body(body),
)
.await
}
}
}
pub(crate) fn decode_and_normalize_response<T: DeserializeOwned>(
@ -135,7 +193,7 @@ pub(crate) fn decode_and_normalize_response<T: DeserializeOwned>(
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(
let decoded = crate::ocr::json::decode_response(
raw_response,
request_format == OcrResponseFormat::Native,
)?;

View file

@ -2,11 +2,12 @@ use serde::{Deserialize, Serialize};
use serde_json::{Map, Value};
use serde_with::serde_as;
use crate::call_arguments::{CallArguments, parse_options};
use crate::constants::{COHERE_API_KEY_ENV, COHERE_PARSE_API_BASE};
use crate::llms::base_llm::ocr::transformation::{BaseOcrConfig, OcrRequestContext};
use crate::ocr::OcrClient;
use crate::ocr::document::InlineDocument;
use crate::ocr::prepare::{credential_env, transform_request_body};
use crate::ocr::prepare::credential_env;
use crate::ocr::types::{
LiteLLMOcrResponse, OcrConnection, OcrDocument, OcrPage, OcrPageImage, OcrUsageInfo,
PreparedOcrRequest,
@ -136,6 +137,14 @@ impl BaseOcrConfig for CohereParseConfig {
&["output_format", "req_format"]
}
fn map_ocr_params(
&self,
arguments: &CallArguments,
_model: &str,
) -> Result<CohereOptions, crate::ocr::Error> {
Ok(parse_options(arguments)?)
}
async fn async_transform_ocr_request(
&self,
model: &str,
@ -160,33 +169,9 @@ impl BaseOcrConfig for CohereParseConfig {
normalize_response,
)
}
}
impl CohereParseConfig {
pub(crate) async fn prepare_request(
&self,
request: &PreparedOcrRequest,
client: &OcrClient,
) -> Result<reqwest::Request, crate::ocr::Error> {
let params = self.map_ocr_params(&request.optional_params, &request.model)?;
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,
request.document.clone(),
&params,
&headers,
OcrRequestContext {
client,
connection: &request.connection,
},
)
.await?;
transform_request_body(client, request, &url, &headers, true, body, |body| {
validate_document(&crate::ocr::prepare::body_document(body)?)
})
.await
fn validate_request_body(&self, body: &Value) -> Result<(), crate::ocr::Error> {
validate_document(&crate::ocr::prepare::body_document(body)?)
}
}
@ -255,7 +240,7 @@ fn page_image(
if let Some(Value::Object(bbox)) = image.get("bounding_box") {
image.insert("bbox".into(), Value::Object(bbox.clone()));
}
crate::ocr::wire::decode_response_value(Value::Object(image), path)
crate::ocr::json::decode_response_value(Value::Object(image), path)
}
fn normalize_page(page: CoherePage, position: usize) -> Result<OcrPage, crate::ocr::Error> {
@ -418,7 +403,11 @@ mod tests {
assert_eq!(arguments["req_format"], "native");
assert_eq!(arguments["extension"], false);
let invalid = serde_json::from_value(json!({"output_format":"html"})).unwrap();
assert!(CohereParseConfig.map_ocr_params(&invalid, "parse").is_err());
assert!(matches!(
CohereParseConfig.map_ocr_params(&invalid, "parse"),
Err(crate::ocr::Error::RequestField { path })
if path == "optional_params.output_format"
));
}
#[test]

View file

@ -1,10 +1,11 @@
use serde::{Deserialize, Serialize};
use serde_json::Value;
use crate::call_arguments::CallArguments;
use crate::constants::MISTRAL_OCR_API_BASE;
use crate::llms::base_llm::ocr::transformation::{BaseOcrConfig, OcrRequestContext};
use crate::ocr::OcrClient;
use crate::ocr::prepare::{credential_env, transform_request_body};
use crate::ocr::prepare::credential_env;
use crate::ocr::types::{
LiteLLMOcrResponse, OcrConnection, OcrDocument, OcrPage, OcrUsageInfo, PreparedOcrRequest,
};
@ -96,6 +97,16 @@ impl BaseOcrConfig for MistralOCRConfig {
]
}
fn map_ocr_params(
&self,
arguments: &CallArguments,
model: &str,
) -> Result<OpaqueParams, crate::ocr::Error> {
Ok(arguments
.select(self.get_supported_ocr_params(model))
.into())
}
async fn async_transform_ocr_request(
&self,
model: &str,
@ -122,31 +133,6 @@ impl BaseOcrConfig for MistralOCRConfig {
}
}
impl MistralOCRConfig {
pub(crate) async fn prepare_request(
&self,
request: &PreparedOcrRequest,
client: &OcrClient,
) -> Result<reqwest::Request, crate::ocr::Error> {
let params = self.map_ocr_params(&request.optional_params, &request.model)?;
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,
request.document.clone(),
&params,
&headers,
OcrRequestContext {
client,
connection: &request.connection,
},
)
.await?;
transform_request_body(client, request, &url, &headers, true, body, |_| Ok(())).await
}
}
pub(crate) fn normalize_response(
model: &str,
response: MistralOcrResponse,
@ -251,7 +237,7 @@ mod tests {
"usage_info.pages_processed",
),
] {
let error = crate::ocr::wire::decode_response::<MistralOcrResponse>(
let error = crate::ocr::json::decode_response::<MistralOcrResponse>(
&serde_json::to_vec(&payload).unwrap(),
false,
)

View file

@ -3,7 +3,7 @@ use std::collections::BTreeMap;
use serde::{Deserialize, Deserializer, Serialize};
use serde_json::{Map, Value, json};
use crate::call_arguments::compose_body;
use crate::call_arguments::{CallArguments, compose_body};
use crate::constants::{REDUCTO_API_BASE, REDUCTO_API_KEY_ENV, REDUCTO_ID_PREFIX};
use crate::llms::base_llm::ocr::transformation::{BaseOcrConfig, OcrRequestContext};
use crate::ocr::OcrClient;
@ -117,6 +117,16 @@ impl BaseOcrConfig for ReductoParseV3Config {
&["formatting", "retrieval", "settings"]
}
fn map_ocr_params(
&self,
arguments: &CallArguments,
model: &str,
) -> Result<ReductoV3Params, crate::ocr::Error> {
Ok(arguments
.select(self.get_supported_ocr_params(model))
.into())
}
#[tracing::instrument(
name = "async_transform_ocr_request",
target = "litellm::function_trace",
@ -151,36 +161,13 @@ impl BaseOcrConfig for ReductoParseV3Config {
normalize_response,
)
}
}
impl ReductoParseV3Config {
pub(crate) async fn prepare_request(
async fn prepare_request(
&self,
request: &PreparedOcrRequest,
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, 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(
&request.model,
document,
&params,
&headers,
OcrRequestContext {
client,
connection: &request.connection,
},
)
.await?;
let body = compose_body(
&request.optional_params,
&body,
self.get_supported_ocr_params(&request.model),
)?;
build_http_request(client, request, &url, &headers, &body)
prepare_upload_request(self, request, client).await
}
}
@ -225,6 +212,16 @@ impl BaseOcrConfig for ReductoParseLegacyConfig {
&["enhance"]
}
fn map_ocr_params(
&self,
arguments: &CallArguments,
model: &str,
) -> Result<ReductoLegacyParams, crate::ocr::Error> {
Ok(arguments
.select(self.get_supported_ocr_params(model))
.into())
}
#[tracing::instrument(
name = "async_transform_ocr_request",
target = "litellm::function_trace",
@ -251,39 +248,48 @@ impl BaseOcrConfig for ReductoParseLegacyConfig {
) -> Result<LiteLLMOcrResponse, crate::ocr::Error> {
ReductoParseV3Config.transform_ocr_response(model, raw_response, request_format)
}
}
impl ReductoParseLegacyConfig {
pub(crate) async fn prepare_request(
async fn prepare_request(
&self,
request: &PreparedOcrRequest,
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, 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(
&request.model,
document,
&params,
&headers,
OcrRequestContext {
client,
connection: &request.connection,
},
)
.await?;
let body = compose_body(
&request.optional_params,
&body,
self.get_supported_ocr_params(&request.model),
)?;
build_http_request(client, request, &url, &headers, &body)
prepare_upload_request(self, request, client).await
}
}
/// Reducto differs from the shared `BaseOcrConfig::prepare_request` flow:
/// guardrails see the *source* document before it is uploaded, because the
/// final body only carries the opaque Reducto file id.
async fn prepare_upload_request<C: BaseOcrConfig<Environment = Vec<(String, String)>>>(
config: &C,
request: &PreparedOcrRequest,
client: &OcrClient,
) -> Result<reqwest::Request, crate::ocr::Error> {
let params = config.map_ocr_params(&request.optional_params, &request.model)?;
let headers = config.validate_environment(request, client).await?;
let url = config.get_complete_url(request, &params, &headers)?;
let (document, headers) = guardrail_document(request, &url, &headers).await?;
let body = config
.async_transform_ocr_request(
&request.model,
document,
&params,
&headers,
OcrRequestContext {
client,
connection: &request.connection,
},
)
.await?;
let body = compose_body(
&request.optional_params,
&body,
config.get_supported_ocr_params(&request.model),
)?;
build_http_request(client, request, &url, &headers, &body)
}
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);

View file

@ -4,9 +4,10 @@ use serde_json::{Map, Value};
use litellm_auth_gcp::{self as vertex, VertexConfig};
use super::transformation::VertexAIOCRConfig;
use crate::call_arguments::CallArguments;
use crate::llms::base_llm::ocr::transformation::{BaseOcrConfig, OcrRequestContext};
use crate::ocr::OcrClient;
use crate::ocr::prepare::{credential_env, transform_request_body};
use crate::ocr::prepare::credential_env;
use crate::ocr::types::{
LiteLLMOcrResponse, OcrDocument, OcrPage, OcrPageDimensions, OcrPageImage, OcrUsageInfo,
PreparedOcrRequest,
@ -106,6 +107,14 @@ impl BaseOcrConfig for VertexAIDeepSeekOCRConfig {
VertexAIOCRConfig.get_api_key_env_var()
}
fn map_ocr_params(
&self,
_arguments: &CallArguments,
_model: &str,
) -> Result<DeepSeekOcrParams, crate::ocr::Error> {
Ok(DeepSeekOcrParams::default())
}
async fn validate_environment(
&self,
request: &PreparedOcrRequest,
@ -183,40 +192,11 @@ impl BaseOcrConfig for VertexAIDeepSeekOCRConfig {
.collect(),
})
}
}
impl VertexAIDeepSeekOCRConfig {
pub(crate) async fn prepare_request(
&self,
request: &PreparedOcrRequest,
client: &OcrClient,
) -> Result<reqwest::Request, crate::ocr::Error> {
let params = self.map_ocr_params(&request.optional_params, &request.model)?;
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,
request.document.clone(),
&params,
&authentication.headers,
OcrRequestContext {
client,
connection: &request.connection,
},
)
.await?;
transform_request_body(
client,
request,
&url,
&authentication.headers,
false,
body,
|_| Ok(()),
)
.await
/// The body carries the document inside `messages`, not a top-level
/// `document` field, so there is nothing for guardrails to retain.
fn retains_document(&self, _document: &OcrDocument) -> bool {
false
}
}
@ -267,7 +247,7 @@ pub(crate) fn normalize_response(
.enumerate()
.filter(|(_, page)| page.is_object())
.map(|(position, page)| {
let page: DeepSeekPage = crate::ocr::wire::decode_response_value(
let page: DeepSeekPage = crate::ocr::json::decode_response_value(
page.clone(),
&format!("choices[0].message.content.pages[{position}]"),
)?;
@ -288,7 +268,7 @@ pub(crate) fn normalize_response(
.or_else(|| (!has_pages).then_some(&response.usage));
let usage_info: Option<OcrUsageInfo> = usage
.filter(|usage| usage.is_object())
.map(|usage| crate::ocr::wire::decode_response_value(usage.clone(), "usage_info"))
.map(|usage| crate::ocr::json::decode_response_value(usage.clone(), "usage_info"))
.transpose()?;
let model = match ocr_data.get("model") {
Some(Value::String(model)) => model.clone(),
@ -683,11 +663,11 @@ mod tests {
#[test]
fn host_registration_selects_deepseek_without_affecting_mistral() {
assert!(crate::ocr::wire::is_supported_request(
assert!(crate::ocr::is_supported_request(
"deepseek-ocr-maas",
Some("vertex_ai")
));
assert!(crate::ocr::wire::is_supported_request(
assert!(crate::ocr::is_supported_request(
"mistral-ocr-maas",
Some("vertex_ai")
));

View file

@ -1,11 +1,15 @@
use litellm_auth_gcp::{self as vertex, VertexConfig};
use serde_json::Value;
use super::common_utils::validate_destination;
use crate::llms::base_llm::ocr::transformation::{BaseOcrConfig, OcrRequestContext};
use crate::call_arguments::CallArguments;
use crate::llms::base_llm::ocr::transformation::{
BaseOcrConfig, OcrEnvironment, OcrRequestContext,
};
use crate::llms::mistral::ocr::transformation::{MistralOCRConfig, MistralOcrRequest};
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::prepare::credential_env;
use crate::ocr::types::{LiteLLMOcrResponse, OcrConnection, OcrDocument, PreparedOcrRequest};
use crate::params::OpaqueParams;
use crate::url_utils::ApiUrl;
@ -71,6 +75,14 @@ impl BaseOcrConfig for VertexAIOCRConfig {
MistralOCRConfig.get_supported_ocr_params(model)
}
fn map_ocr_params(
&self,
arguments: &CallArguments,
model: &str,
) -> Result<OpaqueParams, crate::ocr::Error> {
MistralOCRConfig.map_ocr_params(arguments, model)
}
async fn async_transform_ocr_request(
&self,
model: &str,
@ -96,41 +108,19 @@ impl BaseOcrConfig for VertexAIOCRConfig {
) -> Result<LiteLLMOcrResponse, crate::ocr::Error> {
MistralOCRConfig.transform_ocr_response(model, raw_response, request_format)
}
fn retains_document(&self, document: &OcrDocument) -> bool {
!document.is_remote()
}
fn validate_request_body(&self, body: &Value) -> Result<(), crate::ocr::Error> {
validate_inline_document(&crate::ocr::prepare::body_document(body)?)
}
}
impl VertexAIOCRConfig {
pub(crate) async fn prepare_request(
&self,
request: &PreparedOcrRequest,
client: &OcrClient,
) -> Result<reqwest::Request, crate::ocr::Error> {
let params = self.map_ocr_params(&request.optional_params, &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
.async_transform_ocr_request(
&request.model,
request.document.clone(),
&params,
&authentication.headers,
OcrRequestContext {
client,
connection: &request.connection,
},
)
.await?;
transform_request_body(
client,
request,
&url,
&authentication.headers,
retains_document,
body,
|body| validate_inline_document(&crate::ocr::prepare::body_document(body)?),
)
.await
impl OcrEnvironment for vertex::VertexEnvironment {
fn headers(&self) -> &[(String, String)] {
&self.headers
}
}

View file

@ -0,0 +1,101 @@
use crate::call_arguments::ArgumentSpec;
use super::provider_config::{OcrConfigKind, resolve_provider_config};
const COMMON_OPTION_FIELDS: &[&str] = &["req_format", "extra_body", "max_response_bytes"];
const AZURE_AUTH_OPTION_FIELDS: &[&str] = &[
"azure_ad_token",
"tenant_id",
"client_id",
"client_secret",
"azure_scope",
"azure_authority_host",
"azure_credential",
"azure_federated_token_file",
"enable_azure_ad_token_refresh",
];
const VERTEX_AUTH_OPTION_FIELDS: &[&str] = &[
"vertex_credentials",
"vertex_ai_credentials",
"vertex_project",
"vertex_ai_project",
"vertex_location",
"vertex_ai_location",
];
pub fn is_supported_request(model: &str, custom_llm_provider: Option<&str>) -> bool {
resolve_provider_config(model, custom_llm_provider).is_ok()
}
pub fn consumed_optional_param_names(
model: &str,
custom_llm_provider: Option<&str>,
) -> Result<Vec<&'static str>, super::Error> {
let (model, config) = resolve_provider_config(model, custom_llm_provider)?;
let provider_fields = config.get_supported_ocr_params(&model);
let auth_fields: &[&str] = match config {
OcrConfigKind::AzureAi
| OcrConfigKind::AzureDocumentIntelligence
| OcrConfigKind::AzureCohere => AZURE_AUTH_OPTION_FIELDS,
OcrConfigKind::VertexAi | OcrConfigKind::VertexDeepSeek => VERTEX_AUTH_OPTION_FIELDS,
_ => &[],
};
Ok(COMMON_OPTION_FIELDS
.iter()
.chain(provider_fields)
.chain(auth_fields)
.copied()
.collect())
}
pub fn consumed_optional_params(
model: &str,
custom_llm_provider: Option<&str>,
) -> Result<Vec<ArgumentSpec>, super::Error> {
consumed_optional_param_names(model, custom_llm_provider).map(|names| {
names
.into_iter()
.map(|name| ArgumentSpec {
name,
secret: matches!(
name,
"azure_ad_token"
| "client_secret"
| "azure_federated_token_file"
| "vertex_credentials"
| "vertex_ai_credentials"
),
})
.collect()
})
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn consumed_params_include_provider_options_and_mark_credentials() {
let mistral = consumed_optional_param_names("mistral/model", None).unwrap();
assert!(mistral.contains(&"pages"));
assert!(mistral.contains(&"req_format"));
assert!(!mistral.contains(&"vertex_project"));
let vertex = consumed_optional_param_names("vertex_ai/deepseek-ocr", None).unwrap();
assert!(!vertex.contains(&"temperature"));
assert!(vertex.contains(&"vertex_credentials"));
assert!(!vertex.contains(&"pages"));
let azure = consumed_optional_params("model", Some("azure_ai")).unwrap();
assert!(
azure
.iter()
.any(|spec| spec.name == "client_secret" && spec.secret)
);
assert!(
azure
.iter()
.any(|spec| spec.name == "tenant_id" && !spec.secret)
);
}
}

View file

@ -4,8 +4,8 @@ use std::time::Duration;
use bytes::{Bytes, BytesMut};
use serde::de::DeserializeOwned;
use super::json::{DecodedOcrResponse, decode_response};
use super::types::{LiteLLMOcrRequest, LiteLLMOcrResponse};
use super::wire::{DecodedOcrResponse, decode_response};
use crate::constants::OCR_CONNECT_TIMEOUT_SECS;
use crate::media::MediaFetcher;
use litellm_auth_gcp::VertexAuth;

View file

@ -128,11 +128,11 @@ pub(crate) async fn inline_remote_document(
document: OcrDocument,
connection: &OcrConnection,
) -> Result<OcrDocument, crate::ocr::Error> {
let source = document.source();
if !source.starts_with("http://") && !source.starts_with("https://") {
if !document.is_remote() {
validate_inline_document(&document)?;
return Ok(document);
}
let source = document.source();
let url = Url::parse(source).map_err(|_| crate::ocr::Error::RequestField {
path: "document URL".into(),
})?;

View file

@ -2,18 +2,9 @@ use std::sync::Arc;
use super::OcrClient;
use super::hooks::{OcrHooks, OcrLifecycleHooks, OcrPostCallRequest};
use super::provider_config::OcrConfigKind;
use super::types::{LiteLLMOcrRequest, LiteLLMOcrResponse, PreparedOcrRequest};
use crate::call_lifecycle::{CallLifecycle, CallLifecycleContext};
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;
use crate::llms::base_llm::ocr::transformation::{BaseOcrConfig, OcrResponseContext};
use crate::llms::cohere::ocr::transformation::CohereParseConfig;
use crate::llms::mistral::ocr::transformation::MistralOCRConfig;
use crate::llms::reducto::ocr::transformation::{ReductoParseLegacyConfig, ReductoParseV3Config};
use crate::llms::vertex_ai::ocr::deepseek_transformation::VertexAIDeepSeekOCRConfig;
use crate::llms::vertex_ai::ocr::transformation::VertexAIOCRConfig;
use crate::llms::base_llm::ocr::transformation::OcrResponseContext;
pub(crate) async fn perform_ocr_request(
client: &OcrClient,
@ -55,37 +46,7 @@ impl PreparedOcrCall {
request: LiteLLMOcrRequest,
) -> Result<Self, super::Error> {
let request = super::prepare::prepare_request(request);
let http = match request.config {
OcrConfigKind::Cohere => CohereParseConfig.prepare_request(&request, &client).await?,
OcrConfigKind::Mistral => MistralOCRConfig.prepare_request(&request, &client).await?,
OcrConfigKind::AzureAi => AzureAIOCRConfig.prepare_request(&request, &client).await?,
OcrConfigKind::AzureCohere => {
AzureAICohereParseConfig
.prepare_request(&request, &client)
.await?
}
OcrConfigKind::AzureDocumentIntelligence => {
AzureDocumentIntelligenceOCRConfig
.prepare_request(&request, &client)
.await?
}
OcrConfigKind::ReductoLegacy => {
ReductoParseLegacyConfig
.prepare_request(&request, &client)
.await?
}
OcrConfigKind::ReductoV3 => {
ReductoParseV3Config
.prepare_request(&request, &client)
.await?
}
OcrConfigKind::VertexAi => VertexAIOCRConfig.prepare_request(&request, &client).await?,
OcrConfigKind::VertexDeepSeek => {
VertexAIDeepSeekOCRConfig
.prepare_request(&request, &client)
.await?
}
};
let http = request.config.prepare_request(&request, &client).await?;
Ok(Self {
client,
request,
@ -133,53 +94,10 @@ impl PreparedOcrCall {
url: &url,
headers: &headers,
};
match self.request.config {
OcrConfigKind::Cohere => {
CohereParseConfig
.async_transform_ocr_response(model, response, context)
.await
}
OcrConfigKind::Mistral => {
MistralOCRConfig
.async_transform_ocr_response(model, response, context)
.await
}
OcrConfigKind::AzureAi => {
AzureAIOCRConfig
.async_transform_ocr_response(model, response, context)
.await
}
OcrConfigKind::AzureCohere => {
AzureAICohereParseConfig
.async_transform_ocr_response(model, response, context)
.await
}
OcrConfigKind::AzureDocumentIntelligence => {
AzureDocumentIntelligenceOCRConfig
.async_transform_ocr_response(model, response, context)
.await
}
OcrConfigKind::ReductoLegacy => {
ReductoParseLegacyConfig
.async_transform_ocr_response(model, response, context)
.await
}
OcrConfigKind::ReductoV3 => {
ReductoParseV3Config
.async_transform_ocr_response(model, response, context)
.await
}
OcrConfigKind::VertexAi => {
VertexAIOCRConfig
.async_transform_ocr_response(model, response, context)
.await
}
OcrConfigKind::VertexDeepSeek => {
VertexAIDeepSeekOCRConfig
.async_transform_ocr_response(model, response, context)
.await
}
}
self.request
.config
.async_transform_ocr_response(model, response, context)
.await
}
}

View file

@ -0,0 +1,62 @@
use serde::de::{DeserializeOwned, IntoDeserializer};
use serde_json::{Map, Value};
#[derive(Debug)]
pub struct DecodedOcrResponse<T> {
pub data: T,
pub native: Option<Map<String, Value>>,
pub text: String,
}
pub(crate) fn decode_request_value<T: DeserializeOwned>(
value: Value,
prefix: &str,
) -> Result<T, crate::ocr::Error> {
serde_path_to_error::deserialize(value.into_deserializer()).map_err(|error| {
crate::ocr::Error::RequestField {
path: format!("{prefix}.{}", error.path()),
}
})
}
pub(crate) fn decode_response_value<T: DeserializeOwned>(
value: Value,
prefix: &str,
) -> Result<T, crate::ocr::Error> {
serde_path_to_error::deserialize(value.into_deserializer()).map_err(|error| {
crate::ocr::Error::ResponseField {
path: format!("{prefix}.{}", error.path()),
}
})
}
pub(crate) fn decode_response<T: DeserializeOwned>(
bytes: &[u8],
native: bool,
) -> Result<DecodedOcrResponse<T>, crate::ocr::Error> {
let mut deserializer = serde_json::Deserializer::from_slice(bytes);
let data = serde_path_to_error::deserialize(&mut deserializer).map_err(|error| {
crate::ocr::Error::ResponseField {
path: error.path().to_string(),
}
})?;
deserializer
.end()
.map_err(|_| crate::ocr::Error::ResponseField {
path: "response".into(),
})?;
let native = if native {
Some(
serde_json::from_slice(bytes).map_err(|_| crate::ocr::Error::ResponseField {
path: "response".into(),
})?,
)
} else {
None
};
Ok(DecodedOcrResponse {
data,
native,
text: String::from_utf8_lossy(bytes).into_owned(),
})
}

View file

@ -666,42 +666,37 @@ mod tests {
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,
LiteLLMOcrRequest, NativeOutcome, NoopOcrHost, OcrAdmission, OcrCall, OcrCallStep,
OcrDecline, OcrDocument, 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());
let document = OcrDocument::try_from(
json!({"type":"document_url","document_url":"https://example.com/doc.pdf"}),
)
.unwrap();
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,
})
LiteLLMOcrRequest::new(
"mistral/model".into(),
document.clone(),
None,
json!({"extract_header":true,"unknown":42})
.as_object()
.unwrap()
.clone()
.into(),
)
.is_ok()
);
assert!(
LiteLLMOcrRequest::new(
"model".into(),
document,
Some("unknown"),
Default::default()
)
.is_err()
);
}
@ -1507,11 +1502,20 @@ mod tests {
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 {
let Err(error) = LiteLLMOcrRequest::new(
"mistral/model".into(),
OcrDocument::try_from(json!({
"type": "document_url",
"document_url": "data:application/pdf;base64,YWJj"
}))
.unwrap(),
None,
json!({"max_response_bytes": value})
.as_object()
.unwrap()
.clone()
.into(),
) else {
panic!("invalid response limit accepted")
};
assert!(error.to_string().contains("max_response_bytes"));

View file

@ -1,15 +1,19 @@
mod error;
pub use error::Error;
mod arguments;
pub mod client;
pub(crate) mod document;
pub(crate) mod handler;
pub mod hooks;
pub(crate) mod json;
mod lifecycle;
pub(crate) mod prepare;
mod provider_config;
pub mod types;
pub mod wire;
pub use arguments::{
consumed_optional_param_names, consumed_optional_params, is_supported_request,
};
pub use client::{OcrClient, ocr};
pub use document::{encode_file_document, mime_type_for_name, upload_mime_type};
pub use lifecycle::{
@ -18,8 +22,8 @@ pub use lifecycle::{
};
pub use provider_config::{get_api_key_env_var, get_health_check_document};
pub use types::{
LiteLLMOcrRequest, LiteLLMOcrResponse, OcrConnection, OcrDocument, OcrPage, OcrPageDimensions,
OcrPageImage, OcrUsageInfo,
LiteLLMOcrRequest, LiteLLMOcrResponse, OcrConnection, OcrCredentialInputs, OcrDocument,
OcrPage, OcrPageDimensions, OcrPageImage, OcrTransportConfig, OcrUsageInfo,
};
#[cfg(test)]

View file

@ -104,7 +104,7 @@ pub(crate) async fn guardrail_document(
retained_fields: Vec::new(),
})
.await?;
let document = super::wire::decode_request_value(changed.body, "guardrail.document")?;
let document = super::json::decode_request_value(changed.body, "guardrail.document")?;
Ok((document, changed.headers))
}
@ -120,7 +120,7 @@ pub(crate) fn body_document(body: &Value) -> Result<OcrDocument, super::Error> {
.filter(|(name, _)| matches!(name.as_str(), "type" | "image_url" | "document_url"))
.map(|(name, value)| (name.clone(), value.clone()))
.collect();
super::wire::decode_request_value(Value::Object(source), "body.document")
super::json::decode_request_value(Value::Object(source), "body.document")
}
pub(crate) fn credential_env(name: &str) -> Option<String> {

View file

@ -1,8 +1,12 @@
use super::types::{OcrCredentialInputs, OcrDocument, ResolvedOcrCredentials};
use super::OcrClient;
use super::types::{
LiteLLMOcrResponse, OcrCredentialInputs, OcrDocument, PreparedOcrRequest,
ResolvedOcrCredentials,
};
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;
use crate::llms::base_llm::ocr::transformation::BaseOcrConfig;
use crate::llms::base_llm::ocr::transformation::{BaseOcrConfig, OcrResponseContext};
use crate::llms::cohere::ocr::transformation::CohereParseConfig;
use crate::llms::mistral::ocr::transformation::MistralOCRConfig;
use crate::llms::reducto::ocr::transformation::{ReductoParseLegacyConfig, ReductoParseV3Config};
@ -13,16 +17,22 @@ use strum::{EnumString, IntoStaticStr};
macro_rules! dispatch_config {
($config:expr, $method:ident($($argument:expr),* $(,)?)) => {
dispatch_config!(@arms $config, $method($($argument),*), )
};
($config:expr, $method:ident($($argument:expr),* $(,)?).await) => {
dispatch_config!(@arms $config, $method($($argument),*), .await)
};
(@arms $config:expr, $method:ident($($argument:expr),*), $($suffix:tt)*) => {
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),*),
OcrConfigKind::Cohere => CohereParseConfig.$method($($argument),*)$($suffix)*,
OcrConfigKind::Mistral => MistralOCRConfig.$method($($argument),*)$($suffix)*,
OcrConfigKind::AzureAi => AzureAIOCRConfig.$method($($argument),*)$($suffix)*,
OcrConfigKind::AzureCohere => AzureAICohereParseConfig.$method($($argument),*)$($suffix)*,
OcrConfigKind::AzureDocumentIntelligence => AzureDocumentIntelligenceOCRConfig.$method($($argument),*)$($suffix)*,
OcrConfigKind::ReductoLegacy => ReductoParseLegacyConfig.$method($($argument),*)$($suffix)*,
OcrConfigKind::ReductoV3 => ReductoParseV3Config.$method($($argument),*)$($suffix)*,
OcrConfigKind::VertexAi => VertexAIOCRConfig.$method($($argument),*)$($suffix)*,
OcrConfigKind::VertexDeepSeek => VertexAIDeepSeekOCRConfig.$method($($argument),*)$($suffix)*,
}
};
}
@ -80,6 +90,26 @@ impl OcrConfigKind {
) -> super::Error {
dispatch_config!(self, get_error_class(message, status, headers))
}
pub(crate) async fn prepare_request(
self,
request: &PreparedOcrRequest,
client: &OcrClient,
) -> Result<reqwest::Request, super::Error> {
dispatch_config!(self, prepare_request(request, client).await)
}
pub(crate) async fn async_transform_ocr_response(
self,
model: &str,
raw_response: reqwest::Response,
context: OcrResponseContext<'_>,
) -> Result<LiteLLMOcrResponse, super::Error> {
dispatch_config!(
self,
async_transform_ocr_response(model, raw_response, context).await
)
}
}
pub fn get_api_key_env_var(
@ -157,79 +187,83 @@ fn is_document_intelligence_model(model: &str) -> bool {
mod tests {
use super::*;
use litellm_auth::{InputSource, Sourced};
use rstest::rstest;
#[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()))
);
}
#[rstest]
#[case("cohere")]
#[case("mistral")]
#[case("azure_ai")]
#[case("reducto")]
#[case("vertex_ai")]
fn provider_names_round_trip_exactly(#[case] provider: &str) {
let (_, config) = resolve_provider_config("model", Some(provider)).unwrap();
let resolved: &'static str = config.provider().into();
assert_eq!(resolved, provider);
}
#[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")
);
}
#[rstest]
#[case("Mistral")]
#[case("unknown")]
fn invalid_provider_names_are_rejected(#[case] provider: &str) {
assert_eq!(
resolve_provider_config("model", Some(provider)),
Err(crate::ocr::Error::InvalidProvider(provider.into()))
);
}
#[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}"
);
}
#[rstest]
#[case("mistral/ocr")]
#[case("azure_ai/ocr")]
#[case("azure_ai/doc-intelligence/prebuilt-layout")]
#[case("reducto/parse-v3")]
#[case("vertex_ai/mistral-ocr")]
#[case("vertex_ai/deepseek-ocr")]
fn pdf_health_check_documents_are_valid(#[case] model: &str) {
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-"));
}
#[rstest]
#[case("cohere/parse")]
#[case("azure_ai/cohere-parse")]
fn png_health_check_documents_are_valid(#[case] model: &str) {
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")
);
}
#[rstest]
#[case("mistral/ocr", Some("MISTRAL_API_KEY"))]
#[case("cohere/parse", Some("COHERE_API_KEY"))]
#[case("azure_ai/ocr", Some("AZURE_AI_API_KEY"))]
#[case("azure_ai/cohere-parse", Some("AZURE_AI_API_KEY"))]
#[case(
"azure_ai/doc-intelligence/prebuilt-layout",
Some("AZURE_DOCUMENT_INTELLIGENCE_API_KEY")
)]
#[case("vertex_ai/mistral-ocr", Some("VERTEX_AI_API_KEY"))]
#[case("vertex_ai/deepseek-ocr", Some("VERTEX_AI_API_KEY"))]
#[case("reducto/parse-v3", None)]
#[case("reducto/parse-legacy", None)]
fn api_key_metadata_follows_provider_overrides_and_python_defaults(
#[case] model: &str,
#[case] expected: Option<&str>,
) {
assert_eq!(get_api_key_env_var(model, None).unwrap(), expected);
}
#[test]
@ -268,110 +302,106 @@ mod tests {
connection.api_base.as_ref().map(Sourced::source),
Some(InputSource::Request)
);
for dynamic in [
None,
Some(Sourced::new(String::new(), InputSource::Environment)),
] {
let connection =
OcrConfigKind::Mistral.resolve_connection_params(OcrCredentialInputs {
api_key: Some(Sourced::new("explicit-key".into(), InputSource::Deployment)),
api_base: Some(Sourced::new(
"https://explicit.test".into(),
InputSource::Deployment,
)),
dynamic_api_key: dynamic.clone(),
dynamic_api_base: dynamic,
});
assert_eq!(
connection
.api_key
.as_ref()
.map(|value| value.value().as_str()),
Some("explicit-key")
);
assert_eq!(
connection
.api_base
.as_ref()
.map(|value| value.value().as_str()),
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(
OcrCredentialInputs {
api_key: explicit_key
.map(|value| Sourced::new(value.to_string(), InputSource::Deployment)),
api_base: explicit_base
.map(|value| Sourced::new(value.to_string(), InputSource::Deployment)),
dynamic_api_key: Some(Sourced::new(
"dynamic-key".into(),
InputSource::Environment,
)),
dynamic_api_base: Some(Sourced::new(
"https://dynamic.test".into(),
InputSource::Deployment,
)),
},
);
assert_eq!(
connection
.api_key
.as_ref()
.map(|value| value.value().as_str()),
explicit_key.map(|_| "dynamic-key")
);
assert_eq!(
connection
.api_base
.as_ref()
.map(|value| value.value().as_str()),
explicit_base.map(|_| "https://dynamic.test")
);
}
}
#[test]
fn provider_models_are_preserved_without_a_local_allowlist() {
for (qualified_model, expected_config) in [
("mistral/future-ocr-model", OcrConfigKind::Mistral),
("azure_ai/future-ocr-model", OcrConfigKind::AzureAi),
] {
let expected_model = qualified_model.split_once('/').unwrap().1;
let (model, config) = resolve_provider_config(qualified_model, None).unwrap();
assert_eq!(model, expected_model);
assert_eq!(config, expected_config);
}
}
#[test]
fn provider_specific_models_select_their_config() {
#[rstest]
#[case(None)]
#[case(Some(""))]
fn empty_or_missing_dynamic_credentials_preserve_explicit_values(
#[case] dynamic_value: Option<&str>,
) {
let dynamic =
dynamic_value.map(|value| Sourced::new(value.into(), InputSource::Environment));
let connection = OcrConfigKind::Mistral.resolve_connection_params(OcrCredentialInputs {
api_key: Some(Sourced::new("explicit-key".into(), InputSource::Deployment)),
api_base: Some(Sourced::new(
"https://explicit.test".into(),
InputSource::Deployment,
)),
dynamic_api_key: dynamic.clone(),
dynamic_api_base: dynamic,
});
assert_eq!(
resolve_provider_config("reducto/parse-legacy", None)
.unwrap()
.1,
OcrConfigKind::ReductoLegacy
connection
.api_key
.as_ref()
.map(|value| value.value().as_str()),
Some("explicit-key")
);
assert_eq!(
resolve_provider_config("reducto/future-parse-model", None)
.unwrap()
.1,
OcrConfigKind::ReductoV3
connection
.api_base
.as_ref()
.map(|value| value.value().as_str()),
Some("https://explicit.test")
);
}
#[rstest]
#[case(None, None)]
#[case(Some("key"), None)]
#[case(None, Some("base"))]
#[case(Some("key"), Some("base"))]
fn document_intelligence_only_accepts_dynamic_values_for_explicit_fields(
#[case] explicit_key: Option<&str>,
#[case] explicit_base: Option<&str>,
) {
let connection = OcrConfigKind::AzureDocumentIntelligence.resolve_connection_params(
OcrCredentialInputs {
api_key: explicit_key
.map(|value| Sourced::new(value.into(), InputSource::Deployment)),
api_base: explicit_base
.map(|value| Sourced::new(value.into(), InputSource::Deployment)),
dynamic_api_key: Some(Sourced::new("dynamic-key".into(), InputSource::Environment)),
dynamic_api_base: Some(Sourced::new(
"https://dynamic.test".into(),
InputSource::Deployment,
)),
},
);
assert_eq!(
resolve_provider_config("azure_ai/doc-intelligence/prebuilt-layout", None)
.unwrap()
.1,
OcrConfigKind::AzureDocumentIntelligence
connection
.api_key
.as_ref()
.map(|value| value.value().as_str()),
explicit_key.map(|_| "dynamic-key")
);
assert_eq!(
connection
.api_base
.as_ref()
.map(|value| value.value().as_str()),
explicit_base.map(|_| "https://dynamic.test")
);
}
#[rstest]
#[case("mistral/future-ocr-model", OcrConfigKind::Mistral)]
#[case("azure_ai/future-ocr-model", OcrConfigKind::AzureAi)]
fn provider_models_are_preserved_without_a_local_allowlist(
#[case] qualified_model: &str,
#[case] expected_config: OcrConfigKind,
) {
let expected_model = qualified_model.split_once('/').unwrap().1;
let (model, config) = resolve_provider_config(qualified_model, None).unwrap();
assert_eq!(model, expected_model);
assert_eq!(config, expected_config);
}
#[rstest]
#[case("reducto/parse-legacy", OcrConfigKind::ReductoLegacy)]
#[case("reducto/future-parse-model", OcrConfigKind::ReductoV3)]
#[case(
"azure_ai/doc-intelligence/prebuilt-layout",
OcrConfigKind::AzureDocumentIntelligence
)]
fn provider_specific_models_select_their_config(
#[case] model: &str,
#[case] expected_config: OcrConfigKind,
) {
assert_eq!(
resolve_provider_config(model, None).unwrap().1,
expected_config
);
}
}

View file

@ -4,8 +4,9 @@ use serde_json::{Value, json};
use tokio::io::{AsyncReadExt, AsyncWriteExt};
use tokio::net::TcpListener;
use crate::ocr::wire::{OcrWireRequest, decode_request};
use crate::ocr::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrClient};
use crate::ocr::{
LiteLLMOcrRequest, LiteLLMOcrResponse, OcrClient, OcrCredentialInputs, OcrDocument,
};
pub(crate) fn ocr_client() -> OcrClient {
let document_http = reqwest::Client::builder()
@ -22,18 +23,31 @@ pub(crate) async fn perform_ocr(
}
pub(crate) fn wire_request(model: &str, base: &str, options: Value) -> LiteLLMOcrRequest {
decode_request(OcrWireRequest {
model: model.into(),
document: json!({"type":"document_url","document_url":"data:application/pdf;base64,YWJj"}),
api_key: Some("test-key".into()),
api_base: Some(base.into()),
custom_llm_provider: None,
extra_headers: None,
optional_params: options.as_object().unwrap().clone().into(),
input_sources: Default::default(),
timeout_seconds: Some(2.0),
})
.unwrap()
let request = LiteLLMOcrRequest::new(
model.into(),
OcrDocument::try_from(
json!({"type":"document_url","document_url":"data:application/pdf;base64,YWJj"}),
)
.unwrap(),
None,
options.as_object().unwrap().clone().into(),
)
.unwrap();
let transport = request.transport.clone().with_overrides(
Vec::new(),
Default::default(),
Some(std::time::Duration::from_secs(2)),
);
request.with_connection_inputs(
OcrCredentialInputs::new(
Some("test-key".into()),
Default::default(),
Some(base.into()),
Default::default(),
),
transport,
Default::default(),
)
}
pub(crate) struct MockResponse {

View file

@ -39,6 +39,11 @@ impl OcrDocument {
}
}
pub(crate) fn is_remote(&self) -> bool {
let source = self.source();
source.starts_with("http://") || source.starts_with("https://")
}
pub(crate) fn with_source(self, source: String) -> Self {
match self {
Self::DocumentUrl { extra_fields, .. } => Self::DocumentUrl {
@ -53,6 +58,14 @@ impl OcrDocument {
}
}
impl TryFrom<Value> for OcrDocument {
type Error = super::Error;
fn try_from(value: Value) -> Result<Self, Self::Error> {
super::json::decode_request_value(value, "document")
}
}
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "lowercase")]
pub enum OcrResponseFormat {
@ -69,6 +82,22 @@ pub struct OcrCredentialInputs {
pub dynamic_api_base: Option<Sourced<String>>,
}
impl OcrCredentialInputs {
pub fn new(
api_key: Option<String>,
api_key_source: InputSource,
api_base: Option<String>,
api_base_source: InputSource,
) -> Self {
Self {
api_key: nonblank(api_key).map(|value| Sourced::new(value, api_key_source)),
dynamic_api_key: None,
api_base: nonblank(api_base).map(|value| Sourced::new(value, api_base_source)),
dynamic_api_base: None,
}
}
}
#[derive(Clone)]
pub struct OcrTransportConfig {
pub extra_headers: Vec<(String, String)>,
@ -92,6 +121,28 @@ impl Default for OcrTransportConfig {
}
}
impl OcrTransportConfig {
pub fn with_overrides(
self,
extra_headers: Vec<(String, String)>,
extra_headers_source: InputSource,
timeout: Option<Duration>,
) -> Self {
Self {
extra_headers,
extra_headers_source,
timeout: timeout.unwrap_or(self.timeout),
..self
}
}
}
fn nonblank(value: Option<String>) -> Option<String> {
value
.map(|value| value.trim().to_string())
.filter(|value| !value.is_empty())
}
#[derive(Clone)]
pub struct OcrConnection {
pub api_key: Option<String>,
@ -169,12 +220,34 @@ impl LiteLLMOcrRequest {
optional_params: CallArguments,
) -> Result<Self, super::Error> {
let (model, config) = resolve_provider_config(&model, custom_llm_provider)?;
let default_transport = OcrTransportConfig::default();
let max_response_bytes = optional_params
.get("max_response_bytes")
.map(|value| {
value
.as_u64()
.and_then(|value| usize::try_from(value).ok())
.filter(|value| *value > 0 && *value <= default_transport.max_response_bytes)
.ok_or_else(|| super::Error::RequestField {
path: "max_response_bytes".into(),
})
})
.transpose()?
.unwrap_or(default_transport.max_response_bytes);
let transport = OcrTransportConfig {
max_response_bytes,
..default_transport
};
let optional_params = optional_params
.into_iter()
.filter(|(name, _)| name != "max_response_bytes")
.collect();
Ok(Self {
model,
document,
credentials: OcrCredentialInputs::default(),
transport: OcrTransportConfig::default(),
transport,
hooks: Arc::new(NoopOcrHooks),
litellm_call_id: None,
optional_params,
@ -210,6 +283,20 @@ impl LiteLLMOcrRequest {
..self
}
}
pub fn with_connection_inputs(
self,
credentials: OcrCredentialInputs,
transport: OcrTransportConfig,
input_sources: BTreeMap<String, InputSource>,
) -> Self {
Self {
credentials,
transport,
input_sources,
..self
}
}
}
pub(crate) struct PreparedOcrRequest {

View file

@ -1,345 +0,0 @@
use std::collections::BTreeMap;
use std::time::Duration;
use super::types::{LiteLLMOcrRequest, OcrDocument, OcrTransportConfig};
use crate::call_arguments::{ArgumentSpec, CallArguments};
use litellm_auth::InputSource;
use serde::{
Deserialize,
de::{DeserializeOwned, IntoDeserializer},
};
use serde_json::{Map, Value};
const COMMON_OPTION_FIELDS: &[&str] = &["req_format", "extra_body", "max_response_bytes"];
pub const BOUND_FIELDS: &[&str] = &["model", "document", "timeout", "input_sources"];
const AZURE_AUTH_OPTION_FIELDS: &[&str] = &[
"azure_ad_token",
"tenant_id",
"client_id",
"client_secret",
"azure_scope",
"azure_authority_host",
"azure_credential",
"azure_federated_token_file",
"enable_azure_ad_token_refresh",
];
const VERTEX_AUTH_OPTION_FIELDS: &[&str] = &[
"vertex_credentials",
"vertex_ai_credentials",
"vertex_project",
"vertex_ai_project",
"vertex_location",
"vertex_ai_location",
];
#[derive(Debug)]
pub struct DecodedOcrResponse<T> {
pub data: T,
pub native: Option<Map<String, Value>>,
pub text: String,
}
#[derive(Deserialize)]
#[serde(deny_unknown_fields)]
pub struct OcrWireRequest {
pub model: String,
pub document: Value,
pub api_key: Option<String>,
pub api_base: Option<String>,
pub custom_llm_provider: Option<String>,
pub extra_headers: Option<Map<String, Value>>,
#[serde(default)]
pub optional_params: CallArguments,
#[serde(default)]
pub input_sources: BTreeMap<String, InputSource>,
pub timeout_seconds: Option<f64>,
}
pub fn is_supported_request(model: &str, custom_llm_provider: Option<&str>) -> bool {
super::provider_config::resolve_provider_config(model, custom_llm_provider).is_ok()
}
pub fn consumed_optional_param_names(
model: &str,
custom_llm_provider: Option<&str>,
) -> Result<Vec<&'static str>, crate::ocr::Error> {
use super::provider_config::OcrConfigKind;
let (model, config) =
super::provider_config::resolve_provider_config(model, custom_llm_provider)?;
let provider_fields = config.get_supported_ocr_params(&model);
let auth_fields: &[&str] = match config {
OcrConfigKind::AzureAi
| OcrConfigKind::AzureDocumentIntelligence
| OcrConfigKind::AzureCohere => AZURE_AUTH_OPTION_FIELDS,
OcrConfigKind::VertexAi | OcrConfigKind::VertexDeepSeek => VERTEX_AUTH_OPTION_FIELDS,
_ => &[],
};
Ok(COMMON_OPTION_FIELDS
.iter()
.chain(provider_fields)
.chain(auth_fields)
.copied()
.collect())
}
pub fn consumed_optional_params(
model: &str,
custom_llm_provider: Option<&str>,
) -> Result<Vec<ArgumentSpec>, crate::ocr::Error> {
consumed_optional_param_names(model, custom_llm_provider).map(|names| {
names
.into_iter()
.map(|name| ArgumentSpec {
name,
secret: matches!(
name,
"azure_ad_token"
| "client_secret"
| "azure_federated_token_file"
| "vertex_credentials"
| "vertex_ai_credentials"
),
})
.collect()
})
}
pub fn decode_request(wire: OcrWireRequest) -> Result<LiteLLMOcrRequest, crate::ocr::Error> {
let api_key_source = source_for(&wire.input_sources, "api_key");
let api_base_source = source_for(&wire.input_sources, "api_base");
let extra_headers_source = source_for(&wire.input_sources, "extra_headers");
let document = decode_document(wire.document)?;
let headers = wire
.extra_headers
.unwrap_or_default()
.into_iter()
.map(|(name, value)| {
let value = value
.as_str()
.ok_or_else(|| crate::ocr::Error::RequestField {
path: format!("extra_headers.{name}"),
})?;
Ok((name, value.to_string()))
})
.collect::<Result<Vec<_>, crate::ocr::Error>>()?;
let timeout = wire
.timeout_seconds
.map(|seconds| {
Duration::try_from_secs_f64(seconds).map_err(|_| crate::ocr::Error::RequestField {
path: "timeout_seconds".into(),
})
})
.transpose()?;
let defaults = OcrTransportConfig::default();
let max_response_bytes = wire
.optional_params
.get("max_response_bytes")
.map(|value| {
value
.as_u64()
.and_then(|value| usize::try_from(value).ok())
.filter(|value| *value > 0 && *value <= defaults.max_response_bytes)
.ok_or_else(|| crate::ocr::Error::RequestField {
path: "max_response_bytes".into(),
})
})
.transpose()?
.unwrap_or(defaults.max_response_bytes);
let request = LiteLLMOcrRequest::new(
wire.model,
document,
wire.custom_llm_provider.as_deref(),
wire.optional_params
.into_iter()
.filter(|(name, _)| name != "max_response_bytes")
.collect(),
)?;
let credentials = super::types::OcrCredentialInputs {
api_key: nonblank(wire.api_key)
.map(|value| litellm_auth::Sourced::new(value, api_key_source)),
dynamic_api_key: None,
api_base: nonblank(wire.api_base)
.map(|value| litellm_auth::Sourced::new(value, api_base_source)),
dynamic_api_base: None,
};
let transport = OcrTransportConfig {
extra_headers: headers,
extra_headers_source,
timeout: timeout.unwrap_or(defaults.timeout),
max_download_bytes: defaults.max_download_bytes,
max_response_bytes,
poll_timeout: defaults.poll_timeout,
};
Ok(LiteLLMOcrRequest {
credentials,
transport,
input_sources: wire.input_sources,
..request
})
}
fn decode_document(value: Value) -> Result<OcrDocument, crate::ocr::Error> {
let kind = value.get("type").and_then(Value::as_str);
let missing_url = matches!(kind, Some("document_url")) && value.get("document_url").is_none()
|| matches!(kind, Some("image_url")) && value.get("image_url").is_none();
if missing_url {
return Err(crate::ocr::Error::MissingDocumentUrl);
}
decode_request_value(value, "document")
}
fn source_for(sources: &BTreeMap<String, InputSource>, name: &str) -> InputSource {
sources.get(name).copied().unwrap_or_default()
}
fn nonblank(value: Option<String>) -> Option<String> {
value
.map(|s| s.trim().to_string())
.filter(|s| !s.is_empty())
}
pub fn decode_request_value<T: DeserializeOwned>(
value: Value,
prefix: &str,
) -> Result<T, crate::ocr::Error> {
serde_path_to_error::deserialize(value.into_deserializer()).map_err(|error| {
crate::ocr::Error::RequestField {
path: format!("{prefix}.{}", error.path()),
}
})
}
pub(crate) fn decode_response_value<T: DeserializeOwned>(
value: Value,
prefix: &str,
) -> Result<T, crate::ocr::Error> {
serde_path_to_error::deserialize(value.into_deserializer()).map_err(|error| {
crate::ocr::Error::ResponseField {
path: format!("{prefix}.{}", error.path()),
}
})
}
pub fn decode_response<T: DeserializeOwned>(
bytes: &[u8],
native: bool,
) -> Result<DecodedOcrResponse<T>, crate::ocr::Error> {
let mut deserializer = serde_json::Deserializer::from_slice(bytes);
let data = serde_path_to_error::deserialize(&mut deserializer).map_err(|error| {
crate::ocr::Error::ResponseField {
path: error.path().to_string(),
}
})?;
deserializer
.end()
.map_err(|_| crate::ocr::Error::ResponseField {
path: "response".into(),
})?;
let native = if native {
Some(
serde_json::from_slice(bytes).map_err(|_| crate::ocr::Error::ResponseField {
path: "response".into(),
})?,
)
} else {
None
};
Ok(DecodedOcrResponse {
data,
native,
text: String::from_utf8_lossy(bytes).into_owned(),
})
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn core_selects_consumed_values_without_serializing_host_objects() {
let fields = consumed_optional_params("mistral/model", None).unwrap();
use crate::call_arguments::should_project;
assert!(should_project("future_option", &fields, BOUND_FIELDS));
assert!(should_project("extra_body", &fields, BOUND_FIELDS));
assert!(should_project("id", &fields, BOUND_FIELDS));
assert!(!should_project("metadata", &fields, BOUND_FIELDS));
assert!(!should_project("callbacks", &fields, BOUND_FIELDS));
assert!(!should_project("api_key", &fields, BOUND_FIELDS));
assert!(!should_project("document", &fields, BOUND_FIELDS));
}
#[test]
fn option_projection_is_provider_specific_and_excludes_opaque_fields() {
let mistral = consumed_optional_param_names("mistral/model", None).unwrap();
assert!(mistral.contains(&"pages"));
assert!(mistral.contains(&"req_format"));
assert!(!mistral.contains(&"vertex_project"));
assert!(!mistral.contains(&"opaque_extension"));
let vertex = consumed_optional_param_names("vertex_ai/deepseek-ocr", None).unwrap();
assert!(!vertex.contains(&"temperature"));
assert!(vertex.contains(&"vertex_credentials"));
assert!(!vertex.contains(&"pages"));
}
#[test]
fn optional_param_metadata_marks_only_credentials_as_secret() {
let azure = consumed_optional_params("model", Some("azure_ai")).unwrap();
assert!(
azure
.iter()
.any(|spec| spec.name == "client_secret" && spec.secret)
);
assert!(
azure
.iter()
.any(|spec| spec.name == "tenant_id" && !spec.secret)
);
let vertex = consumed_optional_params("deepseek-ocr", Some("vertex_ai")).unwrap();
assert!(
vertex
.iter()
.any(|spec| spec.name == "vertex_credentials" && spec.secret)
);
assert!(
vertex
.iter()
.any(|spec| spec.name == "vertex_project" && !spec.secret)
);
}
#[test]
fn activation_includes_migrated_providers() {
assert!(is_supported_request("model", Some("mistral")));
assert!(is_supported_request("pixtral-12b", Some("azure_ai")));
assert!(is_supported_request(
"documentintelligence/prebuilt-read",
Some("azure_ai")
));
assert!(is_supported_request("parse-v3", Some("reducto")));
assert!(is_supported_request("parse-legacy", Some("reducto")));
assert!(is_supported_request("mistral-ocr", Some("vertex_ai")));
assert!(is_supported_request("deepseek-ocr", Some("vertex_ai")));
}
#[test]
fn missing_document_source_has_a_typed_public_error() {
for document in [
serde_json::json!({"type": "document_url"}),
serde_json::json!({"type": "image_url"}),
] {
let wire = serde_json::from_value(serde_json::json!({
"model": "mistral/model",
"document": document,
}))
.unwrap();
let error = decode_request(wire).err().expect("missing document URL");
assert_eq!(error, crate::ocr::Error::MissingDocumentUrl);
let error = crate::Error::from(error);
assert!(matches!(
error,
crate::Error::Ocr(crate::ocr::Error::MissingDocumentUrl)
));
}
}
}

View file

@ -9,8 +9,7 @@ use pyo3::types::PyDict;
use pyo3::types::{PyBytes, PyString};
use litellm_core::constants::OCR_INLINE_MAX_BYTES;
use litellm_core::ocr::{OcrDocument, encode_file_document, mime_type_for_name, upload_mime_type};
use litellm_python_interop::to_py_preserving_errors;
use litellm_core::ocr::{OcrDocument, encode_file_document};
enum FileBytes {
Python(PyBackedBytes),
@ -136,44 +135,6 @@ pub(super) fn file_document(py: Python<'_>, document: FileDocumentInput) -> PyRe
.map_err(|error| PyValueError::new_err(error.to_string()))
}
#[pyfunction]
fn _ocr_file_document(py: Python<'_>, document: Bound<'_, PyAny>) -> PyResult<Py<PyAny>> {
to_py_preserving_errors(py, &file_document(py, document.extract()?)?)
}
#[pyfunction]
fn _ocr_mime_type(file_name: &str) -> String {
mime_type_for_name(file_name).into()
}
#[pyfunction]
#[pyo3(signature = (file_content, file_name=None, content_type=None))]
fn _ocr_upload_document(
py: Python<'_>,
file_content: &Bound<'_, PyBytes>,
file_name: Option<&str>,
content_type: Option<&str>,
) -> PyResult<Py<PyAny>> {
let bytes: PyBackedBytes = file_content.extract()?;
let document = py
.detach(|| {
encode_file_document(
&bytes,
None,
Some(upload_mime_type(file_name, content_type)),
)
})
.map_err(|error| PyValueError::new_err(error.to_string()))?;
to_py_preserving_errors(py, &document)
}
pub(super) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> {
module.add("_OCR_MAX_FILE_BYTES", OCR_INLINE_MAX_BYTES)?;
module.add_function(wrap_pyfunction!(_ocr_upload_document, module)?)?;
module.add_function(wrap_pyfunction!(_ocr_file_document, module)?)?;
module.add_function(wrap_pyfunction!(_ocr_mime_type, module)?)
}
#[cfg(test)]
mod tests {
use super::*;

View file

@ -10,7 +10,7 @@ use litellm_python_interop::{
use super::callbacks;
use super::errors::to_pyerr as ocr_error_to_pyerr;
use super::project::{ProjectedOcrFields, admitted_call, project_request};
use super::project::{ProjectedOcrCall, ProjectedOcrFields, PythonOcrInput, admitted_call};
use crate::lifecycle::{
OperationClass, PythonCallState, PythonRoute, missing_state, now, run_call,
};
@ -193,7 +193,10 @@ impl PythonRoute for PythonOcrHost {
let OcrHostData::Unprojected { request } = &self.data else {
return Err(missing_state());
};
let projected = project_request(py, request.bind(py), self.state.kwargs.bind(py))?;
let projected = ProjectedOcrCall::try_from(PythonOcrInput {
request: request.bind(py),
kwargs: self.state.kwargs.bind(py),
})?;
let has_token_provider = projected.fields.azure_ad_token_provider.is_some();
let request = projected.request;
self.data = OcrHostData::Projected(Box::new(ProjectedOcrHost {

View file

@ -3,12 +3,12 @@ mod document;
mod errors;
mod lifecycle;
mod project;
mod request;
mod value;
use pyo3::prelude::*;
pub(super) fn register(module: &Bound<'_, PyModule>) -> PyResult<()> {
value::register(module)?;
document::register(module)?;
lifecycle::register(module)
}

View file

@ -1,7 +1,6 @@
use std::sync::Arc;
use litellm_core::ocr::wire::{OcrWireRequest, consumed_optional_params, decode_request};
use litellm_core::ocr::{LiteLLMOcrRequest, NativeOutcome, OcrCall};
use litellm_core::ocr::{LiteLLMOcrRequest, NativeOutcome, OcrCall, consumed_optional_params};
use litellm_python_interop::{
from_py_preserving_errors as from_py, to_py_preserving_errors as to_py,
};
@ -11,10 +10,13 @@ use serde_json::{Map, Value};
use super::errors::to_pyerr as ocr_error_to_pyerr;
use super::lifecycle::BridgeOcrHooks;
use super::request::BridgeOcrRequest;
use crate::auth::{AZURE_AD_TOKEN_PROVIDER, PythonTokenProvider};
use crate::errors::RustBridgeDeclined;
use crate::marshal::{project_optional_fields, python_timeout_seconds, request_input_sources};
const BOUND_FIELDS: &[&str] = &["model", "document", "timeout", "input_sources"];
pub(super) struct ProjectedOcrFields {
pub boundary_request: Py<PyAny>,
pub document: Py<PyAny>,
@ -29,6 +31,11 @@ pub(super) struct ProjectedOcrCall {
pub fields: ProjectedOcrFields,
}
pub(super) struct PythonOcrInput<'a, 'py> {
pub request: &'a Bound<'py, PyAny>,
pub kwargs: &'a Bound<'py, PyDict>,
}
struct PythonOcrFields<'a, 'py> {
request: &'a Bound<'py, PyAny>,
kwargs: &'a Bound<'py, PyDict>,
@ -110,60 +117,63 @@ impl ProjectedDocument {
}
}
pub(super) fn project_request(
py: Python<'_>,
request: &Bound<'_, PyAny>,
kwargs: &Bound<'_, PyDict>,
) -> PyResult<ProjectedOcrCall> {
let boundary_request = request.clone().unbind();
let arguments = PythonOcrFields { request, kwargs };
let model = arguments.model()?;
let custom_llm_provider = arguments.custom_llm_provider()?;
let (wire_document, retained_document) =
ProjectedDocument::project(py, &arguments.document()?)?.into_parts();
let api_key = arguments.api_key()?;
let specs = consumed_optional_params(&model, custom_llm_provider.as_deref())
impl TryFrom<PythonOcrInput<'_, '_>> for ProjectedOcrCall {
type Error = PyErr;
fn try_from(input: PythonOcrInput<'_, '_>) -> PyResult<Self> {
let request = input.request;
let kwargs = input.kwargs;
let py = request.py();
let boundary_request = request.clone().unbind();
let arguments = PythonOcrFields { request, kwargs };
let model = arguments.model()?;
let custom_llm_provider = arguments.custom_llm_provider()?;
let (wire_document, retained_document) =
ProjectedDocument::project(py, &arguments.document()?)?.into_parts();
let api_key = arguments.api_key()?;
let specs = consumed_optional_params(&model, custom_llm_provider.as_deref())
.map_err(ocr_error_to_pyerr)?;
let optional_params = project_optional_fields(kwargs, &specs, BOUND_FIELDS)?;
let input_sources = request_input_sources(
kwargs,
optional_params.keys().map(String::as_str).chain([
"api_key",
"api_base",
"extra_headers",
]),
)?;
let azure_ad_token_provider = kwargs
.get_item("azure_ad_token_provider")?
.and_then(|provider| PythonTokenProvider::select(provider, AZURE_AD_TOKEN_PROVIDER));
let request = LiteLLMOcrRequest::try_from(BridgeOcrRequest {
model,
document: wire_document,
api_key: api_key.extract()?,
api_base: arguments.api_base()?,
custom_llm_provider,
extra_headers: arguments.extra_headers()?,
optional_params: optional_params.into(),
input_sources,
timeout_seconds: arguments.timeout_seconds()?,
})
.map_err(ocr_error_to_pyerr)?;
let optional_params =
project_optional_fields(kwargs, &specs, litellm_core::ocr::wire::BOUND_FIELDS)?;
let input_sources = request_input_sources(
kwargs,
optional_params
.keys()
.map(String::as_str)
.chain(["api_key", "api_base", "extra_headers"]),
)?;
let azure_ad_token_provider = kwargs
.get_item("azure_ad_token_provider")?
.and_then(|provider| PythonTokenProvider::select(provider, AZURE_AD_TOKEN_PROVIDER));
let wire = OcrWireRequest {
model,
document: wire_document,
api_key: api_key.extract()?,
api_base: arguments.api_base()?,
custom_llm_provider,
extra_headers: arguments.extra_headers()?,
optional_params: optional_params.into(),
input_sources,
timeout_seconds: arguments.timeout_seconds()?,
};
let request = decode_request(wire).map_err(ocr_error_to_pyerr)?;
let provider = request.provider_name();
Ok(ProjectedOcrCall {
request: request.with_host_hooks(Arc::new(BridgeOcrHooks), None),
fields: ProjectedOcrFields {
boundary_request,
document: retained_document,
api_key: api_key.unbind(),
azure_ad_token_provider,
provider,
secret_fields: specs
.into_iter()
.filter(|spec| spec.secret)
.map(|spec| spec.name)
.collect(),
},
})
let provider = request.provider_name();
Ok(Self {
request: request.with_host_hooks(Arc::new(BridgeOcrHooks), None),
fields: ProjectedOcrFields {
boundary_request,
document: retained_document,
api_key: api_key.unbind(),
azure_ad_token_provider,
provider,
secret_fields: specs
.into_iter()
.filter(|spec| spec.secret)
.map(|spec| spec.name)
.collect(),
},
})
}
}
pub(super) fn admitted_call(outcome: NativeOutcome<OcrCall>) -> PyResult<OcrCall> {
@ -183,6 +193,19 @@ mod tests {
use super::*;
#[test]
fn projection_selects_consumed_values_without_serializing_host_objects() {
let fields = consumed_optional_params("mistral/model", None).unwrap();
use litellm_core::call_arguments::should_project;
assert!(should_project("future_option", &fields, BOUND_FIELDS));
assert!(should_project("extra_body", &fields, BOUND_FIELDS));
assert!(should_project("id", &fields, BOUND_FIELDS));
assert!(!should_project("metadata", &fields, BOUND_FIELDS));
assert!(!should_project("callbacks", &fields, BOUND_FIELDS));
assert!(!should_project("api_key", &fields, BOUND_FIELDS));
assert!(!should_project("document", &fields, BOUND_FIELDS));
}
fn eval<'py>(py: Python<'py>, source: &std::ffi::CStr) -> Bound<'py, PyDict> {
let locals = PyDict::new(py);
py.run(source, Some(&locals), Some(&locals)).unwrap();
@ -493,7 +516,7 @@ kwargs = {'api_key': key}
wire_document,
serde_json::json!({"type": "mystery", "mystery": "x"})
);
let error = match decode_request(OcrWireRequest {
let error = match LiteLLMOcrRequest::try_from(BridgeOcrRequest {
model: "mistral/mistral-ocr-latest".into(),
document: wire_document,
api_key: None,

View file

@ -0,0 +1,75 @@
use std::collections::BTreeMap;
use std::time::Duration;
use litellm_auth::InputSource;
use litellm_core::call_arguments::CallArguments;
use litellm_core::ocr::{LiteLLMOcrRequest, OcrCredentialInputs, OcrDocument};
use serde_json::{Map, Value};
pub(super) struct BridgeOcrRequest {
pub model: String,
pub document: Value,
pub api_key: Option<String>,
pub api_base: Option<String>,
pub custom_llm_provider: Option<String>,
pub extra_headers: Option<Map<String, Value>>,
pub optional_params: CallArguments,
pub input_sources: BTreeMap<String, InputSource>,
pub timeout_seconds: Option<f64>,
}
impl TryFrom<BridgeOcrRequest> for LiteLLMOcrRequest {
type Error = litellm_core::ocr::Error;
fn try_from(request: BridgeOcrRequest) -> Result<Self, Self::Error> {
let api_key_source = source_for(&request.input_sources, "api_key");
let api_base_source = source_for(&request.input_sources, "api_base");
let extra_headers_source = source_for(&request.input_sources, "extra_headers");
let timeout = request
.timeout_seconds
.map(|seconds| {
Duration::try_from_secs_f64(seconds).map_err(|_| Self::Error::RequestField {
path: "timeout_seconds".into(),
})
})
.transpose()?;
let headers = request
.extra_headers
.unwrap_or_default()
.into_iter()
.map(|(name, value)| {
value
.as_str()
.map(|value| (name.clone(), value.to_string()))
.ok_or_else(|| Self::Error::RequestField {
path: format!("extra_headers.{name}"),
})
})
.collect::<Result<Vec<_>, _>>()?;
let core_request = LiteLLMOcrRequest::new(
request.model,
OcrDocument::try_from(request.document)?,
request.custom_llm_provider.as_deref(),
request.optional_params,
)?;
let transport =
core_request
.transport
.clone()
.with_overrides(headers, extra_headers_source, timeout);
Ok(core_request.with_connection_inputs(
OcrCredentialInputs::new(
request.api_key,
api_key_source,
request.api_base,
api_base_source,
),
transport,
request.input_sources,
))
}
}
fn source_for(sources: &BTreeMap<String, InputSource>, name: &str) -> InputSource {
sources.get(name).copied().unwrap_or_default()
}

View file

@ -1,11 +1,12 @@
use litellm_core::ocr::Error;
use std::future::Future;
use litellm_core::ocr::wire::{OcrWireRequest, decode_request};
use litellm_core::ocr::LiteLLMOcrRequest;
use pyo3::prelude::*;
use serde_json::Value;
use super::errors::to_pyerr as ocr_error_to_pyerr;
use super::request::BridgeOcrRequest;
use crate::marshal::{RouteOptions, RouteOptionsInputs, object_or_empty};
fn prepare_ocr(
@ -37,7 +38,7 @@ fn prepare_ocr(
extra_headers,
timeout,
} = options;
let request = decode_request(OcrWireRequest {
let request = LiteLLMOcrRequest::try_from(BridgeOcrRequest {
model,
document,
api_key,

View file

@ -12,32 +12,16 @@ Tests that:
import base64
import os
import tempfile
from collections.abc import Generator
from io import BytesIO
from pathlib import Path
from typing import Final
from unittest.mock import AsyncMock, MagicMock, Mock
from unittest.mock import AsyncMock, MagicMock
import orjson
import pytest
from starlette.datastructures import FormData
from litellm.ocr.input import convert_file_document_to_url_document, get_mime_type
@pytest.fixture(autouse=True, params=["native", "disabled", "unavailable"])
def document_runtime(request: pytest.FixtureRequest, monkeypatch: pytest.MonkeyPatch) -> Generator[None]:
from litellm.rust_bridge import bindings, configuration
configuration.reset_rust_configuration()
monkeypatch.delenv("LITELLM_RUST", raising=False)
if request.param == "disabled":
monkeypatch.setenv("LITELLM_RUST", "0")
monkeypatch.setattr(bindings, "get_native_bridge", Mock(side_effect=AssertionError("Rust is disabled")))
elif request.param == "unavailable":
monkeypatch.setattr(bindings, "get_native_bridge", lambda: None)
yield
configuration.reset_rust_configuration()
from litellm.ocr.main import convert_file_document_to_url_document, get_mime_type
class TestGetMimeType:
@ -503,10 +487,9 @@ class TestProxySecurityGuard:
async def test_proxy_upload_stops_reading_at_size_limit() -> None:
from starlette.datastructures import UploadFile
from litellm.ocr.input import get_max_file_bytes
from litellm.proxy.ocr_endpoints.endpoints import _parse_multipart_form
limit: Final = get_max_file_bytes()
limit: Final = 50 * 1024 * 1024
with tempfile.TemporaryFile() as stream:
stream.truncate(limit * 2)
upload: Final = UploadFile(file=stream, filename="large.pdf")

View file

@ -548,111 +548,3 @@ async def test_native_ocr_inherits_named_credentials_without_overwriting_argumen
assert response.pages[0].markdown == "native OCR response"
assert ocr_server.requests[0].headers["authorization"] == f"Bearer {expected_key}"
assert ocr_server.requests[0].body["pages"] == [0, 2]
@pytest.mark.parametrize("source", ["sdk", "proxy"])
@pytest.mark.parametrize(
"filename,mime", [("scan.PNG", "image/png"), ("document.pdf", "application/pdf"), ("note.txt", "text/plain")]
)
def test_ocr_file_helpers_use_native_document_preparation(source: str, filename: str, mime: str) -> None:
from io import BytesIO
from litellm.ocr.input import convert_file_document_to_url_document, get_mime_type
from litellm.proxy.ocr_endpoints.endpoints import _build_document_from_upload
file: Final = BytesIO(b"abc")
file.name = filename
document: Final = (
convert_file_document_to_url_document({"type": "file", "file": file})
if source == "sdk"
else _build_document_from_upload(b"abc", filename, "application/octet-stream; charset=utf-8")
)
field: Final = "image_url" if mime.startswith("image/") else "document_url"
assert get_mime_type(filename) == mime
assert document == {"type": field, field: f"data:{mime};base64,YWJj"}
@pytest.mark.parametrize("attribute", ["read", "name"])
def test_native_file_preparation_preserves_property_errors(attribute: str) -> None:
from litellm.ocr.input import convert_file_document_to_url_document
failure: Final = LookupError("file property failed")
class File:
def __getattribute__(self, name: str):
if name == attribute:
raise failure
return super().__getattribute__(name)
def read(self):
return b"abc"
with pytest.raises(LookupError) as caught:
convert_file_document_to_url_document({"type": "file", "file": File()})
assert caught.value is failure
@pytest.mark.parametrize("kind", ["bytes", "path", "reader"])
def test_native_file_preparation_rejects_oversized_input(kind: str, tmp_path: Path) -> None:
from litellm.ocr.input import FileDocument, convert_file_document_to_url_document, get_max_file_bytes
limit: Final = get_max_file_bytes()
path: Final = tmp_path / "large.pdf"
with path.open("wb") as stream:
stream.truncate(limit + 1)
class Reader:
def read(self) -> bytes:
return b"a" * (limit + 1)
document: Final[FileDocument] = {
"type": "file",
"file": path if kind == "path" else Reader() if kind == "reader" else b"a" * (limit + 1),
}
with pytest.raises(ValueError, match="exceeds the size limit"):
convert_file_document_to_url_document(document)
@pytest.mark.parametrize("kind", ["str", "path", "reader"])
def test_native_upload_binding_rejects_filesystem_inputs(kind: str, tmp_path: Path) -> None:
from io import BytesIO
from typing import cast # noqa: TID251 # deliberately invalid inputs exercise the native runtime boundary
from litellm.ocr.input import convert_upload_to_url_document
path: Final = tmp_path / "secret.pdf"
path.write_bytes(b"server secret")
source: Final = str(path) if kind == "str" else path if kind == "path" else BytesIO(b"abc")
with pytest.raises(TypeError):
convert_upload_to_url_document(cast(bytes, source), "document.pdf", None)
@pytest.mark.parametrize("extra_bytes", [0, 1])
def test_native_upload_enforces_file_size_limit(extra_bytes: int) -> None:
import base64
from litellm.ocr.input import convert_upload_to_url_document, get_max_file_bytes
content: Final = b"a" * (get_max_file_bytes() + extra_bytes)
if extra_bytes:
with pytest.raises(ValueError, match="exceeds the size limit"):
convert_upload_to_url_document(content, "scan.pdf", None)
return
document: Final = convert_upload_to_url_document(content, "scan.pdf", None)
assert document["type"] == "document_url"
assert base64.b64decode(document["document_url"].split(",", 1)[1]) == content
def test_native_file_preparation_preserves_reader_exception() -> None:
from litellm.ocr.input import convert_file_document_to_url_document
failure: Final = RuntimeError("reader failed")
class Reader:
def read(self) -> bytes:
raise failure
with pytest.raises(RuntimeError) as caught:
convert_file_document_to_url_document({"type": "file", "file": Reader()})
assert caught.value is failure