refactor(ocr): mirror Python provider layout and preserve tests

This commit is contained in:
Yujong Lee 2026-09-16 19:48:26 -07:00
parent 351a54e849
commit edfa01da81
88 changed files with 8518 additions and 4121 deletions

147
litellm-rust/Cargo.lock generated
View file

@ -948,8 +948,18 @@ version = "0.20.11"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "fc7f46116c46ff9ab3eb1597a45688b6715c6e628b5c133e288e709a29bcb4ee"
dependencies = [
"darling_core",
"darling_macro",
"darling_core 0.20.11",
"darling_macro 0.20.11",
]
[[package]]
name = "darling"
version = "0.21.3"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "9cdf337090841a411e2a7f3deb9187445851f91b309c0c0a29e05f74a00a48c0"
dependencies = [
"darling_core 0.21.3",
"darling_macro 0.21.3",
]
[[package]]
@ -966,13 +976,38 @@ dependencies = [
"syn 2.0.119",
]
[[package]]
name = "darling_core"
version = "0.21.3"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "1247195ecd7e3c85f83c8d2a366e4210d588e802133e1e355180a9870b517ea4"
dependencies = [
"fnv",
"ident_case",
"proc-macro2",
"quote",
"strsim",
"syn 2.0.119",
]
[[package]]
name = "darling_macro"
version = "0.20.11"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "fc34b93ccb385b40dc71c6fceac4b2ad23662c7eeb248cf10d529b7e055b6ead"
dependencies = [
"darling_core",
"darling_core 0.20.11",
"quote",
"syn 2.0.119",
]
[[package]]
name = "darling_macro"
version = "0.21.3"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "d38308df82d1080de0afee5d069fa14b0326a88c14f15c5ccda35b4a6c414c81"
dependencies = [
"darling_core 0.21.3",
"quote",
"syn 2.0.119",
]
@ -1022,7 +1057,7 @@ version = "0.20.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "2d5bcf7b024d6835cfb3d473887cd966994907effbe9227e8c8219824d06c4e8"
dependencies = [
"darling",
"darling 0.20.11",
"proc-macro2",
"quote",
"syn 2.0.119",
@ -1363,7 +1398,7 @@ dependencies = [
"futures-sink",
"futures-util",
"http 0.2.12",
"indexmap",
"indexmap 2.14.0",
"slab",
"tokio",
"tokio-util",
@ -1382,7 +1417,7 @@ dependencies = [
"futures-core",
"futures-sink",
"http 1.4.2",
"indexmap",
"indexmap 2.14.0",
"slab",
"tokio",
"tokio-util",
@ -1400,6 +1435,12 @@ dependencies = [
"zerocopy",
]
[[package]]
name = "hashbrown"
version = "0.12.3"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "8a9ee70c43aaf417c914396645a0fa852624801b24ebb7ae78fe8272889ac888"
[[package]]
name = "hashbrown"
version = "0.17.1"
@ -1736,6 +1777,17 @@ dependencies = [
"icu_properties",
]
[[package]]
name = "indexmap"
version = "1.9.3"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "bd070e393353796e801d209ad339e89596eb4c8d430d18ede6a1cced8fafbd99"
dependencies = [
"autocfg",
"hashbrown 0.12.3",
"serde",
]
[[package]]
name = "indexmap"
version = "2.14.0"
@ -1743,7 +1795,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "d466e9454f08e4a911e14806c24e16fba1b4c121d1ea474396f396069cf949d9"
dependencies = [
"equivalent",
"hashbrown",
"hashbrown 0.17.1",
"serde",
"serde_core",
]
@ -1971,6 +2023,7 @@ dependencies = [
"serde",
"serde_json",
"serde_path_to_error",
"serde_with",
"sha2 0.10.9",
"strum",
"subtle",
@ -2032,7 +2085,7 @@ version = "0.1.0"
dependencies = [
"base64 0.22.1",
"criterion",
"indexmap",
"indexmap 2.14.0",
"itoa",
"rand 0.8.7",
"rstest",
@ -2753,6 +2806,26 @@ dependencies = [
"bitflags",
]
[[package]]
name = "ref-cast"
version = "1.0.27"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "7e440fb4e4b4147295338efb76001ab9e4efc0e5839df2c47fc5ac2381d365c3"
dependencies = [
"ref-cast-impl",
]
[[package]]
name = "ref-cast-impl"
version = "1.0.27"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "92ecd8964f8453721699a1ed72037b0db49ce2f5a5138486ee89bed6f67cdf3a"
dependencies = [
"proc-macro2",
"quote",
"syn 3.0.0",
]
[[package]]
name = "regex"
version = "1.13.1"
@ -3075,6 +3148,30 @@ dependencies = [
"windows-sys 0.61.2",
]
[[package]]
name = "schemars"
version = "0.9.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "4cd191f9397d57d581cddd31014772520aa448f65ef991055d7f61582c65165f"
dependencies = [
"dyn-clone",
"ref-cast",
"serde",
"serde_json",
]
[[package]]
name = "schemars"
version = "1.2.2"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "687274d293b6cdc6e73e0fee520bf2049650090d7164f87672d212a3c530cf4a"
dependencies = [
"dyn-clone",
"ref-cast",
"serde",
"serde_json",
]
[[package]]
name = "scopeguard"
version = "1.2.0"
@ -3156,6 +3253,7 @@ version = "1.0.150"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "e8014e44b4736ed0538adeecded0fce2a272f22dc9578a7eb6b2d9993c74cfb9"
dependencies = [
"indexmap 2.14.0",
"itoa",
"memchr",
"serde",
@ -3186,6 +3284,37 @@ dependencies = [
"serde",
]
[[package]]
name = "serde_with"
version = "3.16.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "4fa237f2807440d238e0364a218270b98f767a00d3dada77b1c53ae88940e2e7"
dependencies = [
"base64 0.22.1",
"chrono",
"hex",
"indexmap 1.9.3",
"indexmap 2.14.0",
"schemars 0.9.0",
"schemars 1.2.2",
"serde_core",
"serde_json",
"serde_with_macros",
"time",
]
[[package]]
name = "serde_with_macros"
version = "3.16.1"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "52a8e3ca0ca629121f70ab50f95249e5a6f925cc0f6ffe8256c45b728875706c"
dependencies = [
"darling 0.21.3",
"proc-macro2",
"quote",
"syn 2.0.119",
]
[[package]]
name = "sha1"
version = "0.10.7"
@ -3661,7 +3790,7 @@ version = "0.25.13+spec-1.1.0"
source = "registry+https://github.com/rust-lang/crates.io-index"
checksum = "6975367e4d2ef766d86af01ffad14b622fecc8d4357a998fbc4deb6e9bacaf9b"
dependencies = [
"indexmap",
"indexmap 2.14.0",
"toml_datetime",
"toml_parser",
"winnow",

View file

@ -29,6 +29,7 @@ rustls = { version = "0.23", default-features = false, features = ["ring", "std"
rustls-native-certs = "0.8"
serde = { version = "1.0", features = ["derive"] }
serde_json = { version = "1.0", features = ["float_roundtrip"] }
serde_with = { version = "=3.16.1", default-features = false, features = ["std", "macros"] }
sha2 = "0.10"
subtle = "2"
thiserror = "2.0"

View file

@ -22,7 +22,8 @@ reqwest.workspace = true
rustls.workspace = true
rustls-native-certs.workspace = true
serde.workspace = true
serde_json.workspace = true
serde_json = { workspace = true, features = ["preserve_order"] }
serde_with.workspace = true
serde_path_to_error = "0.1"
strum.workspace = true
subtle.workspace = true

View file

@ -0,0 +1,467 @@
use std::ops::Deref;
use serde::{Deserialize, Serialize, de::DeserializeOwned};
use serde_json::{Map, Value};
#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)]
#[serde(transparent)]
pub struct CallArguments(Map<String, Value>);
impl CallArguments {
pub(crate) fn select(&self, names: &[&str]) -> Map<String, Value> {
self.iter()
.filter(|(name, _)| names.contains(&name.as_str()))
.map(|(name, value)| (name.clone(), value.clone()))
.collect()
}
}
#[derive(Clone, Debug, PartialEq, Eq, thiserror::Error)]
#[error("invalid argument: {path}")]
pub struct ArgumentError {
pub path: String,
}
pub fn parse_options<T: DeserializeOwned>(arguments: &CallArguments) -> Result<T, ArgumentError> {
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)]
pub struct ArgumentSpec {
pub name: &'static str,
pub secret: bool,
}
pub fn should_project(name: &str, consumed: &[ArgumentSpec], bound_fields: &[&str]) -> bool {
consumed.iter().any(|field| field.name == name)
|| (!bound_fields.contains(&name) && !is_control(name))
}
pub fn is_control(name: &str) -> bool {
crate::params::is_control_param(name) || HOST_CONTROLS.contains(&name)
}
const HOST_CONTROLS: &[&str] = &[
"_agentic_loop_api_surface",
"_agentic_loop_depth",
"_agentic_loop_fingerprints",
"_code_interpreter_interception_active",
"_code_interpreter_interception_converted_stream",
"_code_interpreter_interception_sandbox_key",
"_code_interpreter_interception_session_scoped",
"_headroom_interception_converted_stream",
"_litellm_strip_stream_usage",
"_router_weights",
"_websearch_interception_converted_stream",
"_websearch_interception_emit_native_blocks",
"acompletion",
"adaptive_router_config",
"adaptive_router_default_model",
"aembedding",
"aimg_generation",
"allm_passthrough_route",
"allow_client_keepalive_override",
"allowed_model_region",
"allowed_openai_params",
"annotation_cost_per_page",
"api_version",
"arize_api_key",
"arize_space_id",
"arize_space_key",
"assistant_continue_message",
"async_call",
"atext_completion",
"attempted_targets",
"auto_router_config",
"auto_router_config_path",
"auto_router_default_model",
"auto_router_embedding_model",
"auto_router_max_input_chars",
"auto_router_model_compression",
"auto_router_routing_compression",
"aws_batch_role_arn",
"azure",
"azure_password",
"azure_username",
"base_model",
"bedrock_tags",
"bos_token",
"budget_duration",
"cache",
"cache_creation_input_audio_token_cost",
"cache_creation_input_token_cost",
"cache_creation_input_token_cost_above_1hr",
"cache_creation_input_token_cost_above_200k_tokens",
"cache_creation_input_token_cost_above_272k_tokens",
"cache_creation_input_token_cost_above_272k_tokens_flex",
"cache_creation_input_token_cost_above_272k_tokens_priority",
"cache_creation_input_token_cost_flex",
"cache_creation_input_token_cost_priority",
"cache_creation_input_token_cost_ultrafast",
"cache_key",
"cache_read_input_audio_token_cost",
"cache_read_input_token_cost",
"cache_read_input_token_cost_above_200k_tokens",
"cache_read_input_token_cost_above_200k_tokens_priority",
"cache_read_input_token_cost_above_272k_tokens",
"cache_read_input_token_cost_above_272k_tokens_flex",
"cache_read_input_token_cost_above_272k_tokens_priority",
"cache_read_input_token_cost_above_512k_tokens",
"cache_read_input_token_cost_flex",
"cache_read_input_token_cost_priority",
"cache_read_input_token_cost_ultrafast",
"caching",
"caching_groups",
"citation_cost_per_token",
"client",
"client_side_timeout",
"complete_response",
"completion_call_id",
"complexity_router_config",
"complexity_router_default_model",
"configurable_clientside_auth_params",
"context_window_fallback_dict",
"cooldown_time",
"cost_per_query",
"custom_prompt_dict",
"data_residency",
"dd_agent_host",
"dd_agent_port",
"dd_api_key",
"dd_site",
"default_api_key_rpm_limit",
"default_api_key_tpm_limit",
"disable_add_transform_inline_image_block",
"enable_json_schema_validation",
"enable_prompt_caching",
"enable_tag_filtering",
"ensure_alternating_roles",
"eos_token",
"fallback_depth",
"fallbacks",
"fastest_response",
"final_prompt_value",
"force_timeout",
"gcs_bucket_name",
"gcs_path_service_account",
"google_maps_grounding_cost_per_query",
"headers",
"hf_model_name",
"humanloop_api_key",
"id",
"input_cost_per_audio_per_second",
"input_cost_per_audio_per_second_above_128k_tokens",
"input_cost_per_audio_token",
"input_cost_per_audio_token_batches",
"input_cost_per_character",
"input_cost_per_character_above_128k_tokens",
"input_cost_per_image",
"input_cost_per_image_above_128k_tokens",
"input_cost_per_image_token",
"input_cost_per_image_token_batches",
"input_cost_per_pixel",
"input_cost_per_query",
"input_cost_per_second",
"input_cost_per_token",
"input_cost_per_token_above_128k_tokens",
"input_cost_per_token_above_200k_tokens",
"input_cost_per_token_above_200k_tokens_priority",
"input_cost_per_token_above_272k_tokens",
"input_cost_per_token_above_272k_tokens_flex",
"input_cost_per_token_above_272k_tokens_priority",
"input_cost_per_token_above_512k_tokens",
"input_cost_per_token_batches",
"input_cost_per_token_cache_hit",
"input_cost_per_token_flex",
"input_cost_per_token_priority",
"input_cost_per_token_ultrafast",
"input_cost_per_video_per_second",
"input_cost_per_video_per_second_above_128k_tokens",
"input_cost_per_video_per_second_above_15s_interval",
"input_cost_per_video_per_second_above_8s_interval",
"input_cost_per_video_token",
"input_cost_per_video_token_batches",
"itpm",
"keepalive_seconds",
"langfuse_environment",
"langfuse_host",
"langfuse_prompt_version",
"langfuse_public_key",
"langfuse_secret",
"langfuse_secret_key",
"langsmith_api_key",
"langsmith_base_url",
"langsmith_project",
"langsmith_sampling_rate",
"langsmith_tenant_id",
"litellm_credential_name",
"litellm_disabled_callbacks",
"litellm_request_debug",
"litellm_session_id",
"litellm_system_prompt",
"litellm_trace_id",
"litellm_trusted_callback_vars",
"logger_fn",
"max_agentic_loops",
"max_budget",
"max_fallbacks",
"max_parallel_requests",
"merge_reasoning_content_in_choices",
"metadata",
"mock_response",
"mock_timeout",
"model_alias_map",
"model_config",
"model_file_id_mapping",
"model_info",
"model_list",
"newrelic_api_key",
"newrelic_region",
"no-log",
"num_retries",
"ocr_cost_per_credit",
"ocr_cost_per_page",
"order",
"otpm",
"output_cost_per_audio_per_second",
"output_cost_per_audio_token",
"output_cost_per_character",
"output_cost_per_character_above_128k_tokens",
"output_cost_per_image",
"output_cost_per_image_token",
"output_cost_per_pixel",
"output_cost_per_reasoning_token",
"output_cost_per_reasoning_token_flex",
"output_cost_per_reasoning_token_priority",
"output_cost_per_second",
"output_cost_per_second_1080p",
"output_cost_per_second_480p",
"output_cost_per_second_4k",
"output_cost_per_second_720p",
"output_cost_per_token",
"output_cost_per_token_above_128k_tokens",
"output_cost_per_token_above_200k_tokens",
"output_cost_per_token_above_200k_tokens_priority",
"output_cost_per_token_above_272k_tokens",
"output_cost_per_token_above_272k_tokens_flex",
"output_cost_per_token_above_272k_tokens_priority",
"output_cost_per_token_above_512k_tokens",
"output_cost_per_token_batches",
"output_cost_per_token_flex",
"output_cost_per_token_priority",
"output_cost_per_token_ultrafast",
"output_cost_per_video_per_second",
"output_cost_per_video_token",
"output_vector_size",
"posthog_api_key",
"posthog_api_url",
"preset_cache_key",
"prompt_environment",
"prompt_id",
"prompt_label",
"prompt_variables",
"prompt_version",
"provider_specific_header",
"quality_router_config",
"quality_router_default_model",
"region_name",
"regional_endpoint_uplift_multiplier",
"regional_processing_uplift_multiplier_eu",
"regional_processing_uplift_multiplier_us",
"retry_policy",
"retry_strategy",
"roles",
"routing_strategy",
"rpm",
"rust",
"s3_bucket_name",
"s3_output_bucket_name",
"s3_region_name",
"search_context_cost_per_query",
"search_tool_name",
"secret_fields",
"self",
"shared_session",
"ssl_verify",
"stream_response",
"stream_timeout",
"supports_system_message",
"tags",
"text_completion",
"tiered_pricing",
"tpm",
"ttl",
"turn_off_message_logging",
"use_chat_completions_api",
"use_client",
"use_in_pass_through",
"use_litellm_proxy",
"use_xai_oauth",
"user_continue_message",
"verbose",
"wandb_api_key",
"weave_project_id",
"weight",
];
pub fn compose_body<B: Serialize>(
arguments: &CallArguments,
body: &B,
consumed: &[&str],
) -> Result<Value, crate::params::Error> {
let Value::Object(fields) =
serde_json::to_value(body).map_err(|_| crate::params::Error::Body)?
else {
return Err(crate::params::Error::Body);
};
let overrides = match arguments.get("extra_body") {
None | Some(Value::Null) => None,
Some(Value::Object(fields)) => Some(fields),
Some(_) => return Err(crate::params::Error::ExtraBody),
};
let extensions = arguments.iter().filter(|(name, _)| {
!consumed.contains(&name.as_str()) && name.as_str() != "extra_body" && !is_control(name)
});
Ok(Value::Object(
fields
.into_iter()
.chain(
extensions
.chain(overrides.into_iter().flatten())
.filter(|(name, _)| {
name.as_str() != "model"
&& name.as_str() != "extra_body"
&& !crate::params::is_control_param(name)
})
.map(|(name, value)| (name.clone(), value.clone())),
)
.collect(),
))
}
impl Deref for CallArguments {
type Target = Map<String, Value>;
fn deref(&self) -> &Self::Target {
&self.0
}
}
impl From<Map<String, Value>> for CallArguments {
fn from(values: Map<String, Value>) -> Self {
Self(values)
}
}
impl From<CallArguments> for Map<String, Value> {
fn from(arguments: CallArguments) -> Self {
arguments.0
}
}
impl FromIterator<(String, Value)> for CallArguments {
fn from_iter<T: IntoIterator<Item = (String, Value)>>(iter: T) -> Self {
Self(iter.into_iter().collect())
}
}
impl IntoIterator for CallArguments {
type Item = (String, Value);
type IntoIter = serde_json::map::IntoIter;
fn into_iter(self) -> Self::IntoIter {
self.0.into_iter()
}
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
#[test]
fn composition_preserves_extensions_and_applies_shallow_explicit_overrides() {
let original = json!({
"known": false, "future": {"old": 1}, "null": null, "zero": 0,
"metadata": {"host": true}, "shared_session": "host", "api_key": "secret",
"extra_body": {
"known": null, "future": {"new": [false, 0, null]},
"metadata": {"provider": true}, "model": "ignored", "api_key": "ignored"
}
});
let arguments = serde_json::from_value(original.clone()).unwrap();
let body = compose_body(
&arguments,
&json!({"model":"resolved", "known":false}),
&["known"],
)
.unwrap();
assert_eq!(
body,
json!({
"model":"resolved", "known":null, "future":{"new":[false,0,null]},
"null":null, "zero":0, "metadata":{"provider":true}
})
);
assert_eq!(serde_json::to_value(arguments).unwrap(), original);
}
#[test]
fn projection_prioritizes_consumed_fields_and_keeps_unknown_names() {
let fields = [ArgumentSpec {
name: "id",
secret: false,
}];
assert!(should_project("id", &fields, &[]));
assert!(!should_project("id", &[], &[]));
assert!(should_project("future_option", &[], &[]));
assert!(!should_project("document", &fields, &["document"]));
assert!(!should_project("metadata", &fields, &[]));
assert!(!should_project("callbacks", &fields, &[]));
assert!(!should_project("ocr_cost_per_page", &fields, &[]));
}
#[test]
fn invalid_extra_body_is_rejected_without_coercing_it_to_empty() {
for value in [json!(false), json!(0), json!([]), json!("")] {
let arguments = serde_json::from_value(json!({"extra_body":value})).unwrap();
assert_eq!(
compose_body(&arguments, &json!({}), &[]),
Err(crate::params::Error::ExtraBody)
);
}
let arguments = serde_json::from_value(json!({"extra_body":null})).unwrap();
assert_eq!(
compose_body(&arguments, &json!({}), &[]).unwrap(),
json!({})
);
}
#[test]
fn typed_views_preserve_missing_and_explicit_null_in_the_source() {
#[derive(Deserialize)]
struct Options {
enabled: Option<bool>,
}
let arguments: CallArguments =
serde_json::from_value(json!({"enabled":null,"future":0})).unwrap();
assert!(
parse_options::<Options>(&arguments)
.unwrap()
.enabled
.is_none()
);
assert_eq!(arguments.get("enabled"), Some(&Value::Null));
assert_eq!(arguments.get("missing"), None);
let invalid = serde_json::from_value(json!({"enabled":0})).unwrap();
assert_eq!(
parse_options::<Options>(&invalid).err().unwrap().path,
"enabled"
);
}
}

View file

@ -1,14 +1,18 @@
pub mod audio_transcription;
pub mod call_arguments;
pub mod call_lifecycle;
pub mod chat_completions;
pub mod constants;
pub mod error;
pub mod http_utils;
pub(crate) mod llms;
mod media;
pub mod messages;
pub mod ocr;
pub mod params;
pub mod providers;
pub mod responses;
mod serde_compat;
pub mod transport;
mod url_utils;

View file

@ -0,0 +1 @@
pub(crate) mod ocr;

View file

@ -0,0 +1,165 @@
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::types::{LiteLLMOcrResponse, OcrDocument, PreparedOcrRequest};
use crate::url_utils::ApiUrl;
use serde_json::Value;
#[derive(Default)]
pub(crate) struct AzureAICohereParseConfig;
impl BaseOcrConfig for AzureAICohereParseConfig {
type OcrParams = CohereOptions;
type ProviderRequest = CohereRequest;
type Environment = Vec<(String, String)>;
fn get_api_key_env_var(&self) -> Option<&'static str> {
super::transformation::AzureAIOCRConfig.get_api_key_env_var()
}
fn get_health_check_document(&self) -> OcrDocument {
CohereParseConfig.get_health_check_document()
}
async fn validate_environment(
&self,
request: &PreparedOcrRequest,
client: &OcrClient,
) -> Result<Self::Environment, crate::ocr::Error> {
BaseOcrConfig::validate_environment(
&super::transformation::AzureAIOCRConfig,
request,
client,
)
.await
}
fn get_complete_url(
&self,
request: &PreparedOcrRequest,
_params: &Self::OcrParams,
_environment: &Self::Environment,
) -> Result<String, crate::ocr::Error> {
let base = super::transformation::AzureAIOCRConfig::resolve_api_base(
request.connection.api_base.as_deref(),
&crate::ocr::prepare::credential_env,
)?;
self.get_complete_url(&base)
}
fn transform_ocr_request(
&self,
model: &str,
document: OcrDocument,
params: &CohereOptions,
headers: &[(String, String)],
) -> Result<CohereRequest, crate::ocr::Error> {
CohereParseConfig.transform_ocr_request(model, document, params, headers)
}
fn get_supported_ocr_params(&self, model: &str) -> &'static [&'static str] {
CohereParseConfig.get_supported_ocr_params(model)
}
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,
document: OcrDocument,
optional_params: &CohereOptions,
headers: &[(String, String)],
context: OcrRequestContext<'_>,
) -> Result<CohereRequest, crate::ocr::Error> {
validate_document(&document)?;
let document = inline_remote_document(
context.client.document_fetcher(),
document,
context.connection,
)
.await?;
self.transform_ocr_request(model, document, optional_params, headers)
}
fn transform_ocr_response(
&self,
model: &str,
raw_response: &[u8],
request_format: crate::ocr::types::OcrResponseFormat,
) -> Result<LiteLLMOcrResponse, crate::ocr::Error> {
CohereParseConfig.transform_ocr_response(model, raw_response, request_format)
}
fn validate_request_body(&self, body: &Value) -> Result<(), crate::ocr::Error> {
let document = crate::ocr::prepare::body_document(body)?;
validate_document(&document)?;
validate_inline_document(&document)
}
}
impl AzureAICohereParseConfig {
fn get_complete_url(&self, base: &str) -> Result<String, crate::ocr::Error> {
let mut url = reqwest::Url::parse(base).map_err(|_| invalid_api_base())?;
if !matches!(url.scheme(), "http" | "https") {
return Err(invalid_api_base());
}
let path = url.path().trim_end_matches('/').to_string();
if path.ends_with("/v2/parse") {
url.set_path(&path);
return Ok(url.into());
}
url.set_path(path.strip_suffix("/models").unwrap_or(&path));
ApiUrl::parse(url.as_str())
.and_then(|url| url.complete_path(&["providers", "cohere", "v2", "parse"]))
.map(|url| url.into_string())
.map_err(|_| invalid_api_base())
}
}
fn invalid_api_base() -> crate::ocr::Error {
crate::ocr::Error::RequestField {
path: "api_base".into(),
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn completes_foundry_urls_without_duplicate_paths_and_preserves_queries() {
for suffix in [
"",
"/models",
"/providers/cohere/v2",
"/providers/cohere/v2/parse",
] {
assert_eq!(
AzureAICohereParseConfig
.get_complete_url(&format!("https://example.com{suffix}?tenant=a"))
.unwrap(),
"https://example.com/providers/cohere/v2/parse?tenant=a"
);
}
assert_eq!(
AzureAICohereParseConfig
.get_complete_url("https://example.com/v2/parse?tenant=a")
.unwrap(),
"https://example.com/v2/parse?tenant=a"
);
assert!(
AzureAICohereParseConfig
.get_complete_url("relative/path")
.is_err()
);
}
}

View file

@ -1,25 +1,13 @@
mod cohere;
mod document_intelligence;
mod mistral;
use std::sync::OnceLock;
use crate::ocr::Error;
use crate::ocr::error::OcrError;
use crate::ocr::types::OcrConnection;
use litellm_auth::{InputSource, Sourced};
use litellm_auth_azure::{AzureAuthInputs, AzureAuthService};
pub(crate) use cohere::AzureCohereAdapter;
pub(crate) use document_intelligence::AzureDocumentIntelligenceAdapter;
pub(crate) use mistral::AzureMistralAdapter;
pub(super) use mistral::validate_environment as validate_ai_environment;
async fn resolve_entra(
pub(super) async fn resolve_entra(
config: &AzureAuthInputs,
env_lookup: &(dyn Fn(&str) -> Option<String> + Sync),
) -> Result<Option<Sourced<String>>, Error> {
) -> Result<Option<Sourced<String>>, crate::ocr::Error> {
static SERVICE: OnceLock<AzureAuthService> = OnceLock::new();
SERVICE
.get_or_init(AzureAuthService::default)
@ -36,18 +24,18 @@ async fn resolve_entra(
Sourced::new(value, source)
})
})
.map_err(Error::from)
.map_err(crate::ocr::Error::from)
}
fn validate_destination(
pub(super) fn validate_destination(
connection: &OcrConnection,
credential_source: InputSource,
) -> Result<(), OcrError> {
) -> Result<(), crate::ocr::Error> {
if connection.api_base.is_some()
&& connection.api_base_source == InputSource::Request
&& credential_source != InputSource::Request
{
return Err(Error::from(litellm_auth::Error::RequestAzureCredentialDestination).into());
return Err(litellm_auth::Error::RequestAzureCredentialDestination.into());
}
Ok(())
}

View file

@ -0,0 +1 @@
pub(crate) mod transformation;

View file

@ -0,0 +1,4 @@
pub(crate) mod cohere_parse_transformation;
pub(crate) mod common_utils;
pub(crate) mod document_intelligence;
pub(crate) mod transformation;

View file

@ -0,0 +1,399 @@
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;
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";
#[derive(Clone, Debug, Default)]
pub(crate) struct AzureAIOCRConfig;
impl BaseOcrConfig for AzureAIOCRConfig {
type OcrParams = OpaqueParams;
type ProviderRequest = MistralOcrRequest;
type Environment = Vec<(String, String)>;
fn get_api_key_env_var(&self) -> Option<&'static str> {
Some(AZURE_AI_API_KEY_ENV)
}
async fn validate_environment(
&self,
request: &PreparedOcrRequest,
_client: &OcrClient,
) -> Result<Self::Environment, crate::ocr::Error> {
let config = AzureAuthInputs {
azure_ad_token_provider: request.azure_ad_token_provider.clone(),
..AzureAuthInputs::from_sourced_optional_params(
&request.optional_params,
&request.input_sources,
)?
};
self.validate_environment(&request.connection, &config, &credential_env)
.await
}
fn get_complete_url(
&self,
request: &PreparedOcrRequest,
_params: &Self::OcrParams,
_environment: &Self::Environment,
) -> Result<String, crate::ocr::Error> {
self.get_complete_url(request.connection.api_base.as_deref(), &credential_env)
}
fn transform_ocr_request(
&self,
model: &str,
document: OcrDocument,
params: &OpaqueParams,
headers: &[(String, String)],
) -> Result<MistralOcrRequest, crate::ocr::Error> {
MistralOCRConfig.transform_ocr_request(model, document, params, headers)
}
fn get_supported_ocr_params(&self, model: &str) -> &'static [&'static str] {
MistralOCRConfig.get_supported_ocr_params(model)
}
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,
document: OcrDocument,
optional_params: &OpaqueParams,
headers: &[(String, String)],
context: OcrRequestContext<'_>,
) -> Result<MistralOcrRequest, crate::ocr::Error> {
let document = inline_remote_document(
context.client.document_fetcher(),
document,
context.connection,
)
.await?;
self.transform_ocr_request(model, document, optional_params, headers)
}
fn transform_ocr_response(
&self,
model: &str,
raw_response: &[u8],
request_format: crate::ocr::types::OcrResponseFormat,
) -> Result<LiteLLMOcrResponse, crate::ocr::Error> {
MistralOCRConfig.transform_ocr_response(model, raw_response, request_format)
}
fn validate_request_body(&self, body: &Value) -> Result<(), crate::ocr::Error> {
validate_inline_document(&crate::ocr::prepare::body_document(body)?)
}
}
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 = 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))
.map(|url| url.into_string())
.map_err(|_| crate::ocr::Error::RequestField {
path: "api_base".into(),
})
}
pub(super) async fn validate_environment(
&self,
connection: &OcrConnection,
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?;
}
super::common_utils::validate_destination(connection, connection.extra_headers_source)?;
return Ok(connection.extra_headers.clone());
}
let key = nonblank(connection.api_key.clone())
.map(|value| Sourced::new(value, connection.api_key_source))
.or_else(|| {
nonblank(self.get_api_key_env_var().and_then(env_lookup))
.map(|value| Sourced::new(value, InputSource::Environment))
});
if let Some(key) = key {
super::common_utils::validate_destination(connection, key.source())?;
return Ok(bearer_headers(connection, key.value()));
}
let key = super::common_utils::resolve_entra(config, env_lookup)
.await?
.ok_or(crate::ocr::Error::MissingAzureAiCredentials)?;
super::common_utils::validate_destination(connection, key.source())?;
Ok(bearer_headers(connection, key.value()))
}
}
fn bearer_headers(connection: &OcrConnection, key: &str) -> Vec<(String, String)> {
std::iter::once(("Authorization".into(), format!("Bearer {key}")))
.chain(connection.extra_headers.clone())
.collect()
}
fn nonblank(value: Option<String>) -> Option<String> {
value
.map(|value| value.trim().to_string())
.filter(|value| !value.is_empty())
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn completes_azure_path_and_preserves_query() {
assert_eq!(
AzureAIOCRConfig
.get_complete_url(Some("https://example.com/?tenant=a"), &|_| None)
.unwrap(),
"https://example.com/providers/mistral/azure/ocr?tenant=a"
);
assert_eq!(
AzureAIOCRConfig
.get_complete_url(
Some("https://example.com/providers/mistral/azure/ocr"),
&|_| None
)
.unwrap(),
"https://example.com/providers/mistral/azure/ocr"
);
}
#[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()
};
assert_eq!(
AzureAIOCRConfig
.validate_environment(&connection, &Default::default(), &|_| {
Some("environment-key".into())
})
.await
.unwrap(),
connection.extra_headers
);
}
#[tokio::test]
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!(
AzureAIOCRConfig
.validate_environment(&connection, &Default::default(), &|_| {
Some("environment-key".into())
})
.await
.unwrap()[0],
("Authorization".into(), "Bearer request-key".into())
);
}
#[tokio::test]
async fn request_endpoint_cannot_receive_environment_key() {
let connection = OcrConnection {
api_base: Some("https://request.example".into()),
api_base_source: InputSource::Request,
..Default::default()
};
let error = AzureAIOCRConfig
.validate_environment(&connection, &Default::default(), &|name| {
(name == AZURE_AI_API_KEY_ENV).then(|| "environment-key".into())
})
.await
.unwrap_err();
assert!(
error
.to_string()
.contains("request-controlled Azure endpoint")
);
}
#[tokio::test]
async fn request_endpoint_accepts_request_owned_key() {
let connection = OcrConnection {
api_key: Some("request-key".into()),
api_key_source: InputSource::Request,
api_base: Some("https://request.example".into()),
api_base_source: InputSource::Request,
..Default::default()
};
let headers = AzureAIOCRConfig
.validate_environment(&connection, &Default::default(), &|_| None)
.await
.unwrap();
assert_eq!(
headers[0],
("Authorization".into(), "Bearer request-key".into())
);
}
use std::sync::Arc;
use serde_json::json;
use crate::ocr::hooks::{OcrDuringCallRequest, OcrHookFuture, OcrHooks};
use crate::ocr::test_support::{MockResponse, mock_server, perform_ocr, wire_request};
#[tokio::test]
async fn facade_executes_azure_mistral_with_prepared_auth() {
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({
"pages":[{"index":0,"markdown":"hello"}],
"usage_info":{"pages_processed":1}
}))])
.await;
let mut request = wire_request(
"azure_ai/model",
&base,
json!({"include_image_base64":true}),
);
request.credentials.api_key = None;
request.transport.extra_headers = vec![(
"Authorization".into(),
"Bearer python-prepared-token".into(),
)];
let result = perform_ocr(request).await.unwrap();
server.await.unwrap();
assert_eq!(result.pages[0].markdown, "hello");
let requests = seen.lock().unwrap();
assert_eq!(requests.len(), 1);
assert!(requests[0].starts_with("POST /providers/mistral/azure/ocr "));
assert!(
requests[0]
.to_ascii_lowercase()
.contains("authorization: bearer python-prepared-token\r\n")
);
let body: Value =
serde_json::from_str(requests[0].split_once("\r\n\r\n").unwrap().1).unwrap();
assert_eq!(
body,
json!({
"model":"model",
"document":{"type":"document_url","document_url":"data:application/pdf;base64,YWJj"},
"include_image_base64":true
})
);
}
#[tokio::test]
async fn facade_acquires_supplied_entra_token_for_final_request() {
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({"pages":[]}))]).await;
let mut request = wire_request(
"azure_ai/model",
&base,
json!({"azure_ad_token":"rust-owned-token"}),
);
request.credentials.api_key = None;
perform_ocr(request).await.unwrap();
server.await.unwrap();
let requests = seen.lock().unwrap();
assert_eq!(requests.len(), 1);
assert!(
requests[0]
.to_ascii_lowercase()
.contains("authorization: bearer rust-owned-token\r\n")
);
}
struct ReplaceBodyDocument;
impl OcrHooks for ReplaceBodyDocument {
fn intercepts_requests(&self) -> bool {
true
}
fn during_call(
&self,
mut request: OcrDuringCallRequest,
) -> OcrHookFuture<'_, OcrDuringCallRequest> {
Box::pin(async move {
request.body["document"] = json!({
"type":"document_url",
"document_url":"https://example.com/not-inline.pdf"
});
Ok(request)
})
}
}
#[tokio::test]
async fn rejects_non_inline_body_after_guardrails() {
let mut request = wire_request("azure_ai/model", "http://127.0.0.1:1", json!({}));
request.hooks = Arc::new(ReplaceBodyDocument);
let error = perform_ocr(request).await.unwrap_err();
assert!(error.to_string().contains("data URI"));
}
}

View file

@ -0,0 +1 @@
pub(crate) mod ocr;

View file

@ -0,0 +1 @@
pub(crate) mod transformation;

View file

@ -0,0 +1,211 @@
use std::future::Future;
use std::sync::Arc;
use serde::Serialize;
use serde::de::DeserializeOwned;
use serde_json::Value;
use crate::call_arguments::CallArguments;
use crate::ocr::OcrClient;
use crate::ocr::hooks::OcrHooks;
use crate::ocr::types::{
LiteLLMOcrResponse, OcrConnection, OcrCredentialInputs, OcrDocument, OcrResponseFormat,
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: Send + Sync;
type ProviderRequest: Serialize + Send;
type Environment: OcrEnvironment;
fn get_api_key_env_var(&self) -> Option<&'static str> {
None
}
fn resolve_connection_params(&self, inputs: OcrCredentialInputs) -> ResolvedOcrCredentials {
ResolvedOcrCredentials {
api_key: inputs
.dynamic_api_key
.filter(|value| !value.value().is_empty())
.or(inputs.api_key),
api_base: inputs
.dynamic_api_base
.filter(|value| !value.value().is_empty())
.or(inputs.api_base),
}
}
fn get_health_check_document(&self) -> OcrDocument {
OcrDocument::DocumentUrl {
document_url: HEALTH_CHECK_PDF_DATA_URI.into(),
extra_fields: Default::default(),
}
}
fn validate_environment(
&self,
request: &PreparedOcrRequest,
client: &OcrClient,
) -> impl Future<Output = Result<Self::Environment, crate::ocr::Error>> + Send;
fn get_complete_url(
&self,
request: &PreparedOcrRequest,
optional_params: &Self::OcrParams,
environment: &Self::Environment,
) -> Result<String, crate::ocr::Error>;
fn get_supported_ocr_params(&self, _model: &str) -> &'static [&'static str] {
&[]
}
fn map_ocr_params(
&self,
arguments: &CallArguments,
model: &str,
) -> Result<Self::OcrParams, crate::ocr::Error>;
fn transform_ocr_request(
&self,
model: &str,
document: OcrDocument,
optional_params: &Self::OcrParams,
headers: &[(String, String)],
) -> Result<Self::ProviderRequest, crate::ocr::Error>;
fn async_transform_ocr_request(
&self,
model: &str,
document: OcrDocument,
optional_params: &Self::OcrParams,
headers: &[(String, String)],
_context: OcrRequestContext<'_>,
) -> impl Future<Output = Result<Self::ProviderRequest, crate::ocr::Error>> + Send {
async move { self.transform_ocr_request(model, document, optional_params, headers) }
}
fn transform_ocr_response(
&self,
model: &str,
raw_response: &[u8],
request_format: OcrResponseFormat,
) -> Result<LiteLLMOcrResponse, crate::ocr::Error>;
fn async_transform_ocr_response(
&self,
model: &str,
raw_response: reqwest::Response,
context: OcrResponseContext<'_>,
) -> impl Future<Output = Result<LiteLLMOcrResponse, crate::ocr::Error>> + Send {
async move {
let bytes = crate::ocr::client::read_response_bytes(
raw_response,
context.connection.max_response_bytes,
)
.await?;
crate::ocr::handler::post_call(context.hooks, &bytes).await?;
self.transform_ocr_response(model, &bytes, context.request_format)
}
}
fn get_error_class(
&self,
error_message: String,
status_code: u16,
headers: Vec<(String, String)>,
) -> crate::ocr::Error {
crate::ocr::Error::Provider {
status: status_code,
body: error_message,
headers,
}
}
/// 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,
body,
|body| self.validate_request_body(body),
)
.await
}
}
}
pub(crate) fn decode_and_normalize_response<T: DeserializeOwned>(
model: &str,
raw_response: &[u8],
request_format: OcrResponseFormat,
normalize: impl FnOnce(&str, T) -> Result<LiteLLMOcrResponse, crate::ocr::Error>,
) -> Result<LiteLLMOcrResponse, crate::ocr::Error> {
let decoded = crate::ocr::json::decode_response(
raw_response,
request_format == OcrResponseFormat::Native,
)?;
Ok(LiteLLMOcrResponse {
provider_native_response: decoded.native,
..normalize(model, decoded.data)?
})
}
#[derive(Clone, Copy)]
pub(crate) struct OcrRequestContext<'a> {
pub client: &'a OcrClient,
pub connection: &'a OcrConnection,
}
#[derive(Clone, Copy)]
pub(crate) struct OcrResponseContext<'a> {
pub client: &'a OcrClient,
pub connection: &'a OcrConnection,
pub hooks: &'a Arc<dyn OcrHooks>,
pub request_format: OcrResponseFormat,
pub url: &'a str,
pub headers: &'a [(String, String)],
}

View file

@ -0,0 +1 @@
pub(crate) mod ocr;

View file

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

View file

@ -0,0 +1,740 @@
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;
use crate::ocr::types::{
LiteLLMOcrResponse, OcrConnection, OcrDocument, OcrPage, OcrPageImage, OcrUsageInfo,
PreparedOcrRequest,
};
use crate::serde_compat::LaxI64;
use crate::url_utils::ApiUrl;
const COHERE_PARSE_HEALTH_CHECK_IMAGE_DATA_URI: &str = "data:image/png;base64,iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAIAAACQd1PeAAAADElEQVR4nGP4//8/AAX+Av4N70a4AAAAAElFTkSuQmCC";
#[derive(Clone, Copy, Debug, Default, Deserialize, Serialize)]
#[serde(rename_all = "lowercase")]
pub(crate) enum OutputFormat {
#[default]
Markdown,
Blocks,
}
#[derive(Default, Deserialize, Serialize)]
pub(crate) struct CohereOptions {
#[serde(skip_serializing_if = "Option::is_none")]
pub output_format: Option<OutputFormat>,
}
#[derive(Deserialize, Serialize)]
pub(crate) struct CohereRequest {
pub model: String,
pub document: CohereParseDocument,
pub output_format: String,
}
#[derive(Deserialize, Serialize)]
#[serde(tag = "type")]
pub(crate) enum CohereParseDocument {
#[serde(rename = "image_url")]
ImageUrl { image_url: String },
}
#[derive(Deserialize)]
pub(crate) struct CohereResponse {
#[serde(default)]
pages: Vec<CoherePage>,
meta: Option<CohereMeta>,
}
#[serde_as]
#[derive(Deserialize)]
struct CoherePage {
#[serde_as(deserialize_as = "Option<LaxI64>")]
index: Option<i64>,
markdown: Option<CohereMarkdown>,
blocks: Option<Vec<Map<String, Value>>>,
}
#[derive(Deserialize, Serialize)]
struct CohereMarkdown {
#[serde(default)]
content: String,
images: Option<Vec<Map<String, Value>>>,
}
#[derive(Deserialize)]
struct CohereMeta {
billed_units: Option<CohereBilledUnits>,
}
#[serde_as]
#[derive(Deserialize)]
struct CohereBilledUnits {
#[serde_as(deserialize_as = "Option<LaxI64>")]
pages: Option<i64>,
}
#[derive(Default)]
pub(crate) struct CohereParseConfig;
impl BaseOcrConfig for CohereParseConfig {
type OcrParams = CohereOptions;
type ProviderRequest = CohereRequest;
type Environment = Vec<(String, String)>;
fn get_api_key_env_var(&self) -> Option<&'static str> {
Some(COHERE_API_KEY_ENV)
}
fn get_health_check_document(&self) -> OcrDocument {
OcrDocument::ImageUrl {
image_url: COHERE_PARSE_HEALTH_CHECK_IMAGE_DATA_URI.into(),
extra_fields: Default::default(),
}
}
async fn validate_environment(
&self,
request: &PreparedOcrRequest,
_client: &OcrClient,
) -> Result<Self::Environment, crate::ocr::Error> {
self.validate_environment(&request.connection, &credential_env)
}
fn get_complete_url(
&self,
request: &PreparedOcrRequest,
_params: &Self::OcrParams,
_environment: &Self::Environment,
) -> Result<String, crate::ocr::Error> {
self.get_complete_url(
request
.connection
.api_base
.as_deref()
.unwrap_or(COHERE_PARSE_API_BASE),
)
}
fn transform_ocr_request(
&self,
model: &str,
document: OcrDocument,
optional_params: &CohereOptions,
_headers: &[(String, String)],
) -> Result<CohereRequest, crate::ocr::Error> {
let image_url = image_url(document)?;
Ok(build_request(model, image_url, optional_params))
}
fn get_supported_ocr_params(&self, _model: &str) -> &'static [&'static str] {
&["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,
document: OcrDocument,
optional_params: &CohereOptions,
headers: &[(String, String)],
_context: OcrRequestContext<'_>,
) -> Result<CohereRequest, crate::ocr::Error> {
self.transform_ocr_request(model, document, optional_params, headers)
}
fn transform_ocr_response(
&self,
model: &str,
raw_response: &[u8],
request_format: crate::ocr::types::OcrResponseFormat,
) -> Result<LiteLLMOcrResponse, crate::ocr::Error> {
crate::llms::base_llm::ocr::transformation::decode_and_normalize_response(
model,
raw_response,
request_format,
normalize_response,
)
}
fn validate_request_body(&self, body: &Value) -> Result<(), crate::ocr::Error> {
validate_document(&crate::ocr::prepare::body_document(body)?)
}
}
pub(crate) fn validate_document(document: &OcrDocument) -> Result<(), crate::ocr::Error> {
let OcrDocument::ImageUrl { image_url, .. } = document else {
return Err(crate::ocr::Error::CohereImageOnly);
};
if image_url.is_empty() {
return Err(crate::ocr::Error::CohereImageOnly);
}
if let Some(inline) = InlineDocument::parse(image_url)? {
if !inline.mime_type().type_.eq_ignore_ascii_case("image") {
return Err(crate::ocr::Error::CohereImageOnly);
}
inline.decode(crate::constants::OCR_INLINE_MAX_BYTES)?;
}
Ok(())
}
pub(crate) fn normalize_response(
model: &str,
response: CohereResponse,
) -> Result<LiteLLMOcrResponse, crate::ocr::Error> {
let pages_processed = billed_pages(&response).map(Ok).unwrap_or_else(|| {
i64::try_from(response.pages.len()).map_err(|_| crate::ocr::Error::NumericRange("pages"))
})?;
let pages = response
.pages
.into_iter()
.enumerate()
.map(|(position, page)| normalize_page(page, position))
.collect::<Result<Vec<_>, crate::ocr::Error>>()?;
Ok(LiteLLMOcrResponse {
usage_info: Some(OcrUsageInfo {
pages_processed: Some(pages_processed),
..Default::default()
}),
..LiteLLMOcrResponse::new(model, pages)
})
}
fn image_url(document: OcrDocument) -> Result<String, crate::ocr::Error> {
validate_document(&document)?;
let OcrDocument::ImageUrl { image_url, .. } = document else {
return Err(crate::ocr::Error::CohereImageOnly);
};
Ok(image_url)
}
fn build_request(model: &str, image_url: String, params: &CohereOptions) -> CohereRequest {
CohereRequest {
model: model.into(),
document: CohereParseDocument::ImageUrl { image_url },
output_format: match params.output_format.unwrap_or_default() {
OutputFormat::Markdown => "markdown",
OutputFormat::Blocks => "blocks",
}
.into(),
}
}
fn page_image(
mut image: Map<String, Value>,
path: &str,
) -> Result<OcrPageImage, crate::ocr::Error> {
if let Some(Value::Object(bbox)) = image.get("bounding_box") {
image.insert("bbox".into(), Value::Object(bbox.clone()));
}
crate::ocr::json::decode_response_value(Value::Object(image), path)
}
fn normalize_page(page: CoherePage, position: usize) -> Result<OcrPage, crate::ocr::Error> {
let index = page.index.map(Ok).unwrap_or_else(|| {
i64::try_from(position).map_err(|_| crate::ocr::Error::NumericRange("page index"))
})?;
let (markdown, images) = match page.markdown {
Some(markdown) => {
let images = markdown
.images
.filter(|images| !images.is_empty())
.map(|images| {
images
.into_iter()
.enumerate()
.map(|(image_index, image)| {
page_image(
image,
&format!("pages[{position}].markdown.images[{image_index}]"),
)
})
.collect::<Result<Vec<_>, _>>()
})
.transpose()?;
(markdown.content, images)
}
None => (String::new(), None),
};
let extra_fields = page
.blocks
.map(|blocks| {
(
"blocks".into(),
Value::Array(blocks.into_iter().map(Value::Object).collect()),
)
})
.into_iter()
.collect();
Ok(OcrPage {
index,
markdown,
images,
extra_fields,
..Default::default()
})
}
fn billed_pages(response: &CohereResponse) -> Option<i64> {
response.meta.as_ref()?.billed_units.as_ref()?.pages
}
impl CohereParseConfig {
fn get_complete_url(&self, base: &str) -> Result<String, crate::ocr::Error> {
let parsed = reqwest::Url::parse(base).map_err(|_| invalid_api_base())?;
if !matches!(parsed.scheme(), "http" | "https") {
return Err(invalid_api_base());
}
ApiUrl::parse(base)
.and_then(|url| url.complete_path(&["v2", "parse"]))
.map(|url| url.into_string())
.map_err(|_| invalid_api_base())
}
fn validate_environment(
&self,
connection: &OcrConnection,
env_lookup: &(dyn Fn(&str) -> Option<String> + Sync),
) -> Result<Vec<(String, String)>, crate::ocr::Error> {
if crate::http_utils::has_header(&connection.extra_headers, "authorization") {
return Ok(connection.extra_headers.clone());
}
let key = connection
.api_key
.as_deref()
.map(str::trim)
.filter(|key| !key.is_empty())
.map(str::to_string)
.or_else(|| {
self.get_api_key_env_var()
.and_then(env_lookup)
.filter(|key| !key.trim().is_empty())
})
.ok_or_else(|| {
crate::ocr::Error::Auth(litellm_auth::Error::ProviderAuthentication(
"Missing COHERE_API_KEY - set it in the environment or pass api_key".into(),
))
})?;
Ok(
std::iter::once(("Authorization".into(), format!("Bearer {key}")))
.chain(connection.extra_headers.clone())
.collect(),
)
}
}
fn invalid_api_base() -> crate::ocr::Error {
crate::ocr::Error::RequestField {
path: "api_base".into(),
}
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
#[tokio::test]
async fn composed_body_preserves_native_document_fields_and_untyped_overrides() {
let request = crate::ocr::test_support::wire_request(
"cohere/parse",
"https://example.com",
json!({
"output_format":"markdown", "metadata":{"host":true},
"extra_body":{
"output_format": {"future":true},
"document":{"type":"image_url","image_url":"https://example.com/a.png",
"provider_options":{"nested":[false,0,null]}}
}
}),
);
let request = request.with_document(
serde_json::from_value(json!({
"type":"image_url","image_url":"https://example.com/original.png"
}))
.unwrap(),
);
let request = crate::ocr::prepare::prepare_request(request);
let http = CohereParseConfig
.prepare_request(&request, &crate::ocr::test_support::ocr_client())
.await
.unwrap();
let body: Value = serde_json::from_slice(http.body().unwrap().as_bytes().unwrap()).unwrap();
assert_eq!(
body,
json!({
"model":"parse", "output_format":{"future":true},
"document":{"type":"image_url","image_url":"https://example.com/a.png",
"provider_options":{"nested":[false,0,null]}}
})
);
}
#[test]
fn options_read_known_fields_without_changing_arguments() {
let arguments = serde_json::from_value(json!({
"output_format":"blocks", "req_format":"native", "extension":false
}))
.unwrap();
for config in [false, true] {
let mapped = if config {
crate::llms::azure_ai::ocr::cohere_parse_transformation::AzureAICohereParseConfig
.map_ocr_params(&arguments, "parse")
} else {
CohereParseConfig.map_ocr_params(&arguments, "parse")
}
.unwrap();
assert_eq!(
serde_json::to_value(mapped).unwrap(),
json!({"output_format":"blocks"})
);
}
assert_eq!(arguments["req_format"], "native");
assert_eq!(arguments["extension"], false);
let invalid = serde_json::from_value(json!({"output_format":"html"})).unwrap();
assert!(matches!(
CohereParseConfig.map_ocr_params(&invalid, "parse"),
Err(crate::ocr::Error::RequestField { path })
if path == "optional_params.output_format"
));
}
#[test]
fn billed_pages_accept_integral_doubles_and_reject_fractional_counts() {
let response = serde_json::from_str::<CohereResponse>(
r#"{"pages":[],"meta":{"billed_units":{"pages":1.0}}}"#,
)
.unwrap();
let normalized = normalize_response("parse", response).unwrap();
assert_eq!(normalized.usage_info.unwrap().pages_processed, Some(1));
assert!(
serde_json::from_str::<CohereResponse>(
r#"{"pages":[],"meta":{"billed_units":{"pages":1.5}}}"#,
)
.is_err()
);
}
#[test]
fn response_preserves_python_mapping_shapes_and_extensions() {
let blocks = json!([
{"type":"text", "text":"Total Due: $4.00"},
{"type":"future", "payload":{"nested":[null,false,0]}}
]);
let response = serde_json::from_value(json!({
"pages":[{
"index":"2",
"markdown":{"content":"receipt", "images":[
{"bounding_box":{"x":1}, "bbox":"replaced", "category":"future", "extension":null},
{"image_base64":"encoded"}
]},
"blocks":blocks
}],
"meta":{"billed_units":{"pages":0}}
})).unwrap();
let response = normalize_response("parse", response).unwrap();
assert_eq!(response.pages[0].index, 2);
assert_eq!(response.usage_info.unwrap().pages_processed, Some(0));
assert_eq!(response.pages[0].extra_fields["blocks"], blocks);
let images = response.pages[0].images.as_ref().unwrap();
assert_eq!(images[0].bbox.as_ref().unwrap()["x"], 1);
assert_eq!(images[0].extra_fields["category"], "future");
assert_eq!(images[0].extra_fields.get("extension"), Some(&Value::Null));
assert_eq!(images[1].image_base64.as_deref(), Some("encoded"));
assert!(images[1].bbox.is_none());
}
#[test]
fn malformed_normalized_image_fields_report_the_original_path() {
let response = serde_json::from_value(json!({
"pages":[{"markdown":{"images":[{"image_base64":42}]}}]
}))
.unwrap();
assert!(matches!(
normalize_response("parse", response).unwrap_err(),
crate::ocr::Error::ResponseField { path }
if path == "pages[0].markdown.images[0].image_base64"
));
}
#[test]
fn provider_options_exclude_response_controls_and_extensions() {
let arguments = serde_json::from_value(
json!({"output_format":"blocks","req_format":"native","unknown":true}),
)
.unwrap();
let params = CohereParseConfig
.map_ocr_params(&arguments, "parse")
.unwrap();
assert_eq!(
serde_json::to_value(&params).unwrap(),
json!({"output_format":"blocks"})
);
let document = serde_json::from_value(
json!({"type":"image_url","image_url":"https://example.com/a.png","ignored":"field"}),
)
.unwrap();
let body = CohereParseConfig
.transform_ocr_request("parse", document, &params, &[])
.unwrap();
assert_eq!(
serde_json::to_value(body).unwrap(),
json!({
"model":"parse", "document":{"type":"image_url","image_url":"https://example.com/a.png"}, "output_format":"blocks"
})
);
}
#[tokio::test]
async fn explicit_null_options_use_defaults_before_http() {
let request = crate::ocr::test_support::wire_request(
"cohere/parse",
"https://example.com",
json!({"output_format":null,"req_format":null}),
);
let request = request.with_document(
serde_json::from_value(
json!({"type":"image_url","image_url":"https://example.com/a.png"}),
)
.unwrap(),
);
assert_eq!(
request.response_format().unwrap(),
crate::ocr::types::OcrResponseFormat::Litellm
);
let request = crate::ocr::prepare::prepare_request(request);
let http = CohereParseConfig
.prepare_request(&request, &crate::ocr::test_support::ocr_client())
.await
.unwrap();
let body: Value = serde_json::from_slice(http.body().unwrap().as_bytes().unwrap()).unwrap();
assert_eq!(body["output_format"], "markdown");
assert!(body.get("req_format").is_none());
}
#[test]
fn response_normalizes_markdown_images_blocks_and_billed_pages() {
let response = serde_json::from_value(json!({
"pages": [
{
"type":"markdown",
"index":4,
"markdown":{
"content":"receipt",
"images":[{
"id":"image",
"bounding_box":{
"top_left_x":1,
"top_left_y":2,
"bottom_right_x":48,
"bottom_right_y":49
},
"bounding_box_normalized":{
"top_left_x":0.04,
"top_left_y":0.05,
"bottom_right_x":0.15,
"bottom_right_y":0.16
},
"description":"scan",
"category":"logo",
"provider_extension":"preserved"
}]
}
},
{"type":"blocks","blocks":[{"type":"text","text":{"content":"total"}}]}
],
"meta":{"api_version":{"version":"2"},"billed_units":{"pages":3}}
}))
.unwrap();
let normalized = normalize_response("parse-v5.0", response).unwrap();
assert_eq!(normalized.pages[0].index, 4);
assert_eq!(normalized.pages[0].markdown, "receipt");
let image = &normalized.pages[0].images.as_ref().unwrap()[0];
assert_eq!(image.bbox.as_ref().unwrap()["top_left_x"], 1);
assert_eq!(
image.extra_fields["bounding_box_normalized"]["bottom_right_x"],
0.15
);
assert_eq!(image.extra_fields["description"], "scan");
assert_eq!(image.extra_fields["category"], "logo");
assert_eq!(image.extra_fields["provider_extension"], "preserved");
assert_eq!(normalized.pages[1].index, 1);
assert_eq!(normalized.pages[1].markdown, "");
assert_eq!(
normalized.pages[1].extra_fields["blocks"][0]["text"]["content"],
"total"
);
assert_eq!(normalized.usage_info.unwrap().pages_processed, Some(3));
}
#[test]
fn response_defaults_and_invalid_fields() {
for value in [
json!({}),
json!({"meta":null}),
json!({"pages":[],"meta":{"billed_units":null}}),
] {
let normalized =
normalize_response("parse", serde_json::from_value(value).unwrap()).unwrap();
assert!(normalized.pages.is_empty());
assert_eq!(normalized.usage_info.unwrap().pages_processed, Some(0));
}
for value in [
json!({"pages":null}),
json!({"pages":[{"markdown":"text"}]}),
json!({"pages":[{"index":"bad"}]}),
] {
assert!(serde_json::from_value::<CohereResponse>(value).is_err());
}
let normalized = normalize_response(
"parse",
serde_json::from_value(json!({"pages":[{"markdown":null}]})).unwrap(),
)
.unwrap();
assert_eq!(normalized.usage_info.unwrap().pages_processed, Some(1));
assert!(normalized.pages[0].images.is_none());
}
#[test]
fn response_types_documented_block_variants() {
let response = serde_json::from_value(json!({
"pages": [{
"type": "blocks",
"index": 0,
"blocks": [
{"type": "text", "text": {"content": "hello"}},
{
"type": "image",
"image": {
"id": "img-0",
"description": "logo",
"category": "logo",
"bounding_box": {
"top_left_x": 1,
"top_left_y": 2,
"bottom_right_x": 3,
"bottom_right_y": 4
},
"bounding_box_normalized": {
"top_left_x": 0.1,
"top_left_y": 0.2,
"bottom_right_x": 0.3,
"bottom_right_y": 0.4
}
}
},
{
"type": "table",
"table": {
"type": "html",
"html": "<table></table>",
"bounding_box": {
"top_left_x": 5,
"top_left_y": 6,
"bottom_right_x": 7,
"bottom_right_y": 8
},
"bounding_box_normalized": {
"top_left_x": 0.5,
"top_left_y": 0.6,
"bottom_right_x": 0.7,
"bottom_right_y": 0.8
},
"title": "Totals"
}
}
]
}]
}))
.unwrap();
let normalized = normalize_response("parse-v5.0", response).unwrap();
let blocks = normalized.pages[0].extra_fields["blocks"]
.as_array()
.unwrap();
assert_eq!(blocks[0]["text"]["content"], "hello");
assert_eq!(blocks[1]["image"]["category"], "logo");
assert_eq!(blocks[2]["table"]["type"], "html");
assert_eq!(blocks[2]["table"]["title"], "Totals");
}
#[test]
fn request_requires_image_and_supported_output_format() {
for value in [
json!({"type":"document_url","document_url":"https://example.com/a.pdf"}),
json!({"type":"image_url","image_url":""}),
json!({"type":"image_url","image_url":"data:application/pdf;base64,YQ=="}),
] {
assert!(matches!(
validate_document(&serde_json::from_value(value).unwrap()),
Err(crate::ocr::Error::CohereImageOnly)
));
}
assert!(serde_json::from_value::<CohereOptions>(json!({"output_format":"html"})).is_err());
for format in ["markdown", "blocks"] {
assert!(
serde_json::from_value::<CohereOptions>(json!({"output_format":format})).is_ok()
);
}
let request = CohereParseConfig
.transform_ocr_request(
"parse-v5.0",
serde_json::from_value(json!({
"type":"image_url",
"image_url":"https://example.com/image.png"
}))
.unwrap(),
&serde_json::from_value(json!({})).unwrap(),
&[],
)
.unwrap();
assert_eq!(
serde_json::to_value(request).unwrap()["output_format"],
"markdown"
);
}
#[test]
fn completes_provider_urls_without_duplicate_paths_and_preserves_queries() {
for suffix in ["", "/v2", "/v2/parse"] {
assert_eq!(
CohereParseConfig
.get_complete_url(&format!("https://example.com{suffix}?tenant=a"))
.unwrap(),
"https://example.com/v2/parse?tenant=a"
);
}
}
#[test]
fn rejects_invalid_urls_and_blank_keys() {
assert!(CohereParseConfig.get_complete_url("relative/path").is_err());
assert!(
CohereParseConfig
.get_complete_url("ftp://example.com")
.is_err()
);
assert!(matches!(
CohereParseConfig.validate_environment(
&OcrConnection {
api_key: Some(" ".into()),
..Default::default()
},
&|_| None,
),
Err(crate::ocr::Error::Auth(_))
));
}
}

View file

@ -0,0 +1 @@
pub(crate) mod ocr;

View file

@ -0,0 +1 @@
pub(crate) mod transformation;

View file

@ -0,0 +1,626 @@
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;
use crate::ocr::types::{
LiteLLMOcrResponse, OcrConnection, OcrDocument, OcrPage, OcrUsageInfo, PreparedOcrRequest,
};
use crate::params::OpaqueParams;
use crate::url_utils::ApiUrl;
const MISTRAL_API_KEY_ENV: &str = "MISTRAL_API_KEY";
#[derive(Clone, Debug, Serialize, Deserialize)]
pub(crate) struct MistralOcrRequest {
pub model: String,
pub document: OcrDocument,
#[serde(flatten)]
pub params: OpaqueParams,
}
#[derive(Clone, Debug, Default, Deserialize)]
pub(crate) struct MistralOcrResponse {
#[serde(flatten)]
pub extra_fields: serde_json::Map<String, Value>,
#[serde(default)]
pub pages: Vec<OcrPage>,
#[serde(
default,
deserialize_with = "serde_with::rust::double_option::deserialize"
)]
pub model: Option<Option<String>>,
pub document_annotation: Option<Value>,
pub usage_info: Option<OcrUsageInfo>,
}
#[derive(Clone, Debug, Default)]
pub(crate) struct MistralOCRConfig;
impl BaseOcrConfig for MistralOCRConfig {
type OcrParams = OpaqueParams;
type ProviderRequest = MistralOcrRequest;
type Environment = Vec<(String, String)>;
fn get_api_key_env_var(&self) -> Option<&'static str> {
Some(MISTRAL_API_KEY_ENV)
}
async fn validate_environment(
&self,
request: &PreparedOcrRequest,
_client: &OcrClient,
) -> Result<Self::Environment, crate::ocr::Error> {
self.validate_environment(&request.connection, &credential_env)
}
fn get_complete_url(
&self,
request: &PreparedOcrRequest,
_params: &Self::OcrParams,
_environment: &Self::Environment,
) -> Result<String, crate::ocr::Error> {
self.get_complete_url(request.connection.api_base.as_deref())
}
fn transform_ocr_request(
&self,
model: &str,
document: OcrDocument,
optional_params: &OpaqueParams,
_headers: &[(String, String)],
) -> Result<MistralOcrRequest, crate::ocr::Error> {
Ok(MistralOcrRequest {
model: model.to_string(),
document,
params: optional_params.clone(),
})
}
fn get_supported_ocr_params(&self, _model: &str) -> &'static [&'static str] {
&[
"pages",
"include_image_base64",
"image_limit",
"image_min_size",
"bbox_annotation_format",
"document_annotation_format",
"document_annotation_prompt",
"extract_header",
"extract_footer",
"table_format",
"confidence_scores_granularity",
"include_blocks",
"id",
]
}
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,
document: OcrDocument,
optional_params: &OpaqueParams,
headers: &[(String, String)],
_context: OcrRequestContext<'_>,
) -> Result<MistralOcrRequest, crate::ocr::Error> {
self.transform_ocr_request(model, document, optional_params, headers)
}
fn transform_ocr_response(
&self,
model: &str,
raw_response: &[u8],
request_format: crate::ocr::types::OcrResponseFormat,
) -> Result<LiteLLMOcrResponse, crate::ocr::Error> {
crate::llms::base_llm::ocr::transformation::decode_and_normalize_response(
model,
raw_response,
request_format,
normalize_response,
)
}
}
pub(crate) fn normalize_response(
model: &str,
response: MistralOcrResponse,
) -> Result<LiteLLMOcrResponse, crate::ocr::Error> {
let model = match response.model {
Some(Some(model)) => model,
Some(None) => {
return Err(crate::ocr::Error::ResponseField {
path: "model".into(),
});
}
None => model.to_string(),
};
Ok(LiteLLMOcrResponse {
extra_fields: response.extra_fields,
document_annotation: response.document_annotation,
usage_info: response.usage_info,
..LiteLLMOcrResponse::new(model, response.pages)
})
}
impl MistralOCRConfig {
fn get_complete_url(&self, api_base: Option<&str>) -> Result<String, crate::ocr::Error> {
let base = api_base
.map(str::trim)
.filter(|base| !base.is_empty())
.unwrap_or(MISTRAL_OCR_API_BASE);
ApiUrl::parse(base)
.and_then(|url| url.complete_path(&["v1", "ocr"]))
.map(|url| url.into_string())
.map_err(|_| crate::ocr::Error::RequestField {
path: "api_base".into(),
})
}
fn validate_environment(
&self,
connection: &OcrConnection,
env_lookup: &(dyn Fn(&str) -> Option<String> + Sync),
) -> Result<Vec<(String, String)>, crate::ocr::Error> {
if crate::http_utils::has_header(&connection.extra_headers, "authorization") {
return Ok(connection.extra_headers.clone());
}
let api_key = connection
.api_key
.as_deref()
.map(str::trim)
.filter(|key| !key.is_empty())
.map(str::to_string)
.or_else(|| {
self.get_api_key_env_var()
.and_then(env_lookup)
.filter(|key| !key.trim().is_empty())
})
.ok_or(litellm_auth::Error::MissingApiKey {
provider: "Mistral",
environment_variable: MISTRAL_API_KEY_ENV,
})?;
Ok(
std::iter::once(("Authorization".into(), format!("Bearer {api_key}")))
.chain(connection.extra_headers.clone())
.collect(),
)
}
}
#[cfg(test)]
mod tests {
use super::*;
use rstest::rstest;
use serde_json::{Value, json};
#[test]
fn explicit_null_model_does_not_use_the_missing_model_default() {
let response = serde_json::from_value(json!({"model":null})).unwrap();
assert!(matches!(
normalize_response("fallback", response).unwrap_err(),
crate::ocr::Error::ResponseField { path } if path == "model"
));
}
#[test]
fn response_validates_normalized_shapes_at_the_provider_boundary() {
for (payload, path) in [
(json!({"pages":[42]}), "pages[0]"),
(json!({"pages":[{"index":0}]}), "pages[0]"),
(
json!({"pages":[{"index":0,"markdown":42}]}),
"pages[0].markdown",
),
(
json!({"pages":[{"index":0,"markdown":"","images":[42]}]}),
"pages[0].images[0]",
),
(
json!({"pages":[{"index":0,"markdown":"","dimensions":{"width":1.5}}]}),
"pages[0].dimensions.width",
),
(
json!({"usage_info":{"pages_processed":"bad"}}),
"usage_info.pages_processed",
),
] {
let error = crate::ocr::json::decode_response::<MistralOcrResponse>(
&serde_json::to_vec(&payload).unwrap(),
false,
)
.unwrap_err();
assert!(matches!(
error,
crate::ocr::Error::ResponseField { path: actual } if actual == path
));
}
}
#[test]
fn response_normalizes_python_numeric_inputs_and_shared_defaults() {
let response = serde_json::from_value(json!({
"pages":[{"index":"2","markdown":"text","dimensions":{"width":1.0},"extension":false}],
"usage_info":{"pages_processed":true,"credits":"1.5","custom":0},
"extra":"ignored"
}))
.unwrap();
let response = normalize_response("model", response).unwrap();
assert_eq!(response.pages[0].index, 2);
assert_eq!(
response.pages[0].dimensions.as_ref().unwrap().width,
Some(1)
);
assert_eq!(
response.usage_info.as_ref().unwrap().pages_processed,
Some(1)
);
assert_eq!(response.usage_info.as_ref().unwrap().credits, Some(1.5));
let serialized = response.into_json();
assert_eq!(serialized["pages"][0]["extension"], false);
assert!(serialized["pages"][0]["images"].is_null());
assert!(serialized["usage_info"]["doc_size_bytes"].is_null());
assert_eq!(serialized["usage_info"]["custom"], 0);
assert!(serialized["content"].is_null());
assert_eq!(serialized["extra"], "ignored");
}
#[test]
fn map_ocr_params_selects_known_fields_without_changing_arguments() {
let input =
serde_json::from_value(json!({"pages":null,"extract_header":false,"unknown":true}))
.unwrap();
let params = MistralOCRConfig.map_ocr_params(&input, "model").unwrap();
assert_eq!(
serde_json::to_value(params).unwrap(),
json!({"pages":null,"extract_header":false})
);
assert_eq!(input["unknown"], true);
assert_eq!(input.get("pages"), Some(&Value::Null));
}
#[test]
fn request_transform_uses_already_mapped_params_without_filtering_again() {
let params = serde_json::from_value(json!({"extension":{"nested":null}})).unwrap();
let body = MistralOCRConfig
.transform_ocr_request("model", document(), &params, &[])
.unwrap();
assert_eq!(
serde_json::to_value(body).unwrap()["extension"],
json!({"nested":null})
);
}
#[test]
fn raw_response_transform_keeps_native_payload_separate_from_typed_normalization() {
let raw = br#"{"pages":[{"index":"2","markdown":"text"}],"provider_extension":false}"#;
let response = MistralOCRConfig
.transform_ocr_response("model", raw, crate::ocr::types::OcrResponseFormat::Native)
.unwrap();
assert_eq!(response.pages[0].index, 2);
let native = response.provider_native_response.unwrap();
assert_eq!(native["pages"][0]["index"], "2");
assert_eq!(native["provider_extension"], false);
assert_eq!(response.extra_fields["provider_extension"], false);
assert!(
MistralOCRConfig
.transform_ocr_response(
"model",
br#"{"pages":[{"index":0}]}"#,
crate::ocr::types::OcrResponseFormat::Litellm
)
.is_err()
);
}
fn mapped_params(value: Value) -> Value {
let params = serde_json::from_value(value).unwrap();
serde_json::to_value(MistralOCRConfig.map_ocr_params(&params, "model").unwrap()).unwrap()
}
fn document() -> OcrDocument {
serde_json::from_value(
json!({"type":"document_url","document_url":"https://example.com/a.pdf"}),
)
.unwrap()
}
#[rstest]
fn extract_header_is_a_supported_ocr_param() {
assert_eq!(
mapped_params(json!({"extract_header":true}))["extract_header"],
true
);
}
#[rstest]
fn extract_footer_is_a_supported_ocr_param() {
assert_eq!(
mapped_params(json!({"extract_footer":false}))["extract_footer"],
false
);
}
#[rstest]
fn existing_ocr_params_remain_supported() {
let mapped = mapped_params(json!({
"pages":[0,2],
"include_image_base64":true,
"image_limit":2,
"image_min_size":100,
"bbox_annotation_format":{"type":"json_schema"},
"document_annotation_format":{"type":"json_schema"}
}));
assert_eq!(mapped["pages"], json!([0, 2]));
assert_eq!(mapped["include_image_base64"], true);
assert_eq!(mapped["image_limit"], 2);
assert_eq!(mapped["image_min_size"], 100);
assert_eq!(mapped["bbox_annotation_format"]["type"], "json_schema");
assert_eq!(mapped["document_annotation_format"]["type"], "json_schema");
}
#[rstest]
fn map_ocr_params_forwards_extract_header() {
assert_eq!(
mapped_params(json!({"extract_header":true}))["extract_header"],
true
);
}
#[rstest]
fn map_ocr_params_forwards_extract_footer() {
assert_eq!(
mapped_params(json!({"extract_footer":true}))["extract_footer"],
true
);
}
#[rstest]
fn map_ocr_params_forwards_extract_header_and_footer() {
let mapped = mapped_params(json!({"extract_header":true,"extract_footer":false}));
assert_eq!(mapped["extract_header"], true);
assert_eq!(mapped["extract_footer"], false);
}
#[rstest]
fn map_ocr_params_excludes_extensions_from_the_provider_options() {
let mapped = mapped_params(json!({"extract_header":true,"unsupported_param":"value"}));
assert_eq!(mapped["extract_header"], true);
assert!(mapped.get("unsupported_param").is_none());
}
#[rstest]
fn map_ocr_params_preserves_unvalidated_values_and_explicit_null() {
let mapped = mapped_params(json!({
"pages":{"future":"shape"},
"include_image_base64":null
}));
assert_eq!(mapped["pages"], json!({"future":"shape"}));
assert!(mapped.get("include_image_base64").unwrap().is_null());
}
#[rstest]
#[case("table_format", json!("html"))]
#[case("confidence_scores_granularity", json!("word"))]
#[case("confidence_scores_granularity", json!("block"))]
#[case("document_annotation_prompt", json!("extract"))]
#[case("include_blocks", json!(true))]
#[case("id", json!("req-123"))]
fn new_ocr_params_are_supported(#[case] name: &str, #[case] value: Value) {
assert_eq!(mapped_params(json!({name:value.clone()}))[name], value);
}
#[rstest]
#[case("table_format", json!("html"))]
#[case("confidence_scores_granularity", json!("word"))]
#[case("document_annotation_prompt", json!("extract"))]
#[case("include_blocks", json!(true))]
#[case("id", json!("req-123"))]
fn map_ocr_params_forwards_new_ocr_params(#[case] name: &str, #[case] value: Value) {
assert_eq!(mapped_params(json!({name:value.clone()}))[name], value);
}
#[rstest]
#[case("pages", json!([0, 2]))]
#[case("pages", json!("0,2-4"))]
#[case("include_image_base64", json!(true))]
#[case("image_limit", json!(2))]
#[case("image_min_size", json!(100))]
#[case("bbox_annotation_format", json!({"type":"json_schema"}))]
#[case("document_annotation_format", json!({"type":"json_schema"}))]
#[case("document_annotation_prompt", json!("extract"))]
#[case("extract_header", json!(true))]
#[case("extract_footer", json!(false))]
#[case("table_format", json!("html"))]
#[case("confidence_scores_granularity", json!("word"))]
#[case("include_blocks", json!(true))]
#[case("id", json!("req-123"))]
fn request_mapping_matches_python(#[case] name: &str, #[case] value: Value) {
let params: OpaqueParams = serde_json::from_value(json!({name: value.clone()})).unwrap();
let result = serde_json::to_value(
MistralOCRConfig
.transform_ocr_request("model", document(), &params, &[])
.unwrap(),
)
.unwrap();
assert_eq!(result["model"], "model");
assert_eq!(result[name], value);
}
#[rstest]
#[case("table_format", json!("html"))]
#[case("confidence_scores_granularity", json!("word"))]
#[case("document_annotation_prompt", json!("extract"))]
#[case("id", json!("req-123"))]
#[case("extract_header", json!(true))]
#[case("include_blocks", json!(true))]
#[case("pages", json!([0,1]))]
fn transform_ocr_request_includes_each_optional_param(
#[case] name: &str,
#[case] value: Value,
) {
let params: OpaqueParams = serde_json::from_value(json!({name:value.clone()})).unwrap();
let result = serde_json::to_value(
MistralOCRConfig
.transform_ocr_request("mistral-ocr-latest", document(), &params, &[])
.unwrap(),
)
.unwrap();
assert_eq!(result[name], value);
assert_eq!(result["model"], "mistral-ocr-latest");
}
#[rstest]
fn transform_ocr_request_includes_multiple_new_params() {
let params: OpaqueParams = serde_json::from_value(json!({
"table_format":"html",
"confidence_scores_granularity":"page",
"extract_header":true
}))
.unwrap();
let result = serde_json::to_value(
MistralOCRConfig
.transform_ocr_request("mistral-ocr-latest", document(), &params, &[])
.unwrap(),
)
.unwrap();
assert_eq!(result["table_format"], "html");
assert_eq!(result["confidence_scores_granularity"], "page");
assert_eq!(result["extract_header"], true);
}
#[rstest]
fn transform_ocr_response_preserves_blocks_and_confidence_scores() {
let response: MistralOcrResponse = serde_json::from_value(json!({
"pages":[{
"index":0,
"markdown":"hello",
"images":[{"id":"img-0","image_base64":"data:image/png;base64,AA=="}],
"dimensions":{"width":612,"height":792,"dpi":72},
"blocks":[{"type":"title","bbox":{"x":1},"confidence_scores":{"mean":0.98}}],
"confidence_scores":{"average_page_confidence_score":0.99,"minimum_page_confidence_score":0.97}
}],
"model":"returned-model",
"document_annotation":"{\"language\":\"en\"}",
"usage_info":{"pages_processed":1}
}))
.unwrap();
let result = normalize_response("model", response).unwrap().into_json();
assert_eq!(result["pages"][0]["blocks"][0]["type"], "title");
assert_eq!(result["pages"][0]["blocks"][0]["bbox"]["x"], 1);
assert_eq!(
result["pages"][0]["blocks"][0]["confidence_scores"]["mean"],
0.98
);
assert_eq!(
result["pages"][0]["confidence_scores"]["average_page_confidence_score"],
0.99
);
assert_eq!(result["pages"][0]["images"][0]["id"], "img-0");
assert_eq!(result["pages"][0]["dimensions"]["dpi"], 72);
assert_eq!(result["model"], "returned-model");
assert_eq!(result["document_annotation"], "{\"language\":\"en\"}");
assert_eq!(result["usage_info"]["pages_processed"], 1);
}
#[rstest]
fn transform_ocr_response_preserves_ocr4_page_fields() {
let page = json!({
"index":0,
"markdown":"table page",
"tables":[{"rows":2,"cols":3}],
"hyperlinks":["https://example.com"],
"header":"header",
"footer":"footer"
});
let response: MistralOcrResponse =
serde_json::from_value(json!({"pages":[page.clone()]})).unwrap();
let result = normalize_response("model", response).unwrap().into_json();
assert_eq!(result["pages"][0]["tables"], page["tables"]);
assert_eq!(result["pages"][0]["hyperlinks"], page["hyperlinks"]);
assert_eq!(result["pages"][0]["header"], page["header"]);
assert_eq!(result["pages"][0]["footer"], page["footer"]);
assert!(result["pages"][0]["images"].is_null());
assert!(result["pages"][0]["dimensions"].is_null());
}
#[test]
fn complete_url_defaults_and_dedupes_v1() {
assert_eq!(
MistralOCRConfig.get_complete_url(None).unwrap(),
"https://api.mistral.ai/v1/ocr"
);
assert_eq!(
MistralOCRConfig
.get_complete_url(Some("https://example.com/v1?tenant=a"))
.unwrap(),
"https://example.com/v1/ocr?tenant=a"
);
assert_eq!(
MistralOCRConfig
.get_complete_url(Some("https://example.com/v1/ocr?tenant=a"))
.unwrap(),
"https://example.com/v1/ocr?tenant=a"
);
}
#[test]
fn environment_prefers_explicit_key_then_environment() {
let explicit = OcrConnection {
api_key: Some("explicit".into()),
..OcrConnection::default()
};
assert_eq!(
MistralOCRConfig
.validate_environment(&explicit, &|_| Some("environment".into()))
.unwrap()[0],
("Authorization".into(), "Bearer explicit".into())
);
assert_eq!(
MistralOCRConfig
.validate_environment(&OcrConnection::default(), &|_| Some("environment".into()))
.unwrap()[0],
("Authorization".into(), "Bearer environment".into())
);
}
#[test]
fn environment_preserves_forwarded_authorization() {
let connection = OcrConnection {
extra_headers: vec![("authorization".into(), "Bearer forwarded".into())],
..OcrConnection::default()
};
assert_eq!(
MistralOCRConfig
.validate_environment(&connection, &|_| None)
.unwrap(),
connection.extra_headers
);
}
#[test]
fn environment_rejects_missing_key() {
assert!(matches!(
MistralOCRConfig.validate_environment(&OcrConnection::default(), &|_| None),
Err(crate::ocr::Error::Auth(
litellm_auth::Error::MissingApiKey {
provider: "Mistral",
environment_variable: MISTRAL_API_KEY_ENV,
}
))
));
}
}

View file

@ -0,0 +1,6 @@
pub(crate) mod azure_ai;
pub(crate) mod base_llm;
pub(crate) mod cohere;
pub(crate) mod mistral;
pub(crate) mod reducto;
pub(crate) mod vertex_ai;

View file

@ -0,0 +1 @@
pub(crate) mod ocr;

View file

@ -0,0 +1 @@
pub(crate) mod transformation;

File diff suppressed because it is too large Load diff

View file

@ -0,0 +1 @@
pub(crate) mod ocr;

View file

@ -0,0 +1,9 @@
use crate::ocr::types::OcrConnection;
use litellm_auth::InputSource;
pub(super) fn validate_destination(connection: &OcrConnection) -> Result<(), crate::ocr::Error> {
if connection.api_base.is_some() && connection.api_base_source == InputSource::Request {
return Err(litellm_auth::Error::RequestVertexCredentialDestination.into());
}
Ok(())
}

View file

@ -0,0 +1,705 @@
use serde::{Deserialize, Serialize};
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;
use crate::ocr::types::{
LiteLLMOcrResponse, OcrDocument, OcrPage, OcrPageDimensions, OcrPageImage, OcrUsageInfo,
PreparedOcrRequest,
};
use crate::params::OpaqueParams;
use crate::providers::model::{ModelNamespace, ProviderModel, RoutedModel};
use crate::url_utils::ApiUrl;
const DEFAULT_API_BASE: &str = "https://aiplatform.googleapis.com";
const MODEL_NAMESPACE: &str = "deepseek-ai";
const DEFAULT_LOCATION: &str = "us-central1";
const DEEPSEEK_OCR_PARAMS: &[&str] = &["stream", "temperature", "max_tokens", "top_p", "n", "stop"];
pub(crate) type DeepSeekOcrParams = OpaqueParams;
#[derive(Clone, Debug, Serialize, Deserialize)]
pub(crate) struct DeepSeekOcrRequest {
pub model: ProviderModel<DeepSeekAi>,
pub messages: Vec<DeepSeekOcrMessage>,
#[serde(flatten)]
pub params: OpaqueParams,
}
#[derive(Clone, Debug, Serialize, Deserialize)]
pub(crate) struct DeepSeekOcrMessage {
pub role: UserRole,
pub content: Vec<DeepSeekDocument>,
}
#[derive(Clone, Debug, Serialize, Deserialize)]
#[serde(tag = "type")]
pub(crate) enum DeepSeekDocument {
#[serde(rename = "image_url")]
ImageUrl { image_url: String },
}
#[derive(Clone, Debug, Serialize, Deserialize)]
#[serde(rename_all = "lowercase")]
pub(crate) enum UserRole {
User,
}
#[derive(Clone, Debug, Deserialize)]
pub(crate) struct DeepSeekOcrResponse {
#[serde(default)]
choices: Vec<DeepSeekChoice>,
#[serde(default = "empty_object")]
usage: Value,
}
#[derive(Clone, Debug, Deserialize)]
struct DeepSeekChoice {
#[serde(default)]
message: DeepSeekResponseMessage,
}
#[derive(Clone, Debug, Default, Deserialize)]
struct DeepSeekResponseMessage {
content: Option<DeepSeekContent>,
}
#[derive(Clone, Debug, Deserialize)]
#[serde(untagged)]
enum DeepSeekContent {
Text(String),
Object(Map<String, Value>),
}
#[serde_with::serde_as]
#[derive(Deserialize)]
struct DeepSeekPage {
#[serde(default)]
#[serde_as(deserialize_as = "crate::serde_compat::LaxI64")]
index: i64,
#[serde(default)]
markdown: String,
images: Option<Vec<OcrPageImage>>,
dimensions: Option<OcrPageDimensions>,
}
#[derive(Clone, Debug)]
pub(crate) struct DeepSeekAi;
impl ModelNamespace for DeepSeekAi {
const NAME: &'static str = MODEL_NAMESPACE;
}
#[derive(Clone, Debug)]
pub(crate) struct VertexAIDeepSeekOCRConfig;
impl BaseOcrConfig for VertexAIDeepSeekOCRConfig {
type OcrParams = DeepSeekOcrParams;
type ProviderRequest = DeepSeekOcrRequest;
type Environment = vertex::VertexEnvironment;
fn get_api_key_env_var(&self) -> Option<&'static str> {
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,
client: &OcrClient,
) -> Result<Self::Environment, crate::ocr::Error> {
BaseOcrConfig::validate_environment(&VertexAIOCRConfig, request, client).await
}
fn get_complete_url(
&self,
request: &PreparedOcrRequest,
_params: &Self::OcrParams,
environment: &Self::Environment,
) -> Result<String, crate::ocr::Error> {
let config = VertexConfig::from_sourced_optional_params(
&request.optional_params,
&request.input_sources,
)?;
let location = vertex::get_vertex_ai_location(&config, &credential_env)
.unwrap_or_else(|| DEFAULT_LOCATION.to_string());
self.get_complete_url(
request.connection.api_base.as_deref(),
&environment.project_id,
&location,
)
}
async fn async_transform_ocr_request(
&self,
model: &str,
document: OcrDocument,
optional_params: &DeepSeekOcrParams,
headers: &[(String, String)],
_context: OcrRequestContext<'_>,
) -> Result<DeepSeekOcrRequest, crate::ocr::Error> {
self.transform_ocr_request(model, document, optional_params, headers)
}
fn transform_ocr_response(
&self,
model: &str,
raw_response: &[u8],
request_format: crate::ocr::types::OcrResponseFormat,
) -> Result<LiteLLMOcrResponse, crate::ocr::Error> {
crate::llms::base_llm::ocr::transformation::decode_and_normalize_response(
model,
raw_response,
request_format,
normalize_response,
)
}
fn transform_ocr_request(
&self,
model: &str,
document: OcrDocument,
optional_params: &DeepSeekOcrParams,
_headers: &[(String, String)],
) -> Result<DeepSeekOcrRequest, crate::ocr::Error> {
if document.source().is_empty() {
return Err(crate::ocr::Error::MissingDocumentUrl);
}
Ok(DeepSeekOcrRequest {
model: provider_model(model)?,
messages: vec![DeepSeekOcrMessage {
role: UserRole::User,
content: vec![DeepSeekDocument::ImageUrl {
image_url: document.source().to_string(),
}],
}],
params: optional_params
.iter()
.filter(|(name, _)| DEEPSEEK_OCR_PARAMS.contains(&name.as_str()))
.map(|(name, value)| (name.clone(), value.clone()))
.collect(),
})
}
}
pub(crate) fn normalize_response(
model: &str,
response: DeepSeekOcrResponse,
) -> Result<LiteLLMOcrResponse, crate::ocr::Error> {
let content = response
.choices
.into_iter()
.next()
.and_then(|choice| choice.message.content)
.ok_or(crate::ocr::Error::EmptyContent)?;
let (ocr_data, fallback_markdown) = match content {
DeepSeekContent::Text(text) if text.is_empty() => {
return Err(crate::ocr::Error::EmptyContent);
}
DeepSeekContent::Text(text) => {
let parsed = text
.trim_start()
.starts_with('{')
.then(|| serde_json::from_str::<Map<String, Value>>(&text).ok())
.flatten();
(parsed.unwrap_or_default(), text)
}
DeepSeekContent::Object(data) if data.is_empty() => {
return Err(crate::ocr::Error::EmptyContent);
}
DeepSeekContent::Object(data) => {
let fallback = if data.contains_key("pages") {
String::new()
} else {
let mut output = Vec::new();
data.serialize(&mut serde_json::Serializer::with_formatter(
&mut output,
PythonJsonFormatter,
))
.map_err(|_| response_field("content"))?;
String::from_utf8(output).map_err(|_| response_field("content"))?
};
(data, fallback)
}
};
let has_pages = ocr_data.contains_key("pages");
let pages = match ocr_data.get("pages") {
Some(Value::Array(pages)) => pages
.iter()
.enumerate()
.filter(|(_, page)| page.is_object())
.map(|(position, page)| {
let page: DeepSeekPage = crate::ocr::json::decode_response_value(
page.clone(),
&format!("choices[0].message.content.pages[{position}]"),
)?;
Ok(OcrPage {
index: page.index,
markdown: page.markdown,
images: page.images,
dimensions: page.dimensions,
..Default::default()
})
})
.collect::<Result<Vec<_>, crate::ocr::Error>>()?,
Some(_) => return Err(response_field("pages")),
None => Vec::new(),
};
let usage = ocr_data
.get("usage_info")
.or_else(|| (!has_pages).then_some(&response.usage));
let usage_info: Option<OcrUsageInfo> = usage
.filter(|usage| usage.is_object())
.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(),
Some(_) => return Err(response_field("model")),
None => model.to_string(),
};
Ok(LiteLLMOcrResponse {
extra_fields: ocr_data
.iter()
.filter(|(name, _)| {
!matches!(
name.as_str(),
"pages"
| "model"
| "document_annotation"
| "usage_info"
| "object"
| "content"
| "tables"
| "keyValuePairs"
| "provider_native_response"
)
})
.map(|(name, value)| (name.clone(), value.clone()))
.collect(),
document_annotation: has_pages
.then(|| ocr_data.get("document_annotation").cloned())
.flatten(),
usage_info,
..LiteLLMOcrResponse::new(
model,
if pages.is_empty() {
vec![OcrPage {
markdown: fallback_markdown,
..Default::default()
}]
} else {
pages
},
)
})
}
fn empty_object() -> Value {
Value::Object(Map::new())
}
struct PythonJsonFormatter;
impl serde_json::ser::Formatter for PythonJsonFormatter {
fn begin_array_value<W: std::io::Write + ?Sized>(
&mut self,
writer: &mut W,
first: bool,
) -> std::io::Result<()> {
if first {
Ok(())
} else {
writer.write_all(b", ")
}
}
fn begin_object_key<W: std::io::Write + ?Sized>(
&mut self,
writer: &mut W,
first: bool,
) -> std::io::Result<()> {
if first {
Ok(())
} else {
writer.write_all(b", ")
}
}
fn begin_object_value<W: std::io::Write + ?Sized>(
&mut self,
writer: &mut W,
) -> std::io::Result<()> {
writer.write_all(b": ")
}
fn write_string_fragment<W: std::io::Write + ?Sized>(
&mut self,
writer: &mut W,
fragment: &str,
) -> std::io::Result<()> {
for character in fragment.chars() {
if character.is_ascii() && character != '\u{7f}' {
writer.write_all(&[character as u8])?;
} else {
for unit in character.encode_utf16(&mut [0; 2]) {
write!(writer, "\\u{unit:04x}")?;
}
}
}
Ok(())
}
}
fn response_field(field: &str) -> crate::ocr::Error {
crate::ocr::Error::ResponseField {
path: format!("choices[0].message.content.{field}"),
}
}
pub(crate) fn provider_model(model: &str) -> Result<ProviderModel<DeepSeekAi>, crate::ocr::Error> {
RoutedModel::new(model)
.and_then(RoutedModel::into_provider::<DeepSeekAi>)
.map_err(|_| crate::ocr::Error::RequestField {
path: "model".into(),
})
}
impl VertexAIDeepSeekOCRConfig {
fn get_complete_url(
&self,
api_base: Option<&str>,
project: &str,
location: &str,
) -> Result<String, crate::ocr::Error> {
let base = api_base
.map(str::trim)
.filter(|base| !base.is_empty())
.unwrap_or(DEFAULT_API_BASE);
ApiUrl::parse(base)
.and_then(|url| {
url.complete_path(&[
"v1",
"projects",
project,
"locations",
location,
"endpoints",
"openapi",
"chat",
"completions",
])
})
.map(|url| url.into_string())
.map_err(|_| crate::ocr::Error::RequestField {
path: "api_base".into(),
})
}
}
#[cfg(test)]
mod tests {
use super::{
DeepSeekOcrParams, DeepSeekOcrResponse, VertexAIDeepSeekOCRConfig, normalize_response,
provider_model,
};
use serde_json::{Value, json};
#[test]
fn unconsumed_options_remain_available_for_body_composition() {
use crate::llms::base_llm::ocr::transformation::BaseOcrConfig;
use serde_json::json;
let arguments =
serde_json::from_value(json!({"temperature":0.5,"extension":null})).unwrap();
assert_eq!(
serde_json::to_value(
VertexAIDeepSeekOCRConfig
.map_ocr_params(&arguments, "deepseek-ocr")
.unwrap()
)
.unwrap(),
json!({})
);
assert_eq!(
crate::call_arguments::compose_body(&arguments, &json!({"model":"deepseek-ocr"}), &[])
.unwrap(),
json!({"model":"deepseek-ocr","temperature":0.5,"extension":null})
);
}
#[test]
fn config_owns_model_namespace_and_endpoint() {
assert_eq!(
provider_model("deepseek-ocr-maas").unwrap().as_str(),
"deepseek-ai/deepseek-ocr-maas"
);
assert_eq!(
provider_model("deepseek-ai/deepseek-ocr-maas")
.unwrap()
.as_str(),
"deepseek-ai/deepseek-ocr-maas"
);
assert_eq!(
VertexAIDeepSeekOCRConfig
.get_complete_url(None, "proj-1", "europe-west4")
.unwrap(),
"https://aiplatform.googleapis.com/v1/projects/proj-1/locations/europe-west4/endpoints/openapi/chat/completions"
);
}
use rstest::rstest;
use crate::llms::base_llm::ocr::transformation::BaseOcrConfig;
use crate::ocr::types::OcrDocument;
fn document() -> OcrDocument {
serde_json::from_value(json!({"type":"image_url","image_url":"gs://bucket/a.png"})).unwrap()
}
#[rstest]
#[case("stream", json!(true))]
#[case("temperature", json!(0.1))]
#[case("max_tokens", json!(1024))]
#[case("top_p", json!(0.9))]
#[case("n", json!(2))]
#[case("stop", json!("done"))]
#[case("stop", json!(["done", "stop"]))]
#[case("temperature", json!(null))]
fn request_mapping_matches_python(#[case] name: &str, #[case] value: Value) {
let params: DeepSeekOcrParams =
serde_json::from_value(json!({name: value.clone(), "ignored": true})).unwrap();
let result = serde_json::to_value(
VertexAIDeepSeekOCRConfig
.transform_ocr_request("deepseek-ai/deepseek-ocr-maas", document(), &params, &[])
.unwrap(),
)
.unwrap();
assert_eq!(result["model"], "deepseek-ai/deepseek-ocr-maas");
assert_eq!(
result["messages"][0]["content"][0],
json!({"type":"image_url","image_url":"gs://bucket/a.png"})
);
assert_eq!(result[name], value);
assert!(result.get("ignored").is_none());
}
#[rstest]
#[case(json!({"type":"image_url","image_url":"data:image/png;base64,AA=="}))]
#[case(json!({"type":"document_url","document_url":"data:application/pdf;base64,AA=="}))]
fn request_maps_both_document_types_to_image_content(#[case] document: Value) {
let source = document
.get("image_url")
.or_else(|| document.get("document_url"))
.unwrap()
.clone();
let request = VertexAIDeepSeekOCRConfig
.transform_ocr_request(
"deepseek-ai/deepseek-ocr-maas",
serde_json::from_value(document).unwrap(),
&DeepSeekOcrParams::default(),
&[],
)
.unwrap();
let result = serde_json::to_value(request).unwrap();
assert_eq!(
result["messages"][0]["content"][0],
json!({"type":"image_url","image_url":source})
);
}
#[rstest]
#[case(json!("# hello"), "# hello")]
#[case(json!("{broken"), "{broken")]
#[case(json!(" {\"pages\":[]} "), " {\"pages\":[]} ")]
#[case(json!({"pages":[]}), "")]
#[case(json!("[]"), "[]")]
#[case(json!("{\"pages\":[{\"markdown\":\"json text\"}]}"), "json text")]
#[case(json!({"pages":[{"markdown":"object"}]}), "object")]
fn response_transform_handles_text_json_and_objects(
#[case] content: Value,
#[case] expected: &str,
) {
let has_pages = content
.as_object()
.is_some_and(|data| data.contains_key("pages"))
|| content
.as_str()
.is_some_and(|text| text.contains("\"pages\""));
let response: DeepSeekOcrResponse = serde_json::from_value(
json!({"choices":[{"message":{"content":content}}],"usage":{"prompt_tokens":1}}),
)
.unwrap();
let result = normalize_response("model", response).unwrap().into_json();
assert_eq!(result["pages"][0]["markdown"], expected);
assert_eq!(result["pages"][0]["index"], 0);
if has_pages {
assert!(result["usage_info"].is_null());
} else {
assert_eq!(result["usage_info"]["prompt_tokens"], 1);
}
}
#[test]
fn structured_result_maps_pages_usage_model_and_annotation() {
let response: DeepSeekOcrResponse = serde_json::from_value(json!({
"choices":[{"message":{"content":{
"pages":[{"index":2,"markdown":"page","images":[{"id":"one"}],"dimensions":{"width":10}}],
"model":"provider-model",
"usage_info":{"pages_processed":1},
"document_annotation":{"language":"en"},
"future":"kept"
}}}]
}))
.unwrap();
let result = normalize_response("requested", response)
.unwrap()
.into_json();
assert_eq!(result["pages"][0]["index"], 2);
assert_eq!(result["pages"][0]["images"][0]["id"], "one");
assert_eq!(result["model"], "provider-model");
assert_eq!(result["usage_info"]["pages_processed"], 1);
assert_eq!(result["document_annotation"]["language"], "en");
assert_eq!(result["future"], "kept");
}
#[test]
fn response_transform_rejects_missing_empty_and_malformed_content() {
for value in [
json!({"choices":[]}),
json!({"choices":[{"message":{"content":{}}}]}),
json!({"choices":[{"message":{"content":""}}]}),
json!({"choices":[{"message":{"content":"{\"pages\":[{\"markdown\":42}]}"}}]}),
json!({"choices":[{"message":{"content":{"pages":[{"markdown":42}]}}}]}),
] {
let result = serde_json::from_value::<DeepSeekOcrResponse>(value)
.map_err(|_| ())
.and_then(|response| normalize_response("model", response).map_err(|_| ()));
assert!(result.is_err());
}
}
#[test]
fn structured_content_preserves_usage_presence_and_shared_page_defaults() {
for (usage, expected) in [(json!(null), None), (json!({"pages_processed":2}), Some(2))] {
let response = serde_json::from_value(json!({
"choices":[{"message":{"content":{
"pages":[42, {"index":"2", "images":[{"id":"kept"}], "ignored":true}],
"usage_info":usage
}}}],
"usage":{"pages_processed":99}
}))
.unwrap();
let normalized = normalize_response("model", response).unwrap();
assert_eq!(normalized.pages.len(), 1);
assert_eq!(normalized.pages[0].index, 2);
assert_eq!(normalized.pages[0].markdown, "");
assert!(normalized.pages[0].extra_fields.is_empty());
assert_eq!(
normalized
.usage_info
.and_then(|usage| usage.pages_processed),
expected
);
}
}
use crate::ocr::test_support::{MockResponse, mock_server, perform_ocr, wire_request};
use litellm_auth::InputSource;
fn request_body(request: &str) -> Value {
serde_json::from_str(request.split_once("\r\n\r\n").unwrap().1).unwrap()
}
#[tokio::test]
async fn facade_executes_vertex_deepseek_at_the_openai_endpoint() {
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({
"choices":[{"message":{"content":"recognized"}}],
"usage":{"prompt_tokens":1}
}))])
.await;
let request = wire_request(
"vertex_ai/deepseek-ocr-maas",
&base,
json!({
"vertex_project":"project-1",
"vertex_location":"europe-west4",
"temperature":0.1,
"future_ocr_option":true,
"extra_body":{"provider_option":"value"}
}),
);
let request = crate::ocr::test_support::with_source(request, "gs://bucket/document.pdf");
let response = perform_ocr(request).await.unwrap();
server.await.unwrap();
assert_eq!(response.pages[0].markdown, "recognized");
assert_eq!(
response.usage_info.unwrap().extra_fields["prompt_tokens"],
1
);
let requests = seen.lock().unwrap();
assert!(requests[0].starts_with(
"POST /v1/projects/project-1/locations/europe-west4/endpoints/openapi/chat/completions "
));
assert!(
requests[0]
.to_ascii_lowercase()
.contains("authorization: bearer test-key")
);
let body = request_body(&requests[0]);
assert_eq!(body["model"], "deepseek-ai/deepseek-ocr-maas");
assert_eq!(body["temperature"], 0.1);
assert_eq!(body["future_ocr_option"], true);
assert_eq!(body["provider_option"], "value");
assert!(body.get("vertex_project").is_none());
assert!(body.get("extra_body").is_none());
assert_eq!(
body["messages"][0]["content"][0],
json!({"type":"image_url","image_url":"gs://bucket/document.pdf"})
);
}
#[test]
fn host_registration_selects_deepseek_without_affecting_mistral() {
assert!(crate::ocr::is_supported_request(
"deepseek-ocr-maas",
Some("vertex_ai")
));
assert!(crate::ocr::is_supported_request(
"mistral-ocr-maas",
Some("vertex_ai")
));
}
#[tokio::test]
async fn request_controlled_api_base_is_rejected_before_vertex_auth() {
let mut request = wire_request(
"vertex_ai/deepseek-ocr-maas",
"https://caller.example",
json!({"vertex_project":"project-1"}),
);
request.credentials.api_base = Some(litellm_auth::Sourced::new(
"https://caller.example".into(),
InputSource::Request,
));
let error = perform_ocr(request).await.unwrap_err();
assert!(
error
.to_string()
.contains("request-controlled Vertex AI endpoint")
);
}
}

View file

@ -0,0 +1,3 @@
pub(crate) mod common_utils;
pub(crate) mod deepseek_transformation;
pub(crate) mod transformation;

View file

@ -0,0 +1,395 @@
use litellm_auth_gcp::{self as vertex, VertexConfig};
use serde_json::Value;
use super::common_utils::validate_destination;
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;
use crate::ocr::types::{LiteLLMOcrResponse, OcrConnection, OcrDocument, PreparedOcrRequest};
use crate::params::OpaqueParams;
use crate::url_utils::ApiUrl;
const DEFAULT_LOCATION: &str = "us-central1";
#[derive(Clone, Debug, Default)]
pub(crate) struct VertexAIOCRConfig;
impl BaseOcrConfig for VertexAIOCRConfig {
type OcrParams = OpaqueParams;
type ProviderRequest = MistralOcrRequest;
type Environment = vertex::VertexEnvironment;
fn get_api_key_env_var(&self) -> Option<&'static str> {
Some("VERTEX_AI_API_KEY")
}
async fn validate_environment(
&self,
request: &PreparedOcrRequest,
client: &OcrClient,
) -> Result<Self::Environment, crate::ocr::Error> {
let config = VertexConfig::from_sourced_optional_params(
&request.optional_params,
&request.input_sources,
)?;
self.validate_environment(&request.connection, &config, client)
.await
}
fn get_complete_url(
&self,
request: &PreparedOcrRequest,
_params: &Self::OcrParams,
environment: &Self::Environment,
) -> Result<String, crate::ocr::Error> {
let config = VertexConfig::from_sourced_optional_params(
&request.optional_params,
&request.input_sources,
)?;
let location = vertex::get_vertex_ai_location(&config, &credential_env)
.unwrap_or_else(|| DEFAULT_LOCATION.to_string());
self.get_complete_url(
request.connection.api_base.as_deref(),
&environment.project_id,
&location,
&request.model,
)
}
fn transform_ocr_request(
&self,
model: &str,
document: OcrDocument,
params: &OpaqueParams,
headers: &[(String, String)],
) -> Result<MistralOcrRequest, crate::ocr::Error> {
MistralOCRConfig.transform_ocr_request(model, document, params, headers)
}
fn get_supported_ocr_params(&self, model: &str) -> &'static [&'static str] {
MistralOCRConfig.get_supported_ocr_params(model)
}
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,
document: OcrDocument,
optional_params: &OpaqueParams,
headers: &[(String, String)],
context: OcrRequestContext<'_>,
) -> Result<MistralOcrRequest, crate::ocr::Error> {
let document = inline_remote_document(
context.client.document_fetcher(),
document,
context.connection,
)
.await?;
self.transform_ocr_request(model, document, optional_params, headers)
}
fn transform_ocr_response(
&self,
model: &str,
raw_response: &[u8],
request_format: crate::ocr::types::OcrResponseFormat,
) -> Result<LiteLLMOcrResponse, crate::ocr::Error> {
MistralOCRConfig.transform_ocr_response(model, raw_response, request_format)
}
fn validate_request_body(&self, body: &Value) -> Result<(), crate::ocr::Error> {
validate_inline_document(&crate::ocr::prepare::body_document(body)?)
}
}
impl OcrEnvironment for vertex::VertexEnvironment {
fn headers(&self) -> &[(String, String)] {
&self.headers
}
}
impl VertexAIOCRConfig {
pub(super) async fn validate_environment(
&self,
connection: &OcrConnection,
config: &VertexConfig,
client: &OcrClient,
) -> Result<vertex::VertexEnvironment, crate::ocr::Error> {
validate_destination(connection)?;
client
.vertex_auth()
.validate_environment(
connection.extra_headers.clone(),
connection.api_key.as_deref(),
config,
&credential_env,
)
.await
.map_err(crate::ocr::Error::from)
}
fn get_complete_url(
&self,
api_base: Option<&str>,
project: &str,
location: &str,
model: &str,
) -> Result<String, crate::ocr::Error> {
validate_location(location)?;
let default_base = format!("https://{location}-aiplatform.googleapis.com");
let base = api_base
.map(str::trim)
.filter(|base| !base.is_empty())
.unwrap_or(&default_base);
let prediction = format!("{model}:rawPredict");
ApiUrl::parse(base)
.and_then(|url| {
url.complete_path(&[
"v1",
"projects",
project,
"locations",
location,
"publishers",
"mistralai",
"models",
&prediction,
])
})
.map(|url| url.into_string())
.map_err(|_| crate::ocr::Error::RequestField {
path: "api_base".into(),
})
}
}
fn validate_location(location: &str) -> Result<(), crate::ocr::Error> {
let valid = !location.is_empty()
&& location
.bytes()
.all(|value| value.is_ascii_lowercase() || value.is_ascii_digit() || value == b'-')
&& location
.as_bytes()
.first()
.is_some_and(u8::is_ascii_alphanumeric)
&& location
.as_bytes()
.last()
.is_some_and(u8::is_ascii_alphanumeric);
if valid {
return Ok(());
}
Err(crate::ocr::Error::RequestField {
path: "vertex_location".into(),
})
}
#[cfg(test)]
mod tests {
use super::VertexAIOCRConfig;
#[test]
fn endpoint_uses_location_project_and_model() {
assert_eq!(
VertexAIOCRConfig
.get_complete_url(None, "proj-1", "europe-west4", "mistral-ocr-maas")
.unwrap(),
"https://europe-west4-aiplatform.googleapis.com/v1/projects/proj-1/locations/europe-west4/publishers/mistralai/models/mistral-ocr-maas:rawPredict"
);
assert!(
VertexAIOCRConfig
.get_complete_url(None, "proj-1", "attacker.example/path", "model")
.is_err()
);
}
use serde_json::{Value, json};
use crate::ocr::test_support::{MockResponse, mock_server, perform_ocr, wire_request};
use litellm_auth::InputSource;
fn request_body(request: &str) -> Value {
serde_json::from_str(request.split_once("\r\n\r\n").unwrap().1).unwrap()
}
#[tokio::test]
async fn facade_executes_vertex_mistral_with_resolved_project_and_location() {
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({
"pages":[{"index":0,"markdown":"hello"}],
"usage_info":{"pages_processed":1}
}))])
.await;
let request = wire_request(
"vertex_ai/mistral-ocr-maas",
&base,
json!({
"vertex_project":"project-1",
"vertex_location":"europe-west4",
"extract_footer":true
}),
);
let response = perform_ocr(request).await.unwrap();
server.await.unwrap();
assert_eq!(response.pages[0].markdown, "hello");
let requests = seen.lock().unwrap();
assert_eq!(requests.len(), 1);
assert!(requests[0].starts_with(
"POST /v1/projects/project-1/locations/europe-west4/publishers/mistralai/models/mistral-ocr-maas:rawPredict "
));
assert!(
requests[0]
.to_ascii_lowercase()
.contains("authorization: bearer test-key")
);
assert_eq!(
request_body(&requests[0]),
json!({
"model":"mistral-ocr-maas",
"document":{"type":"document_url","document_url":"data:application/pdf;base64,YWJj"},
"extract_footer":true
})
);
}
#[tokio::test]
async fn supplied_authorization_is_forwarded_without_a_static_token() {
let (base, seen, server) = mock_server(vec![MockResponse::json(json!({"pages":[]}))]).await;
let mut request = wire_request(
"vertex_ai/model",
&base,
json!({"vertex_project":"project-1"}),
);
request.credentials.api_key = None;
request.transport.extra_headers = vec![("authorization".into(), "Bearer supplied".into())];
perform_ocr(request).await.unwrap();
server.await.unwrap();
assert!(
seen.lock().unwrap()[0]
.to_ascii_lowercase()
.contains("authorization: bearer supplied")
);
}
#[tokio::test]
async fn invalid_credentials_fail_before_provider_http() {
let request = wire_request(
"vertex_ai/model",
"http://127.0.0.1:1",
json!({"vertex_credentials": true}),
);
let error = perform_ocr(request).await.unwrap_err();
assert!(error.to_string().contains("vertex_credentials"));
}
#[tokio::test]
async fn request_controlled_api_base_is_rejected_before_vertex_auth() {
let mut request = wire_request(
"vertex_ai/mistral-ocr-maas",
"https://caller.example",
json!({"vertex_project":"project-1"}),
);
request.credentials.api_base = Some(litellm_auth::Sourced::new(
"https://caller.example".into(),
InputSource::Request,
));
let error = perform_ocr(request).await.unwrap_err();
assert!(
error
.to_string()
.contains("request-controlled Vertex AI endpoint")
);
}
#[tokio::test]
async fn configs_build_complete_requests_and_share_mistral_normalization() {
use std::time::Duration;
use crate::llms::base_llm::ocr::transformation::BaseOcrConfig;
use crate::llms::mistral::ocr::transformation::MistralOCRConfig;
use crate::llms::vertex_ai::ocr::transformation::VertexAIOCRConfig;
use crate::ocr::test_support::ocr_client;
let client = ocr_client();
let options = json!({
"pages": [0, 2],
"include_image_base64": true,
"vertex_project": "project-1",
"vertex_location": "us-central1",
"unknown": "preserved"
});
let direct = wire_request(
"mistral/mistral-ocr-maas",
"https://mistral.test",
options.clone(),
);
let vertex = wire_request("vertex_ai/mistral-ocr-maas", "https://vertex.test", options);
let direct = crate::ocr::prepare::prepare_request(
crate::ocr::test_support::resolved_request(direct),
);
let vertex = crate::ocr::prepare::prepare_request(
crate::ocr::test_support::resolved_request(vertex),
);
let direct_http = MistralOCRConfig
.prepare_request(&direct, &client)
.await
.unwrap();
let vertex_http = VertexAIOCRConfig
.prepare_request(&vertex, &client)
.await
.unwrap();
assert_eq!(direct_http.url().as_str(), "https://mistral.test/v1/ocr");
assert_eq!(
vertex_http.url().as_str(),
"https://vertex.test/v1/projects/project-1/locations/us-central1/publishers/mistralai/models/mistral-ocr-maas:rawPredict"
);
for http in [&direct_http, &vertex_http] {
assert_eq!(http.method(), reqwest::Method::POST);
assert_eq!(http.headers()["authorization"], "Bearer test-key");
assert_eq!(http.headers()["content-type"], "application/json");
assert_eq!(http.timeout(), Some(&Duration::from_secs(2)));
let body: Value =
serde_json::from_slice(http.body().unwrap().as_bytes().unwrap()).unwrap();
assert_eq!(
body,
json!({
"model": "mistral-ocr-maas",
"document": {"type": "document_url", "document_url": "data:application/pdf;base64,YWJj"},
"pages": [0, 2],
"include_image_base64": true,
"unknown": "preserved"
})
);
}
let payload = serde_json::to_vec(
&json!({"pages": [{"index": 0, "markdown": "hello"}], "extra": "preserved"}),
)
.unwrap();
let direct_response = MistralOCRConfig
.transform_ocr_response(&direct.model, &payload, Default::default())
.unwrap()
.into_json();
let vertex_response = VertexAIOCRConfig
.transform_ocr_response(&vertex.model, &payload, Default::default())
.unwrap()
.into_json();
assert_eq!(direct_response, vertex_response);
assert_eq!(direct_response["model"], "mistral-ocr-maas");
assert_eq!(direct_response["object"], "ocr");
assert_eq!(direct_response["extra"], "preserved");
}
}

View file

@ -1,131 +0,0 @@
use super::super::OcrAdapter;
use crate::ocr::Error;
use crate::ocr::OcrClient;
use crate::ocr::codecs::cohere::{
CohereParams, CohereResponse, transform_request, transform_response, validate_document,
};
use crate::ocr::document::{inline_remote_document, validate_inline_document};
use crate::ocr::error::{OcrError, OcrRequestError, OcrResponseError};
use crate::ocr::prepare::{credential_env, transform_request_body};
use crate::ocr::registry::OcrProvider;
use crate::ocr::types::{LiteLLMOcrRequest, LiteLLMOcrResponse};
use crate::url_utils::ApiUrl;
use litellm_auth_azure::AzureAuthInputs;
const AZURE_AI_API_BASE_ENV: &str = "AZURE_AI_API_BASE";
pub(crate) struct AzureCohereAdapter;
impl OcrAdapter for AzureCohereAdapter {
type ProviderResponse = CohereResponse;
const PROVIDER: OcrProvider = OcrProvider::AzureAi;
async fn prepare_request(
&self,
request: &LiteLLMOcrRequest,
client: &OcrClient,
) -> Result<reqwest::Request, OcrError> {
let params = super::super::super::wire::decode_request_value::<CohereParams>(
serde_json::Value::Object(request.optional_params.clone()),
"optional_params",
)?;
let mut config = AzureAuthInputs::from_sourced_optional_params(
&request.optional_params,
&request.input_sources,
)
.map_err(Error::from)?;
config.azure_ad_token_provider = request.azure_ad_token_provider.clone();
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(|| {
Error::Auth(
"Missing Azure AI API Base - Set AZURE_AI_API_BASE or pass api_base".into(),
)
})?;
let headers =
super::validate_ai_environment(&request.connection, &config, &credential_env).await?;
validate_document(&request.document)?;
let remote = request.document.source().starts_with("http://")
|| request.document.source().starts_with("https://");
let document = inline_remote_document(
client.document_fetcher(),
request.document.clone(),
&request.connection,
)
.await?;
let body = transform_request(&request.model, document, params)?;
transform_request_body(
client,
request,
&complete_url(&base)?,
&headers,
!remote,
body,
|body| {
validate_document(&body.document)?;
validate_inline_document(&body.document)
},
)
.await
}
fn transform_ocr_response(
&self,
request: &LiteLLMOcrRequest,
response: Self::ProviderResponse,
) -> Result<LiteLLMOcrResponse, OcrResponseError> {
transform_response(&request.model, response)
}
}
fn complete_url(base: &str) -> Result<String, OcrError> {
let mut url = reqwest::Url::parse(base).map_err(|_| invalid_api_base())?;
if !matches!(url.scheme(), "http" | "https") {
return Err(invalid_api_base().into());
}
let path = url.path().trim_end_matches('/').to_string();
if path.ends_with("/v2/parse") {
url.set_path(&path);
return Ok(url.into());
}
url.set_path(path.strip_suffix("/models").unwrap_or(&path));
ApiUrl::parse(url.as_str())
.and_then(|url| url.complete_path(&["providers", "cohere", "v2", "parse"]))
.map(|url| url.into_string())
.map_err(|_| invalid_api_base().into())
}
fn invalid_api_base() -> OcrRequestError {
OcrRequestError::RequestField {
path: "api_base".into(),
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn completes_foundry_urls_without_duplicate_paths_and_preserves_queries() {
for suffix in [
"",
"/models",
"/providers/cohere/v2",
"/providers/cohere/v2/parse",
] {
assert_eq!(
complete_url(&format!("https://example.com{suffix}?tenant=a")).unwrap(),
"https://example.com/providers/cohere/v2/parse?tenant=a"
);
}
assert_eq!(
complete_url("https://example.com/v2/parse?tenant=a").unwrap(),
"https://example.com/v2/parse?tenant=a"
);
assert!(complete_url("relative/path").is_err());
}
}

View file

@ -1,214 +0,0 @@
use super::super::OcrAdapter;
use crate::constants::{AZURE_DI_API_VERSION, AZURE_DI_SUBSCRIPTION_HEADER};
use crate::ocr::Error;
use crate::ocr::OcrClient;
use crate::ocr::codecs::document_intelligence::{
self, AzureDocumentIntelligenceOperation, DocumentIntelligenceParams,
};
use crate::ocr::error::{OcrError, OcrRequestError, OcrResponseError};
use crate::ocr::prepare::{credential_env, transform_request_body};
use crate::ocr::registry::OcrProvider;
use crate::ocr::types::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrConnection, OcrResponseFormat};
use crate::url_utils::ApiUrl;
use litellm_auth::{InputSource, Sourced};
use litellm_auth_azure::AzureAuthInputs;
mod polling;
const AZURE_DI_API_KEY_ENV: &str = "AZURE_DOCUMENT_INTELLIGENCE_API_KEY";
const AZURE_DI_ENDPOINT_ENV: &str = "AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT";
#[derive(Clone, Debug)]
pub(crate) struct AzureDocumentIntelligenceAdapter;
impl OcrAdapter for AzureDocumentIntelligenceAdapter {
type ProviderResponse = AzureDocumentIntelligenceOperation;
const PROVIDER: OcrProvider = OcrProvider::AzureAi;
async fn prepare_request(
&self,
request: &LiteLLMOcrRequest,
client: &OcrClient,
) -> Result<reqwest::Request, OcrError> {
let params = map_ocr_params(request)?;
let mut config = AzureAuthInputs::from_sourced_optional_params(
&request.optional_params,
&request.input_sources,
)
.map_err(Error::from)?;
config.azure_ad_token_provider = request.azure_ad_token_provider.clone();
let headers = validate_environment(&request.connection, &config, &credential_env).await?;
let endpoint = nonblank(request.connection.api_base.clone())
.or_else(|| nonblank(credential_env(AZURE_DI_ENDPOINT_ENV)))
.ok_or_else(|| Error::Auth("Missing Azure Document Intelligence API Base - Set AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT or pass api_base".into()))?;
let url = get_complete_url(&endpoint, &request.model, &params)?;
let body = document_intelligence::transform_ocr_request(request.document.clone())?;
transform_request_body(client, request, &url, &headers, false, body, |_| Ok(())).await
}
fn transform_ocr_response(
&self,
request: &LiteLLMOcrRequest,
response: Self::ProviderResponse,
) -> Result<LiteLLMOcrResponse, OcrResponseError> {
document_intelligence::transform_ocr_response(&request.model, response)
}
async fn read_response(
&self,
client: &OcrClient,
response: reqwest::Response,
url: &str,
headers: &[(String, String)],
request: &LiteLLMOcrRequest,
) -> Result<crate::ocr::wire::DecodedOcrResponse<Self::ProviderResponse>, OcrError> {
polling::read_operation_response(
client.polling_http(),
response,
url,
headers,
&request.connection,
request.response_format()? == OcrResponseFormat::Native,
&request.hooks,
)
.await
}
}
fn map_ocr_params(
request: &LiteLLMOcrRequest,
) -> Result<DocumentIntelligenceParams, OcrRequestError> {
let params = document_intelligence::decode_input_params(
request.optional_params.clone(),
"optional_params",
)?;
let crate::ocr::prepare::ParsedProviderParams {
known: params,
extra_params: _extra_params,
} = params;
document_intelligence::map_ocr_params(params)
}
fn get_complete_url(
endpoint: &str,
model: &str,
params: &DocumentIntelligenceParams,
) -> Result<String, OcrError> {
let model = format!("{}:analyze", model_id(model)?);
ApiUrl::parse(endpoint)
.and_then(|url| url.complete_path(&["documentintelligence", "documentModels", &model]))
.map(|url| {
url.append_query_pairs(
[("api-version", AZURE_DI_API_VERSION)]
.into_iter()
.chain(params.pages.iter().map(|pages| ("pages", pages.as_str())))
.chain(
params
.features
.iter()
.map(|features| ("features", features.as_str())),
),
)
.into_string()
})
.map_err(|_| OcrRequestError::RequestField {
path: "api_base".into(),
})
.map_err(OcrError::from)
}
async fn validate_environment(
connection: &OcrConnection,
config: &AzureAuthInputs,
env_lookup: &(dyn Fn(&str) -> Option<String> + Sync),
) -> Result<Vec<(String, String)>, OcrError> {
if crate::http_utils::has_header(&connection.extra_headers, "authorization")
|| crate::http_utils::has_header(&connection.extra_headers, AZURE_DI_SUBSCRIPTION_HEADER)
{
super::validate_destination(connection, connection.extra_headers_source)?;
return Ok(connection.extra_headers.clone());
}
let key = nonblank(connection.api_key.clone())
.map(|value| Sourced::new(value, connection.api_key_source))
.or_else(|| {
nonblank(env_lookup(AZURE_DI_API_KEY_ENV))
.map(|value| Sourced::new(value, InputSource::Environment))
});
if let Some(key) = key {
super::validate_destination(connection, key.source())?;
return Ok(
std::iter::once((AZURE_DI_SUBSCRIPTION_HEADER.into(), key.into_value()))
.chain(connection.extra_headers.clone())
.collect(),
);
}
let token = super::resolve_entra(config, env_lookup)
.await?
.ok_or(Error::MissingAzureDocumentIntelligenceCredentials)?;
super::validate_destination(connection, token.source())?;
Ok(
std::iter::once(("Authorization".into(), format!("Bearer {}", token.value())))
.chain(connection.extra_headers.clone())
.collect(),
)
}
fn model_id(model: &str) -> Result<&str, OcrRequestError> {
let model = model.rsplit('/').next().unwrap_or(model);
if matches!(model, "." | "..") {
return Err(OcrRequestError::DotModel);
}
Ok(model)
}
fn nonblank(value: Option<String>) -> Option<String> {
value
.map(|value| value.trim().to_string())
.filter(|value| !value.is_empty())
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn request_endpoint_cannot_receive_environment_key() {
let connection = OcrConnection {
api_base: Some("https://request.example".into()),
api_base_source: InputSource::Request,
..Default::default()
};
let error = validate_environment(&connection, &Default::default(), &|name| {
(name == AZURE_DI_API_KEY_ENV).then(|| "environment-key".into())
})
.await
.unwrap_err();
assert!(
error
.to_string()
.contains("request-controlled Azure endpoint")
);
}
#[tokio::test]
async fn request_endpoint_accepts_request_owned_key() {
let connection = OcrConnection {
api_key: Some("request-key".into()),
api_key_source: InputSource::Request,
api_base: Some("https://request.example".into()),
api_base_source: InputSource::Request,
..Default::default()
};
let headers = validate_environment(&connection, &Default::default(), &|_| None)
.await
.unwrap();
assert_eq!(
headers[0],
(AZURE_DI_SUBSCRIPTION_HEADER.into(), "request-key".into())
);
}
}

View file

@ -1,119 +0,0 @@
use std::sync::Arc;
use std::time::Duration;
use reqwest::Url;
use tokio::time::Instant;
use crate::constants::{AZURE_DI_SUBSCRIPTION_HEADER, OCR_POLL_RETRY_SECS};
use crate::ocr::client::read_json_response;
use crate::ocr::codecs::document_intelligence::{
AzureDocumentIntelligenceOperation, OperationStatus,
};
use crate::ocr::error::{OcrError, OcrPollingError, OcrResponseError};
use crate::ocr::hooks::OcrHooks;
use crate::ocr::types::OcrConnection;
use crate::ocr::wire::DecodedOcrResponse;
pub(super) async fn read_operation_response(
http_client: &reqwest::Client,
response: reqwest::Response,
original_url: &str,
headers: &[(String, String)],
connection: &OcrConnection,
native: bool,
hooks: &Arc<dyn OcrHooks>,
) -> Result<DecodedOcrResponse<AzureDocumentIntelligenceOperation>, OcrError> {
if response.status() != reqwest::StatusCode::ACCEPTED {
let bytes =
crate::ocr::client::read_response_bytes(response, connection.max_response_bytes)
.await?;
crate::ocr::handler::post_call(hooks, &bytes).await?;
return Ok(crate::ocr::wire::decode_response(&bytes, native)?);
}
let location = response
.headers()
.get("operation-location")
.and_then(|value| value.to_str().ok())
.ok_or(OcrPollingError::PollLocation)?
.to_string();
let original = Url::parse(original_url).map_err(|_| OcrPollingError::PollOrigin)?;
let operation = Url::parse(&location).map_err(|_| OcrPollingError::PollOrigin)?;
if original.origin() != operation.origin()
|| !operation.username().is_empty()
|| operation.password().is_some()
{
return Err(OcrPollingError::PollOrigin.into());
}
let bytes =
crate::ocr::client::read_response_bytes(response, connection.max_response_bytes).await?;
crate::ocr::handler::post_call(hooks, &bytes).await?;
poll_operation(http_client, operation, headers, connection, native, hooks).await
}
async fn poll_operation(
http_client: &reqwest::Client,
url: Url,
headers: &[(String, String)],
connection: &OcrConnection,
native: bool,
hooks: &Arc<dyn OcrHooks>,
) -> Result<DecodedOcrResponse<AzureDocumentIntelligenceOperation>, OcrError> {
let deadline = Instant::now()
.checked_add(connection.poll_timeout)
.ok_or(OcrPollingError::PollTimeout)?;
loop {
let remaining = deadline
.checked_duration_since(Instant::now())
.filter(|remaining| !remaining.is_zero())
.ok_or(OcrPollingError::PollTimeout)?;
let builder = http_client
.get(url.clone())
.timeout(remaining.min(connection.timeout));
let builder = crate::http_utils::with_headers(
builder,
headers,
crate::http_utils::HeaderPolicy::Only(&[AZURE_DI_SUBSCRIPTION_HEADER, "authorization"]),
);
let response = tokio::time::timeout_at(deadline, crate::http_utils::http_request(builder))
.await
.map_err(|_| OcrPollingError::PollTimeout)?
.map_err(crate::transport::Error::from)?;
let retry = response
.headers()
.get(reqwest::header::RETRY_AFTER)
.and_then(|value| value.to_str().ok())
.and_then(|value| value.parse::<u64>().ok())
.unwrap_or(OCR_POLL_RETRY_SECS)
.max(1);
let decoded = tokio::time::timeout_at(
deadline,
read_json_response::<AzureDocumentIntelligenceOperation>(
response,
native,
connection.max_response_bytes,
),
)
.await
.map_err(|_| OcrPollingError::PollTimeout)??;
match &decoded.data.status {
Some(OperationStatus::Succeeded) => {
crate::ocr::handler::post_call(hooks, decoded.text.as_bytes()).await?;
return Ok(decoded);
}
Some(OperationStatus::Running | OperationStatus::NotStarted) => {
tokio::time::timeout_at(deadline, tokio::time::sleep(Duration::from_secs(retry)))
.await
.map_err(|_| OcrPollingError::PollTimeout)?;
}
status => {
return Err(OcrResponseError::OperationStatus(
status
.as_ref()
.map(ToString::to_string)
.unwrap_or_else(|| "None".into()),
)
.into());
}
}
}
}

View file

@ -1,229 +0,0 @@
use super::super::OcrAdapter;
use crate::constants::AZURE_AI_OCR_PATH;
use crate::ocr::Error;
use crate::ocr::OcrClient;
use crate::ocr::codecs::mistral::{self, MistralOcrParams, MistralOcrResponse};
use crate::ocr::document::{inline_remote_document, validate_inline_document};
use crate::ocr::error::{OcrError, OcrRequestError, OcrResponseError};
use crate::ocr::prepare::{
_prepare_ocr_request, ParsedProviderParams, credential_env, transform_request_body,
};
use crate::ocr::registry::OcrProvider;
use crate::ocr::types::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrConnection};
use crate::url_utils::ApiUrl;
use litellm_auth::{InputSource, Sourced};
use litellm_auth_azure::AzureAuthInputs;
const AZURE_AI_API_KEY_ENV: &str = "AZURE_AI_API_KEY";
const AZURE_AI_API_BASE_ENV: &str = "AZURE_AI_API_BASE";
#[derive(Clone, Debug)]
pub(crate) struct AzureMistralAdapter;
impl OcrAdapter for AzureMistralAdapter {
type ProviderResponse = MistralOcrResponse;
const PROVIDER: OcrProvider = OcrProvider::AzureAi;
async fn prepare_request(
&self,
request: &LiteLLMOcrRequest,
client: &OcrClient,
) -> Result<reqwest::Request, OcrError> {
let ParsedProviderParams {
known: params,
extra_params: _extra_params,
} = _prepare_ocr_request::<MistralOcrParams>(request)?;
let mut config = AzureAuthInputs::from_sourced_optional_params(
&request.optional_params,
&request.input_sources,
)
.map_err(Error::from)?;
config.azure_ad_token_provider = request.azure_ad_token_provider.clone();
let url = get_complete_url(request.connection.api_base.as_deref(), &credential_env)?;
let headers = validate_environment(&request.connection, &config, &credential_env).await?;
let retains_document = !request.document.source().starts_with("http://")
&& !request.document.source().starts_with("https://");
let document = inline_remote_document(
client.document_fetcher(),
request.document.clone(),
&request.connection,
)
.await?;
let body = mistral::transform_ocr_request(&request.model, document, &params)?;
transform_request_body(
client,
request,
&url,
&headers,
retains_document,
body,
|body| validate_inline_document(&body.document),
)
.await
}
fn transform_ocr_response(
&self,
request: &LiteLLMOcrRequest,
response: Self::ProviderResponse,
) -> Result<LiteLLMOcrResponse, OcrResponseError> {
mistral::transform_ocr_response(&request.model, response)
}
}
fn get_complete_url(
api_base: Option<&str>,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> Result<String, OcrError> {
let base = nonblank(api_base.map(str::to_string))
.or_else(|| nonblank(env_lookup(AZURE_AI_API_BASE_ENV)))
.ok_or_else(|| Error::Auth(
"Missing Azure AI API Base - Set AZURE_AI_API_BASE environment variable or pass api_base parameter".into(),
))?;
let path: Vec<&str> = AZURE_AI_OCR_PATH.trim_matches('/').split('/').collect();
ApiUrl::parse(&base)
.and_then(|url| url.complete_path(&path))
.map(|url| url.into_string())
.map_err(|_| {
OcrRequestError::RequestField {
path: "api_base".into(),
}
.into()
})
}
pub(in crate::ocr::adapters) async fn validate_environment(
connection: &OcrConnection,
config: &AzureAuthInputs,
env_lookup: &(dyn Fn(&str) -> Option<String> + Sync),
) -> Result<Vec<(String, String)>, OcrError> {
if crate::http_utils::has_header(&connection.extra_headers, "authorization") {
if config.azure_ad_token_provider.is_some() {
super::resolve_entra(config, env_lookup).await?;
}
super::validate_destination(connection, connection.extra_headers_source)?;
return Ok(connection.extra_headers.clone());
}
let key = nonblank(connection.api_key.clone())
.map(|value| Sourced::new(value, connection.api_key_source))
.or_else(|| {
nonblank(env_lookup(AZURE_AI_API_KEY_ENV))
.map(|value| Sourced::new(value, InputSource::Environment))
});
if let Some(key) = key {
super::validate_destination(connection, key.source())?;
return Ok(bearer_headers(connection, key.value()));
}
let key = super::resolve_entra(config, env_lookup)
.await?
.ok_or(Error::MissingAzureAiCredentials)?;
super::validate_destination(connection, key.source())?;
Ok(bearer_headers(connection, key.value()))
}
fn bearer_headers(connection: &OcrConnection, key: &str) -> Vec<(String, String)> {
std::iter::once(("Authorization".into(), format!("Bearer {key}")))
.chain(connection.extra_headers.clone())
.collect()
}
fn nonblank(value: Option<String>) -> Option<String> {
value
.map(|value| value.trim().to_string())
.filter(|value| !value.is_empty())
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn completes_azure_path_and_preserves_query() {
assert_eq!(
get_complete_url(Some("https://example.com/?tenant=a"), &|_| None).unwrap(),
"https://example.com/providers/mistral/azure/ocr?tenant=a"
);
assert_eq!(
get_complete_url(
Some("https://example.com/providers/mistral/azure/ocr"),
&|_| None
)
.unwrap(),
"https://example.com/providers/mistral/azure/ocr"
);
}
#[tokio::test]
async fn supplied_authorization_precedes_keys() {
let connection = OcrConnection {
api_key: Some("request-key".into()),
extra_headers: vec![("authorization".into(), "Bearer prepared".into())],
..Default::default()
};
assert_eq!(
validate_environment(&connection, &Default::default(), &|_| {
Some("environment-key".into())
})
.await
.unwrap(),
connection.extra_headers
);
}
#[tokio::test]
async fn request_key_precedes_environment_key() {
let connection = OcrConnection {
api_key: Some("request-key".into()),
..Default::default()
};
assert_eq!(
validate_environment(&connection, &Default::default(), &|_| {
Some("environment-key".into())
})
.await
.unwrap()[0],
("Authorization".into(), "Bearer request-key".into())
);
}
#[tokio::test]
async fn request_endpoint_cannot_receive_environment_key() {
let connection = OcrConnection {
api_base: Some("https://request.example".into()),
api_base_source: InputSource::Request,
..Default::default()
};
let error = validate_environment(&connection, &Default::default(), &|name| {
(name == AZURE_AI_API_KEY_ENV).then(|| "environment-key".into())
})
.await
.unwrap_err();
assert!(
error
.to_string()
.contains("request-controlled Azure endpoint")
);
}
#[tokio::test]
async fn request_endpoint_accepts_request_owned_key() {
let connection = OcrConnection {
api_key: Some("request-key".into()),
api_key_source: InputSource::Request,
api_base: Some("https://request.example".into()),
api_base_source: InputSource::Request,
..Default::default()
};
let headers = validate_environment(&connection, &Default::default(), &|_| None)
.await
.unwrap();
assert_eq!(
headers[0],
("Authorization".into(), "Bearer request-key".into())
);
}
}

View file

@ -1,123 +0,0 @@
use super::OcrAdapter;
use crate::constants::{COHERE_API_KEY_ENV, COHERE_PARSE_API_BASE};
use crate::ocr::Error;
use crate::ocr::OcrClient;
use crate::ocr::codecs::cohere::{
CohereParams, CohereResponse, transform_request, transform_response, validate_document,
};
use crate::ocr::error::{OcrError, OcrRequestError, OcrResponseError};
use crate::ocr::prepare::{credential_env, transform_request_body};
use crate::ocr::registry::OcrProvider;
use crate::ocr::types::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrConnection};
use crate::url_utils::ApiUrl;
pub(crate) struct CohereAdapter;
impl OcrAdapter for CohereAdapter {
type ProviderResponse = CohereResponse;
const PROVIDER: OcrProvider = OcrProvider::Cohere;
async fn prepare_request(
&self,
request: &LiteLLMOcrRequest,
client: &OcrClient,
) -> Result<reqwest::Request, OcrError> {
let params = super::super::wire::decode_request_value::<CohereParams>(
serde_json::Value::Object(request.optional_params.clone()),
"optional_params",
)?;
let headers = validate_environment(&request.connection, &credential_env)?;
let url = complete_url(
request
.connection
.api_base
.as_deref()
.unwrap_or(COHERE_PARSE_API_BASE),
)?;
let body = transform_request(&request.model, request.document.clone(), params)?;
transform_request_body(client, request, &url, &headers, true, body, |body| {
validate_document(&body.document)
})
.await
}
fn transform_ocr_response(
&self,
request: &LiteLLMOcrRequest,
response: Self::ProviderResponse,
) -> Result<LiteLLMOcrResponse, OcrResponseError> {
transform_response(&request.model, response)
}
}
fn complete_url(base: &str) -> Result<String, OcrError> {
let parsed = reqwest::Url::parse(base).map_err(|_| invalid_api_base())?;
if !matches!(parsed.scheme(), "http" | "https") {
return Err(invalid_api_base().into());
}
ApiUrl::parse(base)
.and_then(|url| url.complete_path(&["v2", "parse"]))
.map(|url| url.into_string())
.map_err(|_| invalid_api_base().into())
}
fn invalid_api_base() -> OcrRequestError {
OcrRequestError::RequestField {
path: "api_base".into(),
}
}
fn validate_environment(
connection: &OcrConnection,
env_lookup: &(dyn Fn(&str) -> Option<String> + Sync),
) -> Result<Vec<(String, String)>, OcrError> {
if crate::http_utils::has_header(&connection.extra_headers, "authorization") {
return Ok(connection.extra_headers.clone());
}
let key = connection
.api_key
.as_deref()
.map(str::trim)
.filter(|key| !key.is_empty())
.map(str::to_string)
.or_else(|| env_lookup(COHERE_API_KEY_ENV).filter(|key| !key.trim().is_empty()))
.ok_or_else(|| {
Error::Auth("Missing COHERE_API_KEY - set it in the environment or pass api_key".into())
})?;
Ok(
std::iter::once(("Authorization".into(), format!("Bearer {key}")))
.chain(connection.extra_headers.clone())
.collect(),
)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn completes_provider_urls_without_duplicate_paths_and_preserves_queries() {
for suffix in ["", "/v2", "/v2/parse"] {
assert_eq!(
complete_url(&format!("https://example.com{suffix}?tenant=a")).unwrap(),
"https://example.com/v2/parse?tenant=a"
);
}
}
#[test]
fn rejects_invalid_urls_and_blank_keys() {
assert!(complete_url("relative/path").is_err());
assert!(complete_url("ftp://example.com").is_err());
assert!(matches!(
validate_environment(
&OcrConnection {
api_key: Some(" ".into()),
..Default::default()
},
&|_| None,
),
Err(OcrError::Public(Error::Auth(_)))
));
}
}

View file

@ -1,147 +0,0 @@
use super::OcrAdapter;
use crate::constants::MISTRAL_OCR_API_BASE;
use crate::ocr::Error;
use crate::ocr::OcrClient;
use crate::ocr::codecs::mistral::{self, MistralOcrParams, MistralOcrResponse};
use crate::ocr::error::{OcrError, OcrRequestError, OcrResponseError};
use crate::ocr::prepare::{
_prepare_ocr_request, ParsedProviderParams, credential_env, transform_request_body,
};
use crate::ocr::registry::OcrProvider;
use crate::ocr::types::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrConnection};
use crate::url_utils::ApiUrl;
const MISTRAL_API_KEY_ENV: &str = "MISTRAL_API_KEY";
#[derive(Clone, Debug)]
pub(crate) struct MistralAdapter;
impl OcrAdapter for MistralAdapter {
type ProviderResponse = MistralOcrResponse;
const PROVIDER: OcrProvider = OcrProvider::Mistral;
async fn prepare_request(
&self,
request: &LiteLLMOcrRequest,
client: &OcrClient,
) -> Result<reqwest::Request, OcrError> {
let ParsedProviderParams {
known: params,
extra_params: _extra_params,
} = _prepare_ocr_request::<MistralOcrParams>(request)?;
let headers = validate_environment(&request.connection, &credential_env)?;
let url = get_complete_url(request.connection.api_base.as_deref())?;
let body =
mistral::transform_ocr_request(&request.model, request.document.clone(), &params)?;
transform_request_body(client, request, &url, &headers, true, body, |_| Ok(())).await
}
fn transform_ocr_response(
&self,
request: &LiteLLMOcrRequest,
response: Self::ProviderResponse,
) -> Result<LiteLLMOcrResponse, OcrResponseError> {
mistral::transform_ocr_response(&request.model, response)
}
}
pub(crate) fn get_complete_url(api_base: Option<&str>) -> Result<String, OcrError> {
let base = api_base
.map(str::trim)
.filter(|base| !base.is_empty())
.unwrap_or(MISTRAL_OCR_API_BASE);
ApiUrl::parse(base)
.and_then(|url| url.complete_path(&["v1", "ocr"]))
.map(|url| url.into_string())
.map_err(|_| {
OcrRequestError::RequestField {
path: "api_base".into(),
}
.into()
})
}
fn validate_environment(
connection: &OcrConnection,
env_lookup: &(dyn Fn(&str) -> Option<String> + Sync),
) -> Result<Vec<(String, String)>, OcrError> {
if crate::http_utils::has_header(&connection.extra_headers, "authorization") {
return Ok(connection.extra_headers.clone());
}
let api_key = connection
.api_key
.as_deref()
.map(str::trim)
.filter(|key| !key.is_empty())
.map(str::to_string)
.or_else(|| env_lookup(MISTRAL_API_KEY_ENV).filter(|key| !key.trim().is_empty()))
.ok_or(Error::MissingApiKey {
provider: "Mistral",
})?;
Ok(
std::iter::once(("Authorization".into(), format!("Bearer {api_key}")))
.chain(connection.extra_headers.clone())
.collect(),
)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn complete_url_defaults_and_dedupes_v1() {
assert_eq!(
get_complete_url(None).unwrap(),
"https://api.mistral.ai/v1/ocr"
);
assert_eq!(
get_complete_url(Some("https://example.com/v1?tenant=a")).unwrap(),
"https://example.com/v1/ocr?tenant=a"
);
assert_eq!(
get_complete_url(Some("https://example.com/v1/ocr?tenant=a")).unwrap(),
"https://example.com/v1/ocr?tenant=a"
);
}
#[test]
fn environment_prefers_explicit_key_then_environment() {
let explicit = OcrConnection {
api_key: Some("explicit".into()),
..OcrConnection::default()
};
assert_eq!(
validate_environment(&explicit, &|_| Some("environment".into())).unwrap()[0],
("Authorization".into(), "Bearer explicit".into())
);
assert_eq!(
validate_environment(&OcrConnection::default(), &|_| Some("environment".into()))
.unwrap()[0],
("Authorization".into(), "Bearer environment".into())
);
}
#[test]
fn environment_preserves_forwarded_authorization() {
let connection = OcrConnection {
extra_headers: vec![("authorization".into(), "Bearer forwarded".into())],
..OcrConnection::default()
};
assert_eq!(
validate_environment(&connection, &|_| None).unwrap(),
connection.extra_headers
);
}
#[test]
fn environment_rejects_missing_key() {
assert!(matches!(
validate_environment(&OcrConnection::default(), &|_| None),
Err(OcrError::Public(Error::MissingApiKey {
provider: "Mistral"
}))
));
}
}

View file

@ -1,91 +0,0 @@
use std::future::Future;
use serde::de::DeserializeOwned;
use super::OcrClient;
use super::error::{OcrError, OcrResponseError};
use super::registry::OcrProvider;
use super::types::{LiteLLMOcrRequest, LiteLLMOcrResponse};
mod azure;
mod cohere;
mod mistral;
mod reducto;
mod vertex;
pub(crate) use azure::{AzureCohereAdapter, AzureDocumentIntelligenceAdapter, AzureMistralAdapter};
pub(crate) use cohere::CohereAdapter;
pub(crate) use mistral::MistralAdapter;
pub(crate) use reducto::{ReductoLegacyAdapter, ReductoV3Adapter};
pub(crate) use vertex::{VertexDeepSeekAdapter, VertexMistralAdapter};
/// Converts a complete LiteLLM OCR call to provider HTTP and normalizes its response.
pub(crate) trait OcrAdapter: Send + Sync + Sized + 'static {
/// Provider JSON schema; direct and Vertex Mistral share `MistralOcrResponse`.
type ProviderResponse: DeserializeOwned + Send;
const PROVIDER: OcrProvider;
/// Prepares the complete provider HTTP request.
/// `request` contains the model, document, connection, and unmapped caller options.
/// `client` supplies reusable provider and document HTTP clients.
/// Returns the complete HTTP request, whereas Python returns body data.
fn prepare_request(
&self,
request: &LiteLLMOcrRequest,
client: &OcrClient,
) -> impl Future<Output = Result<reqwest::Request, OcrError>> + Send;
/// Python: `transform_ocr_response`.
/// `request` supplies caller context, including the fallback model.
/// `response` is the decoded provider payload; the output is the shared LiteLLM schema.
fn transform_ocr_response(
&self,
request: &LiteLLMOcrRequest,
response: Self::ProviderResponse,
) -> Result<LiteLLMOcrResponse, OcrResponseError>;
/// Decodes provider HTTP; adapters may override this to poll asynchronous operations.
/// Python performs that polling inside `async_transform_ocr_response`.
/// `client` is reused for polling; `response` is the initial HTTP response.
/// `url` and `headers` describe the submitted call; `request` supplies limits and format.
fn read_response(
&self,
_client: &OcrClient,
response: reqwest::Response,
_url: &str,
_headers: &[(String, String)],
request: &LiteLLMOcrRequest,
) -> impl Future<
Output = Result<super::wire::DecodedOcrResponse<Self::ProviderResponse>, OcrError>,
> + Send {
async move {
let bytes =
super::client::read_response_bytes(response, request.connection.max_response_bytes)
.await?;
super::handler::post_call(&request.hooks, &bytes).await?;
Ok(super::wire::decode_response(
&bytes,
request.response_format()? == super::types::OcrResponseFormat::Native,
)?)
}
}
}
macro_rules! for_each_ocr_adapter {
($callback:ident) => {
$callback! {
Cohere, $crate::ocr::adapters::CohereAdapter, $crate::ocr::adapters::CohereAdapter, Cohere;
AzureCohere, $crate::ocr::adapters::AzureCohereAdapter, $crate::ocr::adapters::AzureCohereAdapter, AzureAi;
Mistral, $crate::ocr::adapters::MistralAdapter, $crate::ocr::adapters::MistralAdapter, Mistral;
AzureMistral, $crate::ocr::adapters::AzureMistralAdapter, $crate::ocr::adapters::AzureMistralAdapter, AzureAi;
AzureDocumentIntelligence, $crate::ocr::adapters::AzureDocumentIntelligenceAdapter, $crate::ocr::adapters::AzureDocumentIntelligenceAdapter, AzureAi;
ReductoLegacy, $crate::ocr::adapters::ReductoLegacyAdapter, $crate::ocr::adapters::ReductoLegacyAdapter, Reducto;
ReductoV3, $crate::ocr::adapters::ReductoV3Adapter, $crate::ocr::adapters::ReductoV3Adapter, Reducto;
VertexMistral, $crate::ocr::adapters::VertexMistralAdapter, $crate::ocr::adapters::VertexMistralAdapter, VertexAi;
VertexDeepSeek, $crate::ocr::adapters::VertexDeepSeekAdapter, $crate::ocr::adapters::VertexDeepSeekAdapter, VertexAi;
}
};
}
pub(crate) use for_each_ocr_adapter;

View file

@ -1,45 +0,0 @@
use super::super::OcrAdapter;
use crate::ocr::OcrClient;
use crate::ocr::codecs::reducto::{self, ReductoLegacyParams, ReductoResponse};
use crate::ocr::error::{OcrError, OcrResponseError};
use crate::ocr::prepare::{
_prepare_ocr_request, ParsedProviderParams, build_http_request, credential_env,
guardrail_document, merge_extra_params,
};
use crate::ocr::registry::OcrProvider;
use crate::ocr::types::{LiteLLMOcrRequest, LiteLLMOcrResponse};
#[derive(Clone, Debug)]
pub(crate) struct ReductoLegacyAdapter;
impl OcrAdapter for ReductoLegacyAdapter {
type ProviderResponse = ReductoResponse;
const PROVIDER: OcrProvider = OcrProvider::Reducto;
async fn prepare_request(
&self,
request: &LiteLLMOcrRequest,
client: &OcrClient,
) -> Result<reqwest::Request, OcrError> {
let ParsedProviderParams {
known: params,
extra_params,
} = _prepare_ocr_request::<ReductoLegacyParams>(request)?;
let headers = super::validate_environment(&request.connection, &credential_env)?;
let url = super::get_complete_url(request.connection.api_base.as_deref(), "parse")?;
let (document, headers) = guardrail_document(request, &url, &headers).await?;
let document =
super::prepare_document(client, document, &request.connection, &headers).await?;
let body = reducto::transform_legacy_ocr_request(&request.model, document, &params)?;
let body = merge_extra_params(&body, extra_params)?;
build_http_request(client, request, &url, &headers, &body)
}
fn transform_ocr_response(
&self,
request: &LiteLLMOcrRequest,
response: Self::ProviderResponse,
) -> Result<LiteLLMOcrResponse, OcrResponseError> {
reducto::transform_ocr_response(&request.model, response)
}
}

View file

@ -1,148 +0,0 @@
mod legacy;
mod v3;
use crate::constants::{REDUCTO_API_BASE, REDUCTO_API_KEY_ENV, REDUCTO_ID_PREFIX};
use crate::ocr::Error;
use crate::ocr::document::InlineDocument;
use crate::ocr::error::{OcrError, OcrRequestError, OcrResponseError};
use crate::ocr::types::{OcrConnection, OcrDocument};
use crate::url_utils::ApiUrl;
pub(crate) use legacy::ReductoLegacyAdapter;
pub(crate) use v3::ReductoV3Adapter;
pub(super) fn get_complete_url(api_base: Option<&str>, path: &str) -> Result<String, OcrError> {
let base = api_base
.map(str::trim)
.filter(|base| !base.is_empty())
.unwrap_or(REDUCTO_API_BASE);
ApiUrl::parse(base)
.and_then(|url| url.complete_path(&[path]))
.map(|url| url.into_string())
.map_err(|_| {
OcrRequestError::RequestField {
path: "api_base".into(),
}
.into()
})
}
pub(super) fn validate_environment(
connection: &OcrConnection,
env_lookup: &(dyn Fn(&str) -> Option<String> + Sync),
) -> Result<Vec<(String, String)>, OcrError> {
if crate::http_utils::has_header(&connection.extra_headers, "authorization") {
return Ok(connection.extra_headers.clone());
}
let api_key = connection
.api_key
.as_deref()
.map(str::trim)
.filter(|key| !key.is_empty())
.map(str::to_string)
.or_else(|| {
env_lookup(REDUCTO_API_KEY_ENV)
.map(|key| key.trim().to_string())
.filter(|key| !key.is_empty())
})
.ok_or(Error::MissingReductoApiKey)?;
Ok(
std::iter::once(("Authorization".into(), format!("Bearer {api_key}")))
.chain(connection.extra_headers.clone())
.collect(),
)
}
pub(super) async fn prepare_document(
client: &crate::ocr::OcrClient,
document: OcrDocument,
connection: &OcrConnection,
headers: &[(String, String)],
) -> Result<OcrDocument, OcrError> {
if document.source().starts_with(REDUCTO_ID_PREFIX) {
if document.source()[REDUCTO_ID_PREFIX.len()..]
.trim()
.is_empty()
{
return Err(OcrRequestError::RequestField {
path: "document file id".into(),
}
.into());
}
return Ok(document);
}
let inline = InlineDocument::parse(document.source())?.ok_or(OcrRequestError::ReductoSource)?;
let mime = inline.mime_type().to_string();
let bytes = inline.decode(crate::constants::OCR_INLINE_MAX_BYTES)?;
let part = reqwest::multipart::Part::bytes(bytes)
.file_name("document")
.mime_str(&mime)
.map_err(|_| OcrRequestError::InvalidDataUri)?;
let builder = client
.provider_http()
.post(get_complete_url(connection.api_base.as_deref(), "upload")?)
.multipart(reqwest::multipart::Form::new().part("file", part))
.timeout(connection.timeout);
let builder = crate::http_utils::with_headers(
builder,
headers,
crate::http_utils::HeaderPolicy::Except(&["content-type", "content-length"]),
);
let response = crate::http_utils::http_request(builder)
.await
.map_err(crate::transport::Error::from)?;
let uploaded = crate::ocr::client::read_json_response::<
crate::ocr::codecs::reducto::ReductoUploadResponse,
>(response, false, connection.max_response_bytes)
.await?
.data;
let file_id = uploaded
.file_id
.as_deref()
.map(str::trim)
.filter(|id| !id.is_empty());
let Some(file_id) = file_id else {
return Err(OcrResponseError::ResponseField {
path: "file_id".into(),
}
.into());
};
Ok(document.with_source(file_id.to_string()))
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn explicit_key_precedes_environment_key() {
let connection = OcrConnection {
api_key: Some("passed-key".into()),
..Default::default()
};
let headers = validate_environment(&connection, &|_| Some("env-key".into())).unwrap();
assert_eq!(headers[0].1, "Bearer passed-key");
}
#[test]
fn blank_explicit_key_uses_environment_key() {
let connection = OcrConnection {
api_key: Some(" ".into()),
..Default::default()
};
let headers = validate_environment(&connection, &|_| Some(" env-key ".into())).unwrap();
assert_eq!(headers[0].1, "Bearer env-key");
}
#[test]
fn existing_authorization_skips_key_lookup() {
let connection = OcrConnection {
extra_headers: vec![("authorization".into(), "Bearer existing".into())],
..Default::default()
};
assert_eq!(
validate_environment(&connection, &|_| None).unwrap(),
connection.extra_headers
);
}
}

View file

@ -1,45 +0,0 @@
use super::super::OcrAdapter;
use crate::ocr::OcrClient;
use crate::ocr::codecs::reducto::{self, ReductoResponse, ReductoV3Params};
use crate::ocr::error::{OcrError, OcrResponseError};
use crate::ocr::prepare::{
_prepare_ocr_request, ParsedProviderParams, build_http_request, credential_env,
guardrail_document, merge_extra_params,
};
use crate::ocr::registry::OcrProvider;
use crate::ocr::types::{LiteLLMOcrRequest, LiteLLMOcrResponse};
#[derive(Clone, Debug)]
pub(crate) struct ReductoV3Adapter;
impl OcrAdapter for ReductoV3Adapter {
type ProviderResponse = ReductoResponse;
const PROVIDER: OcrProvider = OcrProvider::Reducto;
async fn prepare_request(
&self,
request: &LiteLLMOcrRequest,
client: &OcrClient,
) -> Result<reqwest::Request, OcrError> {
let ParsedProviderParams {
known: params,
extra_params,
} = _prepare_ocr_request::<ReductoV3Params>(request)?;
let headers = super::validate_environment(&request.connection, &credential_env)?;
let url = super::get_complete_url(request.connection.api_base.as_deref(), "parse")?;
let (document, headers) = guardrail_document(request, &url, &headers).await?;
let document =
super::prepare_document(client, document, &request.connection, &headers).await?;
let body = reducto::transform_v3_ocr_request(&request.model, document, &params)?;
let body = merge_extra_params(&body, extra_params)?;
build_http_request(client, request, &url, &headers, &body)
}
fn transform_ocr_response(
&self,
request: &LiteLLMOcrRequest,
response: Self::ProviderResponse,
) -> Result<LiteLLMOcrResponse, OcrResponseError> {
reducto::transform_ocr_response(&request.model, response)
}
}

View file

@ -1,140 +0,0 @@
use super::super::OcrAdapter;
use super::validate_destination;
use crate::ocr::Error;
use crate::ocr::OcrClient;
use crate::ocr::codecs::deepseek::{self, DeepSeekOcrParams, DeepSeekOcrResponse};
use crate::ocr::error::{OcrError, OcrRequestError, OcrResponseError};
use crate::ocr::prepare::{
_prepare_ocr_request, ParsedProviderParams, credential_env, transform_request_body,
};
use crate::ocr::registry::OcrProvider;
use crate::ocr::types::{LiteLLMOcrRequest, LiteLLMOcrResponse};
use crate::url_utils::ApiUrl;
use litellm_auth_gcp::{self as vertex, VertexConfig};
const DEFAULT_API_BASE: &str = "https://aiplatform.googleapis.com";
const MODEL_NAMESPACE: &str = "deepseek-ai";
const DEFAULT_LOCATION: &str = "us-central1";
#[derive(Clone, Debug)]
pub(crate) struct VertexDeepSeekAdapter;
impl OcrAdapter for VertexDeepSeekAdapter {
type ProviderResponse = DeepSeekOcrResponse;
const PROVIDER: OcrProvider = OcrProvider::VertexAi;
async fn prepare_request(
&self,
request: &LiteLLMOcrRequest,
client: &OcrClient,
) -> Result<reqwest::Request, OcrError> {
validate_destination(&request.connection)?;
let ParsedProviderParams {
known: params,
extra_params: _extra_params,
} = _prepare_ocr_request::<DeepSeekOcrParams>(request)?;
let config = VertexConfig::from_sourced_optional_params(
&request.optional_params,
&request.input_sources,
)
.map_err(Error::from)?;
let authentication = client
.vertex_auth()
.validate_environment(
request.connection.extra_headers.clone(),
request.connection.api_key.as_deref(),
&config,
&credential_env,
)
.await
.map_err(Error::from)?;
let location = vertex::get_vertex_ai_location(&config, &credential_env)
.unwrap_or_else(|| DEFAULT_LOCATION.to_string());
let url = get_complete_url(
request.connection.api_base.as_deref(),
&authentication.project_id,
&location,
)?;
let document = request.document.clone();
let body =
deepseek::transform_ocr_request(&provider_model(&request.model), document, &params)?;
transform_request_body(
client,
request,
&url,
&authentication.headers,
false,
body,
|_| Ok(()),
)
.await
}
fn transform_ocr_response(
&self,
request: &LiteLLMOcrRequest,
response: Self::ProviderResponse,
) -> Result<LiteLLMOcrResponse, OcrResponseError> {
deepseek::transform_ocr_response(&request.model, response)
}
}
fn provider_model(model: &str) -> String {
if model.starts_with(&format!("{MODEL_NAMESPACE}/")) {
model.to_string()
} else {
format!("{MODEL_NAMESPACE}/{model}")
}
}
fn get_complete_url(
api_base: Option<&str>,
project: &str,
location: &str,
) -> Result<String, OcrError> {
let base = api_base
.map(str::trim)
.filter(|base| !base.is_empty())
.unwrap_or(DEFAULT_API_BASE);
ApiUrl::parse(base)
.and_then(|url| {
url.complete_path(&[
"v1",
"projects",
project,
"locations",
location,
"endpoints",
"openapi",
"chat",
"completions",
])
})
.map(|url| url.into_string())
.map_err(|_| {
OcrRequestError::RequestField {
path: "api_base".into(),
}
.into()
})
}
#[cfg(test)]
mod tests {
use super::{get_complete_url, provider_model};
#[test]
fn adapter_owns_model_namespace_and_endpoint() {
assert_eq!(
provider_model("deepseek-ocr-maas"),
"deepseek-ai/deepseek-ocr-maas"
);
assert_eq!(
provider_model("deepseek-ai/deepseek-ocr-maas"),
"deepseek-ai/deepseek-ocr-maas"
);
assert_eq!(
get_complete_url(None, "proj-1", "europe-west4").unwrap(),
"https://aiplatform.googleapis.com/v1/projects/proj-1/locations/europe-west4/endpoints/openapi/chat/completions"
);
}
}

View file

@ -1,157 +0,0 @@
use super::super::OcrAdapter;
use super::validate_destination;
use crate::ocr::Error;
use crate::ocr::OcrClient;
use crate::ocr::codecs::mistral::{self, MistralOcrParams, MistralOcrResponse};
use crate::ocr::document::{inline_remote_document, validate_inline_document};
use crate::ocr::error::{OcrError, OcrRequestError, OcrResponseError};
use crate::ocr::prepare::{
_prepare_ocr_request, ParsedProviderParams, credential_env, transform_request_body,
};
use crate::ocr::registry::OcrProvider;
use crate::ocr::types::{LiteLLMOcrRequest, LiteLLMOcrResponse};
use crate::url_utils::ApiUrl;
use litellm_auth_gcp::{self as vertex, VertexConfig};
const DEFAULT_LOCATION: &str = "us-central1";
#[derive(Clone, Debug)]
pub(crate) struct VertexMistralAdapter;
impl OcrAdapter for VertexMistralAdapter {
type ProviderResponse = MistralOcrResponse;
const PROVIDER: OcrProvider = OcrProvider::VertexAi;
async fn prepare_request(
&self,
request: &LiteLLMOcrRequest,
client: &OcrClient,
) -> Result<reqwest::Request, OcrError> {
validate_destination(&request.connection)?;
let ParsedProviderParams {
known: params,
extra_params: _extra_params,
} = _prepare_ocr_request::<MistralOcrParams>(request)?;
let config = VertexConfig::from_sourced_optional_params(
&request.optional_params,
&request.input_sources,
)
.map_err(Error::from)?;
let authentication = client
.vertex_auth()
.validate_environment(
request.connection.extra_headers.clone(),
request.connection.api_key.as_deref(),
&config,
&credential_env,
)
.await
.map_err(Error::from)?;
let location = vertex::get_vertex_ai_location(&config, &credential_env)
.unwrap_or_else(|| DEFAULT_LOCATION.to_string());
let url = get_complete_url(
request.connection.api_base.as_deref(),
&authentication.project_id,
&location,
&request.model,
)?;
let retains_document = !request.document.source().starts_with("http://")
&& !request.document.source().starts_with("https://");
let document = inline_remote_document(
client.document_fetcher(),
request.document.clone(),
&request.connection,
)
.await?;
let body = mistral::transform_ocr_request(&request.model, document, &params)?;
transform_request_body(
client,
request,
&url,
&authentication.headers,
retains_document,
body,
|body| validate_inline_document(&body.document),
)
.await
}
fn transform_ocr_response(
&self,
request: &LiteLLMOcrRequest,
response: Self::ProviderResponse,
) -> Result<LiteLLMOcrResponse, OcrResponseError> {
mistral::transform_ocr_response(&request.model, response)
}
}
fn get_complete_url(
api_base: Option<&str>,
project: &str,
location: &str,
model: &str,
) -> Result<String, OcrError> {
validate_location(location)?;
let default_base = format!("https://{location}-aiplatform.googleapis.com");
let base = api_base
.map(str::trim)
.filter(|base| !base.is_empty())
.unwrap_or(&default_base);
let prediction = format!("{model}:rawPredict");
ApiUrl::parse(base)
.and_then(|url| {
url.complete_path(&[
"v1",
"projects",
project,
"locations",
location,
"publishers",
"mistralai",
"models",
&prediction,
])
})
.map(|url| url.into_string())
.map_err(|_| {
OcrRequestError::RequestField {
path: "api_base".into(),
}
.into()
})
}
fn validate_location(location: &str) -> Result<(), OcrError> {
let valid = !location.is_empty()
&& location
.bytes()
.all(|value| value.is_ascii_lowercase() || value.is_ascii_digit() || value == b'-')
&& location
.as_bytes()
.first()
.is_some_and(u8::is_ascii_alphanumeric)
&& location
.as_bytes()
.last()
.is_some_and(u8::is_ascii_alphanumeric);
if valid {
return Ok(());
}
Err(OcrRequestError::RequestField {
path: "vertex_location".into(),
}
.into())
}
#[cfg(test)]
mod tests {
use super::get_complete_url;
#[test]
fn endpoint_uses_location_project_and_model() {
assert_eq!(
get_complete_url(None, "proj-1", "europe-west4", "mistral-ocr-maas").unwrap(),
"https://europe-west4-aiplatform.googleapis.com/v1/projects/proj-1/locations/europe-west4/publishers/mistralai/models/mistral-ocr-maas:rawPredict"
);
assert!(get_complete_url(None, "proj-1", "attacker.example/path", "model").is_err());
}
}

View file

@ -1,18 +0,0 @@
mod deepseek;
mod mistral;
use crate::ocr::Error;
use litellm_auth::InputSource;
use crate::ocr::error::OcrError;
use crate::ocr::types::OcrConnection;
pub(crate) use deepseek::VertexDeepSeekAdapter;
pub(crate) use mistral::VertexMistralAdapter;
fn validate_destination(connection: &OcrConnection) -> Result<(), OcrError> {
if connection.api_base.is_some() && connection.api_base_source == InputSource::Request {
return Err(Error::from(litellm_auth::Error::RequestVertexCredentialDestination).into());
}
Ok(())
}

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,12 +4,10 @@ use std::time::Duration;
use bytes::{Bytes, BytesMut};
use serde::de::DeserializeOwned;
use super::error::{Error, OcrError, OcrResponseError};
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 crate::transport::Error as TransportError;
use litellm_auth_gcp::VertexAuth;
#[derive(Clone)]
@ -21,8 +19,8 @@ pub struct OcrClient {
}
impl OcrClient {
pub fn new(provider_http: reqwest::Client) -> Result<Self, TransportError> {
let document_fetcher = MediaFetcher::new().map_err(TransportError::from)?;
pub fn new(provider_http: reqwest::Client) -> Result<Self, crate::transport::Error> {
let document_fetcher = MediaFetcher::new().map_err(crate::transport::Error::from)?;
Ok(Self {
provider_http,
polling_http: no_redirect_http()?,
@ -31,11 +29,14 @@ impl OcrClient {
})
}
pub fn shared() -> Result<Self, Error> {
pub fn shared() -> Result<Self, crate::ocr::Error> {
shared_client()
}
pub async fn perform(&self, request: LiteLLMOcrRequest) -> Result<LiteLLMOcrResponse, Error> {
pub async fn perform(
&self,
request: LiteLLMOcrRequest,
) -> Result<LiteLLMOcrResponse, crate::ocr::Error> {
use super::{
NativeOutcome, OcrAdmission, OcrCall, OcrCallStep, OcrHookHost, OcrHost,
OcrHostOperation, OcrHostResult,
@ -45,7 +46,7 @@ impl OcrClient {
let mut request = Some(request);
let NativeOutcome::Completed(mut call) = OcrCall::admit(self.clone(), OcrAdmission::all())
else {
return Err(Error::InvalidRequest(
return Err(crate::ocr::Error::InvalidRequest(
"native OCR host admission declined".into(),
));
};
@ -54,16 +55,11 @@ impl OcrClient {
match call.resume(result.take()).await? {
OcrCallStep::Host(OcrHostOperation::ProjectRequest) => {
result = Some(OcrHostResult::Request(Ok((
Box::new(
request
.take()
.ok_or_else(|| {
Error::InvalidRequest(
"OCR request was already projected".into(),
)
})?
.into(),
),
Box::new(request.take().ok_or_else(|| {
crate::ocr::Error::InvalidRequest(
"OCR request was already projected".into(),
)
})?),
false,
))))
}
@ -100,29 +96,29 @@ impl OcrClient {
}
}
fn no_redirect_http() -> Result<reqwest::Client, TransportError> {
fn no_redirect_http() -> Result<reqwest::Client, crate::transport::Error> {
reqwest::Client::builder()
.connect_timeout(Duration::from_secs(OCR_CONNECT_TIMEOUT_SECS))
.redirect(reqwest::redirect::Policy::none())
.build()
.map_err(TransportError::from)
.map_err(crate::transport::Error::from)
}
pub(crate) fn shared_client() -> Result<OcrClient, Error> {
static CLIENT: OnceLock<Result<OcrClient, TransportError>> = OnceLock::new();
pub(crate) fn shared_client() -> Result<OcrClient, crate::ocr::Error> {
static CLIENT: OnceLock<Result<OcrClient, crate::transport::Error>> = OnceLock::new();
let client = CLIENT
.get_or_init(|| {
reqwest::Client::builder()
.connect_timeout(Duration::from_secs(OCR_CONNECT_TIMEOUT_SECS))
.build()
.map_err(TransportError::from)
.map_err(crate::transport::Error::from)
.and_then(OcrClient::new)
})
.clone()?;
Ok(client)
}
pub async fn ocr(request: LiteLLMOcrRequest) -> Result<LiteLLMOcrResponse, Error> {
pub async fn ocr(request: LiteLLMOcrRequest) -> Result<LiteLLMOcrResponse, crate::ocr::Error> {
shared_client()?.perform(request).await
}
@ -130,15 +126,15 @@ pub async fn read_json_response<T: DeserializeOwned>(
response: reqwest::Response,
native: bool,
max_response_bytes: usize,
) -> Result<DecodedOcrResponse<T>, OcrError> {
) -> Result<DecodedOcrResponse<T>, crate::ocr::Error> {
let bytes = read_response_bytes(response, max_response_bytes).await?;
Ok(decode_response(&bytes, native)?)
decode_response(&bytes, native)
}
pub(crate) async fn read_response_bytes(
mut response: reqwest::Response,
max_response_bytes: usize,
) -> Result<Bytes, OcrError> {
) -> Result<Bytes, crate::ocr::Error> {
let status = response.status();
let limit = if status.is_success() {
max_response_bytes
@ -150,13 +146,13 @@ pub(crate) async fn read_response_bytes(
.content_length()
.is_some_and(|length| length > limit as u64)
{
return Err(OcrResponseError::TooLarge { limit }.into());
return Err(crate::ocr::Error::TooLarge { limit });
}
let mut bytes = BytesMut::new();
while let Some(chunk) = response.chunk().await.map_err(transport_error)? {
let remaining = limit.saturating_sub(bytes.len());
if status.is_success() && chunk.len() > remaining {
return Err(OcrResponseError::TooLarge { limit }.into());
return Err(crate::ocr::Error::TooLarge { limit });
}
bytes.extend_from_slice(&chunk[..chunk.len().min(remaining)]);
if !status.is_success() && bytes.len() == limit {
@ -173,12 +169,12 @@ pub(crate) async fn read_response_bytes(
Ok(bytes.freeze())
}
pub(crate) fn transport_error(error: reqwest::Error) -> Error {
pub(crate) fn transport_error(error: reqwest::Error) -> crate::ocr::Error {
if error.is_timeout() {
return Error::Http {
return crate::ocr::Error::Transport(crate::transport::Error::Http {
status: 408,
body: "OCR request timed out".into(),
};
});
}
crate::transport::Error::from(error).into()
}
@ -203,7 +199,7 @@ mod tests {
.unwrap_err();
assert!(matches!(
transport_error(error),
Error::Http { status: 408, .. }
crate::ocr::Error::Transport(crate::transport::Error::Http { status: 408, .. })
));
server.abort();
}

View file

@ -1,254 +0,0 @@
use serde::{Deserialize, Serialize};
use serde_json::{Map, Value, json};
use crate::ocr::document::InlineDocument;
use crate::ocr::error::{OcrRequestError, OcrResponseError};
use crate::ocr::types::{LiteLLMOcrResponse, OcrDocument};
#[derive(Clone, Copy, Debug, Default, Deserialize, Serialize)]
#[serde(rename_all = "lowercase")]
pub(crate) enum OutputFormat {
#[default]
Markdown,
Blocks,
}
#[derive(Deserialize)]
pub(crate) struct CohereParams {
#[serde(default)]
pub output_format: OutputFormat,
}
#[derive(Deserialize, Serialize)]
pub(crate) struct CohereRequest {
pub model: String,
pub document: OcrDocument,
pub output_format: OutputFormat,
}
pub(crate) fn validate_document(document: &OcrDocument) -> Result<(), OcrRequestError> {
let OcrDocument::ImageUrl { image_url, .. } = document else {
return Err(OcrRequestError::CohereImageOnly);
};
if image_url.is_empty() {
return Err(OcrRequestError::CohereImageOnly);
}
if let Some(inline) = InlineDocument::parse(image_url)? {
if !inline.mime_type().type_.eq_ignore_ascii_case("image") {
return Err(OcrRequestError::CohereImageOnly);
}
inline.decode(crate::constants::OCR_INLINE_MAX_BYTES)?;
}
Ok(())
}
#[derive(Deserialize)]
pub(crate) struct CohereResponse {
#[serde(default)]
pages: Vec<CoherePage>,
meta: Option<CohereMeta>,
}
#[derive(Deserialize)]
struct CoherePage {
index: Option<i64>,
markdown: Option<CohereMarkdown>,
blocks: Option<Vec<Map<String, Value>>>,
}
#[derive(Deserialize)]
struct CohereMarkdown {
#[serde(default)]
content: String,
images: Option<Vec<Map<String, Value>>>,
}
#[derive(Deserialize)]
struct CohereMeta {
billed_units: Option<CohereBilledUnits>,
}
#[derive(Deserialize)]
struct CohereBilledUnits {
pages: Option<i64>,
}
pub(crate) fn transform_response(
model: &str,
response: CohereResponse,
) -> Result<LiteLLMOcrResponse, OcrResponseError> {
let pages_processed = response
.meta
.and_then(|meta| meta.billed_units)
.and_then(|units| units.pages)
.map(Ok)
.unwrap_or_else(|| {
i64::try_from(response.pages.len()).map_err(|_| OcrResponseError::NumericRange("pages"))
})?;
let pages = response
.pages
.into_iter()
.enumerate()
.map(|(position, page)| {
let index = page.index.map(Ok).unwrap_or_else(|| {
i64::try_from(position).map_err(|_| OcrResponseError::NumericRange("page index"))
})?;
let (content, images) = page
.markdown
.map(|markdown| {
let images =
markdown
.images
.filter(|images| !images.is_empty())
.map(|images| {
images
.into_iter()
.map(|mut image| {
if let Some(Value::Object(bbox)) =
image.get("bounding_box").cloned()
{
image.insert("bbox".into(), Value::Object(bbox));
}
Value::Object(image)
})
.collect::<Vec<_>>()
});
(markdown.content, images)
})
.unwrap_or_default();
let mut normalized = json!({"index": index, "markdown": content, "images": images});
if let Some(blocks) = page.blocks {
normalized["blocks"] = json!(blocks);
}
Ok(normalized)
})
.collect::<Result<Vec<_>, OcrResponseError>>()?;
Ok(LiteLLMOcrResponse {
pages,
model: model.into(),
document_annotation: None,
usage_info: Some(json!({"pages_processed": pages_processed})),
object: "ocr".into(),
extra_fields: Map::new(),
provider_native_response: None,
})
}
pub(crate) fn transform_request(
model: &str,
document: OcrDocument,
params: CohereParams,
) -> Result<CohereRequest, OcrRequestError> {
validate_document(&document)?;
Ok(CohereRequest {
model: model.into(),
document,
output_format: params.output_format,
})
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn response_normalizes_markdown_images_blocks_and_billed_pages() {
let response = serde_json::from_value(json!({
"pages": [
{
"type":"markdown",
"index":4,
"markdown":{
"content":"receipt",
"images":[{
"id":"image",
"bounding_box":{"top_left_x":1,"bottom_right_x":48},
"bounding_box_normalized":{"top_left_x":0.04,"bottom_right_x":0.15},
"description":"scan",
"category":"logo"
}]
}
},
{"type":"blocks","blocks":[{"type":"text","text":{"content":"total"}}]}
],
"meta":{"api_version":{"version":"2"},"billed_units":{"pages":3}}
}))
.unwrap();
let normalized = transform_response("parse-v5.0", response).unwrap();
assert_eq!(normalized.pages[0]["index"], 4);
assert_eq!(normalized.pages[0]["markdown"], "receipt");
assert_eq!(normalized.pages[0]["images"][0]["bbox"]["top_left_x"], 1);
assert_eq!(
normalized.pages[0]["images"][0]["bounding_box_normalized"]["bottom_right_x"],
0.15
);
assert_eq!(normalized.pages[0]["images"][0]["description"], "scan");
assert_eq!(normalized.pages[0]["images"][0]["category"], "logo");
assert_eq!(normalized.pages[1]["index"], 1);
assert_eq!(normalized.pages[1]["markdown"], "");
assert_eq!(normalized.pages[1]["blocks"][0]["text"]["content"], "total");
assert_eq!(normalized.usage_info.unwrap()["pages_processed"], 3);
}
#[test]
fn response_defaults_and_invalid_fields() {
for value in [
json!({}),
json!({"meta":null}),
json!({"pages":[],"meta":{"billed_units":null}}),
] {
let normalized =
transform_response("parse", serde_json::from_value(value).unwrap()).unwrap();
assert!(normalized.pages.is_empty());
assert_eq!(normalized.usage_info.unwrap()["pages_processed"], 0);
}
for value in [
json!({"pages":null}),
json!({"pages":[{"markdown":"text"}]}),
json!({"pages":[{"index":"bad"}]}),
] {
assert!(serde_json::from_value::<CohereResponse>(value).is_err());
}
let normalized = transform_response(
"parse",
serde_json::from_value(json!({"pages":[{"markdown":null}]})).unwrap(),
)
.unwrap();
assert_eq!(normalized.usage_info.unwrap()["pages_processed"], 1);
assert!(normalized.pages[0]["images"].is_null());
}
#[test]
fn request_requires_image_and_supported_output_format() {
for value in [
json!({"type":"document_url","document_url":"https://example.com/a.pdf"}),
json!({"type":"image_url","image_url":""}),
json!({"type":"image_url","image_url":"data:application/pdf;base64,YQ=="}),
] {
assert_eq!(
validate_document(&serde_json::from_value(value).unwrap()),
Err(OcrRequestError::CohereImageOnly)
);
}
assert!(serde_json::from_value::<CohereParams>(json!({"output_format":"html"})).is_err());
for format in ["markdown", "blocks"] {
assert!(
serde_json::from_value::<CohereParams>(json!({"output_format":format})).is_ok()
);
}
let request = transform_request(
"parse-v5.0",
serde_json::from_value(json!({
"type":"image_url",
"image_url":"https://example.com/image.png"
}))
.unwrap(),
serde_json::from_value(json!({})).unwrap(),
)
.unwrap();
assert_eq!(
serde_json::to_value(request).unwrap()["output_format"],
"markdown"
);
}
}

View file

@ -1,5 +0,0 @@
mod transformation;
mod types;
pub(crate) use transformation::{transform_ocr_request, transform_ocr_response};
pub(crate) use types::{DeepSeekOcrParams, DeepSeekOcrResponse};

View file

@ -1,101 +0,0 @@
use serde::de::IntoDeserializer;
use serde_json::{Value, json};
use super::types::*;
use crate::ocr::error::{OcrRequestError, OcrResponseError};
use crate::ocr::types::{LiteLLMOcrResponse, OcrDocument};
pub(crate) fn transform_ocr_request(
provider_model: &str,
document: OcrDocument,
params: &DeepSeekOcrParams,
) -> Result<DeepSeekOcrRequest, OcrRequestError> {
if document.source().is_empty() {
return Err(OcrRequestError::MissingDocumentUrl);
}
let content = OcrDocument::ImageUrl {
image_url: document.source().to_string(),
extra_fields: serde_json::Map::new(),
};
Ok(DeepSeekOcrRequest {
model: provider_model.to_string(),
messages: vec![DeepSeekOcrMessage {
role: UserRole::User,
content: vec![content],
}],
params: params.clone(),
})
}
pub(crate) fn transform_ocr_response(
model: &str,
response: DeepSeekOcrResponse,
) -> Result<LiteLLMOcrResponse, OcrResponseError> {
let content = response
.choices
.into_iter()
.next()
.and_then(|choice| choice.message.content)
.ok_or(OcrResponseError::EmptyContent)?;
let decoded = decode_content(content)?;
let pages = match decoded.result.pages {
Some(pages) if !pages.is_empty() => pages
.into_iter()
.map(|page| serde_json::to_value(page).expect("DeepSeek page serializes"))
.collect(),
_ => vec![json!({
"index":0,
"markdown":decoded.fallback_markdown,
"images":null
})],
};
Ok(LiteLLMOcrResponse {
pages,
model: decoded.result.model.unwrap_or_else(|| model.to_string()),
document_annotation: decoded.result.document_annotation,
usage_info: decoded.result.usage_info.or(response.usage),
object: "ocr".into(),
extra_fields: decoded.result.extra_fields,
provider_native_response: None,
})
}
struct DecodedContent {
result: DeepSeekOcrResult,
fallback_markdown: String,
}
fn decode_content(content: DeepSeekContent) -> Result<DecodedContent, OcrResponseError> {
let (result, fallback_markdown) = match content {
DeepSeekContent::Text(text) if text.is_empty() => {
return Err(OcrResponseError::EmptyContent);
}
DeepSeekContent::Text(text) => (decode_json_content(&text)?, text),
DeepSeekContent::Object(object) => {
let fallback =
serde_json::to_string(&object).map_err(|_| OcrResponseError::ResponseField {
path: "choices[0].message.content".into(),
})?;
(Some(object), fallback)
}
};
Ok(DecodedContent {
result: result.unwrap_or_default(),
fallback_markdown,
})
}
fn decode_json_content(text: &str) -> Result<Option<DeepSeekOcrResult>, OcrResponseError> {
if !text.trim_start().starts_with('{') {
return Ok(None);
}
let value = match serde_json::from_str::<Value>(text) {
Ok(value) => value,
Err(_) => return Ok(None),
};
serde_path_to_error::deserialize(value.into_deserializer())
.map(Some)
.map_err(|error| OcrResponseError::ResponseField {
path: format!("choices[0].message.content.{}", error.path()),
})
}

View file

@ -1,95 +0,0 @@
use serde::{Deserialize, Serialize};
use serde_json::{Map, Value};
#[derive(Clone, Debug, Default, Serialize, Deserialize)]
pub(crate) struct DeepSeekOcrParams {
#[serde(skip_serializing_if = "Option::is_none")]
pub stream: Option<bool>,
#[serde(skip_serializing_if = "Option::is_none")]
pub temperature: Option<f64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub max_tokens: Option<i64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub top_p: Option<f64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub n: Option<i64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub stop: Option<StopSequences>,
}
#[derive(Clone, Debug, Serialize, Deserialize)]
#[serde(untagged)]
pub(crate) enum StopSequences {
One(String),
Many(Vec<String>),
}
#[derive(Clone, Debug, Serialize, Deserialize)]
pub(crate) struct DeepSeekOcrRequest {
pub model: String,
pub messages: Vec<DeepSeekOcrMessage>,
#[serde(flatten)]
pub params: DeepSeekOcrParams,
}
#[derive(Clone, Debug, Serialize, Deserialize)]
pub(crate) struct DeepSeekOcrMessage {
pub role: UserRole,
pub content: Vec<crate::ocr::types::OcrDocument>,
}
#[derive(Clone, Debug, Serialize, Deserialize)]
#[serde(rename_all = "lowercase")]
pub(crate) enum UserRole {
User,
}
#[derive(Clone, Debug, Deserialize)]
pub(crate) struct DeepSeekOcrResponse {
#[serde(default)]
pub choices: Vec<DeepSeekChoice>,
pub usage: Option<Value>,
}
#[derive(Clone, Debug, Deserialize)]
pub(crate) struct DeepSeekChoice {
pub message: DeepSeekResponseMessage,
}
#[derive(Clone, Debug, Deserialize)]
pub(crate) struct DeepSeekResponseMessage {
pub content: Option<DeepSeekContent>,
}
#[derive(Clone, Debug, Deserialize)]
#[serde(untagged)]
pub(crate) enum DeepSeekContent {
Text(String),
Object(DeepSeekOcrResult),
}
#[derive(Clone, Debug, Default, Serialize, Deserialize)]
pub(crate) struct DeepSeekOcrResult {
#[serde(skip_serializing_if = "Option::is_none")]
pub pages: Option<Vec<DeepSeekPage>>,
#[serde(skip_serializing_if = "Option::is_none")]
pub model: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub usage_info: Option<Value>,
#[serde(skip_serializing_if = "Option::is_none")]
pub document_annotation: Option<Value>,
#[serde(flatten)]
pub extra_fields: Map<String, Value>,
}
#[derive(Clone, Debug, Serialize, Deserialize)]
pub(crate) struct DeepSeekPage {
#[serde(default)]
pub index: i64,
#[serde(default)]
pub markdown: String,
pub images: Option<Value>,
pub dimensions: Option<Value>,
#[serde(flatten)]
pub extra_fields: Map<String, Value>,
}

View file

@ -1,9 +0,0 @@
mod params;
mod transformation;
mod types;
pub(crate) use params::{decode_input_params, map_ocr_params};
pub(crate) use transformation::{transform_ocr_request, transform_ocr_response};
pub(crate) use types::{
AzureDocumentIntelligenceOperation, DocumentIntelligenceParams, OperationStatus,
};

View file

@ -1,219 +0,0 @@
use std::collections::BTreeSet;
use serde_json::{Map, Value};
use super::types::{
DocumentIntelligenceInputParams, DocumentIntelligenceParams, FeaturesInput, PagesInput,
};
use crate::ocr::error::OcrRequestError;
use crate::ocr::prepare::ParsedProviderParams;
pub(crate) fn decode_input_params(
params: Map<String, Value>,
prefix: &str,
) -> Result<ParsedProviderParams<DocumentIntelligenceInputParams>, OcrRequestError> {
if let Some(Value::Array(pages)) = params.get("pages") {
if pages.iter().any(Value::is_boolean) {
return Err(OcrRequestError::Pages("boolean page index".into()));
}
if pages
.iter()
.any(|page| page.is_number() && page.as_i64().is_none())
{
return Err(OcrRequestError::Pages("page index is out of range".into()));
}
if !pages.iter().all(Value::is_i64) && !pages.iter().all(Value::is_string) {
return Err(OcrRequestError::Pages("mixed page element types".into()));
}
}
crate::ocr::wire::decode_request_value(Value::Object(params), prefix)
}
pub(crate) fn map_ocr_params(
params: DocumentIntelligenceInputParams,
) -> Result<DocumentIntelligenceParams, OcrRequestError> {
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>, OcrRequestError> {
let normalized = match pages {
PagesInput::ZeroBasedIndices(indices) => {
if indices.is_empty() {
return Ok(None);
}
indices
.into_iter()
.map(|page| {
if page < 0 {
return Err(OcrRequestError::Pages("negative page index".into()));
}
page.checked_add(1)
.ok_or_else(|| OcrRequestError::Pages("page index is out of range".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
.split(',')
.map(str::trim)
.collect::<Vec<_>>()
.join(","),
};
if !normalized.split(',').all(valid_page_token) {
return Err(OcrRequestError::Pages("invalid native page range".into()));
}
Ok(Some(normalized))
}
fn valid_page_token(token: &str) -> bool {
let mut parts = token.split('-');
let start = parts.next().unwrap_or_default();
if start.is_empty() || !start.chars().all(|character| character.is_ascii_digit()) {
return false;
}
match parts.next() {
None => true,
Some(end) => {
!end.is_empty()
&& end.chars().all(|character| character.is_ascii_digit())
&& parts.next().is_none()
}
}
}
fn normalize_features(features: FeaturesInput) -> Result<Option<String>, OcrRequestError> {
let tokens = match features {
FeaturesInput::Names(names) => names,
FeaturesInput::CommaSeparated(names) => names.split(',').map(str::to_string).collect(),
};
if tokens.is_empty() {
return Ok(None);
}
let normalized = tokens.iter().map(|token| token.trim()).collect::<Vec<_>>();
if !normalized.iter().all(|token| {
let Some((first, rest)) = token.as_bytes().split_first() else {
return false;
};
first.is_ascii_alphabetic() && rest.iter().all(u8::is_ascii_alphanumeric)
}) {
return Err(OcrRequestError::Features);
}
Ok(Some(normalized.join(",")))
}
#[cfg(test)]
mod tests {
use rstest::rstest;
use serde_json::{Value, json};
use super::*;
fn map(value: Value) -> Result<DocumentIntelligenceParams, OcrRequestError> {
let fields = value.as_object().unwrap().clone();
map_ocr_params(decode_input_params(fields, "optional_params")?.known)
}
#[test]
fn input_params_retain_unknown_fields() {
let parsed = decode_input_params(
json!({
"pages": [0],
"future_ocr_option": true,
"extra_body": {"provider_option": "value"}
})
.as_object()
.unwrap()
.clone(),
"optional_params",
)
.unwrap();
assert_eq!(
parsed.known.pages,
Some(PagesInput::ZeroBasedIndices(vec![0]))
);
assert_eq!(parsed.extra_params["future_ocr_option"], true);
assert_eq!(
parsed.extra_params["extra_body"],
json!({"provider_option": "value"})
);
assert_eq!(
serde_json::to_value(map_ocr_params(parsed.known).unwrap()).unwrap(),
json!({"pages": "1", "features": null})
);
}
#[rstest]
#[case(json!([0, 1, 2]), Some("1,2,3"))]
#[case(json!([2, 0, 0, 1]), Some("1,2,3"))]
#[case(json!([]), None)]
#[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"))]
fn page_mapping_matches_python(#[case] input: Value, #[case] expected: Option<&str>) {
assert_eq!(
map(json!({"pages": input})).unwrap().pages.as_deref(),
expected
);
}
#[rstest]
#[case(json!("a,b"))]
#[case(json!([-1]))]
#[case(json!([true, false]))]
#[case(json!([1, "2"]))]
#[case(json!(5))]
fn invalid_page_mapping_matches_python(#[case] input: Value) {
assert!(map(json!({"pages": input})).is_err());
}
#[rstest]
#[case(json!(["keyValuePairs"]), "keyValuePairs")]
#[case(json!(["keyValuePairs", "languages"]), "keyValuePairs,languages")]
#[case(json!("keyValuePairs"), "keyValuePairs")]
#[case(json!("keyValuePairs,languages"), "keyValuePairs,languages")]
#[case(json!("keyValuePairs, languages"), "keyValuePairs,languages")]
fn feature_mapping_matches_python(#[case] input: Value, #[case] expected: &str) {
assert_eq!(
map(json!({"features": input})).unwrap().features.as_deref(),
Some(expected)
);
}
#[rstest]
#[case(json!("keyValuePairs&pages=9"))]
#[case(json!("key value pairs"))]
#[case(json!(""))]
#[case(json!([1, 2]))]
#[case(json!([["keyValuePairs"]]))]
#[case(json!({"feature":"keyValuePairs"}))]
#[case(json!(5))]
fn invalid_feature_mapping_matches_python(#[case] input: Value) {
assert!(map(json!({"features": input})).is_err());
}
#[test]
fn empty_feature_list_is_omitted() {
assert_eq!(map(json!({"features": []})).unwrap().features, None);
}
}

View file

@ -1,107 +0,0 @@
use base64::{Engine, engine::general_purpose::STANDARD};
use serde_json::{Map, Value, json};
use super::types::*;
use crate::constants::{AZURE_DI_DEFAULT_DPI, AZURE_DI_DEFAULT_HEIGHT, AZURE_DI_DEFAULT_WIDTH};
use crate::ocr::document::InlineDocument;
use crate::ocr::error::{OcrRequestError, OcrResponseError};
use crate::ocr::types::{LiteLLMOcrResponse, OcrDocument};
pub(crate) fn transform_ocr_request(
document: OcrDocument,
) -> Result<DocumentIntelligenceRequest, OcrRequestError> {
let source = document.source();
if source.is_empty() {
return Err(OcrRequestError::MissingDocumentUrl);
}
Ok(if let Some(document) = InlineDocument::parse(source)? {
DocumentIntelligenceRequest::Base64Source(
STANDARD.encode(document.decode(crate::constants::OCR_INLINE_MAX_BYTES)?),
)
} else {
DocumentIntelligenceRequest::UrlSource(source.to_string())
})
}
pub(crate) fn transform_ocr_response(
model: &str,
response: AzureDocumentIntelligenceOperation,
) -> Result<LiteLLMOcrResponse, OcrResponseError> {
if response.status != Some(OperationStatus::Succeeded) {
return Err(OcrResponseError::OperationStatus(
response
.status
.map(|status| status.to_string())
.unwrap_or_else(|| "None".into()),
));
}
let result = response.analyze_result.unwrap_or_default();
let pages = result
.pages
.into_iter()
.map(normalize_page)
.collect::<Result<Vec<_>, _>>()?;
let pages_processed = pages.len();
let mut extra_fields = Map::new();
extra_fields.insert("content".into(), option_value(result.content));
extra_fields.insert("tables".into(), option_value(result.tables));
extra_fields.insert("keyValuePairs".into(), option_value(result.key_value_pairs));
Ok(LiteLLMOcrResponse {
pages,
model: model.into(),
document_annotation: None,
usage_info: Some(json!({"pages_processed":pages_processed})),
object: "ocr".into(),
extra_fields,
provider_native_response: None,
})
}
fn normalize_page(page: AzureDocumentIntelligencePage) -> Result<Value, OcrResponseError> {
let index = page
.page_number
.unwrap_or(1)
.checked_sub(1)
.ok_or(OcrResponseError::NumericRange("page.pageNumber"))?;
let scale = if page.unit.as_deref().unwrap_or("inch") == "inch" {
AZURE_DI_DEFAULT_DPI as f64
} else {
1.0
};
let width = pixel_dimension(
page.width.unwrap_or(AZURE_DI_DEFAULT_WIDTH),
scale,
"page.width",
)?;
let height = pixel_dimension(
page.height.unwrap_or(AZURE_DI_DEFAULT_HEIGHT),
scale,
"page.height",
)?;
let markdown = page
.lines
.iter()
.map(|line| line.content.as_deref().unwrap_or_default())
.collect::<Vec<_>>()
.join("\n");
Ok(json!({
"index":index,
"markdown":markdown,
"images":null,
"dimensions":{"width":width,"height":height,"dpi":AZURE_DI_DEFAULT_DPI}
}))
}
fn pixel_dimension(value: f64, scale: f64, field: &'static str) -> Result<i64, OcrResponseError> {
let value = value * scale;
if !value.is_finite() || value < i64::MIN as f64 || value > i64::MAX as f64 {
return Err(OcrResponseError::NumericRange(field));
}
Ok(value.trunc() as i64)
}
fn option_value<T: serde::Serialize>(value: Option<T>) -> Value {
value
.and_then(|value| serde_json::to_value(value).ok())
.unwrap_or(Value::Null)
}

View file

@ -1,138 +0,0 @@
use serde::{Deserialize, Deserializer, Serialize};
use serde_json::{Map, Value};
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
#[serde(untagged)]
pub(crate) enum PagesInput {
ZeroBasedIndices(Vec<i64>),
NativeTokens(Vec<String>),
NativeRange(String),
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
#[serde(untagged)]
pub(crate) enum FeaturesInput {
Names(Vec<String>),
CommaSeparated(String),
}
#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)]
pub(crate) struct DocumentIntelligenceInputParams {
pub pages: Option<PagesInput>,
pub features: Option<FeaturesInput>,
}
#[derive(Clone, Debug, PartialEq, Serialize)]
pub(crate) struct DocumentIntelligenceParams {
pub pages: Option<String>,
pub features: Option<String>,
}
#[derive(Clone, Debug, Serialize, Deserialize)]
pub(crate) enum DocumentIntelligenceRequest {
#[serde(rename = "urlSource")]
UrlSource(String),
#[serde(rename = "base64Source")]
Base64Source(String),
}
#[derive(Clone, Debug, PartialEq)]
pub(crate) enum OperationStatus {
Succeeded,
Running,
NotStarted,
Failed,
Unknown(String),
}
impl<'de> Deserialize<'de> for OperationStatus {
fn deserialize<D: Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
Ok(match String::deserialize(deserializer)?.as_str() {
"succeeded" => Self::Succeeded,
"running" => Self::Running,
"notStarted" => Self::NotStarted,
"failed" => Self::Failed,
value => Self::Unknown(value.to_string()),
})
}
}
impl std::fmt::Display for OperationStatus {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter.write_str(match self {
Self::Succeeded => "succeeded",
Self::Running => "running",
Self::NotStarted => "notStarted",
Self::Failed => "failed",
Self::Unknown(value) => value,
})
}
}
#[derive(Clone, Debug, Deserialize)]
pub(crate) struct AzureDocumentIntelligenceOperation {
pub status: Option<OperationStatus>,
#[serde(rename = "analyzeResult")]
pub analyze_result: Option<AzureDocumentIntelligenceAnalyzeResult>,
}
#[derive(Clone, Debug, Default, Deserialize)]
pub(crate) struct AzureDocumentIntelligenceAnalyzeResult {
pub content: Option<String>,
#[serde(default)]
pub pages: Vec<AzureDocumentIntelligencePage>,
pub tables: Option<Vec<Map<String, Value>>>,
#[serde(rename = "keyValuePairs")]
pub key_value_pairs: Option<Vec<Map<String, Value>>>,
}
#[derive(Clone, Debug, Deserialize)]
pub(crate) struct AzureDocumentIntelligencePage {
#[serde(rename = "pageNumber", default, deserialize_with = "optional_i64")]
pub page_number: Option<i64>,
#[serde(default, deserialize_with = "optional_f64")]
pub width: Option<f64>,
#[serde(default, deserialize_with = "optional_f64")]
pub height: Option<f64>,
pub unit: Option<String>,
#[serde(default)]
pub lines: Vec<AzureDocumentIntelligenceLine>,
}
#[derive(Clone, Debug, Deserialize)]
pub(crate) struct AzureDocumentIntelligenceLine {
pub content: Option<String>,
}
fn optional_i64<'de, D: Deserializer<'de>>(deserializer: D) -> Result<Option<i64>, D::Error> {
match Option::<Value>::deserialize(deserializer)? {
None | Some(Value::Null) => Ok(None),
Some(Value::Number(number)) => number
.as_i64()
.map(Some)
.ok_or_else(|| serde::de::Error::custom("expected an integer")),
Some(Value::String(value)) => value
.parse::<i64>()
.map(Some)
.map_err(|_| serde::de::Error::custom("expected an integer")),
Some(_) => Err(serde::de::Error::custom("expected an integer")),
}
}
fn optional_f64<'de, D: Deserializer<'de>>(deserializer: D) -> Result<Option<f64>, D::Error> {
match Option::<Value>::deserialize(deserializer)? {
None | Some(Value::Null) => Ok(None),
Some(Value::Number(number)) => number
.as_f64()
.filter(|value| value.is_finite())
.map(Some)
.ok_or_else(|| serde::de::Error::custom("expected a finite number")),
Some(Value::String(value)) => value
.parse::<f64>()
.ok()
.filter(|value| value.is_finite())
.map(Some)
.ok_or_else(|| serde::de::Error::custom("expected a finite number")),
Some(_) => Err(serde::de::Error::custom("expected a number")),
}
}

View file

@ -1,5 +0,0 @@
mod transformation;
mod types;
pub(crate) use transformation::{transform_ocr_request, transform_ocr_response};
pub(crate) use types::{MistralOcrParams, MistralOcrRequest, MistralOcrResponse};

View file

@ -1,250 +0,0 @@
use super::{MistralOcrParams, MistralOcrRequest, MistralOcrResponse};
use crate::ocr::error::{OcrRequestError, OcrResponseError};
use crate::ocr::types::{LiteLLMOcrResponse, OcrDocument};
pub(crate) fn transform_ocr_request(
model: &str,
document: OcrDocument,
params: &MistralOcrParams,
) -> Result<MistralOcrRequest, OcrRequestError> {
Ok(MistralOcrRequest {
model: model.to_string(),
document,
params: params.clone(),
})
}
pub(crate) fn transform_ocr_response(
model: &str,
response: MistralOcrResponse,
) -> Result<LiteLLMOcrResponse, OcrResponseError> {
Ok(LiteLLMOcrResponse {
pages: response.pages,
model: response.model.unwrap_or_else(|| model.to_string()),
document_annotation: response.document_annotation,
usage_info: response.usage_info,
object: "ocr".to_string(),
extra_fields: response.extra_fields,
provider_native_response: None,
})
}
#[cfg(test)]
mod tests {
use super::*;
use rstest::rstest;
use serde_json::{Value, json};
fn mapped_params(value: Value) -> Value {
serde_json::to_value(serde_json::from_value::<MistralOcrParams>(value).unwrap()).unwrap()
}
fn document() -> OcrDocument {
serde_json::from_value(
json!({"type":"document_url","document_url":"https://example.com/a.pdf"}),
)
.unwrap()
}
#[rstest]
fn extract_header_is_a_supported_ocr_param() {
assert_eq!(
mapped_params(json!({"extract_header":true}))["extract_header"],
true
);
}
#[rstest]
fn extract_footer_is_a_supported_ocr_param() {
assert_eq!(
mapped_params(json!({"extract_footer":false}))["extract_footer"],
false
);
}
#[rstest]
fn existing_ocr_params_remain_supported() {
let mapped = mapped_params(json!({
"pages":[0,2],
"include_image_base64":true,
"image_limit":2,
"image_min_size":100,
"bbox_annotation_format":{"type":"json_schema"},
"document_annotation_format":{"type":"json_schema"}
}));
assert_eq!(mapped["pages"], json!([0, 2]));
assert_eq!(mapped["include_image_base64"], true);
assert_eq!(mapped["image_limit"], 2);
assert_eq!(mapped["image_min_size"], 100);
assert_eq!(mapped["bbox_annotation_format"]["type"], "json_schema");
assert_eq!(mapped["document_annotation_format"]["type"], "json_schema");
}
#[rstest]
fn map_ocr_params_forwards_extract_header() {
assert_eq!(
mapped_params(json!({"extract_header":true}))["extract_header"],
true
);
}
#[rstest]
fn map_ocr_params_forwards_extract_footer() {
assert_eq!(
mapped_params(json!({"extract_footer":true}))["extract_footer"],
true
);
}
#[rstest]
fn map_ocr_params_forwards_extract_header_and_footer() {
let mapped = mapped_params(json!({"extract_header":true,"extract_footer":false}));
assert_eq!(mapped["extract_header"], true);
assert_eq!(mapped["extract_footer"], false);
}
#[rstest]
fn map_ocr_params_drops_unknown_params() {
let mapped = mapped_params(json!({"extract_header":true,"unsupported_param":"value"}));
assert_eq!(mapped["extract_header"], true);
assert!(mapped.get("unsupported_param").is_none());
}
#[rstest]
#[case("table_format", json!("html"))]
#[case("confidence_scores_granularity", json!("word"))]
#[case("confidence_scores_granularity", json!("block"))]
#[case("document_annotation_prompt", json!("extract"))]
#[case("include_blocks", json!(true))]
#[case("id", json!("req-123"))]
fn new_ocr_params_are_supported(#[case] name: &str, #[case] value: Value) {
assert_eq!(mapped_params(json!({name:value.clone()}))[name], value);
}
#[rstest]
#[case("table_format", json!("html"))]
#[case("confidence_scores_granularity", json!("word"))]
#[case("document_annotation_prompt", json!("extract"))]
#[case("include_blocks", json!(true))]
#[case("id", json!("req-123"))]
fn map_ocr_params_forwards_new_ocr_params(#[case] name: &str, #[case] value: Value) {
assert_eq!(mapped_params(json!({name:value.clone()}))[name], value);
}
#[rstest]
#[case("pages", json!([0, 2]))]
#[case("pages", json!("0,2-4"))]
#[case("include_image_base64", json!(true))]
#[case("image_limit", json!(2))]
#[case("image_min_size", json!(100))]
#[case("bbox_annotation_format", json!({"type":"json_schema"}))]
#[case("document_annotation_format", json!({"type":"json_schema"}))]
#[case("document_annotation_prompt", json!("extract"))]
#[case("extract_header", json!(true))]
#[case("extract_footer", json!(false))]
#[case("table_format", json!("html"))]
#[case("confidence_scores_granularity", json!("word"))]
#[case("include_blocks", json!(true))]
#[case("id", json!("req-123"))]
fn request_mapping_matches_python(#[case] name: &str, #[case] value: Value) {
let params: MistralOcrParams =
serde_json::from_value(json!({name: value.clone()})).unwrap();
let result =
serde_json::to_value(transform_ocr_request("model", document(), &params).unwrap())
.unwrap();
assert_eq!(result["model"], "model");
assert_eq!(result[name], value);
}
#[rstest]
#[case("table_format", json!("html"))]
#[case("confidence_scores_granularity", json!("word"))]
#[case("document_annotation_prompt", json!("extract"))]
#[case("id", json!("req-123"))]
#[case("extract_header", json!(true))]
#[case("include_blocks", json!(true))]
#[case("pages", json!([0,1]))]
fn transform_ocr_request_includes_each_optional_param(
#[case] name: &str,
#[case] value: Value,
) {
let params: MistralOcrParams = serde_json::from_value(json!({name:value.clone()})).unwrap();
let result = serde_json::to_value(
transform_ocr_request("mistral-ocr-latest", document(), &params).unwrap(),
)
.unwrap();
assert_eq!(result[name], value);
assert_eq!(result["model"], "mistral-ocr-latest");
}
#[rstest]
fn transform_ocr_request_includes_multiple_new_params() {
let params: MistralOcrParams = serde_json::from_value(json!({
"table_format":"html",
"confidence_scores_granularity":"page",
"extract_header":true
}))
.unwrap();
let result = serde_json::to_value(
transform_ocr_request("mistral-ocr-latest", document(), &params).unwrap(),
)
.unwrap();
assert_eq!(result["table_format"], "html");
assert_eq!(result["confidence_scores_granularity"], "page");
assert_eq!(result["extract_header"], true);
}
#[rstest]
fn transform_ocr_response_preserves_blocks_and_confidence_scores() {
let response: MistralOcrResponse = serde_json::from_value(json!({
"pages":[{
"index":0,
"markdown":"hello",
"images":[{"id":"img-0","image_base64":"data:image/png;base64,AA=="}],
"dimensions":{"width":612,"height":792,"dpi":72},
"blocks":[{"type":"title","bbox":{"x":1},"confidence_scores":{"mean":0.98}}],
"confidence_scores":{"average_page_confidence_score":0.99,"minimum_page_confidence_score":0.97}
}],
"model":"returned-model",
"document_annotation":"{\"language\":\"en\"}",
"usage_info":{"pages_processed":1}
}))
.unwrap();
let result = transform_ocr_response("model", response)
.unwrap()
.into_json();
assert_eq!(result["pages"][0]["blocks"][0]["type"], "title");
assert_eq!(result["pages"][0]["blocks"][0]["bbox"]["x"], 1);
assert_eq!(
result["pages"][0]["blocks"][0]["confidence_scores"]["mean"],
0.98
);
assert_eq!(
result["pages"][0]["confidence_scores"]["average_page_confidence_score"],
0.99
);
assert_eq!(result["pages"][0]["images"][0]["id"], "img-0");
assert_eq!(result["pages"][0]["dimensions"]["dpi"], 72);
assert_eq!(result["model"], "returned-model");
assert_eq!(result["document_annotation"], "{\"language\":\"en\"}");
assert_eq!(result["usage_info"]["pages_processed"], 1);
}
#[rstest]
fn transform_ocr_response_preserves_ocr4_page_fields() {
let page = json!({
"index":0,
"markdown":"table page",
"tables":[{"rows":2,"cols":3}],
"hyperlinks":["https://example.com"],
"header":"header",
"footer":"footer"
});
let response: MistralOcrResponse =
serde_json::from_value(json!({"pages":[page.clone()]})).unwrap();
let result = transform_ocr_response("model", response)
.unwrap()
.into_json();
assert_eq!(result["pages"][0], page);
}
}

View file

@ -1,60 +0,0 @@
use serde::{Deserialize, Serialize};
use serde_json::{Map, Value};
use crate::ocr::types::OcrDocument;
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
#[serde(untagged)]
pub(crate) enum MistralOcrPages {
Range(String),
Indices(Vec<i64>),
}
#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)]
pub(crate) struct MistralOcrParams {
#[serde(skip_serializing_if = "Option::is_none")]
pub pages: Option<MistralOcrPages>,
#[serde(skip_serializing_if = "Option::is_none")]
pub include_image_base64: Option<bool>,
#[serde(skip_serializing_if = "Option::is_none")]
pub image_limit: Option<i64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub image_min_size: Option<i64>,
#[serde(skip_serializing_if = "Option::is_none")]
pub bbox_annotation_format: Option<Value>,
#[serde(skip_serializing_if = "Option::is_none")]
pub document_annotation_format: Option<Value>,
#[serde(skip_serializing_if = "Option::is_none")]
pub document_annotation_prompt: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub extract_header: Option<bool>,
#[serde(skip_serializing_if = "Option::is_none")]
pub extract_footer: Option<bool>,
#[serde(skip_serializing_if = "Option::is_none")]
pub table_format: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub confidence_scores_granularity: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub include_blocks: Option<bool>,
#[serde(skip_serializing_if = "Option::is_none")]
pub id: Option<String>,
}
#[derive(Clone, Debug, Serialize, Deserialize)]
pub(crate) struct MistralOcrRequest {
pub model: String,
pub document: OcrDocument,
#[serde(flatten)]
pub params: MistralOcrParams,
}
#[derive(Clone, Debug, Default, Deserialize)]
pub(crate) struct MistralOcrResponse {
#[serde(default)]
pub pages: Vec<Value>,
pub model: Option<String>,
pub document_annotation: Option<Value>,
pub usage_info: Option<Value>,
#[serde(flatten)]
pub extra_fields: Map<String, Value>,
}

View file

@ -1,5 +0,0 @@
pub(crate) mod cohere;
pub(crate) mod deepseek;
pub(crate) mod document_intelligence;
pub(crate) mod mistral;
pub(crate) mod reducto;

View file

@ -1,9 +0,0 @@
mod transformation;
mod types;
pub(crate) use transformation::{
transform_legacy_ocr_request, transform_ocr_response, transform_v3_ocr_request,
};
pub(crate) use types::{
ReductoLegacyParams, ReductoResponse, ReductoUploadResponse, ReductoV3Params,
};

View file

@ -1,103 +0,0 @@
use std::collections::BTreeMap;
use serde_json::{Value, json};
use super::types::*;
use crate::ocr::error::{OcrRequestError, OcrResponseError};
use crate::ocr::types::{LiteLLMOcrResponse, OcrDocument};
pub(crate) fn transform_v3_ocr_request(
_model: &str,
document: OcrDocument,
params: &ReductoV3Params,
) -> Result<ReductoV3Request, OcrRequestError> {
Ok(ReductoV3Request {
input: document.source().to_string(),
params: params.clone(),
})
}
pub(crate) fn transform_legacy_ocr_request(
_model: &str,
document: OcrDocument,
params: &ReductoLegacyParams,
) -> Result<ReductoLegacyRequest, OcrRequestError> {
Ok(ReductoLegacyRequest {
document_url: document.source().to_string(),
options: params.enhance.as_ref().map(|_| params.clone()),
})
}
pub(crate) fn transform_ocr_response(
model: &str,
response: ReductoResponse,
) -> Result<LiteLLMOcrResponse, OcrResponseError> {
let result = match response.result {
Some(result) => result.unwrap_or_default(),
None => ReductoResult {
chunks: response.chunks,
},
};
let usage = response.usage.unwrap_or_default();
Ok(LiteLLMOcrResponse {
pages: build_pages(result.chunks.unwrap_or_default()),
model: model.to_string(),
document_annotation: None,
usage_info: Some(json!({
"pages_processed": usage.num_pages,
"credits": usage.credits,
})),
object: "ocr".to_string(),
extra_fields: serde_json::Map::new(),
provider_native_response: None,
})
}
fn build_pages(chunks: Vec<ReductoChunk>) -> Vec<Value> {
let blocks_by_page = chunks
.iter()
.flat_map(|chunk| chunk.blocks.iter().flatten())
.filter_map(|block| block.bbox.as_ref()?.page.map(|page| (page, block)))
.fold(
BTreeMap::<i64, Vec<&ReductoBlock>>::new(),
|mut pages, (page, block)| {
pages.entry(page).or_default().push(block);
pages
},
);
if blocks_by_page.is_empty() {
let markdown = join_content(chunks.iter().map(|chunk| chunk.content.as_deref()));
return if markdown.is_empty() {
Vec::new()
} else {
vec![page(0, markdown, None)]
};
}
blocks_by_page
.into_iter()
.map(|(index, blocks)| {
let markdown = join_content(blocks.iter().map(|block| block.content.as_deref()));
page(
index.saturating_sub(1).max(0),
markdown,
Some(json!(blocks)),
)
})
.collect()
}
fn join_content<'a>(content: impl Iterator<Item = Option<&'a str>>) -> String {
content
.flatten()
.filter(|text| !text.is_empty())
.collect::<Vec<_>>()
.join("\n\n")
}
fn page(index: i64, markdown: String, blocks: Option<Value>) -> Value {
let mut result = json!({"index":index,"markdown":markdown,"images":null});
if let (Value::Object(fields), Some(blocks)) = (&mut result, blocks) {
fields.insert("blocks".into(), blocks);
}
result
}

View file

@ -1,128 +0,0 @@
use serde::{Deserialize, Deserializer, Serialize};
use serde_json::{Map, Value};
#[derive(Clone, Debug, Default, Serialize, Deserialize)]
pub(crate) struct ReductoV3Params {
#[serde(skip_serializing_if = "Option::is_none")]
pub formatting: Option<Map<String, Value>>,
#[serde(skip_serializing_if = "Option::is_none")]
pub retrieval: Option<Map<String, Value>>,
#[serde(skip_serializing_if = "Option::is_none")]
pub settings: Option<Map<String, Value>>,
}
#[derive(Clone, Debug, Default, Serialize, Deserialize)]
pub(crate) struct ReductoLegacyParams {
#[serde(skip_serializing_if = "Option::is_none")]
pub enhance: Option<Map<String, Value>>,
}
#[derive(Clone, Debug, Serialize, Deserialize)]
pub(crate) struct ReductoV3Request {
pub input: String,
#[serde(flatten)]
pub params: ReductoV3Params,
}
#[derive(Clone, Debug, Serialize, Deserialize)]
pub(crate) struct ReductoLegacyRequest {
pub document_url: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub options: Option<ReductoLegacyParams>,
}
#[derive(Deserialize)]
pub(crate) struct ReductoUploadResponse {
pub file_id: Option<String>,
}
#[derive(Clone, Debug, Deserialize)]
pub(crate) struct ReductoResponse {
#[serde(default, deserialize_with = "present_nullable")]
pub result: Option<Option<ReductoResult>>,
pub usage: Option<ReductoUsage>,
#[serde(default)]
pub chunks: Option<Vec<ReductoChunk>>,
}
fn present_nullable<'de, D: Deserializer<'de>, T: Deserialize<'de>>(
deserializer: D,
) -> Result<Option<Option<T>>, D::Error> {
Option::<T>::deserialize(deserializer).map(Some)
}
#[derive(Clone, Debug, Default, Deserialize)]
pub(crate) struct ReductoResult {
pub chunks: Option<Vec<ReductoChunk>>,
}
#[derive(Clone, Debug, Default, Deserialize)]
pub(crate) struct ReductoUsage {
#[serde(default, deserialize_with = "optional_i64")]
pub num_pages: Option<i64>,
#[serde(default, deserialize_with = "optional_f64")]
pub credits: Option<f64>,
}
#[derive(Clone, Debug, Deserialize)]
pub(crate) struct ReductoChunk {
pub content: Option<String>,
pub blocks: Option<Vec<ReductoBlock>>,
}
#[derive(Clone, Debug, Serialize, Deserialize)]
pub(crate) struct ReductoBlock {
#[serde(skip_serializing_if = "Option::is_none")]
pub content: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub bbox: Option<ReductoBoundingBox>,
#[serde(flatten)]
pub extra_fields: Map<String, Value>,
}
#[derive(Clone, Debug, Serialize, Deserialize)]
pub(crate) struct ReductoBoundingBox {
#[serde(default, deserialize_with = "optional_i64")]
pub page: Option<i64>,
#[serde(flatten)]
pub extra_fields: Map<String, Value>,
}
fn optional_i64<'de, D: Deserializer<'de>>(deserializer: D) -> Result<Option<i64>, D::Error> {
match Option::<Value>::deserialize(deserializer)? {
None | Some(Value::Null) => Ok(None),
Some(Value::Number(number)) => number
.as_i64()
.or_else(|| number.as_f64().and_then(checked_truncated_i64))
.map(Some)
.ok_or_else(|| serde::de::Error::custom("expected an integer")),
Some(Value::String(value)) => value
.trim()
.parse::<i64>()
.map(Some)
.map_err(|_| serde::de::Error::custom("expected an integer")),
Some(Value::Bool(value)) => Ok(Some(i64::from(value))),
Some(_) => Ok(None),
}
}
fn optional_f64<'de, D: Deserializer<'de>>(deserializer: D) -> Result<Option<f64>, D::Error> {
match Option::<Value>::deserialize(deserializer)? {
None | Some(Value::Null) => Ok(None),
Some(Value::Number(number)) => number
.as_f64()
.map(Some)
.ok_or_else(|| serde::de::Error::custom("expected a number")),
Some(Value::String(value)) => value
.trim()
.parse::<f64>()
.map(Some)
.map_err(|_| serde::de::Error::custom("expected a number")),
Some(_) => Ok(None),
}
}
fn checked_truncated_i64(value: f64) -> Option<i64> {
(value.is_finite() && value >= i64::MIN as f64 && value <= i64::MAX as f64)
.then(|| value.trunc() as i64)
}

View file

@ -5,9 +5,11 @@ use base64::{Engine, engine::general_purpose::STANDARD};
use data_url::mime::Mime;
use data_url::{DataUrl, DataUrlError, forgiving_base64::DecodeError};
use reqwest::Url;
use serde_json::Map;
use std::collections::BTreeMap as Map;
use super::error::{OcrError, OcrRequestError, OcrResponseError};
use super::Error as OcrError;
use super::Error as OcrRequestError;
use super::Error as OcrResponseError;
use super::types::{OcrConnection, OcrDocument, OcrDocumentInput};
use crate::constants::{OCR_INLINE_MAX_BYTES, OCR_MAX_FETCH_REDIRECTS};
use crate::media::Error as MediaError;
@ -47,11 +49,10 @@ pub fn read_path_document(
})
.map_err(|source| super::Error::FileRead {
path: path.to_owned(),
kind: source.kind(),
message: source.to_string(),
source: std::sync::Arc::new(source),
})?;
let name = path.file_name().map(|name| name.to_string_lossy());
Ok(encode_file_document(&bytes, name.as_deref(), mime_type)?)
encode_file_document(&bytes, name.as_deref(), mime_type)
}
pub fn encode_file_document(
@ -164,7 +165,7 @@ pub(crate) async fn inline_remote_document(
connection: &OcrConnection,
) -> Result<OcrDocument, OcrError> {
let source = document.source();
if !source.starts_with("http://") && !source.starts_with("https://") {
if !document.is_remote() {
validate_inline_document(&document)?;
return Ok(document);
}
@ -193,12 +194,12 @@ pub(crate) async fn inline_remote_document(
fn map_media_error(error: MediaError) -> OcrError {
match error {
MediaError::BlockedUrl => OcrRequestError::BlockedDocumentUrl.into(),
MediaError::DownloadDisabled => OcrRequestError::DownloadDisabled.into(),
MediaError::DownloadTooLarge => OcrRequestError::DownloadTooLarge.into(),
MediaError::TooManyRedirects => OcrRequestError::TooManyRedirects.into(),
MediaError::MissingRedirectLocation => OcrResponseError::MissingRedirectLocation.into(),
MediaError::InvalidRedirect => OcrResponseError::InvalidRedirect.into(),
MediaError::BlockedUrl => OcrRequestError::BlockedDocumentUrl,
MediaError::DownloadDisabled => OcrRequestError::DownloadDisabled,
MediaError::DownloadTooLarge => OcrRequestError::DownloadTooLarge,
MediaError::TooManyRedirects => OcrRequestError::TooManyRedirects,
MediaError::MissingRedirectLocation => OcrResponseError::MissingRedirectLocation,
MediaError::InvalidRedirect => OcrResponseError::InvalidRedirect,
MediaError::Http(status) => TransportError::Http {
status,
body: "OCR document download failed".into(),
@ -216,7 +217,7 @@ fn map_media_error(error: MediaError) -> OcrError {
#[cfg(test)]
mod tests {
use super::*;
use serde_json::Map;
use std::collections::BTreeMap as Map;
fn document(source: &str) -> OcrDocument {
OcrDocument::DocumentUrl {
@ -286,17 +287,17 @@ mod tests {
document("data:application/pdf;base64,YWJj")
);
std::fs::write(&path, vec![b'a'; OCR_INLINE_MAX_BYTES + 1]).unwrap();
assert_eq!(
assert!(matches!(
prepare_document(OcrDocumentInput::Path {
path: path.clone(),
mime_type: None,
}),
Err(OcrRequestError::InlineDocumentTooLarge.into())
);
Err(OcrRequestError::InlineDocumentTooLarge)
));
std::fs::remove_dir_all(&dir).unwrap();
let missing = dir.join("missing.pdf");
let Err(super::super::Error::FileRead { path, kind, .. }) =
let Err(super::super::Error::FileRead { path, source, .. }) =
prepare_document(OcrDocumentInput::Path {
path: missing.clone(),
mime_type: None,
@ -305,7 +306,7 @@ mod tests {
panic!("missing paths must surface a file read error");
};
assert_eq!(path, missing);
assert_eq!(kind, std::io::ErrorKind::NotFound);
assert_eq!(source.kind(), std::io::ErrorKind::NotFound);
}
#[test]
@ -325,10 +326,10 @@ mod tests {
#[test]
fn file_encoding_enforces_decoded_size_limit() {
let bytes = vec![b'a'; OCR_INLINE_MAX_BYTES + 1];
assert_eq!(
assert!(matches!(
encode_file_document(&bytes, None, None),
Err(OcrRequestError::InlineDocumentTooLarge)
);
));
let document = encode_file_document(&bytes[..OCR_INLINE_MAX_BYTES], None, None).unwrap();
let inline = InlineDocument::parse(document.source()).unwrap().unwrap();
assert_eq!(
@ -359,10 +360,10 @@ mod tests {
] {
let inline = InlineDocument::parse(source).unwrap().unwrap();
assert_eq!(inline.decode(expected.len()).unwrap(), expected);
assert_eq!(
assert!(matches!(
inline.decode(expected.len() - 1),
Err(OcrRequestError::InlineDocumentTooLarge)
);
));
}
}
@ -427,7 +428,7 @@ mod tests {
client.document_fetcher(),
OcrDocument::ImageUrl {
image_url: format!("http://{address}/image"),
extra_fields: Map::from_iter([("detail".into(), serde_json::json!("high"))]),
extra_fields: Map::from_iter([("detail".into(), "high".into())]),
},
&OcrConnection::default(),
)
@ -439,7 +440,7 @@ mod tests {
converted,
OcrDocument::ImageUrl {
image_url: "data:image/png;base64,YWJj".into(),
extra_fields: Map::from_iter([("detail".into(), serde_json::json!("high"))]),
extra_fields: Map::from_iter([("detail".into(), "high".into())]),
}
);
assert!(!request.to_ascii_lowercase().contains("authorization"));

View file

@ -1,117 +1,21 @@
use thiserror::Error;
use crate::transport::Error as TransportError;
#[derive(Clone, Debug, Error, PartialEq, Eq)]
#[derive(Clone, Debug, thiserror::Error)]
pub enum Error {
#[error("expected {expected}, got {actual}")]
InvalidType {
expected: &'static str,
actual: &'static str,
#[error("upstream OCR error ({status}): {body}")]
Provider {
status: u16,
body: String,
headers: Vec<(String, String)>,
},
#[error("missing required field: {0}")]
MissingField(&'static str),
#[error("Document URL is required")]
MissingDocumentUrl,
#[error("invalid response: {0}")]
InvalidResponse(String),
#[error("invalid provider: {0}")]
InvalidProvider(String),
#[error("invalid request: {0}")]
InvalidRequest(String),
#[error("{0}")]
Auth(String),
#[error(
"Missing {provider} API Key - A call is being made to {provider} but no key is set either in the environment variables or via params"
)]
MissingApiKey { provider: &'static str },
#[error(
"invalid authentication configuration: Missing Azure AI credentials - set AZURE_AI_API_KEY or configure Entra ID"
)]
MissingAzureAiCredentials,
#[error(
"invalid authentication configuration: Missing Azure Document Intelligence credentials - set AZURE_DOCUMENT_INTELLIGENCE_API_KEY or configure Entra ID"
)]
MissingAzureDocumentIntelligenceCredentials,
#[error(
"Missing REDUCTO_API_KEY - set it in the environment or pass api_key to litellm.ocr()/litellm.aocr()"
)]
MissingReductoApiKey,
#[error("upstream request failed with status {status}: {body}")]
Http { status: u16, body: String },
#[error("upstream network error: {0}")]
Network(String),
/// The provider was never reached: DNS, TCP, TLS or proxy setup failed
/// before any byte of the request went out. Nothing was billed, so a host
/// that keeps a reference implementation can serve the request itself.
/// A timeout is deliberately not this, since the provider may have received
/// and answered the request already.
#[error("could not reach the provider: {0}")]
Connect(String),
#[error("routing error: {0}")]
Routing(String),
#[error("Failed to read OCR file {}: {message}", path.display())]
FileRead {
path: std::path::PathBuf,
kind: std::io::ErrorKind,
message: String,
},
/// The request is outside the surface this route covers in Rust. Hosts that
/// keep a reference implementation treat this as "fall back", not "fail".
#[error("unsupported by the rust path: {0}")]
Unsupported(&'static str),
}
impl Error {
pub const fn http_status_code(&self) -> Option<u16> {
match self {
Self::InvalidRequest(_) => Some(400),
Self::MissingDocumentUrl => Some(500),
Self::Http { status, .. } => Some(*status),
_ => None,
}
}
}
impl From<OcrRequestError> for Error {
fn from(error: OcrRequestError) -> Self {
match error {
OcrRequestError::MissingField(field) => Self::MissingField(field),
OcrRequestError::MissingDocumentUrl => Self::MissingDocumentUrl,
error => Self::InvalidRequest(error.to_string()),
}
}
}
impl From<OcrResponseError> for Error {
fn from(error: OcrResponseError) -> Self {
Self::InvalidResponse(error.to_string())
}
}
impl From<TransportError> for Error {
fn from(error: TransportError) -> Self {
match error {
TransportError::Http { status, body } => Self::Http { status, body },
TransportError::Network(message) => Self::Network(message),
TransportError::Connect(message) => Self::Connect(message),
}
}
}
impl From<litellm_auth::Error> for Error {
fn from(error: litellm_auth::Error) -> Self {
match error {
litellm_auth::Error::MissingApiKey { provider, .. } => Self::MissingApiKey { provider },
error => Self::Auth(error.to_string()),
}
}
}
#[derive(Debug, Clone, PartialEq, Eq, Error)]
pub enum OcrRequestError {
#[error("File is empty or could not be read")]
EmptyFile,
#[error("Failed to read OCR file {}: {source}", path.display())]
FileRead {
path: std::path::PathBuf,
#[source]
source: std::sync::Arc<std::io::Error>,
},
#[error("OCR document preparation task failed: {0}")]
DocumentTask(#[source] std::sync::Arc<tokio::task::JoinError>),
#[error("Invalid MIME type: {0}")]
InvalidMimeType(String),
#[error(
@ -148,10 +52,6 @@ pub enum OcrRequestError {
Features,
#[error("OCR model cannot be a dot segment")]
DotModel,
}
#[derive(Debug, Clone, PartialEq, Eq, Error)]
pub enum OcrResponseError {
#[error("OCR response exceeds the size limit of {limit} bytes")]
TooLarge { limit: usize },
#[error("invalid OCR response field: {path}")]
@ -166,40 +66,101 @@ pub enum OcrResponseError {
OperationStatus(String),
#[error("OCR response numeric value is out of range: {0}")]
NumericRange(&'static str),
}
#[derive(Debug, Clone, PartialEq, Eq, Error)]
pub enum OcrPollingError {
#[error("OCR accepted response is missing a valid operation-location")]
PollLocation,
#[error("OCR operation-location must use the submission origin without credentials")]
PollOrigin,
#[error("OCR polling timed out")]
PollTimeout,
#[error("unsupported by the rust path: {0}")]
Unsupported(&'static str),
#[error("invalid provider: {0}")]
InvalidProvider(String),
#[error("invalid request: {0}")]
InvalidRequest(String),
#[error("invalid response: {0}")]
InvalidResponse(String),
#[error(
"invalid authentication configuration: Missing Azure AI credentials - set AZURE_AI_API_KEY or configure Entra ID"
)]
MissingAzureAiCredentials,
#[error(
"invalid authentication configuration: Missing Azure Document Intelligence credentials - set AZURE_DOCUMENT_INTELLIGENCE_API_KEY or configure Entra ID"
)]
MissingAzureDocumentIntelligenceCredentials,
#[error(
"Missing REDUCTO_API_KEY - set it in the environment or pass api_key to litellm.ocr()/litellm.aocr()"
)]
MissingReductoApiKey,
#[error(transparent)]
Auth(#[from] litellm_auth::Error),
#[error(transparent)]
Transport(#[from] crate::transport::Error),
#[error(transparent)]
Params(#[from] crate::params::Error),
#[error(transparent)]
Headers(#[from] crate::http_utils::HeaderError),
}
#[derive(Debug, Error)]
pub enum OcrError {
#[error("{0}")]
Request(#[from] OcrRequestError),
#[error("{0}")]
Response(#[from] OcrResponseError),
#[error("{0}")]
Transport(#[from] TransportError),
#[error("{0}")]
Polling(#[from] OcrPollingError),
#[error("{0}")]
Public(#[from] Error),
}
impl From<OcrError> for Error {
fn from(error: OcrError) -> Self {
match error {
OcrError::Request(error) => error.into(),
OcrError::Response(error) => error.into(),
OcrError::Transport(error) => error.into(),
OcrError::Polling(error) => Error::InvalidResponse(error.to_string()),
OcrError::Public(error) => error,
impl From<crate::call_arguments::ArgumentError> for Error {
fn from(error: crate::call_arguments::ArgumentError) -> Self {
Self::RequestField {
path: format!("optional_params.{}", error.path),
}
}
}
impl Error {
pub fn http_status_code(&self) -> Option<u16> {
match self {
Self::MissingDocumentUrl => Some(500),
Self::Provider { status, .. }
| Self::Transport(crate::transport::Error::Http { status, .. }) => Some(*status),
error if error.is_request() => Some(400),
_ => None,
}
}
pub fn is_request(&self) -> bool {
matches!(
self,
Self::EmptyFile
| Self::InvalidMimeType(_)
| Self::CohereImageOnly
| Self::RequestFormat
| Self::RequestField { .. }
| Self::MissingField(_)
| Self::MissingDocumentUrl
| Self::InvalidDataUri
| Self::ReductoSource
| Self::InlineDocumentTooLarge
| Self::BlockedDocumentUrl
| Self::DownloadDisabled
| Self::DownloadTooLarge
| Self::TooManyRedirects
| Self::Pages(_)
| Self::Features
| Self::DotModel
| Self::InvalidRequest(_)
| Self::Params(_)
| Self::Headers(_)
)
}
pub fn is_response(&self) -> bool {
matches!(
self,
Self::TooLarge { .. }
| Self::ResponseField { .. }
| Self::EmptyContent
| Self::MissingRedirectLocation
| Self::InvalidRedirect
| Self::OperationStatus(_)
| Self::NumericRange(_)
| Self::PollLocation
| Self::PollOrigin
| Self::PollTimeout
| Self::InvalidResponse(_)
)
}
}

View file

@ -1,21 +1,20 @@
use super::OcrClient;
use super::adapters::OcrAdapter;
use super::hooks::{OcrHooks, OcrLifecycleHooks, OcrPostCallRequest};
use super::registry::OcrAdapterKind;
use super::types::{LiteLLMOcrRequest, LiteLLMOcrResponse};
use crate::call_lifecycle::{CallLifecycle, CallLifecycleContext};
use crate::ocr::Error;
use std::sync::Arc;
use super::OcrClient;
use super::hooks::{OcrHooks, OcrLifecycleHooks, OcrPostCallRequest};
use super::types::{LiteLLMOcrResponse, PreparedOcrRequest, ResolvedOcrRequest};
use crate::call_lifecycle::{CallLifecycle, CallLifecycleContext};
use crate::llms::base_llm::ocr::transformation::OcrResponseContext;
pub(crate) async fn perform_ocr_request(
client: &OcrClient,
request: LiteLLMOcrRequest,
) -> Result<LiteLLMOcrResponse, Error> {
request: ResolvedOcrRequest,
) -> Result<LiteLLMOcrResponse, super::Error> {
request.response_format()?;
let context = CallLifecycleContext::new(
"ocr",
request.model.clone(),
request.adapter.provider().as_str(),
request.provider_name(),
request
.litellm_call_id
.clone()
@ -30,31 +29,24 @@ pub(crate) async fn perform_ocr_request(
PreparedOcrCall::prepare(client.clone(), request)
.await?
.execute()
.await?
.normalize()
.await
})
.await
}
pub(crate) struct PreparedOcrCall {
client: OcrClient,
request: LiteLLMOcrRequest,
request: PreparedOcrRequest,
http: reqwest::Request,
}
impl PreparedOcrCall {
pub(crate) async fn prepare(
client: OcrClient,
request: LiteLLMOcrRequest,
) -> Result<Self, Error> {
macro_rules! prepare_adapter {
($( $variant:ident, $adapter:ty, $instance:expr, $provider:ident; )+) => {
match request.adapter {
$( OcrAdapterKind::$variant => $instance.prepare_request(&request, &client).await?, )+
}
};
}
let http = super::adapters::for_each_ocr_adapter!(prepare_adapter);
request: ResolvedOcrRequest,
) -> Result<Self, super::Error> {
let request = super::prepare::prepare_request(request);
let http = request.config.prepare_request(&request, &client).await?;
Ok(Self {
client,
request,
@ -62,33 +54,54 @@ impl PreparedOcrCall {
})
}
pub(crate) async fn execute(self) -> Result<OcrProviderResponse, Error> {
pub(crate) async fn execute(self) -> Result<LiteLLMOcrResponse, super::Error> {
let url = self.http.url().to_string();
let headers = request_headers(&self.http)?;
let response = crate::http_utils::http_request(reqwest::RequestBuilder::from_parts(
self.client.provider_http().clone(),
self.http,
))
.await
.map_err(super::client::transport_error)?;
macro_rules! read_adapter {
($( $variant:ident, $adapter:ty, $instance:expr, $provider:ident; )+) => {
match self.request.adapter {
$( OcrAdapterKind::$variant => {
let decoded = $instance.read_response(&self.client, response, &url, &headers, &self.request).await?;
Ok(OcrProviderResponse {
request: self.request,
data: OcrProviderData::$variant(decoded),
})
}, )+
let response =
crate::http_utils::execute_http_request(self.client.provider_http(), self.http)
.await
.map_err(super::client::transport_error)?;
if !response.status().is_success() {
let headers = response
.headers()
.iter()
.filter_map(|(name, value)| {
value
.to_str()
.ok()
.map(|value| (name.to_string(), value.to_string()))
})
.collect();
return match super::client::read_response_bytes(
response,
self.request.connection.max_response_bytes,
)
.await
{
Err(super::Error::Transport(crate::transport::Error::Http { status, body })) => {
Err(self.request.config.get_error_class(body, status, headers))
}
Err(error) => Err(error),
Ok(_) => unreachable!("non-success response produces an HTTP error"),
};
}
super::adapters::for_each_ocr_adapter!(read_adapter)
let model = &self.request.model;
let context = OcrResponseContext {
client: &self.client,
connection: &self.request.connection,
hooks: &self.request.hooks,
request_format: self.request.response_format()?,
url: &url,
headers: &headers,
};
self.request
.config
.async_transform_ocr_response(model, response, context)
.await
}
}
fn request_headers(request: &reqwest::Request) -> Result<Vec<(String, String)>, Error> {
fn request_headers(request: &reqwest::Request) -> Result<Vec<(String, String)>, super::Error> {
request
.headers()
.iter()
@ -96,44 +109,17 @@ fn request_headers(request: &reqwest::Request) -> Result<Vec<(String, String)>,
value
.to_str()
.map(|value| (name.to_string(), value.to_string()))
.map_err(|_| super::error::OcrRequestError::RequestField {
.map_err(|_| super::Error::RequestField {
path: "headers".into(),
})
.map_err(Error::from)
})
.collect()
}
macro_rules! provider_data {
($( $variant:ident, $adapter:ty, $instance:expr, $provider:ident; )+) => {
enum OcrProviderData {
$( $variant(super::wire::DecodedOcrResponse<<$adapter as OcrAdapter>::ProviderResponse>), )+
}
impl OcrProviderResponse {
pub(crate) fn normalize(self) -> Result<LiteLLMOcrResponse, Error> {
match self.data {
$( OcrProviderData::$variant(decoded) => {
let response = $instance.transform_ocr_response(&self.request, decoded.data)?;
Ok(LiteLLMOcrResponse { provider_native_response: decoded.native, ..response })
}, )+
}
}
}
};
}
pub(crate) struct OcrProviderResponse {
request: LiteLLMOcrRequest,
data: OcrProviderData,
}
pub(crate) async fn post_call(hooks: &Arc<dyn OcrHooks>, bytes: &[u8]) -> Result<(), Error> {
pub(crate) async fn post_call(hooks: &Arc<dyn OcrHooks>, bytes: &[u8]) -> Result<(), super::Error> {
let original_response = serde_json::Value::String(String::from_utf8_lossy(bytes).into_owned());
hooks
.post_call(OcrPostCallRequest { original_response })
.await?;
Ok(())
}
super::adapters::for_each_ocr_adapter!(provider_data);

View file

@ -2,7 +2,7 @@ use std::future::Future;
use std::pin::Pin;
use std::sync::Arc;
use super::types::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrDocument};
use super::types::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrDocument, ResolvedOcrRequest};
use crate::call_lifecycle::{CallLifecycleContext, CallLifecycleHooks, CallLifecycleTiming};
use crate::ocr::Error;
use serde::Serialize;
@ -23,6 +23,7 @@ pub struct OcrPreCallRequest {
pub struct OcrDuringCallRequest {
pub model: String,
pub custom_llm_provider: String,
pub api_key: Option<String>,
pub url: String,
pub headers: Vec<(String, String)>,
pub body: Value,
@ -77,19 +78,19 @@ pub(crate) struct OcrLifecycleHooks {
pub provider_name: String,
}
impl CallLifecycleHooks<LiteLLMOcrRequest, LiteLLMOcrRequest, LiteLLMOcrResponse>
impl CallLifecycleHooks<ResolvedOcrRequest, ResolvedOcrRequest, LiteLLMOcrResponse>
for OcrLifecycleHooks
{
type Error = crate::ocr::Error;
type PreCallFuture<'a> = OcrHookFuture<'a, LiteLLMOcrRequest>;
type DuringCallFuture<'a> = OcrHookFuture<'a, LiteLLMOcrRequest>;
type PreCallFuture<'a> = OcrHookFuture<'a, ResolvedOcrRequest>;
type DuringCallFuture<'a> = OcrHookFuture<'a, ResolvedOcrRequest>;
type SuccessFuture<'a> = OcrLogFuture<'a>;
type FailureFuture<'a> = OcrLogFuture<'a>;
fn async_pre_call_hook<'a>(
&'a self,
_context: &'a CallLifecycleContext,
request: LiteLLMOcrRequest,
request: ResolvedOcrRequest,
) -> Self::PreCallFuture<'a> {
Box::pin(async move {
if !self.hooks.intercepts_requests() {
@ -101,18 +102,17 @@ impl CallLifecycleHooks<LiteLLMOcrRequest, LiteLLMOcrRequest, LiteLLMOcrResponse
model: request.model.clone(),
custom_llm_provider: self.provider_name.clone(),
document: request.document,
optional_params: Value::Object(request.optional_params),
optional_params: Value::Object(request.optional_params.into()),
})
.await?;
let Value::Object(optional_params) = changed.optional_params else {
return Err(super::error::OcrRequestError::RequestField {
return Err(super::Error::RequestField {
path: "guardrail.optional_params".into(),
}
.into());
});
};
Ok(LiteLLMOcrRequest {
document: changed.document,
optional_params,
optional_params: optional_params.into(),
..request
})
})
@ -121,7 +121,7 @@ impl CallLifecycleHooks<LiteLLMOcrRequest, LiteLLMOcrRequest, LiteLLMOcrResponse
fn async_during_call_hook<'a>(
&'a self,
_context: &'a CallLifecycleContext,
request: LiteLLMOcrRequest,
request: ResolvedOcrRequest,
) -> Self::DuringCallFuture<'a> {
Box::pin(async move { Ok(request) })
}

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

@ -380,7 +380,7 @@ impl OcrExecution {
self.execution = None;
self.completed = true;
result
.map_err(|error| Error::Network(format!("OCR execution task failed: {error}")))?
.map_err(|error| Error::Transport(crate::transport::Error::Network(format!("OCR execution task failed: {error}"))))?
.map(OcrCallStep::Complete)
}
}
@ -431,7 +431,7 @@ impl OcrExecution {
async fn prepare_request_document(
request: LiteLLMOcrRequest<OcrDocumentInput>,
hooks: &ProtocolHooks,
) -> Result<LiteLLMOcrRequest, Error> {
) -> Result<super::types::ResolvedOcrRequest, Error> {
let request = match &request.document {
OcrDocumentInput::HostReader { mime_type } => {
let mime_type = mime_type.clone();

View file

@ -1,26 +1,31 @@
mod adapters;
mod arguments;
pub mod client;
mod codecs;
mod document;
pub(crate) mod document;
pub mod error;
pub use error::Error;
mod handler;
pub(crate) mod handler;
pub mod hooks;
pub(crate) mod json;
mod lifecycle;
mod prepare;
mod registry;
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, read_path_document};
pub use lifecycle::{
NativeOutcome, NativeResult, NoopOcrHost, OcrAdmission, OcrCall, OcrCallStep, OcrDecline,
OcrHookHost, OcrHost, OcrHostOperation, OcrHostResult,
};
pub use provider_config::{get_api_key_env_var, get_health_check_document};
pub use types::{
LiteLLMOcrRequest, LiteLLMOcrResponse, OcrConnection, OcrDocument, OcrDocumentInput,
OcrFileContent,
LiteLLMOcrRequest, LiteLLMOcrResponse, OcrConnection, OcrConnectionInputs, OcrCredentialInputs,
OcrDocument, OcrDocumentInput, OcrFileContent, OcrPage, OcrPageDimensions, OcrPageImage,
OcrTransportConfig, OcrUsageInfo,
};
#[cfg(test)]

View file

@ -1,117 +1,72 @@
use serde::{Deserialize, Serialize, de::DeserializeOwned};
use serde_json::{Map, Value};
use serde::Serialize;
use serde_json::Value;
use super::OcrClient;
use super::error::{OcrError, OcrRequestError};
use super::hooks::OcrDuringCallRequest;
use super::types::{LiteLLMOcrRequest, OcrDocument};
#[derive(Debug, Deserialize)]
pub(crate) struct ParsedProviderParams<T> {
#[serde(flatten)]
pub known: T,
#[serde(default, flatten)]
pub extra_params: Map<String, Value>,
}
pub(crate) fn _prepare_ocr_request<T: DeserializeOwned>(
request: &LiteLLMOcrRequest,
) -> Result<ParsedProviderParams<T>, OcrRequestError> {
super::wire::decode_request_value(
Value::Object(request.optional_params.clone()),
"optional_params",
)
}
pub(crate) fn merge_extra_params<B: Serialize>(
body: &B,
extra_params: Map<String, Value>,
) -> Result<Value, OcrRequestError> {
let Value::Object(fields) =
serde_json::to_value(body).map_err(|_| OcrRequestError::RequestField {
path: "body".into(),
})?
else {
return Err(OcrRequestError::RequestField {
path: "body".into(),
});
};
let extra_body = extra_params
.get("extra_body")
.and_then(Value::as_object)
.cloned()
.unwrap_or_default()
.into_iter()
.collect::<Map<String, Value>>();
Ok(Value::Object(
fields
.into_iter()
.chain(
extra_params
.into_iter()
.filter(|(name, _)| name != "extra_body"),
)
.chain(extra_body)
.collect(),
))
}
use super::types::{OcrConnection, OcrDocument, PreparedOcrRequest, ResolvedOcrRequest};
pub(crate) async fn transform_request_body<B>(
client: &OcrClient,
request: &LiteLLMOcrRequest,
request: &PreparedOcrRequest,
url: &str,
headers: &[(String, String)],
retains_document: bool,
body: B,
validate: impl FnOnce(&B) -> Result<(), OcrRequestError>,
) -> Result<reqwest::Request, OcrError>
validate: impl Fn(&Value) -> Result<(), super::Error>,
) -> Result<reqwest::Request, super::Error>
where
B: Serialize + DeserializeOwned,
B: Serialize,
{
let composed = crate::call_arguments::compose_body(
&request.optional_params,
&body,
request.config.get_supported_ocr_params(&request.model),
)?;
validate(&composed)?;
let retained_fields = request
.optional_params
.keys()
.filter(|name| composed.get(*name).is_some())
.cloned()
.chain(
composed
.get("document")
.is_some()
.then(|| "document".to_string()),
)
.collect();
let (body, headers) = if request.hooks.intercepts_requests() {
let body = serde_json::to_value(body).map_err(|_| OcrRequestError::RequestField {
path: "body".into(),
})?;
let retained_fields = request
.optional_params
.keys()
.filter(|name| body.get(*name).is_some())
.cloned()
.chain(retains_document.then(|| "document".to_string()))
.collect();
let changed = request
.hooks
.during_call(OcrDuringCallRequest {
model: request.model.clone(),
custom_llm_provider: request.adapter.provider().as_str().into(),
custom_llm_provider: request.provider_name().into(),
api_key: request.connection.api_key.clone(),
url: url.into(),
headers: headers.to_vec(),
body,
body: composed,
retained_fields,
})
.await?;
let body = OcrWireBody::<B>::decode(changed.body)?;
validate(&body.body)?;
(body, changed.headers)
if !changed.body.is_object() {
return Err(super::Error::RequestField {
path: "guardrail.body".into(),
});
}
validate(&changed.body)?;
(changed.body, changed.headers)
} else {
(
OcrWireBody {
body,
extra: Map::new(),
},
headers.to_vec(),
)
(composed, headers.to_vec())
};
build_http_request(client, request, url, &headers, &body)
}
pub(crate) fn build_http_request<B: Serialize>(
client: &OcrClient,
request: &LiteLLMOcrRequest,
request: &PreparedOcrRequest,
url: &str,
headers: &[(String, String)],
body: &B,
) -> Result<reqwest::Request, OcrError> {
) -> Result<reqwest::Request, super::Error> {
let builder = client
.provider_http()
.post(url)
@ -120,14 +75,14 @@ pub(crate) fn build_http_request<B: Serialize>(
crate::http_utils::with_headers(builder, headers, crate::http_utils::HeaderPolicy::All)
.build()
.map_err(crate::transport::Error::from)
.map_err(OcrError::from)
.map_err(super::Error::from)
}
pub(crate) async fn guardrail_document(
request: &LiteLLMOcrRequest,
request: &PreparedOcrRequest,
url: &str,
headers: &[(String, String)],
) -> Result<(OcrDocument, Vec<(String, String)>), OcrError> {
) -> Result<(OcrDocument, Vec<(String, String)>), super::Error> {
if !request.hooks.intercepts_requests() {
return Ok((request.document.clone(), headers.to_vec()));
}
@ -135,80 +90,113 @@ pub(crate) async fn guardrail_document(
.hooks
.during_call(OcrDuringCallRequest {
model: request.model.clone(),
custom_llm_provider: request.adapter.provider().as_str().into(),
custom_llm_provider: request.provider_name().into(),
api_key: request.connection.api_key.clone(),
url: url.into(),
headers: headers.to_vec(),
body: serde_json::to_value(&request.document).map_err(|_| {
OcrRequestError::RequestField {
super::Error::RequestField {
path: "document".into(),
}
})?,
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))
}
#[derive(Serialize)]
struct OcrWireBody<B> {
#[serde(flatten)]
body: B,
#[serde(flatten)]
extra: Map<String, Value>,
}
impl<B: Serialize + DeserializeOwned> OcrWireBody<B> {
fn decode(value: Value) -> Result<Self, OcrRequestError> {
let body: B = super::wire::decode_request_value(value.clone(), "guardrail.body")?;
let Value::Object(fields) = value else {
return Err(OcrRequestError::RequestField {
path: "guardrail.body".into(),
});
};
let known = serde_json::to_value(&body).map_err(|_| OcrRequestError::RequestField {
path: "guardrail.body".into(),
pub(crate) fn body_document(body: &Value) -> Result<OcrDocument, super::Error> {
let document = body
.get("document")
.and_then(Value::as_object)
.ok_or_else(|| super::Error::RequestField {
path: "body.document".into(),
})?;
let extra = fields
.into_iter()
.filter(|(key, _)| known.get(key).is_none())
.collect();
Ok(Self { body, extra })
}
let source = document
.iter()
.filter(|(name, _)| matches!(name.as_str(), "type" | "image_url" | "document_url"))
.map(|(name, value)| (name.clone(), value.clone()))
.collect();
super::json::decode_request_value(Value::Object(source), "body.document")
}
pub(crate) fn credential_env(name: &str) -> Option<String> {
std::env::var(name).ok()
}
pub(crate) fn prepare_request(request: ResolvedOcrRequest) -> PreparedOcrRequest {
use litellm_auth::{InputSource, Sourced};
let credentials = request.credentials.clone();
let api_base_env = match request.config.provider() {
super::provider_config::OcrProvider::Mistral => Some("MISTRAL_API_BASE"),
super::provider_config::OcrProvider::AzureAi => Some("AZURE_AI_API_BASE"),
super::provider_config::OcrProvider::Cohere
| super::provider_config::OcrProvider::Reducto
| super::provider_config::OcrProvider::VertexAi => None,
};
let dynamic_api_key = credentials.dynamic_api_key.or_else(|| {
credentials.api_key.clone().or_else(|| {
request
.config
.get_api_key_env_var()
.and_then(credential_env)
.map(|value| Sourced::new(value, InputSource::Environment))
})
});
let dynamic_api_base = credentials.dynamic_api_base.or_else(|| {
credentials.api_base.clone().or_else(|| {
api_base_env
.and_then(credential_env)
.map(|value| Sourced::new(value, InputSource::Environment))
})
});
let resolved = request
.config
.resolve_connection_params(super::types::OcrCredentialInputs {
dynamic_api_key,
dynamic_api_base,
..credentials
});
let transport = request.transport.clone();
PreparedOcrRequest::new(request, OcrConnection::new(resolved, transport))
}
#[cfg(test)]
mod tests {
use crate::call_arguments::{CallArguments, compose_body, parse_options};
use serde_json::json;
use super::*;
#[derive(Debug, Deserialize, PartialEq)]
#[derive(serde::Deserialize)]
struct KnownParams {
pages: Option<Vec<i64>>,
}
#[test]
fn parsed_provider_params_separates_known_and_extra_params() {
let parsed: ParsedProviderParams<KnownParams> = super::super::wire::decode_request_value(
json!({
"pages": [0, 2],
"future_ocr_option": true,
"extra_body": {"provider_option": "value"}
}),
"optional_params",
)
let arguments: CallArguments = serde_json::from_value(json!({
"pages": [0, 2],
"future_ocr_option": true,
"extra_body": {"provider_option": "value"}
}))
.unwrap();
assert_eq!(parsed.known.pages, Some(vec![0, 2]));
assert_eq!(parsed.extra_params["future_ocr_option"], true);
let known: KnownParams = parse_options(&arguments).unwrap();
assert_eq!(known.pages, Some(vec![0, 2]));
assert_eq!(arguments["future_ocr_option"], true);
assert_eq!(arguments["extra_body"], json!({"provider_option": "value"}));
assert_eq!(
parsed.extra_params["extra_body"],
json!({"provider_option": "value"})
arguments
.iter()
.filter(|(name, _)| name.as_str() != "pages")
.count(),
2
);
assert_eq!(
compose_body(&arguments, &json!({"pages": known.pages}), &["pages"]).unwrap(),
json!({
"pages": [0, 2], "future_ocr_option": true, "provider_option": "value"
})
);
assert_eq!(parsed.extra_params.len(), 2);
}
}

View file

@ -0,0 +1,411 @@
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, 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::providers::custom_llm_provider::{CustomLlmProvider, get_custom_llm_provider};
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),*)$($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)*,
}
};
}
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub(crate) enum OcrConfigKind {
Cohere,
Mistral,
AzureAi,
AzureCohere,
AzureDocumentIntelligence,
ReductoLegacy,
ReductoV3,
VertexAi,
VertexDeepSeek,
}
impl OcrConfigKind {
pub(crate) const fn provider(self) -> OcrProvider {
match self {
Self::Cohere => OcrProvider::Cohere,
Self::Mistral => OcrProvider::Mistral,
Self::AzureAi | Self::AzureCohere | Self::AzureDocumentIntelligence => {
OcrProvider::AzureAi
}
Self::ReductoLegacy | Self::ReductoV3 => OcrProvider::Reducto,
Self::VertexAi | Self::VertexDeepSeek => OcrProvider::VertexAi,
}
}
pub(crate) fn get_supported_ocr_params(self, model: &str) -> &'static [&'static str] {
dispatch_config!(self, get_supported_ocr_params(model))
}
pub(crate) fn get_api_key_env_var(self) -> Option<&'static str> {
dispatch_config!(self, get_api_key_env_var())
}
pub(crate) fn get_health_check_document(self) -> OcrDocument {
dispatch_config!(self, get_health_check_document())
}
pub(crate) fn resolve_connection_params(
self,
inputs: OcrCredentialInputs,
) -> ResolvedOcrCredentials {
dispatch_config!(self, resolve_connection_params(inputs))
}
pub(crate) fn get_error_class(
self,
message: String,
status: u16,
headers: Vec<(String, String)>,
) -> super::Error {
dispatch_config!(self, get_error_class(message, status, headers))
}
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(
model: &str,
custom_llm_provider: Option<&str>,
) -> Result<Option<&'static str>, super::Error> {
Ok(resolve_provider_config(model, custom_llm_provider)?
.1
.get_api_key_env_var())
}
pub fn get_health_check_document(
model: &str,
custom_llm_provider: Option<&str>,
) -> Result<OcrDocument, super::Error> {
Ok(resolve_provider_config(model, custom_llm_provider)?
.1
.get_health_check_document())
}
#[derive(Clone, Copy, Debug, EnumString, IntoStaticStr, PartialEq, Eq)]
#[strum(serialize_all = "snake_case")]
pub(crate) enum OcrProvider {
Cohere,
Mistral,
AzureAi,
Reducto,
VertexAi,
}
pub(crate) fn resolve_provider_config(
model: &str,
custom_llm_provider: Option<&str>,
) -> Result<(String, OcrConfigKind), super::Error> {
let provider =
get_custom_llm_provider(model, custom_llm_provider).unwrap_or(CustomLlmProvider {
model,
custom_llm_provider: OcrProvider::Mistral.into(),
});
let ocr_provider = provider
.custom_llm_provider
.parse::<OcrProvider>()
.map_err(|_| super::Error::InvalidProvider(provider.custom_llm_provider.to_string()))?;
let config = match ocr_provider {
OcrProvider::Cohere => OcrConfigKind::Cohere,
OcrProvider::Mistral => OcrConfigKind::Mistral,
OcrProvider::AzureAi if is_document_intelligence_model(provider.model) => {
OcrConfigKind::AzureDocumentIntelligence
}
OcrProvider::AzureAi
if provider.model.to_ascii_lowercase().contains("cohere")
&& provider.model.to_ascii_lowercase().contains("parse") =>
{
OcrConfigKind::AzureCohere
}
OcrProvider::AzureAi => OcrConfigKind::AzureAi,
OcrProvider::Reducto if provider.model.eq_ignore_ascii_case("parse-legacy") => {
OcrConfigKind::ReductoLegacy
}
OcrProvider::Reducto => OcrConfigKind::ReductoV3,
OcrProvider::VertexAi if provider.model.to_ascii_lowercase().contains("deepseek") => {
OcrConfigKind::VertexDeepSeek
}
OcrProvider::VertexAi => OcrConfigKind::VertexAi,
};
Ok((provider.model.to_string(), config))
}
fn is_document_intelligence_model(model: &str) -> bool {
let model = model.to_ascii_lowercase();
model.contains("doc-intelligence") || model.contains("documentintelligence")
}
#[cfg(test)]
mod tests {
use super::*;
use litellm_auth::{InputSource, Sourced};
use rstest::rstest;
#[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);
}
#[rstest]
#[case("Mistral")]
#[case("unknown")]
fn invalid_provider_names_are_rejected(#[case] provider: &str) {
assert!(matches!(
resolve_provider_config("model", Some(provider)),
Err(crate::ocr::Error::InvalidProvider(value)) if value == provider
));
}
#[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]
fn connection_resolution_preserves_dynamic_precedence_and_input_sources() {
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: Some(Sourced::new("dynamic-key".into(), InputSource::Environment)),
dynamic_api_base: Some(Sourced::new(
"https://dynamic.test".into(),
InputSource::Request,
)),
});
assert_eq!(
connection
.api_key
.as_ref()
.map(|value| value.value().as_str()),
Some("dynamic-key")
);
assert_eq!(
connection
.api_base
.as_ref()
.map(|value| value.value().as_str()),
Some("https://dynamic.test")
);
assert_eq!(
connection.api_key.as_ref().map(Sourced::source),
Some(InputSource::Environment)
);
assert_eq!(
connection.api_base.as_ref().map(Sourced::source),
Some(InputSource::Request)
);
}
#[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!(
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")
);
}
#[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!(
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
);
assert_eq!(
resolve_provider_config(model, None).unwrap().0,
model.split_once('/').unwrap().1
);
}
}

View file

@ -1,132 +0,0 @@
use super::adapters::OcrAdapter;
use crate::ocr::Error;
use crate::providers::custom_llm_provider::{CustomLlmProvider, get_custom_llm_provider};
macro_rules! define_adapter_types {
($( $variant:ident, $adapter:ty, $instance:expr, $provider:ident; )+) => {
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub(crate) enum OcrAdapterKind {
$( $variant, )+
}
impl OcrAdapterKind {
pub(crate) const fn provider(self) -> OcrProvider {
match self {
$( Self::$variant => <$adapter>::PROVIDER, )+
}
}
}
};
}
super::adapters::for_each_ocr_adapter!(define_adapter_types);
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub(crate) enum OcrProvider {
Cohere,
Mistral,
AzureAi,
Reducto,
VertexAi,
}
impl OcrProvider {
pub(crate) const fn as_str(self) -> &'static str {
match self {
Self::Cohere => "cohere",
Self::Mistral => "mistral",
Self::AzureAi => "azure_ai",
Self::Reducto => "reducto",
Self::VertexAi => "vertex_ai",
}
}
}
pub(crate) fn resolve_wire_adapter(
model: &str,
custom_llm_provider: Option<&str>,
) -> Result<(String, OcrAdapterKind), Error> {
let provider =
get_custom_llm_provider(model, custom_llm_provider).unwrap_or(CustomLlmProvider {
model,
custom_llm_provider: OcrProvider::Mistral.as_str(),
});
let typed_provider = match provider.custom_llm_provider {
"cohere" => OcrProvider::Cohere,
"mistral" => OcrProvider::Mistral,
"azure_ai" => OcrProvider::AzureAi,
"reducto" => OcrProvider::Reducto,
"vertex_ai" => OcrProvider::VertexAi,
value => return Err(Error::InvalidProvider(value.to_string())),
};
let adapter = match typed_provider {
OcrProvider::Cohere => OcrAdapterKind::Cohere,
OcrProvider::Mistral => OcrAdapterKind::Mistral,
OcrProvider::AzureAi if is_document_intelligence_model(provider.model) => {
OcrAdapterKind::AzureDocumentIntelligence
}
OcrProvider::AzureAi
if provider.model.to_ascii_lowercase().contains("cohere")
&& provider.model.to_ascii_lowercase().contains("parse") =>
{
OcrAdapterKind::AzureCohere
}
OcrProvider::AzureAi => OcrAdapterKind::AzureMistral,
OcrProvider::Reducto if provider.model.eq_ignore_ascii_case("parse-legacy") => {
OcrAdapterKind::ReductoLegacy
}
OcrProvider::Reducto if provider.model.eq_ignore_ascii_case("parse-v3") => {
OcrAdapterKind::ReductoV3
}
OcrProvider::Reducto => OcrAdapterKind::ReductoV3,
OcrProvider::VertexAi if provider.model.to_ascii_lowercase().contains("deepseek") => {
OcrAdapterKind::VertexDeepSeek
}
OcrProvider::VertexAi => OcrAdapterKind::VertexMistral,
};
Ok((provider.model.to_string(), adapter))
}
fn is_document_intelligence_model(model: &str) -> bool {
let model = model.to_ascii_lowercase();
model.contains("doc-intelligence") || model.contains("documentintelligence")
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn provider_models_are_preserved_without_a_local_allowlist() {
let cases = [
("mistral/future-ocr-model", OcrAdapterKind::Mistral),
("azure_ai/future-ocr-model", OcrAdapterKind::AzureMistral),
];
for (qualified_model, expected_adapter) in cases {
let expected_model = qualified_model.split_once('/').unwrap().1;
let (model, adapter) = resolve_wire_adapter(qualified_model, None).unwrap();
assert_eq!(model, expected_model);
assert_eq!(adapter, expected_adapter);
}
}
#[test]
fn unknown_reducto_models_use_the_current_protocol() {
let (model, adapter) = resolve_wire_adapter("reducto/future-parse-model", None).unwrap();
assert_eq!(model, "future-parse-model");
assert_eq!(adapter, OcrAdapterKind::ReductoV3);
}
#[test]
fn known_protocol_models_still_select_specialized_adapters() {
let (model, adapter) = resolve_wire_adapter("reducto/parse-legacy", None).unwrap();
assert_eq!(model, "parse-legacy");
assert_eq!(adapter, OcrAdapterKind::ReductoLegacy);
let (model, adapter) =
resolve_wire_adapter("azure_ai/doc-intelligence/prebuilt-layout", None).unwrap();
assert_eq!(model, "doc-intelligence/prebuilt-layout");
assert_eq!(adapter, OcrAdapterKind::AzureDocumentIntelligence);
}
}

View file

@ -1,5 +1,4 @@
use std::collections::BTreeMap;
use std::convert::Infallible;
use std::path::PathBuf;
use std::sync::Arc;
use std::time::Duration;
@ -7,12 +6,15 @@ use std::time::Duration;
use bytes::Bytes;
use serde::{Deserialize, Serialize};
use serde_json::{Map, Value};
use serde_with::serde_as;
use litellm_auth::{InputSource, Sourced, TokenProviderHandle};
use super::hooks::{NoopOcrHooks, OcrHooks};
use super::registry::{OcrAdapterKind, resolve_wire_adapter};
use super::provider_config::{OcrConfigKind, resolve_provider_config};
use crate::call_arguments::CallArguments;
use crate::constants::OCR_HTTP_TIMEOUT_SECS;
use crate::ocr::Error;
use litellm_auth::{InputSource, TokenProviderHandle};
use crate::serde_compat::{FiniteF64, LaxI64};
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
#[serde(tag = "type")]
@ -21,13 +23,13 @@ pub enum OcrDocument {
DocumentUrl {
document_url: String,
#[serde(flatten)]
extra_fields: Map<String, Value>,
extra_fields: BTreeMap<String, String>,
},
#[serde(rename = "image_url")]
ImageUrl {
image_url: String,
#[serde(flatten)]
extra_fields: Map<String, Value>,
extra_fields: BTreeMap<String, String>,
},
}
@ -39,6 +41,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 +60,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, Debug, PartialEq)]
pub enum OcrDocumentInput {
Document(OcrDocument),
@ -76,6 +91,15 @@ impl From<OcrDocument> for OcrDocumentInput {
}
}
impl From<PathBuf> for OcrDocumentInput {
fn from(path: PathBuf) -> Self {
Self::Path {
path,
mime_type: None,
}
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub struct OcrFileContent {
pub bytes: Bytes,
@ -90,6 +114,107 @@ pub enum OcrResponseFormat {
Native,
}
#[derive(Clone, Default)]
pub struct OcrCredentialInputs {
pub api_key: Option<Sourced<String>>,
pub dynamic_api_key: Option<Sourced<String>>,
pub api_base: Option<Sourced<String>>,
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)>,
pub extra_headers_source: InputSource,
pub timeout: Duration,
pub max_download_bytes: u64,
pub max_response_bytes: usize,
pub poll_timeout: Duration,
}
impl Default for OcrTransportConfig {
fn default() -> Self {
Self {
extra_headers: Vec::new(),
extra_headers_source: InputSource::Deployment,
timeout: Duration::from_secs(OCR_HTTP_TIMEOUT_SECS),
max_download_bytes: crate::constants::OCR_DOWNLOAD_MAX_BYTES,
max_response_bytes: crate::constants::OCR_RESPONSE_MAX_BYTES,
poll_timeout: Duration::from_secs(crate::constants::OCR_POLL_TIMEOUT_SECS),
}
}
}
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())
}
/// Caller-supplied connection overrides for a [`LiteLLMOcrRequest`], in the
/// shape hosts receive them: JSON-ish headers, optional timeout, optional
/// credentials, and per-field provenance in `input_sources`.
#[derive(Clone, Debug, Default)]
pub struct OcrConnectionInputs {
pub api_key: Option<String>,
pub api_base: Option<String>,
pub extra_headers: Map<String, Value>,
pub timeout: Option<Duration>,
pub input_sources: BTreeMap<String, InputSource>,
}
impl OcrConnectionInputs {
fn source(&self, name: &str) -> InputSource {
self.input_sources.get(name).copied().unwrap_or_default()
}
fn header_pairs(&self) -> Result<Vec<(String, String)>, super::Error> {
self.extra_headers
.iter()
.map(|(name, value)| {
value
.as_str()
.map(|value| (name.clone(), value.to_string()))
.ok_or_else(|| super::Error::RequestField {
path: format!("extra_headers.{name}"),
})
})
.collect()
}
}
#[derive(Clone)]
pub struct OcrConnection {
pub api_key: Option<String>,
@ -104,72 +229,154 @@ pub struct OcrConnection {
pub poll_timeout: Duration,
}
impl Default for OcrConnection {
fn default() -> Self {
impl OcrConnection {
pub(crate) fn new(credentials: ResolvedOcrCredentials, transport: OcrTransportConfig) -> Self {
let api_key_source = credentials
.api_key
.as_ref()
.map(Sourced::source)
.unwrap_or(InputSource::Deployment);
let api_base_source = credentials
.api_base
.as_ref()
.map(Sourced::source)
.unwrap_or(InputSource::Deployment);
Self {
api_key: None,
api_key_source: InputSource::Deployment,
api_base: None,
api_base_source: InputSource::Deployment,
extra_headers: Vec::new(),
extra_headers_source: InputSource::Deployment,
timeout: Duration::from_secs(OCR_HTTP_TIMEOUT_SECS),
max_download_bytes: crate::constants::OCR_DOWNLOAD_MAX_BYTES,
max_response_bytes: crate::constants::OCR_RESPONSE_MAX_BYTES,
poll_timeout: Duration::from_secs(crate::constants::OCR_POLL_TIMEOUT_SECS),
api_key: credentials.api_key.map(Sourced::into_value),
api_key_source,
api_base: credentials.api_base.map(Sourced::into_value),
api_base_source,
extra_headers: transport.extra_headers,
extra_headers_source: transport.extra_headers_source,
timeout: transport.timeout,
max_download_bytes: transport.max_download_bytes,
max_response_bytes: transport.max_response_bytes,
poll_timeout: transport.poll_timeout,
}
}
}
pub struct LiteLLMOcrRequest<D = OcrDocument> {
pub model: String,
pub document: D,
pub connection: OcrConnection,
pub hooks: Arc<dyn OcrHooks>,
pub litellm_call_id: Option<String>,
pub optional_params: Map<String, Value>,
pub input_sources: BTreeMap<String, InputSource>,
pub azure_ad_token_provider: Option<TokenProviderHandle>,
pub(crate) adapter: OcrAdapterKind,
impl Default for OcrConnection {
fn default() -> Self {
Self::new(
ResolvedOcrCredentials::default(),
OcrTransportConfig::default(),
)
}
}
impl<D> LiteLLMOcrRequest<D> {
#[derive(Clone, Default)]
pub(crate) struct ResolvedOcrCredentials {
pub api_key: Option<Sourced<String>>,
pub api_base: Option<Sourced<String>>,
}
pub struct LiteLLMOcrRequest<D = OcrDocumentInput> {
pub model: String,
pub document: D,
pub credentials: OcrCredentialInputs,
pub transport: OcrTransportConfig,
pub hooks: Arc<dyn OcrHooks>,
pub litellm_call_id: Option<String>,
pub optional_params: CallArguments,
pub input_sources: BTreeMap<String, InputSource>,
pub azure_ad_token_provider: Option<TokenProviderHandle>,
pub(crate) config: OcrConfigKind,
}
impl LiteLLMOcrRequest {
pub fn new(
model: String,
document: D,
document: impl Into<OcrDocumentInput>,
custom_llm_provider: Option<&str>,
optional_params: Map<String, Value>,
) -> Result<Self, Error> {
let (model, adapter_kind) = resolve_wire_adapter(&model, custom_llm_provider)?;
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,
connection: OcrConnection::default(),
document: document.into(),
credentials: OcrCredentialInputs::default(),
transport,
hooks: Arc::new(NoopOcrHooks),
litellm_call_id: None,
optional_params,
input_sources: BTreeMap::new(),
azure_ad_token_provider: None,
adapter: adapter_kind,
config,
})
}
}
impl<D> LiteLLMOcrRequest<D> {
pub fn map_document<T, E>(
self,
map: impl FnOnce(D) -> Result<T, E>,
) -> Result<LiteLLMOcrRequest<T>, E> {
Ok(LiteLLMOcrRequest {
model: self.model,
document: map(self.document)?,
credentials: self.credentials,
transport: self.transport,
hooks: self.hooks,
litellm_call_id: self.litellm_call_id,
optional_params: self.optional_params,
input_sources: self.input_sources,
azure_ad_token_provider: self.azure_ad_token_provider,
config: self.config,
})
}
pub(crate) fn response_format(
&self,
) -> Result<OcrResponseFormat, super::error::OcrRequestError> {
pub fn with_document<T>(self, document: T) -> LiteLLMOcrRequest<T> {
LiteLLMOcrRequest {
model: self.model,
document,
credentials: self.credentials,
transport: self.transport,
hooks: self.hooks,
litellm_call_id: self.litellm_call_id,
optional_params: self.optional_params,
input_sources: self.input_sources,
azure_ad_token_provider: self.azure_ad_token_provider,
config: self.config,
}
}
pub(crate) fn response_format(&self) -> Result<OcrResponseFormat, super::Error> {
self.optional_params
.get("req_format")
.filter(|value| !value.is_null())
.map(|value| {
serde_json::from_value(value.clone())
.map_err(|_| super::error::OcrRequestError::RequestFormat)
serde_json::from_value(value.clone()).map_err(|_| super::Error::RequestFormat)
})
.transpose()
.map(|format| format.unwrap_or_default())
}
pub fn provider_name(&self) -> &'static str {
self.adapter.provider().as_str()
self.config.provider().into()
}
pub fn with_host_hooks(
@ -184,61 +391,335 @@ impl<D> LiteLLMOcrRequest<D> {
}
}
pub fn map_document<T, E>(
pub fn with_connection_inputs(
self,
map: impl FnOnce(D) -> Result<T, E>,
) -> Result<LiteLLMOcrRequest<T>, E> {
Ok(LiteLLMOcrRequest {
model: self.model,
document: map(self.document)?,
connection: self.connection,
hooks: self.hooks,
litellm_call_id: self.litellm_call_id,
optional_params: self.optional_params,
input_sources: self.input_sources,
azure_ad_token_provider: self.azure_ad_token_provider,
adapter: self.adapter,
})
}
pub fn with_document<T>(self, document: T) -> LiteLLMOcrRequest<T> {
let Ok(request) = self.map_document(|_| Ok::<T, Infallible>(document));
request
credentials: OcrCredentialInputs,
transport: OcrTransportConfig,
input_sources: BTreeMap<String, InputSource>,
) -> Self {
Self {
credentials,
transport,
input_sources,
..self
}
}
}
impl From<LiteLLMOcrRequest> for LiteLLMOcrRequest<OcrDocumentInput> {
fn from(request: LiteLLMOcrRequest) -> Self {
let Ok(request) = request
.map_document(|document| Ok::<_, Infallible>(OcrDocumentInput::Document(document)));
request
impl LiteLLMOcrRequest {
/// Builds a request from host-shaped inputs in one step: provider
/// resolution, optional-param validation, header/timeout overrides and
/// sourced credentials. Hosts should prefer this over sequencing
/// [`Self::new`], [`OcrTransportConfig::with_overrides`] and
/// [`Self::with_connection_inputs`] by hand.
pub fn from_inputs(
model: String,
document: impl Into<OcrDocumentInput>,
custom_llm_provider: Option<&str>,
optional_params: CallArguments,
connection: OcrConnectionInputs,
) -> Result<Self, super::Error> {
let request = Self::new(model, document, custom_llm_provider, optional_params)?;
let transport = request.transport.clone().with_overrides(
connection.header_pairs()?,
connection.source("extra_headers"),
connection.timeout,
);
let (api_key_source, api_base_source) =
(connection.source("api_key"), connection.source("api_base"));
let credentials = OcrCredentialInputs::new(
connection.api_key,
api_key_source,
connection.api_base,
api_base_source,
);
Ok(request.with_connection_inputs(credentials, transport, connection.input_sources))
}
}
pub(crate) type ResolvedOcrRequest = LiteLLMOcrRequest<OcrDocument>;
pub(crate) struct PreparedOcrRequest {
pub model: String,
pub document: OcrDocument,
pub connection: OcrConnection,
pub hooks: Arc<dyn OcrHooks>,
pub optional_params: CallArguments,
pub input_sources: BTreeMap<String, InputSource>,
pub azure_ad_token_provider: Option<TokenProviderHandle>,
pub(crate) config: OcrConfigKind,
}
impl PreparedOcrRequest {
pub(crate) fn new(request: ResolvedOcrRequest, connection: OcrConnection) -> Self {
let LiteLLMOcrRequest {
model,
document,
credentials: _,
transport: _,
hooks,
litellm_call_id: _,
optional_params,
input_sources,
azure_ad_token_provider,
config,
} = request;
Self {
model,
document,
connection,
hooks,
optional_params,
input_sources,
azure_ad_token_provider,
config,
}
}
pub(crate) fn response_format(&self) -> Result<OcrResponseFormat, super::Error> {
self.optional_params
.get("req_format")
.filter(|value| !value.is_null())
.map(|value| {
serde_json::from_value(value.clone()).map_err(|_| super::Error::RequestFormat)
})
.transpose()
.map(|format| format.unwrap_or_default())
}
pub(crate) fn provider_name(&self) -> &'static str {
self.config.provider().into()
}
}
#[serde_as]
#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)]
pub struct OcrPageDimensions {
#[serde_as(deserialize_as = "Option<LaxI64>")]
pub dpi: Option<i64>,
#[serde_as(deserialize_as = "Option<LaxI64>")]
pub height: Option<i64>,
#[serde_as(deserialize_as = "Option<LaxI64>")]
pub width: Option<i64>,
}
#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)]
pub struct OcrPageImage {
pub image_base64: Option<String>,
pub bbox: Option<Map<String, Value>>,
#[serde(flatten)]
pub extra_fields: Map<String, Value>,
}
#[serde_as]
#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)]
pub struct OcrPage {
#[serde_as(deserialize_as = "LaxI64")]
pub index: i64,
pub markdown: String,
pub images: Option<Vec<OcrPageImage>>,
pub dimensions: Option<OcrPageDimensions>,
#[serde(flatten)]
pub extra_fields: Map<String, Value>,
}
#[serde_as]
#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)]
pub struct OcrUsageInfo {
#[serde_as(deserialize_as = "Option<LaxI64>")]
pub pages_processed: Option<i64>,
#[serde_as(deserialize_as = "Option<LaxI64>")]
pub pages_processed_annotation: Option<i64>,
#[serde_as(deserialize_as = "Option<FiniteF64>")]
pub credits: Option<f64>,
#[serde_as(deserialize_as = "Option<LaxI64>")]
pub doc_size_bytes: Option<i64>,
#[serde(flatten)]
pub extra_fields: Map<String, Value>,
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
pub struct LiteLLMOcrResponse {
pub pages: Vec<Value>,
pub pages: Vec<OcrPage>,
pub model: String,
pub document_annotation: Option<Value>,
pub usage_info: Option<Value>,
pub usage_info: Option<OcrUsageInfo>,
pub content: Option<String>,
pub tables: Option<Vec<Map<String, Value>>>,
#[serde(rename = "keyValuePairs")]
pub key_value_pairs: Option<Vec<Map<String, Value>>>,
#[serde(default = "ocr_object")]
pub object: String,
#[serde(flatten)]
pub extra_fields: Map<String, Value>,
#[serde(skip_serializing_if = "Option::is_none")]
pub provider_native_response: Option<Value>,
pub provider_native_response: Option<Map<String, Value>>,
}
impl LiteLLMOcrResponse {
pub fn new(model: impl Into<String>, pages: Vec<OcrPage>) -> Self {
Self {
pages,
model: model.into(),
document_annotation: None,
usage_info: None,
content: None,
tables: None,
key_value_pairs: None,
object: ocr_object(),
extra_fields: Map::new(),
provider_native_response: None,
}
}
pub fn into_json(self) -> Value {
serde_json::to_value(self).expect("OCR response fields are JSON-compatible")
}
}
fn ocr_object() -> String {
"ocr".into()
}
#[cfg(test)]
mod tests {
use super::*;
use serde_json::json;
fn document() -> OcrDocument {
OcrDocument::try_from(
json!({"type":"document_url","document_url":"data:application/pdf;base64,YWJj"}),
)
.unwrap()
}
#[test]
fn from_inputs_applies_connection_overrides_with_field_sources() {
let request = LiteLLMOcrRequest::from_inputs(
"mistral/model".into(),
document(),
None,
Default::default(),
OcrConnectionInputs {
api_key: Some(" key ".into()),
api_base: Some("".into()),
extra_headers: json!({"x-a": "1"}).as_object().unwrap().clone(),
timeout: Some(Duration::from_secs(7)),
input_sources: [
("api_key".to_string(), InputSource::Request),
("extra_headers".to_string(), InputSource::Request),
]
.into(),
},
)
.unwrap();
let api_key = request.credentials.api_key.as_ref().unwrap();
assert_eq!(api_key.clone().into_value(), "key");
assert_eq!(api_key.source(), InputSource::Request);
assert!(request.credentials.api_base.is_none());
assert_eq!(
request.transport.extra_headers,
vec![("x-a".to_string(), "1".to_string())]
);
assert_eq!(request.transport.extra_headers_source, InputSource::Request);
assert_eq!(request.transport.timeout, Duration::from_secs(7));
assert_eq!(request.input_sources.len(), 2);
let defaulted = LiteLLMOcrRequest::from_inputs(
"mistral/model".into(),
document(),
None,
Default::default(),
OcrConnectionInputs::default(),
)
.unwrap();
assert_eq!(
defaulted.transport.timeout,
OcrTransportConfig::default().timeout
);
assert_eq!(
defaulted.transport.extra_headers_source,
InputSource::Deployment
);
}
#[test]
fn from_inputs_rejects_non_string_header_values_by_path() {
let Err(error) = LiteLLMOcrRequest::from_inputs(
"mistral/model".into(),
document(),
None,
Default::default(),
OcrConnectionInputs {
extra_headers: json!({"x-a": 1}).as_object().unwrap().clone(),
..Default::default()
},
) else {
panic!("non-string header value accepted");
};
assert!(matches!(
error,
super::super::Error::RequestField { ref path } if path == "extra_headers.x-a"
));
}
#[test]
fn normalized_response_rejects_invalid_shared_fields() {
for fields in [
json!({"pages":[{}]}),
json!({"pages":[{"index":0,"markdown":false}]}),
json!({"pages":[{"index":0,"markdown":"","images":[{"bbox":[]}]}]}),
json!({"usage_info":{"pages_processed":1.5}}),
json!({"tables":[false]}),
json!({"keyValuePairs":[[]]}),
json!({"provider_native_response":[]}),
] {
let payload: Map<String, Value> = json!({"model":"model", "pages":[]})
.as_object()
.unwrap()
.iter()
.chain(fields.as_object().unwrap())
.map(|(key, value)| (key.clone(), value.clone()))
.collect();
assert!(serde_json::from_value::<LiteLLMOcrResponse>(Value::Object(payload)).is_err());
}
assert!(
serde_json::from_value::<OcrDocument>(json!({
"type":"image_url", "image_url":"https://example.com/image", "detail":42
}))
.is_err()
);
}
#[test]
fn numeric_coercion_preserves_integer_precision_and_rejects_fractional_values() {
for (value, expected) in [
(json!("9007199254740993.0"), 9_007_199_254_740_993),
(json!("+2.000"), 2),
(json!("1_000"), 1000),
(json!(true), 1),
(json!(2.0), 2),
] {
let page: OcrPage =
serde_json::from_value(json!({"index":value,"markdown":""})).unwrap();
assert_eq!(page.index, expected);
}
for value in [
json!("1e2"),
json!(".0"),
json!("2."),
json!("_2"),
json!("2__0"),
json!(2.5),
json!(null),
] {
assert!(
serde_json::from_value::<OcrPage>(json!({"index":value,"markdown":""})).is_err()
);
}
}
#[test]
fn document_variants_preserve_provider_fields_when_rewriting_sources() {
for (value, original, replacement, expected) in [
@ -283,16 +764,11 @@ mod tests {
#[test]
fn response_serialization_flattens_extra_fields_and_omits_absent_native_response() {
let response = LiteLLMOcrResponse {
pages: vec![],
model: "model".into(),
document_annotation: None,
usage_info: None,
object: "ocr".into(),
extra_fields: json!({"provider_field":"kept"})
.as_object()
.unwrap()
.clone(),
provider_native_response: None,
..LiteLLMOcrResponse::new("model", vec![])
};
let serialized = response.into_json();
assert_eq!(serialized["provider_field"], "kept");

View file

@ -1,69 +1,40 @@
use crate::ocr::error::OcrRequestError;
use crate::ocr::error::OcrResponseError;
use std::collections::BTreeMap;
use std::time::Duration;
use super::types::{LiteLLMOcrRequest, OcrConnection, OcrDocument};
use crate::ocr::Error;
use litellm_auth::InputSource;
use serde::{
Deserialize,
de::{DeserializeOwned, IntoDeserializer},
};
use serde::Deserialize;
use serde_json::{Map, Value};
const COMMON_OPTION_FIELDS: &[&str] = &["req_format", "extra_body", "max_response_bytes"];
const MISTRAL_OPTION_FIELDS: &[&str] = &[
"pages",
"include_image_base64",
"image_limit",
"image_min_size",
"bbox_annotation_format",
"document_annotation_format",
"document_annotation_prompt",
"extract_header",
"extract_footer",
"table_format",
"confidence_scores_granularity",
"include_blocks",
"id",
];
const DEEPSEEK_OPTION_FIELDS: &[&str] =
&["stream", "temperature", "max_tokens", "top_p", "n", "stop"];
const DOCUMENT_INTELLIGENCE_OPTION_FIELDS: &[&str] = &["pages", "features"];
const REDUCTO_V3_OPTION_FIELDS: &[&str] = &["formatting", "retrieval", "settings"];
const REDUCTO_LEGACY_OPTION_FIELDS: &[&str] = &["enhance"];
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 use super::is_supported_request;
use super::{Error, LiteLLMOcrRequest, OcrConnectionInputs, OcrDocument, OcrDocumentInput};
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct OptionalParamSpec {
pub name: &'static str,
pub secret: bool,
pub fn consumed_optional_params(
model: &str,
provider: Option<&str>,
) -> Result<Vec<crate::call_arguments::ArgumentSpec>, Error> {
let specs = super::consumed_optional_params(model, provider)?;
Ok(consumed_optional_param_names(model, provider)?
.into_iter()
.map(|name| crate::call_arguments::ArgumentSpec {
name,
secret: specs.iter().any(|spec| spec.name == name && spec.secret),
})
.collect())
}
#[derive(Debug)]
pub struct DecodedOcrResponse<T> {
pub data: T,
pub native: Option<Value>,
pub text: String,
pub fn consumed_optional_param_names(
model: &str,
provider: Option<&str>,
) -> Result<Vec<&'static str>, Error> {
let names = super::consumed_optional_param_names(model, provider)?;
let (_, config) = super::provider_config::resolve_provider_config(model, provider)?;
if config == super::provider_config::OcrConfigKind::VertexDeepSeek {
return Ok(names
.into_iter()
.chain(["stream", "temperature", "max_tokens", "top_p", "n", "stop"])
.collect());
}
Ok(names)
}
#[derive(Deserialize)]
@ -82,216 +53,54 @@ pub struct OcrWireRequest<D = Value> {
pub timeout_seconds: Option<f64>,
}
pub fn is_supported_request(model: &str, custom_llm_provider: Option<&str>) -> bool {
super::registry::resolve_wire_adapter(model, custom_llm_provider).is_ok()
}
pub fn consumed_optional_param_names(
model: &str,
custom_llm_provider: Option<&str>,
) -> Result<Vec<&'static str>, Error> {
use super::registry::OcrAdapterKind;
let (_, adapter) = super::registry::resolve_wire_adapter(model, custom_llm_provider)?;
let provider_fields: &[&str] = match adapter {
OcrAdapterKind::Cohere | OcrAdapterKind::AzureCohere => &["output_format"],
OcrAdapterKind::Mistral | OcrAdapterKind::AzureMistral | OcrAdapterKind::VertexMistral => {
MISTRAL_OPTION_FIELDS
}
OcrAdapterKind::AzureDocumentIntelligence => DOCUMENT_INTELLIGENCE_OPTION_FIELDS,
OcrAdapterKind::ReductoV3 => REDUCTO_V3_OPTION_FIELDS,
OcrAdapterKind::ReductoLegacy => REDUCTO_LEGACY_OPTION_FIELDS,
OcrAdapterKind::VertexDeepSeek => DEEPSEEK_OPTION_FIELDS,
};
let auth_fields: &[&str] = match adapter {
OcrAdapterKind::AzureMistral
| OcrAdapterKind::AzureDocumentIntelligence
| OcrAdapterKind::AzureCohere => AZURE_AUTH_OPTION_FIELDS,
OcrAdapterKind::VertexMistral | OcrAdapterKind::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<OptionalParamSpec>, Error> {
consumed_optional_param_names(model, custom_llm_provider).map(|names| {
names
.into_iter()
.map(|name| OptionalParamSpec {
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, Error> {
let OcrWireRequest {
model,
document,
api_key,
api_base,
custom_llm_provider,
extra_headers,
optional_params,
input_sources,
timeout_seconds,
} = wire;
decode_request_input(OcrWireRequest {
model,
document: decode_document(document)?,
api_key,
api_base,
custom_llm_provider,
extra_headers,
optional_params,
input_sources,
timeout_seconds,
model: wire.model,
document: decode_document(wire.document)?,
api_key: wire.api_key,
api_base: wire.api_base,
custom_llm_provider: wire.custom_llm_provider,
extra_headers: wire.extra_headers,
optional_params: wire.optional_params,
input_sources: wire.input_sources,
timeout_seconds: wire.timeout_seconds,
})
}
pub fn decode_request_input<D>(wire: OcrWireRequest<D>) -> Result<LiteLLMOcrRequest<D>, 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 headers = wire
.extra_headers
.unwrap_or_default()
.into_iter()
.map(|(name, value)| {
let value = value
.as_str()
.ok_or_else(|| OcrRequestError::RequestField {
path: format!("extra_headers.{name}"),
})?;
Ok((name, value.to_string()))
})
.collect::<Result<Vec<_>, OcrRequestError>>()?;
pub fn decode_request_input<D: Into<OcrDocumentInput>>(
wire: OcrWireRequest<D>,
) -> Result<LiteLLMOcrRequest, Error> {
let timeout = wire
.timeout_seconds
.map(|seconds| {
Duration::try_from_secs_f64(seconds).map_err(|_| OcrRequestError::RequestField {
Duration::try_from_secs_f64(seconds).map_err(|_| Error::RequestField {
path: "timeout_seconds".into(),
})
})
.transpose()?;
let defaults = OcrConnection::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(|| OcrRequestError::RequestField {
path: "max_response_bytes".into(),
})
})
.transpose()?
.unwrap_or(defaults.max_response_bytes);
let request = LiteLLMOcrRequest::new(
LiteLLMOcrRequest::from_inputs(
wire.model,
wire.document,
wire.custom_llm_provider.as_deref(),
wire.optional_params
.into_iter()
.filter(|(name, _)| name != "max_response_bytes")
.collect(),
)?;
let connection = OcrConnection {
api_key: nonblank(wire.api_key),
api_key_source,
api_base: nonblank(wire.api_base),
api_base_source,
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 {
connection,
input_sources: wire.input_sources,
..request
})
wire.optional_params.into(),
OcrConnectionInputs {
api_key: wire.api_key,
api_base: wire.api_base,
extra_headers: wire.extra_headers.unwrap_or_default(),
timeout,
input_sources: wire.input_sources,
},
)
}
pub fn decode_document(value: Value) -> Result<OcrDocument, 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(OcrRequestError::MissingDocumentUrl.into());
if matches!(kind, Some("document_url")) && value.get("document_url").is_none()
|| matches!(kind, Some("image_url")) && value.get("image_url").is_none()
{
return Err(Error::MissingDocumentUrl);
}
Ok(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, OcrRequestError> {
serde_path_to_error::deserialize(value.into_deserializer()).map_err(|error| {
OcrRequestError::RequestField {
path: format!("{prefix}.{}", error.path()),
}
})
}
pub fn decode_response<T: DeserializeOwned>(
bytes: &[u8],
native: bool,
) -> Result<DecodedOcrResponse<T>, OcrResponseError> {
let mut deserializer = serde_json::Deserializer::from_slice(bytes);
let data = serde_path_to_error::deserialize(&mut deserializer).map_err(|error| {
OcrResponseError::ResponseField {
path: error.path().to_string(),
}
})?;
deserializer
.end()
.map_err(|_| OcrResponseError::ResponseField {
path: "response".into(),
})?;
let native = if native {
Some(
serde_json::from_slice(bytes).map_err(|_| OcrResponseError::ResponseField {
path: "response".into(),
})?,
)
} else {
None
};
Ok(DecodedOcrResponse {
data,
native,
text: String::from_utf8_lossy(bytes).into_owned(),
})
super::json::decode_request_value(value, "document")
}
#[cfg(test)]
@ -305,7 +114,6 @@ mod tests {
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"));
@ -358,7 +166,10 @@ mod tests {
serde_json::json!({"type": "document_url"}),
serde_json::json!({"type": "image_url"}),
] {
assert_eq!(decode_document(document), Err(Error::MissingDocumentUrl));
assert!(matches!(
decode_document(document),
Err(Error::MissingDocumentUrl)
));
}
}
}

View file

@ -0,0 +1,231 @@
#[derive(Clone, Debug, PartialEq, Eq, thiserror::Error)]
pub enum Error {
#[error("invalid request: extra_body must be an object")]
ExtraBody,
#[error("invalid request: body must be a JSON object")]
Body,
}
use std::ops::{Deref, DerefMut};
use serde::{Deserialize, Serialize};
use serde_json::{Map, Value};
#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)]
#[serde(transparent)]
pub struct OpaqueParams(Map<String, Value>);
pub fn is_control_param(name: &str) -> bool {
matches!(
name,
"api_key"
| "api_base"
| "custom_llm_provider"
| "extra_headers"
| "timeout"
| "timeout_seconds"
| "request_timeout"
| "max_retries"
| "req_format"
| "max_response_bytes"
| "litellm_call_id"
| "litellm_logging_obj"
| "litellm_metadata"
| "proxy_server_request"
| "callbacks"
| "success_callback"
| "failure_callback"
| "guardrails"
| "azure_ad_token"
| "azure_ad_token_provider"
| "tenant_id"
| "client_id"
| "client_secret"
| "azure_scope"
| "azure_authority_host"
| "azure_credential"
| "azure_federated_token_file"
| "enable_azure_ad_token_refresh"
| "vertex_credentials"
| "vertex_ai_credentials"
| "vertex_project"
| "vertex_ai_project"
| "vertex_location"
| "vertex_ai_location"
| "aws_access_key_id"
| "aws_secret_access_key"
| "aws_session_token"
| "aws_region_name"
| "aws_session_name"
| "aws_profile_name"
| "aws_role_name"
| "aws_web_identity_token"
| "aws_sts_endpoint"
| "aws_external_id"
| "aws_bedrock_runtime_endpoint"
)
}
impl OpaqueParams {
pub fn into_inner(self) -> Map<String, Value> {
self.0
}
pub fn without(&self, names: &[&str]) -> Self {
self.iter()
.filter(|(name, _)| !names.contains(&name.as_str()))
.map(|(name, value)| (name.clone(), value.clone()))
.collect()
}
pub fn provider_params(&self) -> Self {
self.iter()
.filter(|(name, _)| !is_control_param(name))
.map(|(name, value)| (name.clone(), value.clone()))
.collect()
}
pub fn into_provider_body(self) -> Result<Map<String, Value>, Error> {
let mut fields = self.0;
let overrides = match fields.remove("extra_body") {
None | Some(Value::Null) => Map::new(),
Some(Value::Object(fields)) => fields,
Some(_) => {
return Err(Error::ExtraBody);
}
};
Ok(fields
.into_iter()
.chain(overrides)
.filter(|(name, _)| name != "extra_body" && !is_control_param(name))
.collect())
}
}
#[cfg(test)]
fn merge_extra_params<B: Serialize>(body: &B, extra_params: OpaqueParams) -> Result<Value, Error> {
let Value::Object(fields) = serde_json::to_value(body).map_err(|_| Error::Body)? else {
return Err(Error::Body);
};
Ok(Value::Object(
fields
.into_iter()
.chain(
extra_params
.into_provider_body()?
.into_iter()
.filter(|(name, _)| name != "model"),
)
.collect(),
))
}
impl Deref for OpaqueParams {
type Target = Map<String, Value>;
fn deref(&self) -> &Self::Target {
&self.0
}
}
impl DerefMut for OpaqueParams {
fn deref_mut(&mut self) -> &mut Self::Target {
&mut self.0
}
}
impl From<Map<String, Value>> for OpaqueParams {
fn from(value: Map<String, Value>) -> Self {
Self(value)
}
}
impl From<OpaqueParams> for Map<String, Value> {
fn from(value: OpaqueParams) -> Self {
value.0
}
}
impl FromIterator<(String, Value)> for OpaqueParams {
fn from_iter<T: IntoIterator<Item = (String, Value)>>(iter: T) -> Self {
Self(iter.into_iter().collect())
}
}
impl IntoIterator for OpaqueParams {
type Item = (String, Value);
type IntoIter = serde_json::map::IntoIter;
fn into_iter(self) -> Self::IntoIter {
self.0.into_iter()
}
}
#[cfg(test)]
mod tests {
use serde_json::json;
use super::*;
#[test]
fn extras_merge_shallowly_and_preserve_values_without_leaking_controls() {
let extras: OpaqueParams = serde_json::from_value(json!({
"future": {"nested": [false, 0, null]},
"explicit_null": null,
"azure_ad_token": "secret",
"req_format": "native",
"extra_body": {
"future": {"replacement": true},
"temperature": 0.5,
"model": "override",
"aws_secret_access_key": "secret"
}
}))
.unwrap();
let body =
merge_extra_params(&json!({"model":"resolved", "temperature":0.1}), extras).unwrap();
assert_eq!(
body,
json!({
"model":"resolved", "temperature":0.5,
"future":{"replacement":true}, "explicit_null":null
})
);
}
#[test]
fn invalid_extra_body_is_rejected_and_null_is_empty() {
for value in [json!(false), json!([]), json!("value"), json!(1)] {
let params: OpaqueParams = serde_json::from_value(json!({"extra_body":value})).unwrap();
assert!(params.into_provider_body().is_err());
}
let params: OpaqueParams =
serde_json::from_value(json!({"extra_body":null,"future":null})).unwrap();
assert_eq!(
Value::Object(params.into_provider_body().unwrap()),
json!({"future":null})
);
}
#[test]
fn provider_params_preserve_opaque_values() {
let params: OpaqueParams = serde_json::from_value(json!({
"object": {"future": [1, null]},
"null": null,
"azure_ad_token": "secret"
}))
.unwrap();
let retained = params.provider_params();
assert_eq!(
serde_json::to_value(retained).unwrap(),
json!({"object": {"future": [1, null]}, "null": null})
);
}
#[test]
fn outer_value_must_be_an_object() {
assert!(serde_json::from_value::<OpaqueParams>(json!(["value"])).is_err());
}
}

View file

@ -2,4 +2,5 @@ pub mod anthropic;
pub mod azure_ai;
pub mod bedrock;
pub mod custom_llm_provider;
pub(crate) mod model;
pub mod openai;

View file

@ -0,0 +1,219 @@
use std::marker::PhantomData;
use serde::{Deserialize, Deserializer, Serialize, Serializer};
#[derive(Clone, Copy, Debug, Eq, PartialEq, thiserror::Error)]
pub(crate) enum ModelNameError {
#[error("model name cannot be empty")]
EmptyModel,
#[error("model namespace must be one non-empty path segment: {0}")]
InvalidNamespace(&'static str),
}
pub(crate) trait ModelNamespace {
const NAME: &'static str;
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub(crate) struct RoutedModel<'a>(&'a str);
impl<'a> RoutedModel<'a> {
pub(crate) fn new(value: &'a str) -> Result<Self, ModelNameError> {
if value.is_empty() {
return Err(ModelNameError::EmptyModel);
}
Ok(Self(value))
}
pub(crate) fn into_provider<N: ModelNamespace>(
self,
) -> Result<ProviderModel<N>, ModelNameError> {
let namespace = N::NAME;
if namespace.is_empty() || namespace.contains('/') {
return Err(ModelNameError::InvalidNamespace(namespace));
}
let prefix = format!("{namespace}/");
let local_model = self.0.trim_start_matches(prefix.as_str());
if local_model.is_empty() {
return Err(ModelNameError::EmptyModel);
}
Ok(ProviderModel {
value: format!("{prefix}{local_model}"),
namespace: PhantomData,
})
}
}
#[derive(Clone, Debug, Eq, PartialEq)]
pub(crate) struct ProviderModel<N> {
value: String,
namespace: PhantomData<N>,
}
impl<N> ProviderModel<N> {
#[cfg(test)]
pub(crate) fn as_str(&self) -> &str {
&self.value
}
}
impl<N> Serialize for ProviderModel<N> {
fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
where
S: Serializer,
{
self.value.serialize(serializer)
}
}
impl<'de, N: ModelNamespace> Deserialize<'de> for ProviderModel<N> {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: Deserializer<'de>,
{
let value = String::deserialize(deserializer)?;
RoutedModel::new(&value)
.and_then(RoutedModel::into_provider::<N>)
.map_err(<D::Error as serde::de::Error>::custom)
}
}
#[cfg(test)]
mod tests {
use serde_json::json;
use super::*;
#[derive(Clone, Debug, Eq, PartialEq)]
struct DeepSeekAi;
impl ModelNamespace for DeepSeekAi {
const NAME: &'static str = "deepseek-ai";
}
#[derive(Clone, Debug, Eq, PartialEq)]
struct FalAi;
impl ModelNamespace for FalAi {
const NAME: &'static str = "fal-ai";
}
#[test]
fn qualifies_a_bare_model() {
let model = RoutedModel::new("deepseek-ocr-maas")
.and_then(RoutedModel::into_provider::<DeepSeekAi>)
.unwrap();
assert_eq!(model.as_str(), "deepseek-ai/deepseek-ocr-maas");
}
#[test]
fn preserves_an_already_qualified_model() {
let model = RoutedModel::new("deepseek-ai/deepseek-ocr-maas")
.and_then(RoutedModel::into_provider::<DeepSeekAi>)
.unwrap();
assert_eq!(model.as_str(), "deepseek-ai/deepseek-ocr-maas");
}
#[test]
fn collapses_repeated_owned_namespaces() {
let model = RoutedModel::new("deepseek-ai/deepseek-ai/deepseek-ai/deepseek-ocr-maas")
.and_then(RoutedModel::into_provider::<DeepSeekAi>)
.unwrap();
assert_eq!(model.as_str(), "deepseek-ai/deepseek-ocr-maas");
}
#[test]
fn matches_the_namespace_as_a_complete_segment() {
let model = RoutedModel::new("deepseek-ai-v2/model")
.and_then(RoutedModel::into_provider::<DeepSeekAi>)
.unwrap();
assert_eq!(model.as_str(), "deepseek-ai/deepseek-ai-v2/model");
}
#[test]
fn preserves_nested_provider_model_paths() {
let model = RoutedModel::new("publishers/vendor/models/model-v1")
.and_then(RoutedModel::into_provider::<DeepSeekAi>)
.unwrap();
assert_eq!(
model.as_str(),
"deepseek-ai/publishers/vendor/models/model-v1"
);
}
#[test]
fn namespace_markers_select_different_wire_names() {
let routed = RoutedModel::new("model-v1").unwrap();
let deepseek = routed.into_provider::<DeepSeekAi>().unwrap();
let fal = routed.into_provider::<FalAi>().unwrap();
assert_eq!(deepseek.as_str(), "deepseek-ai/model-v1");
assert_eq!(fal.as_str(), "fal-ai/model-v1");
}
#[test]
fn rejects_empty_routed_models() {
assert_eq!(RoutedModel::new(""), Err(ModelNameError::EmptyModel));
}
#[test]
fn rejects_a_namespace_without_a_model() {
let result =
RoutedModel::new("deepseek-ai/").and_then(RoutedModel::into_provider::<DeepSeekAi>);
assert_eq!(result, Err(ModelNameError::EmptyModel));
}
#[test]
fn rejects_invalid_namespace_markers() {
struct Empty;
impl ModelNamespace for Empty {
const NAME: &'static str = "";
}
struct MultipleSegments;
impl ModelNamespace for MultipleSegments {
const NAME: &'static str = "one/two";
}
assert!(matches!(
RoutedModel::new("model").and_then(RoutedModel::into_provider::<Empty>),
Err(ModelNameError::InvalidNamespace(""))
));
assert!(matches!(
RoutedModel::new("model").and_then(RoutedModel::into_provider::<MultipleSegments>),
Err(ModelNameError::InvalidNamespace("one/two"))
));
}
#[test]
fn provider_models_serialize_as_plain_strings() {
let model = RoutedModel::new("deepseek-ocr-maas")
.and_then(RoutedModel::into_provider::<DeepSeekAi>)
.unwrap();
assert_eq!(
serde_json::to_value(model).unwrap(),
json!("deepseek-ai/deepseek-ocr-maas")
);
}
#[test]
fn deserialization_reestablishes_the_namespace_invariant() {
let model: ProviderModel<DeepSeekAi> =
serde_json::from_value(json!("deepseek-ai/deepseek-ai/model-v1")).unwrap();
assert_eq!(model.as_str(), "deepseek-ai/model-v1");
}
#[test]
fn deserialization_rejects_missing_model_names() {
let result = serde_json::from_value::<ProviderModel<DeepSeekAi>>(json!("deepseek-ai/"));
assert!(result.is_err());
}
}

View file

@ -0,0 +1,151 @@
use serde::{Deserialize, Deserializer, de::Error};
use serde_json::Value;
use serde_with::DeserializeAs;
pub(crate) struct LaxI64;
pub(crate) struct FiniteF64;
impl<'de> DeserializeAs<'de, i64> for LaxI64 {
fn deserialize_as<D: Deserializer<'de>>(deserializer: D) -> Result<i64, D::Error> {
match Value::deserialize(deserializer)? {
Value::Number(number) if number.is_f64() => number.as_f64().and_then(integral_float),
Value::Number(number) => number.as_i64(),
Value::String(value) => integer_string(value.trim()),
Value::Bool(value) => Some(i64::from(value)),
_ => None,
}
.ok_or_else(|| D::Error::custom("expected an integer in the i64 range"))
}
}
impl<'de> DeserializeAs<'de, f64> for FiniteF64 {
fn deserialize_as<D: Deserializer<'de>>(deserializer: D) -> Result<f64, D::Error> {
match Value::deserialize(deserializer)? {
Value::Number(number) => number.as_f64(),
Value::String(value) => value.trim().parse::<f64>().ok(),
Value::Bool(value) => Some(f64::from(value)),
_ => None,
}
.filter(|value| value.is_finite())
.ok_or_else(|| D::Error::custom("expected a finite number"))
}
}
fn integer_string(value: &str) -> Option<i64> {
let integer = match value.split_once('.') {
Some((integer, fraction)) => {
if fraction.is_empty() || !fraction.bytes().all(|byte| byte == b'0') {
return None;
}
integer
}
None => value,
};
if integer.starts_with('_') || integer.ends_with('_') || integer.contains("__") {
return None;
}
let digits = integer.strip_prefix(['+', '-']).unwrap_or(integer);
if digits.is_empty()
|| digits.starts_with('_')
|| !digits
.bytes()
.all(|byte| byte.is_ascii_digit() || byte == b'_')
{
return None;
}
integer.replace('_', "").parse().ok()
}
fn integral_float(value: f64) -> Option<i64> {
(value.is_finite()
&& value.fract() == 0.0
&& value >= i64::MIN as f64
&& value < -(i64::MIN as f64))
.then_some(value as i64)
}
#[cfg(test)]
mod tests {
use super::*;
use serde::Serialize;
use serde_json::json;
use serde_with::serde_as;
#[serde_as]
#[derive(Debug, Deserialize, Serialize, PartialEq)]
struct Numbers {
#[serde_as(deserialize_as = "Option<Vec<LaxI64>>")]
integers: Option<Vec<i64>>,
#[serde_as(deserialize_as = "Option<FiniteF64>")]
float: Option<f64>,
}
#[test]
fn adapters_compose_and_serialize_as_numbers() {
let numbers: Numbers = serde_json::from_value(json!({
"integers": ["9007199254740993.0", "1_000", " +2.000 ", 3.0, true],
"float": " 1.5 "
}))
.unwrap();
assert_eq!(
serde_json::to_value(numbers).unwrap(),
json!({
"integers": [9_007_199_254_740_993_i64, 1000, 2, 3, 1], "float": 1.5
})
);
for input in [json!({}), json!({"integers": null, "float": null})] {
assert_eq!(
serde_json::from_value::<Numbers>(input).unwrap(),
Numbers {
integers: None,
float: None,
}
);
}
}
#[test]
fn integer_bounds_and_invalid_values_are_checked() {
for input in [
json!(i64::MIN),
json!(i64::MAX),
json!(i64::MAX.to_string()),
] {
assert!(serde_json::from_value::<Numbers>(json!({"integers": [input]})).is_ok());
}
for input in [
json!(u64::MAX),
json!(9_223_372_036_854_775_808_u64),
json!(9_223_372_036_854_775_808.0),
json!("-9223372036854775809"),
json!("1.0000000000000001"),
json!("1e3"),
json!("2."),
json!(".0"),
json!("_2"),
json!("2__0"),
json!(2.5),
json!(null),
json!({}),
] {
assert!(serde_json::from_value::<Numbers>(json!({"integers": [input]})).is_err());
}
}
#[test]
fn floats_reject_nonfinite_and_invalid_values() {
for input in [
json!("NaN"),
json!("inf"),
json!("-inf"),
json!("1e999"),
json!([]),
] {
assert!(serde_json::from_value::<Numbers>(json!({"float": input})).is_err());
}
for (input, expected) in [(json!(2), 2.0), (json!(2.5), 2.5), (json!(true), 1.0)] {
let numbers: Numbers = serde_json::from_value(json!({"float": input})).unwrap();
assert_eq!(numbers.float, Some(expected));
}
}
}

View file

@ -17,15 +17,15 @@ async fn facade_executes_azure_mistral_with_prepared_auth() {
&base,
json!({"include_image_base64":true}),
);
request.connection.api_key = None;
request.connection.extra_headers = vec![(
request.credentials.api_key = None;
request.transport.extra_headers = vec![(
"Authorization".into(),
"Bearer python-prepared-token".into(),
)];
let result = perform_ocr(request).await.unwrap();
server.await.unwrap();
assert_eq!(result.pages[0]["markdown"], "hello");
assert_eq!(result.pages[0].markdown, "hello");
let requests = seen.lock().unwrap();
assert_eq!(requests.len(), 1);
assert!(requests[0].starts_with("POST /providers/mistral/azure/ocr "));
@ -53,7 +53,7 @@ async fn facade_acquires_supplied_entra_token_for_final_request() {
&base,
json!({"azure_ad_token":"rust-owned-token"}),
);
request.connection.api_key = None;
request.credentials.api_key = None;
perform_ocr(request).await.unwrap();
server.await.unwrap();

View file

@ -23,11 +23,12 @@ async fn facade_maps_pages_features_and_url_document() {
&base,
json!({"pages":[2,0,0,1],"features":["keyValuePairs","languages"]}),
);
request.document = serde_json::from_value(json!({
request.document = serde_json::from_value::<super::OcrDocument>(json!({
"type":"document_url",
"document_url":"https://example.com/document.pdf"
}))
.unwrap();
.unwrap()
.into();
perform_ocr(request).await.unwrap();
server.await.unwrap();
@ -118,13 +119,13 @@ async fn immediate_response_normalizes_pages_and_preserves_native() {
.unwrap();
server.await.unwrap();
assert_eq!(result.pages[0]["index"], 1);
assert_eq!(result.pages[0]["markdown"], "A\n\nB");
assert_eq!(result.pages[0].index, 1);
assert_eq!(result.pages[0].markdown, "A\n\nB");
assert_eq!(
result.pages[0]["dimensions"],
serde_json::to_value(&result.pages[0].dimensions).unwrap(),
json!({"width":816,"height":1056,"dpi":96})
);
assert_eq!(result.usage_info, Some(json!({"pages_processed":1})));
assert_eq!(result.usage_info.as_ref().unwrap().pages_processed, Some(1));
let serialized = result.clone().into_json();
assert_eq!(serialized["content"], "A\n\nB");
assert_eq!(serialized["tables"], json!([{"cells":[]}]));
@ -133,7 +134,10 @@ async fn immediate_response_normalizes_pages_and_preserves_native() {
json!([{"key":{"content":"A"}}])
);
assert!(serialized.get("key_value_pairs").is_none());
assert_eq!(result.provider_native_response, Some(operation));
assert_eq!(
result.provider_native_response.map(Value::Object),
Some(operation)
);
}
#[tokio::test]
@ -159,13 +163,16 @@ async fn accepted_response_polls_to_success_with_only_credentials() {
json!({"req_format":"native"}),
);
request
.connection
.transport
.extra_headers
.push(("X-Trace".into(), "initial-only".into()));
let result = perform_ocr(request).await.unwrap();
server.await.unwrap();
assert_eq!(result.provider_native_response, Some(operation));
assert_eq!(
result.provider_native_response.map(Value::Object),
Some(operation)
);
let requests = seen.lock().unwrap();
assert_eq!(requests.len(), 3);
assert!(requests[0].to_ascii_lowercase().contains("x-trace:"));
@ -239,8 +246,8 @@ async fn polling_forwards_bearer_credentials() {
])
.await;
let mut request = wire_request("azure_ai/doc-intelligence/prebuilt-read", &base, json!({}));
request.connection.api_key = None;
request.connection.extra_headers = vec![("Authorization".into(), "Bearer token".into())];
request.credentials.api_key = None;
request.transport.extra_headers = vec![("Authorization".into(), "Bearer token".into())];
perform_ocr(request).await.unwrap();
server.await.unwrap();
@ -375,7 +382,7 @@ async fn polling_deadline_bounds_retry_delay() {
])
.await;
let mut request = wire_request("azure_ai/doc-intelligence/prebuilt-read", &base, json!({}));
request.connection.poll_timeout = std::time::Duration::from_millis(100);
request.transport.poll_timeout = std::time::Duration::from_millis(100);
let error = tokio::time::timeout(std::time::Duration::from_secs(1), perform_ocr(request))
.await

View file

@ -1,8 +1,10 @@
use rstest::rstest;
use serde_json::{Value, json};
use crate::ocr::codecs::deepseek::{
DeepSeekOcrParams, DeepSeekOcrResponse, transform_ocr_request, transform_ocr_response,
use crate::llms::base_llm::ocr::transformation::BaseOcrConfig;
use crate::llms::vertex_ai::ocr::deepseek_transformation::{
DeepSeekOcrParams, DeepSeekOcrResponse, VertexAIDeepSeekOCRConfig,
normalize_response as transform_ocr_response,
};
use crate::ocr::types::OcrDocument;
@ -22,7 +24,9 @@ fn request_mapping_matches_python(#[case] name: &str, #[case] value: Value) {
let params: DeepSeekOcrParams =
serde_json::from_value(json!({name: value.clone(), "ignored": true})).unwrap();
let result = serde_json::to_value(
transform_ocr_request("deepseek-ai/deepseek-ocr-maas", document(), &params).unwrap(),
VertexAIDeepSeekOCRConfig
.transform_ocr_request("deepseek-ai/deepseek-ocr-maas", document(), &params, &[])
.unwrap(),
)
.unwrap();
assert_eq!(result["model"], "deepseek-ai/deepseek-ocr-maas");
@ -43,12 +47,14 @@ fn request_maps_both_document_types_to_image_content(#[case] document: Value) {
.or_else(|| document.get("document_url"))
.unwrap()
.clone();
let request = transform_ocr_request(
"deepseek-ai/deepseek-ocr-maas",
serde_json::from_value(document).unwrap(),
&DeepSeekOcrParams::default(),
)
.unwrap();
let request = VertexAIDeepSeekOCRConfig
.transform_ocr_request(
"deepseek-ai/deepseek-ocr-maas",
serde_json::from_value(document).unwrap(),
&DeepSeekOcrParams::default(),
&[],
)
.unwrap();
let result = serde_json::to_value(request).unwrap();
assert_eq!(
result["messages"][0]["content"][0],
@ -60,12 +66,17 @@ fn request_maps_both_document_types_to_image_content(#[case] document: Value) {
#[case(json!("# hello"), "# hello")]
#[case(json!("{broken"), "{broken")]
#[case(json!(" {\"pages\":[]} "), " {\"pages\":[]} ")]
#[case(json!({"pages":[]}), "{\"pages\":[]}")]
#[case(json!({}), "{}")]
#[case(json!({"pages":[]}), "")]
#[case(json!("[]"), "[]")]
#[case(json!("{\"pages\":[{\"markdown\":\"json text\"}]}"), "json text")]
#[case(json!({"pages":[{"markdown":"object"}]}), "object")]
fn response_codec_handles_text_json_and_objects(#[case] content: Value, #[case] expected: &str) {
let structured = content
.as_object()
.is_some_and(|object| object.contains_key("pages"))
|| content
.as_str()
.is_some_and(|text| text.contains("\"pages\""));
let response: DeepSeekOcrResponse = serde_json::from_value(
json!({"choices":[{"message":{"content":content}}],"usage":{"prompt_tokens":1}}),
)
@ -75,7 +86,11 @@ fn response_codec_handles_text_json_and_objects(#[case] content: Value, #[case]
.into_json();
assert_eq!(result["pages"][0]["markdown"], expected);
assert_eq!(result["pages"][0]["index"], 0);
assert_eq!(result["usage_info"]["prompt_tokens"], 1);
if structured {
assert!(result["usage_info"].is_null());
} else {
assert_eq!(result["usage_info"]["prompt_tokens"], 1);
}
}
#[test]
@ -104,6 +119,7 @@ fn structured_result_maps_pages_usage_model_and_annotation() {
#[test]
fn response_codec_rejects_missing_empty_and_malformed_content() {
for value in [
json!({"choices":[{"message":{"content":{}}}]}),
json!({"choices":[]}),
json!({"choices":[{"message":{"content":""}}]}),
json!({"choices":[{"message":{"content":"{\"pages\":[{\"markdown\":42}]}"}}]}),

View file

@ -83,10 +83,10 @@ fn failure_handler_errors_do_not_replace_selected_failure_or_suppress_async_disp
lifecycle.accept::<Error>(Ok(()));
}
let selected = Error::InvalidRequest("provider".into());
assert_eq!(
assert!(matches!(
lifecycle.accept(Err(HostFailure::Error(selected.clone()))),
Some(selected)
);
Some(Error::InvalidRequest(message)) if message == "provider"
));
lifecycle.accept::<Error>(Ok(()));
for phase in [
HostPhase::DeploymentFailure,
@ -94,11 +94,12 @@ fn failure_handler_errors_do_not_replace_selected_failure_or_suppress_async_disp
HostPhase::AsyncFailure,
] {
assert_eq!(lifecycle.phase(), phase);
assert_eq!(
lifecycle.accept(Err(HostFailure::Error(Error::InvalidRequest(
"callback".into()
)))),
None
assert!(
lifecycle
.accept(Err(HostFailure::Error(Error::InvalidRequest(
"callback".into()
))))
.is_none()
);
}
assert_eq!(lifecycle.phase(), HostPhase::Complete);
@ -108,9 +109,9 @@ fn failure_handler_errors_do_not_replace_selected_failure_or_suppress_async_disp
fn cancellation_skips_terminal_dispatch() {
let mut lifecycle = HostLifecycle::new(true);
let error = Error::InvalidRequest("cancelled".into());
assert_eq!(
assert!(matches!(
lifecycle.accept(Err(HostFailure::Cancelled(error.clone()))),
Some(error)
);
Some(Error::InvalidRequest(message)) if message == "cancelled"
));
assert_eq!(lifecycle.phase(), HostPhase::Complete);
}

View file

@ -63,8 +63,8 @@ async fn facade_executes_direct_mistral_once() {
.await
.unwrap();
server.await.unwrap();
assert_eq!(result.pages[0]["markdown"], "hello");
assert_eq!(result.pages[0]["custom"], "preserved");
assert_eq!(result.pages[0].markdown, "hello");
assert_eq!(result.pages[0].extra_fields["custom"], "preserved");
let requests = seen.lock().unwrap();
assert_eq!(requests.len(), 1);
assert!(requests[0].starts_with("POST /v1/ocr "));
@ -80,7 +80,8 @@ async fn facade_executes_direct_mistral_once() {
"model":"model",
"document":{"type":"document_url","document_url":"data:application/pdf;base64,YWJj"},
"pages":"0,2-4",
"extract_header":true
"extract_header":true,
"unknown":"ignored"
})
);
}
@ -102,7 +103,10 @@ async fn facade_retains_native_response_when_requested() {
.unwrap();
server.await.unwrap();
assert_eq!(response.provider_native_response, Some(provider_response));
assert_eq!(
response.provider_native_response.map(Value::Object),
Some(provider_response)
);
}
#[tokio::test]
@ -348,7 +352,7 @@ async fn fallible_host_phases_do_not_replay_or_reach_transport() {
}
OcrHostOperation::ProjectRequest => {
result = Some(OcrHostResult::Request(Ok((
Box::new(request.take().unwrap().into()),
Box::new(request.take().unwrap()),
false,
))))
}
@ -406,7 +410,7 @@ async fn invalid_provider_response_runs_post_call_before_normalization_failure()
match call.resume(result.take()).await {
Ok(OcrCallStep::Host(OcrHostOperation::ProjectRequest)) => {
result = Some(OcrHostResult::Request(Ok((
Box::new(request.take().unwrap().into()),
Box::new(request.take().unwrap()),
false,
))));
}
@ -421,7 +425,7 @@ async fn invalid_provider_response_runs_post_call_before_normalization_failure()
}
};
server.await.unwrap();
assert!(matches!(error, crate::ocr::Error::InvalidResponse(_)));
assert!(matches!(error, crate::ocr::Error::ResponseField { .. }));
assert_eq!(seen.lock().unwrap().len(), 1);
assert_eq!(post_calls, [json!(r#"{"pages":"invalid"}"#)]);
}
@ -462,16 +466,15 @@ async fn direct_native_host_drives_the_same_state_machine() {
OcrHostOperation::PostCall(_) => "PostCall".into(),
OcrHostOperation::ConstructResponse(_) => "ConstructResponse".into(),
OcrHostOperation::Success { response, .. } => {
assert_eq!(response.pages[0]["markdown"], "native");
assert_eq!(response.pages[0].markdown, "native");
"Success".into()
}
_ => panic!("unexpected OCR operation"),
});
result = Some(match operation {
OcrHostOperation::ProjectRequest => OcrHostResult::Request(Ok((
Box::new(request.take().unwrap().into()),
false,
))),
OcrHostOperation::ProjectRequest => {
OcrHostResult::Request(Ok((Box::new(request.take().unwrap()), false)))
}
operation => host.invoke(operation).await,
});
}
@ -479,7 +482,7 @@ async fn direct_native_host_drives_the_same_state_machine() {
}
};
server.await.unwrap();
assert_eq!(response.pages[0]["markdown"], "native");
assert_eq!(response.pages[0].markdown, "native");
assert_eq!(seen.lock().unwrap().len(), 1);
assert_eq!(
operations,
@ -556,7 +559,7 @@ async fn host_reader_documents_are_read_once_at_the_core_selected_point_and_enco
)
.await;
server.await.unwrap();
assert_eq!(response.unwrap().pages[0]["markdown"], "file");
assert_eq!(response.unwrap().pages[0].markdown, "file");
assert_eq!(reads, 1);
assert!(seen.lock().unwrap()[0].contains("data:application/pdf;base64,YWJj"));
}
@ -571,7 +574,9 @@ async fn host_reader_failures_and_empty_files_fail_before_the_provider_is_called
Err(failure.clone()),
)
.await;
assert_eq!(response.unwrap_err(), failure);
assert!(
matches!(response.unwrap_err(), crate::ocr::Error::InvalidRequest(message) if message == "reader exploded")
);
assert_eq!(reads, 1);
let request = wire_request("mistral/model", &base, json!({}));
@ -585,7 +590,7 @@ async fn host_reader_failures_and_empty_files_fail_before_the_provider_is_called
.await;
assert!(matches!(
response.unwrap_err(),
crate::ocr::Error::InvalidRequest(_)
crate::ocr::Error::EmptyFile
));
assert!(seen.lock().unwrap().is_empty());
}
@ -613,7 +618,7 @@ async fn path_documents_are_read_by_core_without_a_host_operation() {
.await;
server.await.unwrap();
std::fs::remove_dir_all(&dir).unwrap();
assert_eq!(response.unwrap().pages[0]["markdown"], "path");
assert_eq!(response.unwrap().pages[0].markdown, "path");
assert_eq!(reads, 0);
assert!(seen.lock().unwrap()[0].contains("data:image/png;base64,YWJj"));
@ -629,7 +634,7 @@ async fn path_documents_are_read_by_core_without_a_host_operation() {
.await;
assert!(matches!(
response.unwrap_err(),
crate::ocr::Error::FileRead { path: failed, kind: std::io::ErrorKind::NotFound, .. } if failed == path
crate::ocr::Error::FileRead { path: failed, source } if failed == path && source.kind() == std::io::ErrorKind::NotFound
));
assert!(seen.lock().unwrap().is_empty());
}
@ -661,7 +666,9 @@ async fn public_finalization_failure_never_dispatches_success_or_replays_provide
OcrHostResult::Lifecycle(Err(HostFailure::Error(selected.clone())))
}
OcrHostOperation::Failure { error, .. } => {
assert_eq!(error, selected);
assert!(
matches!(error, crate::ocr::Error::InvalidRequest(message) if message == "public metadata failed")
);
failures.push("sync");
OcrHostResult::Lifecycle(Err(HostFailure::Error(
crate::ocr::Error::InvalidRequest("failure callback failed".into()),
@ -676,10 +683,9 @@ async fn public_finalization_failure_never_dispatches_success_or_replays_provide
| OcrHostOperation::Lifecycle(HostPhase::DeploymentFailure) => {
panic!("finalization failure used provider/success dispatch")
}
OcrHostOperation::ProjectRequest => OcrHostResult::Request(Ok((
Box::new(request.take().unwrap().into()),
false,
))),
OcrHostOperation::ProjectRequest => {
OcrHostResult::Request(Ok((Box::new(request.take().unwrap()), false)))
}
operation => host.invoke(operation).await,
});
}
@ -688,7 +694,9 @@ async fn public_finalization_failure_never_dispatches_success_or_replays_provide
}
};
server.await.unwrap();
assert_eq!(error, selected);
assert!(
matches!(error, crate::ocr::Error::InvalidRequest(message) if message == "public metadata failed")
);
assert_eq!(failures, ["sync", "async"]);
assert_eq!(seen.lock().unwrap().len(), 1);
}
@ -716,7 +724,7 @@ async fn cancellation_at_provider_hook_prevents_execution_and_further_resumption
OcrCallStep::Host(OcrHostOperation::PreCall(_)) => break,
OcrCallStep::Host(OcrHostOperation::ProjectRequest) => {
result = Some(OcrHostResult::Request(Ok((
Box::new(request.take().unwrap().into()),
Box::new(request.take().unwrap()),
false,
))))
}
@ -727,7 +735,7 @@ async fn cancellation_at_provider_hook_prevents_execution_and_further_resumption
let selected = crate::ocr::Error::InvalidRequest("cancelled".into());
assert!(matches!(
call.interrupt(HostFailure::Cancelled(selected.clone())).await,
Err(error) if error == selected
Err(crate::ocr::Error::InvalidRequest(message)) if message == "cancelled"
));
assert!(
call.resume(Some(OcrHostResult::Lifecycle(Ok(()))))
@ -761,7 +769,7 @@ async fn missing_host_result_preserves_pending_operation() {
async fn read_bounded_response(
response: Vec<u8>,
limit: usize,
) -> Result<bytes::Bytes, super::error::OcrError> {
) -> Result<bytes::Bytes, super::Error> {
use tokio::io::{AsyncReadExt, AsyncWriteExt};
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
@ -790,7 +798,7 @@ async fn read_bounded_response(
#[tokio::test]
async fn response_limit_accepts_exact_size_and_rejects_declared_and_chunked_overflow() {
use super::error::{OcrError, OcrResponseError};
use super::Error;
for response in [
"HTTP/1.1 200 OK\r\nContent-Length: 8\r\n\r\nabcdefgh",
@ -809,7 +817,7 @@ async fn response_limit_accepts_exact_size_and_rejects_declared_and_chunked_over
] {
assert!(matches!(
read_bounded_response(response.as_bytes().to_vec(), 8).await,
Err(OcrError::Response(OcrResponseError::TooLarge { limit: 8 }))
Err(Error::TooLarge { limit: 8 })
));
}
}
@ -828,7 +836,7 @@ async fn oversized_error_retains_http_status_and_bounded_diagnostics_without_dra
.await
.unwrap_err();
match error {
super::error::OcrError::Transport(crate::transport::Error::Http { status, body }) => {
super::Error::Transport(crate::transport::Error::Http { status, body }) => {
assert_eq!(status, 429);
assert_eq!(
body,
@ -850,7 +858,7 @@ fn response_limit_is_validated_and_not_forwarded_to_the_provider() {
"http://localhost",
json!({"max_response_bytes": 123}),
);
assert_eq!(request.connection.max_response_bytes, 123);
assert_eq!(request.transport.max_response_bytes, 123);
assert!(!request.optional_params.contains_key("max_response_bytes"));
for value in [
json!(0),
@ -908,9 +916,9 @@ async fn cancellation_waits_for_provider_capture_drop_even_when_acknowledgement_
let dropped = Arc::new(AtomicBool::new(false));
let request = wire_request("azure_ai/mistral-ocr", "https://example.invalid", json!({}));
let request = super::LiteLLMOcrRequest {
connection: super::OcrConnection {
transport: super::OcrTransportConfig {
extra_headers: vec![("authorization".into(), "Bearer test-key".into())],
..request.connection
..request.transport
},
azure_ad_token_provider: Some(litellm_auth::TokenProviderHandle::new(Arc::new(
PendingToken {
@ -933,7 +941,7 @@ async fn cancellation_waits_for_provider_capture_drop_even_when_acknowledgement_
_ = entered.notified() => break,
step = call.resume(result.take()) => {
result = Some(match step.unwrap() {
OcrCallStep::Host(OcrHostOperation::ProjectRequest) => OcrHostResult::Request(Ok((Box::new(request.take().unwrap().into()), false))),
OcrCallStep::Host(OcrHostOperation::ProjectRequest) => OcrHostResult::Request(Ok((Box::new(request.take().unwrap()), false))),
OcrCallStep::Host(operation) => NoopOcrHost.invoke(operation).await,
OcrCallStep::Complete(_) => panic!("pending provider completed"),
});
@ -960,7 +968,9 @@ async fn cancellation_waits_for_provider_capture_drop_even_when_acknowledgement_
)
.await
.unwrap();
assert!(matches!(result, Err(error) if error == selected));
assert!(
matches!(result, Err(crate::ocr::Error::InvalidRequest(message)) if message == "cancelled")
);
assert!(
dropped.load(Ordering::SeqCst),
"cancellation returned while provider captures were still alive"

View file

@ -36,6 +36,20 @@ pub(crate) fn wire_request(model: &str, base: &str, options: Value) -> LiteLLMOc
.unwrap()
}
pub(crate) fn resolved_request(
request: LiteLLMOcrRequest,
) -> crate::ocr::types::ResolvedOcrRequest {
request
.map_document(crate::ocr::document::prepare_document)
.unwrap()
}
pub(crate) fn with_source(request: LiteLLMOcrRequest, source: &str) -> LiteLLMOcrRequest {
let request = resolved_request(request);
let document = request.document.clone().with_source(source.into());
request.with_document(document.into())
}
pub(crate) struct MockResponse {
pub status: u16,
pub headers: Vec<(&'static str, String)>,

View file

@ -56,8 +56,7 @@ async fn request_mapping_matches_python(
"result":{"chunks":[]}
}))])
.await;
let mut request = wire_request(model, &base, options);
request.document = request.document.with_source(source.into());
let request = super::test_support::with_source(wire_request(model, &base, options), source);
perform_ocr(request).await.unwrap();
server.await.unwrap();
@ -78,14 +77,14 @@ async fn data_uri_upload_preserves_multipart_headers(#[case] model: &str) {
])
.await;
let mut request = wire_request(&format!("reducto/{model}"), &base, json!({}));
request.connection.extra_headers = vec![
request.transport.extra_headers = vec![
("Content-Type".into(), "application/json".into()),
("X-Trace".into(), "upload-test".into()),
];
let response = perform_ocr(request).await.unwrap();
server.await.unwrap();
assert_eq!(response.pages[0]["markdown"], "hello");
assert_eq!(response.pages[0].markdown, "hello");
let requests = seen.lock().unwrap();
assert_eq!(requests.len(), 2);
assert!(requests[0].starts_with("POST /upload "));
@ -175,14 +174,18 @@ async fn upload_failure_stops_before_parse() {
#[case("data:application/pdf;base64,INVALID!")]
#[tokio::test]
async fn rejects_invalid_document_sources_before_network(#[case] source: &str) {
let mut request = wire_request("reducto/parse-v3", "http://127.0.0.1:1", json!({}));
request.document = request.document.with_source(source.into());
let request = super::test_support::with_source(
wire_request("reducto/parse-v3", "http://127.0.0.1:1", json!({})),
source,
);
assert!(perform_ocr(request).await.is_err());
}
#[test]
fn response_normalization_groups_blocks_and_distinguishes_null_result() {
use crate::ocr::codecs::reducto::{ReductoResponse, transform_ocr_response};
use crate::llms::reducto::ocr::transformation::{
ReductoResponse, normalize_response as transform_ocr_response,
};
let raw = json!({"usage":{"num_pages":"2","credits":"3"},"result":{"type":"full","chunks":[
{"blocks":[{
@ -218,7 +221,7 @@ fn response_normalization_groups_blocks_and_distinguishes_null_result() {
let missing: ReductoResponse =
serde_json::from_value(json!({"chunks":[{"content":"text"}]})).unwrap();
let missing = transform_ocr_response("parse-v3", missing).unwrap();
assert_eq!(missing.pages[0]["markdown"], "text");
assert_eq!(missing.pages[0].markdown, "text");
let null: ReductoResponse = serde_json::from_value(
json!({"result":null,"chunks":[{"content":"ignored"}],"usage":null}),
)
@ -231,9 +234,11 @@ fn response_normalization_groups_blocks_and_distinguishes_null_result() {
async fn facade_omits_native_response_by_default_and_preserves_auth_priority() {
let raw = json!({"job_id":"job-1","result":{"chunks":[]}});
let (base, seen, server) = mock_server(vec![MockResponse::json(raw)]).await;
let mut request = wire_request("reducto/parse-v3", &base, json!({}));
request.document = request.document.with_source("reducto://ready.pdf".into());
request.connection.extra_headers = vec![("authorization".into(), "Bearer existing".into())];
let mut request = super::test_support::with_source(
wire_request("reducto/parse-v3", &base, json!({})),
"reducto://ready.pdf",
);
request.transport.extra_headers = vec![("authorization".into(), "Bearer existing".into())];
let response = perform_ocr(request).await.unwrap();
server.await.unwrap();

View file

@ -14,7 +14,7 @@ async fn facade_executes_vertex_deepseek_at_the_openai_endpoint() {
"usage":{"prompt_tokens":1}
}))])
.await;
let mut request = wire_request(
let request = wire_request(
"vertex_ai/deepseek-ocr-maas",
&base,
json!({
@ -25,14 +25,15 @@ async fn facade_executes_vertex_deepseek_at_the_openai_endpoint() {
"extra_body":{"provider_option":"value"}
}),
);
request.document = request
.document
.with_source("gs://bucket/document.pdf".into());
let request = super::test_support::with_source(request, "gs://bucket/document.pdf");
let response = perform_ocr(request).await.unwrap();
server.await.unwrap();
assert_eq!(response.pages[0]["markdown"], "recognized");
assert_eq!(response.usage_info.unwrap()["prompt_tokens"], 1);
assert_eq!(response.pages[0].markdown, "recognized");
assert_eq!(
response.usage_info.unwrap().extra_fields["prompt_tokens"],
1
);
let requests = seen.lock().unwrap();
assert!(requests[0].starts_with(
"POST /v1/projects/project-1/locations/europe-west4/endpoints/openapi/chat/completions "
@ -45,7 +46,7 @@ async fn facade_executes_vertex_deepseek_at_the_openai_endpoint() {
let body = request_body(&requests[0]);
assert_eq!(body["model"], "deepseek-ai/deepseek-ocr-maas");
assert_eq!(body["temperature"], 0.1);
assert!(body.get("future_ocr_option").is_none());
assert_eq!(body["future_ocr_option"], true);
assert!(body.get("extra_body").is_none());
assert_eq!(
body["messages"][0]["content"][0],
@ -72,7 +73,10 @@ async fn request_controlled_api_base_is_rejected_before_vertex_auth() {
"https://caller.example",
json!({"vertex_project":"project-1"}),
);
request.connection.api_base_source = InputSource::Request;
request.credentials.api_base = Some(litellm_auth::Sourced::new(
"https://caller.example".into(),
InputSource::Request,
));
let error = perform_ocr(request).await.unwrap_err();
assert!(

View file

@ -26,7 +26,7 @@ async fn facade_executes_vertex_mistral_with_resolved_project_and_location() {
let response = perform_ocr(request).await.unwrap();
server.await.unwrap();
assert_eq!(response.pages[0]["markdown"], "hello");
assert_eq!(response.pages[0].markdown, "hello");
let requests = seen.lock().unwrap();
assert_eq!(requests.len(), 1);
assert!(requests[0].starts_with(
@ -55,8 +55,8 @@ async fn supplied_authorization_is_forwarded_without_a_static_token() {
&base,
json!({"vertex_project":"project-1"}),
);
request.connection.api_key = None;
request.connection.extra_headers = vec![("authorization".into(), "Bearer supplied".into())];
request.credentials.api_key = None;
request.transport.extra_headers = vec![("authorization".into(), "Bearer supplied".into())];
perform_ocr(request).await.unwrap();
server.await.unwrap();
@ -85,7 +85,10 @@ async fn request_controlled_api_base_is_rejected_before_vertex_auth() {
"https://caller.example",
json!({"vertex_project":"project-1"}),
);
request.connection.api_base_source = InputSource::Request;
request.credentials.api_base = Some(litellm_auth::Sourced::new(
"https://caller.example".into(),
InputSource::Request,
));
let error = perform_ocr(request).await.unwrap_err();
assert!(
@ -99,7 +102,9 @@ async fn request_controlled_api_base_is_rejected_before_vertex_auth() {
async fn adapters_build_complete_requests_and_share_mistral_normalization() {
use std::time::Duration;
use crate::ocr::adapters::{MistralAdapter, OcrAdapter, VertexMistralAdapter};
use crate::llms::base_llm::ocr::transformation::BaseOcrConfig;
use crate::llms::mistral::ocr::transformation::MistralOCRConfig;
use crate::llms::vertex_ai::ocr::transformation::VertexAIOCRConfig;
use crate::ocr::test_support::ocr_client;
let client = ocr_client();
@ -116,11 +121,15 @@ async fn adapters_build_complete_requests_and_share_mistral_normalization() {
options.clone(),
);
let vertex = wire_request("vertex_ai/mistral-ocr-maas", "https://vertex.test", options);
let direct_http = MistralAdapter
let direct =
crate::ocr::prepare::prepare_request(super::test_support::resolved_request(direct));
let vertex =
crate::ocr::prepare::prepare_request(super::test_support::resolved_request(vertex));
let direct_http = MistralOCRConfig
.prepare_request(&direct, &client)
.await
.unwrap();
let vertex_http = VertexMistralAdapter
let vertex_http = VertexAIOCRConfig
.prepare_request(&vertex, &client)
.await
.unwrap();
@ -141,17 +150,27 @@ async fn adapters_build_complete_requests_and_share_mistral_normalization() {
"model": "mistral-ocr-maas",
"document": {"type": "document_url", "document_url": "data:application/pdf;base64,YWJj"},
"pages": [0, 2],
"include_image_base64": true
"include_image_base64": true,
"unknown": "ignored"
})
);
}
let payload = json!({"pages": [{"index": 0, "markdown": "hello"}], "extra": "preserved"});
let direct_response = MistralAdapter
.transform_ocr_response(&direct, serde_json::from_value(payload.clone()).unwrap())
let raw = serde_json::to_vec(&payload).unwrap();
let direct_response = MistralOCRConfig
.transform_ocr_response(
&direct.model,
&raw,
crate::ocr::types::OcrResponseFormat::Litellm,
)
.unwrap()
.into_json();
let vertex_response = VertexMistralAdapter
.transform_ocr_response(&vertex, serde_json::from_value(payload).unwrap())
let vertex_response = VertexAIOCRConfig
.transform_ocr_response(
&vertex.model,
&raw,
crate::ocr::types::OcrResponseFormat::Litellm,
)
.unwrap()
.into_json();
assert_eq!(direct_response, vertex_response);

View file

@ -35,15 +35,17 @@ pub(crate) fn responses_error_to_pyerr(error: responses::Error) -> PyErr {
pub(crate) fn core_error_to_pyerr(error: Error) -> PyErr {
let value_error = match &error {
Error::Ocr(error) => matches!(
error,
ocr::Error::Auth(_)
| ocr::Error::InvalidProvider(_)
| ocr::Error::InvalidRequest(_)
| ocr::Error::InvalidType { .. }
| ocr::Error::MissingField(_)
| ocr::Error::MissingDocumentUrl
),
Error::Ocr(error) => {
error.is_request()
|| matches!(
error,
ocr::Error::Auth(_)
| ocr::Error::InvalidProvider(_)
| ocr::Error::InvalidRequest(_)
| ocr::Error::MissingField(_)
| ocr::Error::MissingDocumentUrl
)
}
Error::Messages(error) => match error {
messages::Error::Auth(source) => auth_is_value_error(source),
messages::Error::InvalidProvider(_)

View file

@ -7,13 +7,14 @@ use crate::errors::{RustUpstreamError, core_error_to_pyerr};
pub(super) fn to_pyerr(error: Error) -> PyErr {
let status = error.http_status_code();
let mapped = match error {
Error::Http { status, body } => RustUpstreamError::new_err((status, body)),
Error::FileRead {
path,
kind: std::io::ErrorKind::NotFound,
..
} => PyFileNotFoundError::new_err(format!("File not found: {}", path.display())),
Error::FileRead { message, .. } => PyOSError::new_err(message),
Error::Provider { status, body, .. }
| Error::Transport(litellm_core::transport::Error::Http { status, body }) => {
RustUpstreamError::new_err((status, body))
}
Error::FileRead { path, source } if source.kind() == std::io::ErrorKind::NotFound => {
PyFileNotFoundError::new_err(format!("File not found: {}", path.display()))
}
Error::FileRead { source, .. } => PyOSError::new_err(source.to_string()),
other => core_error_to_pyerr(other.into()),
};
attach_status(mapped, status)
@ -51,9 +52,10 @@ mod tests {
.unwrap(),
500
);
let mapped = to_pyerr(Error::Http {
let mapped = to_pyerr(Error::Provider {
status: 429,
body: r#"{"message":"rate limited"}"#.to_string(),
headers: Vec::new(),
});
assert!(mapped.is_instance_of::<RustUpstreamError>(py));
let args: (u16, String) = mapped

View file

@ -215,7 +215,7 @@ mod tests {
fn url_document(url: &str) -> OcrDocumentInput {
litellm_core::ocr::OcrDocument::DocumentUrl {
document_url: url.into(),
extra_fields: Map::new(),
extra_fields: Default::default(),
}
.into()
}