mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-04 02:31:27 +00:00
refactor(ocr): mirror Python provider layout and preserve tests
This commit is contained in:
parent
351a54e849
commit
edfa01da81
88 changed files with 8518 additions and 4121 deletions
147
litellm-rust/Cargo.lock
generated
147
litellm-rust/Cargo.lock
generated
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
467
litellm-rust/crates/core/src/call_arguments.rs
Normal file
467
litellm-rust/crates/core/src/call_arguments.rs
Normal 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"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
|
@ -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;
|
||||
|
||||
|
|
|
|||
1
litellm-rust/crates/core/src/llms/azure_ai/mod.rs
Normal file
1
litellm-rust/crates/core/src/llms/azure_ai/mod.rs
Normal file
|
|
@ -0,0 +1 @@
|
|||
pub(crate) mod ocr;
|
||||
|
|
@ -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()
|
||||
);
|
||||
}
|
||||
}
|
||||
|
|
@ -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(())
|
||||
}
|
||||
|
|
@ -0,0 +1 @@
|
|||
pub(crate) mod transformation;
|
||||
File diff suppressed because it is too large
Load diff
4
litellm-rust/crates/core/src/llms/azure_ai/ocr/mod.rs
Normal file
4
litellm-rust/crates/core/src/llms/azure_ai/ocr/mod.rs
Normal 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;
|
||||
399
litellm-rust/crates/core/src/llms/azure_ai/ocr/transformation.rs
Normal file
399
litellm-rust/crates/core/src/llms/azure_ai/ocr/transformation.rs
Normal 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"));
|
||||
}
|
||||
}
|
||||
1
litellm-rust/crates/core/src/llms/base_llm/mod.rs
Normal file
1
litellm-rust/crates/core/src/llms/base_llm/mod.rs
Normal file
|
|
@ -0,0 +1 @@
|
|||
pub(crate) mod ocr;
|
||||
1
litellm-rust/crates/core/src/llms/base_llm/ocr/mod.rs
Normal file
1
litellm-rust/crates/core/src/llms/base_llm/ocr/mod.rs
Normal file
|
|
@ -0,0 +1 @@
|
|||
pub(crate) mod transformation;
|
||||
211
litellm-rust/crates/core/src/llms/base_llm/ocr/transformation.rs
Normal file
211
litellm-rust/crates/core/src/llms/base_llm/ocr/transformation.rs
Normal 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, ¶ms, &environment)?;
|
||||
let headers = environment.headers();
|
||||
let body = self
|
||||
.async_transform_ocr_request(
|
||||
&request.model,
|
||||
request.document.clone(),
|
||||
¶ms,
|
||||
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)],
|
||||
}
|
||||
1
litellm-rust/crates/core/src/llms/cohere/mod.rs
Normal file
1
litellm-rust/crates/core/src/llms/cohere/mod.rs
Normal file
|
|
@ -0,0 +1 @@
|
|||
pub(crate) mod ocr;
|
||||
3
litellm-rust/crates/core/src/llms/cohere/ocr/mod.rs
Normal file
3
litellm-rust/crates/core/src/llms/cohere/ocr/mod.rs
Normal file
|
|
@ -0,0 +1,3 @@
|
|||
pub(crate) mod transformation;
|
||||
|
||||
pub(crate) use transformation::{CohereOptions, validate_document};
|
||||
740
litellm-rust/crates/core/src/llms/cohere/ocr/transformation.rs
Normal file
740
litellm-rust/crates/core/src/llms/cohere/ocr/transformation.rs
Normal 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(¶ms).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, ¶ms, &[])
|
||||
.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(_))
|
||||
));
|
||||
}
|
||||
}
|
||||
1
litellm-rust/crates/core/src/llms/mistral/mod.rs
Normal file
1
litellm-rust/crates/core/src/llms/mistral/mod.rs
Normal file
|
|
@ -0,0 +1 @@
|
|||
pub(crate) mod ocr;
|
||||
1
litellm-rust/crates/core/src/llms/mistral/ocr/mod.rs
Normal file
1
litellm-rust/crates/core/src/llms/mistral/ocr/mod.rs
Normal file
|
|
@ -0,0 +1 @@
|
|||
pub(crate) mod transformation;
|
||||
626
litellm-rust/crates/core/src/llms/mistral/ocr/transformation.rs
Normal file
626
litellm-rust/crates/core/src/llms/mistral/ocr/transformation.rs
Normal 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(), ¶ms, &[])
|
||||
.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(¶ms, "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(), ¶ms, &[])
|
||||
.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(), ¶ms, &[])
|
||||
.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(), ¶ms, &[])
|
||||
.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,
|
||||
}
|
||||
))
|
||||
));
|
||||
}
|
||||
}
|
||||
6
litellm-rust/crates/core/src/llms/mod.rs
Normal file
6
litellm-rust/crates/core/src/llms/mod.rs
Normal 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;
|
||||
1
litellm-rust/crates/core/src/llms/reducto/mod.rs
Normal file
1
litellm-rust/crates/core/src/llms/reducto/mod.rs
Normal file
|
|
@ -0,0 +1 @@
|
|||
pub(crate) mod ocr;
|
||||
1
litellm-rust/crates/core/src/llms/reducto/ocr/mod.rs
Normal file
1
litellm-rust/crates/core/src/llms/reducto/ocr/mod.rs
Normal file
|
|
@ -0,0 +1 @@
|
|||
pub(crate) mod transformation;
|
||||
1018
litellm-rust/crates/core/src/llms/reducto/ocr/transformation.rs
Normal file
1018
litellm-rust/crates/core/src/llms/reducto/ocr/transformation.rs
Normal file
File diff suppressed because it is too large
Load diff
1
litellm-rust/crates/core/src/llms/vertex_ai/mod.rs
Normal file
1
litellm-rust/crates/core/src/llms/vertex_ai/mod.rs
Normal file
|
|
@ -0,0 +1 @@
|
|||
pub(crate) mod ocr;
|
||||
|
|
@ -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(())
|
||||
}
|
||||
|
|
@ -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(), ¶ms, &[])
|
||||
.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")
|
||||
);
|
||||
}
|
||||
}
|
||||
3
litellm-rust/crates/core/src/llms/vertex_ai/ocr/mod.rs
Normal file
3
litellm-rust/crates/core/src/llms/vertex_ai/ocr/mod.rs
Normal file
|
|
@ -0,0 +1,3 @@
|
|||
pub(crate) mod common_utils;
|
||||
pub(crate) mod deepseek_transformation;
|
||||
pub(crate) mod transformation;
|
||||
|
|
@ -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");
|
||||
}
|
||||
}
|
||||
|
|
@ -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());
|
||||
}
|
||||
}
|
||||
|
|
@ -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, ¶ms)?;
|
||||
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())
|
||||
);
|
||||
}
|
||||
}
|
||||
|
|
@ -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());
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
@ -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, ¶ms)?;
|
||||
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())
|
||||
);
|
||||
}
|
||||
}
|
||||
|
|
@ -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(_)))
|
||||
));
|
||||
}
|
||||
}
|
||||
|
|
@ -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(), ¶ms)?;
|
||||
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"
|
||||
}))
|
||||
));
|
||||
}
|
||||
}
|
||||
|
|
@ -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;
|
||||
|
|
@ -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, ¶ms)?;
|
||||
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)
|
||||
}
|
||||
}
|
||||
|
|
@ -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
|
||||
);
|
||||
}
|
||||
}
|
||||
|
|
@ -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, ¶ms)?;
|
||||
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)
|
||||
}
|
||||
}
|
||||
|
|
@ -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, ¶ms)?;
|
||||
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"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
|
@ -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, ¶ms)?;
|
||||
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());
|
||||
}
|
||||
}
|
||||
|
|
@ -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(())
|
||||
}
|
||||
101
litellm-rust/crates/core/src/ocr/arguments.rs
Normal file
101
litellm-rust/crates/core/src/ocr/arguments.rs
Normal 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)
|
||||
);
|
||||
}
|
||||
}
|
||||
|
|
@ -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();
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
|
@ -1,5 +0,0 @@
|
|||
mod transformation;
|
||||
mod types;
|
||||
|
||||
pub(crate) use transformation::{transform_ocr_request, transform_ocr_response};
|
||||
pub(crate) use types::{DeepSeekOcrParams, DeepSeekOcrResponse};
|
||||
|
|
@ -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()),
|
||||
})
|
||||
}
|
||||
|
|
@ -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>,
|
||||
}
|
||||
|
|
@ -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,
|
||||
};
|
||||
|
|
@ -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);
|
||||
}
|
||||
}
|
||||
|
|
@ -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)
|
||||
}
|
||||
|
|
@ -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")),
|
||||
}
|
||||
}
|
||||
|
|
@ -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};
|
||||
|
|
@ -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(), ¶ms).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(), ¶ms).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(), ¶ms).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);
|
||||
}
|
||||
}
|
||||
|
|
@ -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>,
|
||||
}
|
||||
|
|
@ -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;
|
||||
|
|
@ -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,
|
||||
};
|
||||
|
|
@ -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
|
||||
}
|
||||
|
|
@ -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)
|
||||
}
|
||||
|
|
@ -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"));
|
||||
|
|
|
|||
|
|
@ -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(_)
|
||||
)
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
|
|
|
|||
|
|
@ -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) })
|
||||
}
|
||||
|
|
|
|||
62
litellm-rust/crates/core/src/ocr/json.rs
Normal file
62
litellm-rust/crates/core/src/ocr/json.rs
Normal 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(),
|
||||
})
|
||||
}
|
||||
|
|
@ -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();
|
||||
|
|
|
|||
|
|
@ -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)]
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
}
|
||||
}
|
||||
|
|
|
|||
411
litellm-rust/crates/core/src/ocr/provider_config.rs
Normal file
411
litellm-rust/crates/core/src/ocr/provider_config.rs
Normal 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
|
||||
);
|
||||
}
|
||||
}
|
||||
|
|
@ -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);
|
||||
}
|
||||
}
|
||||
|
|
@ -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");
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
|
|||
231
litellm-rust/crates/core/src/params.rs
Normal file
231
litellm-rust/crates/core/src/params.rs
Normal 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());
|
||||
}
|
||||
}
|
||||
|
|
@ -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;
|
||||
|
|
|
|||
219
litellm-rust/crates/core/src/providers/model.rs
Normal file
219
litellm-rust/crates/core/src/providers/model.rs
Normal 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());
|
||||
}
|
||||
}
|
||||
151
litellm-rust/crates/core/src/serde_compat.rs
Normal file
151
litellm-rust/crates/core/src/serde_compat.rs
Normal 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));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
@ -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();
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(), ¶ms).unwrap(),
|
||||
VertexAIDeepSeekOCRConfig
|
||||
.transform_ocr_request("deepseek-ai/deepseek-ocr-maas", document(), ¶ms, &[])
|
||||
.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}]}"}}]}),
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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)>,
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
|
|
|
|||
|
|
@ -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!(
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
|
|
|
|||
|
|
@ -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(_)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
}
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue