This commit is contained in:
Yujong Lee 2026-09-15 17:49:18 -07:00
parent 31b48f6191
commit 5de196af63
21 changed files with 702 additions and 458 deletions

View file

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

View file

@ -1,4 +1,5 @@
pub mod audio_transcription;
pub mod call_arguments;
pub mod call_lifecycle;
pub mod chat_completions;
pub mod constants;

View file

@ -1,7 +1,6 @@
use crate::llms::base_llm::ocr::transformation::{BaseOcrConfig, OcrRequestContext};
use crate::llms::cohere::ocr::transformation::{CohereParseConfig, CohereRequest};
use crate::llms::cohere::ocr::{CohereParams, CohereResponse, validate_document};
use crate::ocr::OcrArguments;
use crate::llms::cohere::ocr::{CohereOptions, CohereResponse, validate_document};
use crate::ocr::OcrClient;
use crate::ocr::document::{inline_remote_document, validate_inline_document};
use crate::ocr::prepare::{credential_env, transform_request_body};
@ -15,7 +14,7 @@ const AZURE_AI_API_BASE_ENV: &str = "AZURE_AI_API_BASE";
pub(crate) struct AzureAICohereParseConfig;
impl BaseOcrConfig for AzureAICohereParseConfig {
type OcrParams = CohereParams;
type OcrParams = CohereOptions;
type ProviderRequest = CohereRequest;
type ProviderResponse = CohereResponse;
@ -23,20 +22,11 @@ impl BaseOcrConfig for AzureAICohereParseConfig {
CohereParseConfig.get_supported_ocr_params(model)
}
fn map_ocr_params(
&self,
non_default_params: &OcrArguments,
optional_params: &OcrArguments,
model: &str,
) -> Result<OcrArguments, crate::ocr::Error> {
CohereParseConfig.map_ocr_params(non_default_params, optional_params, model)
}
async fn async_transform_ocr_request(
&self,
model: &str,
document: OcrDocument,
optional_params: &CohereParams,
optional_params: &CohereOptions,
headers: &[(String, String)],
context: OcrRequestContext<'_>,
) -> Result<CohereRequest, crate::ocr::Error> {
@ -65,7 +55,7 @@ impl AzureAICohereParseConfig {
request: &LiteLLMOcrRequest,
client: &OcrClient,
) -> Result<reqwest::Request, crate::ocr::Error> {
let params = self.parse_options(&request.optional_params, &request.model)?;
let params = self.map_ocr_params(&request.optional_params, &request.model)?;
let config = AzureAuthInputs {
azure_ad_token_provider: request.azure_ad_token_provider.clone(),
..AzureAuthInputs::from_sourced_optional_params(
@ -110,8 +100,9 @@ impl AzureAICohereParseConfig {
!remote,
body,
|body| {
validate_document(&body.document.as_document())?;
validate_inline_document(&body.document.as_document())
let document = crate::ocr::prepare::body_document(body)?;
validate_document(&document)?;
validate_inline_document(&document)
},
)
.await

View file

@ -12,6 +12,7 @@ use tokio::time::Instant;
use litellm_auth::{InputSource, Sourced};
use litellm_auth_azure::AzureAuthInputs;
use crate::call_arguments::CallArguments;
use crate::constants::{
AZURE_DI_API_VERSION, AZURE_DI_DEFAULT_DPI, AZURE_DI_DEFAULT_HEIGHT, AZURE_DI_DEFAULT_WIDTH,
AZURE_DI_SUBSCRIPTION_HEADER, OCR_POLL_RETRY_SECS,
@ -19,7 +20,6 @@ use crate::constants::{
use crate::llms::base_llm::ocr::transformation::{
BaseOcrConfig, OcrRequestContext, OcrResponseContext,
};
use crate::ocr::OcrArguments;
use crate::ocr::OcrClient;
use crate::ocr::client::read_json_response;
use crate::ocr::document::InlineDocument;
@ -30,7 +30,6 @@ use crate::ocr::types::{
OcrResponseFormat, OcrUsageInfo,
};
use crate::ocr::wire::DecodedOcrResponse;
use crate::params::OpaqueParams;
use crate::serde_compat::{FiniteF64, LaxI64};
use crate::url_utils::ApiUrl;
@ -61,10 +60,6 @@ pub(crate) struct DocumentIntelligenceParams {
pub pages: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub features: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub req_format: Option<OcrResponseFormat>,
#[serde(flatten)]
pub extra_fields: OpaqueParams,
}
#[derive(Clone, Debug, Serialize, Deserialize)]
@ -183,8 +178,6 @@ fn normalize_ocr_params(
.map(normalize_features)
.transpose()?
.flatten(),
req_format: None,
extra_fields: OpaqueParams::default(),
})
}
@ -485,53 +478,13 @@ impl BaseOcrConfig for AzureDocumentIntelligenceOCRConfig {
fn map_ocr_params(
&self,
non_default_params: &OcrArguments,
optional_params: &OcrArguments,
arguments: &CallArguments,
_model: &str,
) -> Result<OcrArguments, crate::ocr::Error> {
let mapped = normalize_ocr_params(decode_input_params(
non_default_params
.iter()
.filter(|(name, _)| matches!(name.as_str(), "pages" | "features"))
.map(|(name, value)| (name.clone(), value.clone()))
.collect(),
) -> Result<DocumentIntelligenceParams, crate::ocr::Error> {
normalize_ocr_params(decode_input_params(
arguments.select(&["pages", "features"]),
"optional_params",
)?)?;
let request_format = non_default_params
.get("req_format")
.filter(|value| !value.is_null())
.map(|value| {
serde_json::from_value::<OcrResponseFormat>(value.clone())
.map_err(|_| crate::ocr::Error::RequestFormat)
})
.transpose()?;
let fields = optional_params
.iter()
.map(|(name, value)| (name.clone(), value.clone()))
.chain(
mapped
.pages
.map(|value| ("pages".into(), Value::String(value))),
)
.chain(
mapped
.features
.map(|value| ("features".into(), Value::String(value))),
)
.chain(request_format.map(|value| {
(
"req_format".into(),
Value::String(
match value {
OcrResponseFormat::Litellm => "litellm",
OcrResponseFormat::Native => "native",
}
.into(),
),
)
}))
.collect();
Ok(fields)
)?)
}
async fn async_transform_ocr_request(
@ -592,7 +545,7 @@ impl AzureDocumentIntelligenceOCRConfig {
request: &LiteLLMOcrRequest,
client: &OcrClient,
) -> Result<reqwest::Request, crate::ocr::Error> {
let params = self.parse_options(&request.optional_params, &request.model)?;
let params = self.map_ocr_params(&request.optional_params, &request.model)?;
let config = AzureAuthInputs {
azure_ad_token_provider: request.azure_ad_token_provider.clone(),
..AzureAuthInputs::from_sourced_optional_params(
@ -728,35 +681,27 @@ mod tests {
serde_json::from_value(json!({"pages":[], "features":null, "req_format":"native"}))
.unwrap();
let mapped = AzureDocumentIntelligenceOCRConfig
.parse_options(&overrides, "model")
.map_ocr_params(&overrides, "model")
.unwrap();
assert_eq!(
serde_json::to_value(mapped).unwrap(),
json!({
"req_format":"native"
})
);
assert_eq!(serde_json::to_value(mapped).unwrap(), json!({}));
}
#[test]
fn mapping_preserves_supplied_options_when_overrides_are_empty() {
let supplied = serde_json::from_value(json!({
fn options_normalize_query_fields_without_consuming_extensions() {
let arguments = serde_json::from_value(json!({
"pages":"4", "features":"languages", "extension":true
}))
.unwrap();
let overrides = serde_json::from_value(json!({
"pages":[], "features":null, "req_format":"native", "ignored":true
}))
.unwrap();
let mapped = AzureDocumentIntelligenceOCRConfig
.map_ocr_params(&overrides, &supplied, "model")
.map_ocr_params(&arguments, "model")
.unwrap();
assert_eq!(
serde_json::to_value(mapped).unwrap(),
json!({
"pages":"4", "features":"languages", "extension":true, "req_format":"native"
"pages":"4", "features":"languages"
})
);
assert_eq!(arguments["extension"], true);
}
#[test]

View file

@ -2,7 +2,6 @@ use crate::constants::AZURE_AI_OCR_PATH;
use crate::llms::base_llm::ocr::transformation::{BaseOcrConfig, OcrRequestContext};
use crate::llms::mistral::ocr::MistralOcrResponse;
use crate::llms::mistral::ocr::transformation::{MistralOCRConfig, MistralOcrRequest};
use crate::ocr::OcrArguments;
use crate::ocr::OcrClient;
use crate::ocr::document::{inline_remote_document, validate_inline_document};
use crate::ocr::prepare::{credential_env, transform_request_body};
@ -27,15 +26,6 @@ impl BaseOcrConfig for AzureAIOCRConfig {
MistralOCRConfig.get_supported_ocr_params(model)
}
fn map_ocr_params(
&self,
non_default_params: &OcrArguments,
optional_params: &OcrArguments,
model: &str,
) -> Result<OcrArguments, crate::ocr::Error> {
MistralOCRConfig.map_ocr_params(non_default_params, optional_params, model)
}
async fn async_transform_ocr_request(
&self,
model: &str,
@ -68,7 +58,7 @@ impl AzureAIOCRConfig {
request: &LiteLLMOcrRequest,
client: &OcrClient,
) -> Result<reqwest::Request, crate::ocr::Error> {
let params = self.parse_options(&request.optional_params, &request.model)?;
let params = self.map_ocr_params(&request.optional_params, &request.model)?;
let config = AzureAuthInputs {
azure_ad_token_provider: request.azure_ad_token_provider.clone(),
..AzureAuthInputs::from_sourced_optional_params(
@ -102,7 +92,7 @@ impl AzureAIOCRConfig {
&headers,
retains_document,
body,
|body| validate_inline_document(&body.document),
|body| validate_inline_document(&crate::ocr::prepare::body_document(body)?),
)
.await
}

View file

@ -4,7 +4,7 @@ use std::sync::Arc;
use serde::Serialize;
use serde::de::DeserializeOwned;
use crate::ocr::OcrArguments;
use crate::call_arguments::{CallArguments, parse_options};
use crate::ocr::OcrClient;
use crate::ocr::hooks::OcrHooks;
use crate::ocr::types::{LiteLLMOcrResponse, OcrConnection, OcrDocument, OcrResponseFormat};
@ -20,20 +20,14 @@ pub(crate) trait BaseOcrConfig: Send + Sync + Sized + 'static {
fn map_ocr_params(
&self,
_non_default_params: &OcrArguments,
optional_params: &OcrArguments,
_model: &str,
) -> Result<OcrArguments, crate::ocr::Error> {
Ok(optional_params.clone())
}
fn parse_options(
&self,
arguments: &OcrArguments,
arguments: &CallArguments,
model: &str,
) -> Result<Self::OcrParams, crate::ocr::Error> {
self.map_ocr_params(arguments, &OcrArguments::default(), model)?
.parse()
Ok(parse_options(
&arguments
.select(self.get_supported_ocr_params(model))
.into(),
)?)
}
fn async_transform_ocr_request(

View file

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

View file

@ -5,7 +5,6 @@ use serde_with::serde_as;
use crate::constants::{COHERE_API_KEY_ENV, COHERE_PARSE_API_BASE};
use crate::llms::base_llm::ocr::transformation::{BaseOcrConfig, OcrRequestContext};
use crate::ocr::OcrArguments;
use crate::ocr::OcrClient;
use crate::ocr::document::InlineDocument;
use crate::ocr::prepare::{credential_env, transform_request_body};
@ -24,7 +23,7 @@ pub(crate) enum OutputFormat {
}
#[derive(Default, Deserialize, Serialize)]
pub(crate) struct CohereParams {
pub(crate) struct CohereOptions {
#[serde(skip_serializing_if = "Option::is_none")]
pub output_format: Option<OutputFormat>,
}
@ -43,16 +42,6 @@ pub(crate) enum CohereParseDocument {
ImageUrl { image_url: String },
}
impl CohereParseDocument {
pub(crate) fn as_document(&self) -> OcrDocument {
let Self::ImageUrl { image_url } = self;
OcrDocument::ImageUrl {
image_url: image_url.clone(),
extra_fields: Default::default(),
}
}
}
#[derive(Deserialize)]
pub(crate) struct CohereResponse {
#[serde(default)]
@ -96,7 +85,7 @@ impl CohereParseConfig {
&self,
model: &str,
document: OcrDocument,
optional_params: &CohereParams,
optional_params: &CohereOptions,
_headers: &[(String, String)],
) -> Result<CohereRequest, crate::ocr::Error> {
let image_url = image_url(document)?;
@ -105,7 +94,7 @@ impl CohereParseConfig {
}
impl BaseOcrConfig for CohereParseConfig {
type OcrParams = CohereParams;
type OcrParams = CohereOptions;
type ProviderRequest = CohereRequest;
type ProviderResponse = CohereResponse;
@ -113,34 +102,11 @@ impl BaseOcrConfig for CohereParseConfig {
&["output_format", "req_format"]
}
fn map_ocr_params(
&self,
non_default_params: &OcrArguments,
optional_params: &OcrArguments,
_model: &str,
) -> Result<OcrArguments, crate::ocr::Error> {
let overrides: OcrArguments = non_default_params
.select(&["output_format", "req_format"])
.into_iter()
.filter(|(_, value)| !value.is_null())
.collect();
overrides.parse::<CohereParams>()?;
if let Some(value) = overrides.get("req_format") {
serde_json::from_value::<crate::ocr::types::OcrResponseFormat>(value.clone())
.map_err(|_| crate::ocr::Error::RequestFormat)?;
}
Ok(optional_params
.iter()
.chain(overrides.iter())
.map(|(name, value)| (name.clone(), value.clone()))
.collect())
}
async fn async_transform_ocr_request(
&self,
model: &str,
document: OcrDocument,
optional_params: &CohereParams,
optional_params: &CohereOptions,
headers: &[(String, String)],
_context: OcrRequestContext<'_>,
) -> Result<CohereRequest, crate::ocr::Error> {
@ -162,7 +128,7 @@ impl CohereParseConfig {
request: &LiteLLMOcrRequest,
client: &OcrClient,
) -> Result<reqwest::Request, crate::ocr::Error> {
let params = self.parse_options(&request.optional_params, &request.model)?;
let params = self.map_ocr_params(&request.optional_params, &request.model)?;
let headers = self.validate_environment(&request.connection, &credential_env)?;
let url = self.get_complete_url(
request
@ -184,7 +150,7 @@ impl CohereParseConfig {
)
.await?;
transform_request_body(client, request, &url, &headers, true, body, |body| {
validate_document(&body.document.as_document())
validate_document(&crate::ocr::prepare::body_document(body)?)
})
.await
}
@ -236,7 +202,7 @@ fn image_url(document: OcrDocument) -> Result<String, crate::ocr::Error> {
Ok(image_url)
}
fn build_request(model: &str, image_url: String, params: &CohereParams) -> CohereRequest {
fn build_request(model: &str, image_url: String, params: &CohereOptions) -> CohereRequest {
CohereRequest {
model: model.into(),
document: CohereParseDocument::ImageUrl { image_url },
@ -358,43 +324,62 @@ mod tests {
use super::*;
use serde_json::json;
#[test]
fn mapping_merges_non_null_supported_overrides_and_preserves_supplied_options() {
let supplied = serde_json::from_value(json!({
"output_format":"blocks", "req_format":"native", "extension":false
#[tokio::test]
async fn composed_body_preserves_native_document_fields_and_untyped_overrides() {
let mut 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]}}
}
}),
);
request.document = serde_json::from_value(json!({
"type":"image_url","image_url":"https://example.com/original.png"
}))
.unwrap();
let overrides = serde_json::from_value(json!({
"output_format":null, "req_format":null, "ignored":true
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(&overrides, &supplied, "parse")
.map_ocr_params(&arguments, "parse")
} else {
CohereParseConfig.map_ocr_params(&overrides, &supplied, "parse")
CohereParseConfig.map_ocr_params(&arguments, "parse")
}
.unwrap();
assert_eq!(mapped, supplied);
}
for overrides in [
json!({"output_format":"html"}),
json!({"req_format":"invalid"}),
] {
let overrides = serde_json::from_value(overrides).unwrap();
assert!(
CohereParseConfig
.map_ocr_params(&overrides, &supplied, "parse")
.is_err()
assert_eq!(
serde_json::to_value(mapped).unwrap(),
json!({"output_format":"blocks"})
);
}
let overrides = serde_json::from_value(json!({"output_format":"markdown"})).unwrap();
let mapped = CohereParseConfig
.map_ocr_params(&overrides, &supplied, "parse")
.unwrap();
assert_eq!(mapped["output_format"], "markdown");
assert_eq!(mapped["extension"], false);
assert_eq!(arguments["req_format"], "native");
assert_eq!(arguments["extension"], false);
let invalid = serde_json::from_value(json!({"output_format":"html"})).unwrap();
assert!(CohereParseConfig.map_ocr_params(&invalid, "parse").is_err());
}
#[test]
@ -463,7 +448,7 @@ mod tests {
)
.unwrap();
let params = CohereParseConfig
.parse_options(&arguments, "parse")
.map_ocr_params(&arguments, "parse")
.unwrap();
assert_eq!(
serde_json::to_value(&params).unwrap(),
@ -668,10 +653,10 @@ mod tests {
Err(crate::ocr::Error::CohereImageOnly)
);
}
assert!(serde_json::from_value::<CohereParams>(json!({"output_format":"html"})).is_err());
assert!(serde_json::from_value::<CohereOptions>(json!({"output_format":"html"})).is_err());
for format in ["markdown", "blocks"] {
assert!(
serde_json::from_value::<CohereParams>(json!({"output_format":format})).is_ok()
serde_json::from_value::<CohereOptions>(json!({"output_format":format})).is_ok()
);
}
let request = CohereParseConfig

View file

@ -3,7 +3,6 @@ use serde_json::Value;
use crate::constants::MISTRAL_OCR_API_BASE;
use crate::llms::base_llm::ocr::transformation::{BaseOcrConfig, OcrRequestContext};
use crate::ocr::OcrArguments;
use crate::ocr::OcrClient;
use crate::ocr::prepare::{credential_env, transform_request_body};
use crate::ocr::types::{
@ -78,17 +77,6 @@ impl BaseOcrConfig for MistralOCRConfig {
]
}
fn map_ocr_params(
&self,
non_default_params: &OcrArguments,
_optional_params: &OcrArguments,
model: &str,
) -> Result<OcrArguments, crate::ocr::Error> {
Ok(non_default_params
.select(self.get_supported_ocr_params(model))
.into())
}
async fn async_transform_ocr_request(
&self,
model: &str,
@ -115,7 +103,7 @@ impl MistralOCRConfig {
request: &LiteLLMOcrRequest,
client: &OcrClient,
) -> Result<reqwest::Request, crate::ocr::Error> {
let params = self.parse_options(&request.optional_params, &request.model)?;
let params = self.map_ocr_params(&request.optional_params, &request.model)?;
let headers = self.validate_environment(&request.connection, &credential_env)?;
let url = self.get_complete_url(request.connection.api_base.as_deref())?;
let body = self
@ -275,18 +263,17 @@ mod tests {
}
#[test]
fn parameter_mapping_uses_only_supported_non_default_values() {
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 supplied = serde_json::from_value(json!({"pages":[9],"extension":true})).unwrap();
let params = MistralOCRConfig
.map_ocr_params(&input, &supplied, "model")
.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]
@ -328,13 +315,8 @@ mod tests {
}
fn mapped_params(value: Value) -> Value {
let params = serde_json::from_value::<OcrArguments>(value).unwrap();
serde_json::to_value(
MistralOCRConfig
.map_ocr_params(&params, &OcrArguments::default(), "model")
.unwrap(),
)
.unwrap()
let params = serde_json::from_value(value).unwrap();
serde_json::to_value(MistralOCRConfig.map_ocr_params(&params, "model").unwrap()).unwrap()
}
fn document() -> OcrDocument {
@ -402,7 +384,7 @@ mod tests {
}
#[rstest]
fn map_ocr_params_filters_unsupported_non_default_params() {
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());

View file

@ -3,9 +3,9 @@ use std::collections::BTreeMap;
use serde::{Deserialize, Deserializer, Serialize};
use serde_json::{Map, Value, json};
use crate::call_arguments::compose_body;
use crate::constants::{REDUCTO_API_BASE, REDUCTO_API_KEY_ENV, REDUCTO_ID_PREFIX};
use crate::llms::base_llm::ocr::transformation::{BaseOcrConfig, OcrRequestContext};
use crate::ocr::OcrArguments;
use crate::ocr::OcrClient;
use crate::ocr::document::InlineDocument;
use crate::ocr::prepare::{build_http_request, credential_env, guardrail_document};
@ -87,19 +87,6 @@ impl BaseOcrConfig for ReductoParseV3Config {
&["formatting", "retrieval", "settings"]
}
fn map_ocr_params(
&self,
non_default_params: &OcrArguments,
optional_params: &OcrArguments,
model: &str,
) -> Result<OcrArguments, crate::ocr::Error> {
Ok(map_ocr_params(
non_default_params,
optional_params,
self.get_supported_ocr_params(model),
))
}
#[tracing::instrument(
name = "async_transform_ocr_request",
target = "litellm::function_trace",
@ -136,7 +123,7 @@ impl ReductoParseV3Config {
request: &LiteLLMOcrRequest,
client: &OcrClient,
) -> Result<reqwest::Request, crate::ocr::Error> {
let params = self.parse_options(&request.optional_params, &request.model)?;
let params = self.map_ocr_params(&request.optional_params, &request.model)?;
let headers = validate_environment(&request.connection, &credential_env)?;
let url = get_complete_url(request.connection.api_base.as_deref())?;
let (document, headers) = guardrail_document(request, &url, &headers).await?;
@ -152,9 +139,11 @@ impl ReductoParseV3Config {
},
)
.await?;
let body = request
.optional_params
.compose_body(&body, self.get_supported_ocr_params(&request.model))?;
let body = compose_body(
&request.optional_params,
&body,
self.get_supported_ocr_params(&request.model),
)?;
build_http_request(client, request, &url, &headers, &body)
}
}
@ -171,19 +160,6 @@ impl BaseOcrConfig for ReductoParseLegacyConfig {
&["enhance"]
}
fn map_ocr_params(
&self,
non_default_params: &OcrArguments,
optional_params: &OcrArguments,
model: &str,
) -> Result<OcrArguments, crate::ocr::Error> {
Ok(map_ocr_params(
non_default_params,
optional_params,
self.get_supported_ocr_params(model),
))
}
#[tracing::instrument(
name = "async_transform_ocr_request",
target = "litellm::function_trace",
@ -217,7 +193,7 @@ impl ReductoParseLegacyConfig {
request: &LiteLLMOcrRequest,
client: &OcrClient,
) -> Result<reqwest::Request, crate::ocr::Error> {
let params = self.parse_options(&request.optional_params, &request.model)?;
let params = self.map_ocr_params(&request.optional_params, &request.model)?;
let headers = validate_environment(&request.connection, &credential_env)?;
let url = get_complete_url(request.connection.api_base.as_deref())?;
let (document, headers) = guardrail_document(request, &url, &headers).await?;
@ -233,29 +209,15 @@ impl ReductoParseLegacyConfig {
},
)
.await?;
let body = request
.optional_params
.compose_body(&body, self.get_supported_ocr_params(&request.model))?;
let body = compose_body(
&request.optional_params,
&body,
self.get_supported_ocr_params(&request.model),
)?;
build_http_request(client, request, &url, &headers, &body)
}
}
fn map_ocr_params(
non_default_params: &OcrArguments,
optional_params: &OcrArguments,
supported_params: &[&str],
) -> OcrArguments {
optional_params
.iter()
.chain(
non_default_params
.iter()
.filter(|(name, _)| supported_params.contains(&name.as_str())),
)
.map(|(name, value)| (name.clone(), value.clone()))
.collect()
}
fn present_nullable<'de, D: Deserializer<'de>, T: Deserialize<'de>>(
deserializer: D,
) -> Result<Option<Option<T>>, D::Error> {
@ -499,31 +461,27 @@ mod tests {
use super::*;
#[test]
fn mapping_merges_supported_overrides_including_null_with_supplied_options() {
let supplied = serde_json::from_value(json!({
"formatting":{"old":true}, "enhance":true, "extension":false
}))
.unwrap();
fn options_preserve_null_and_select_the_provider_fields() {
let overrides = serde_json::from_value(json!({
"formatting":null, "enhance":null, "ignored":true
}))
.unwrap();
let v3 = ReductoParseV3Config
.map_ocr_params(&overrides, &supplied, "parse-v3")
.map_ocr_params(&overrides, "parse-v3")
.unwrap();
assert_eq!(
serde_json::to_value(v3).unwrap(),
json!({
"formatting":null, "enhance":true, "extension":false
"formatting":null
})
);
let legacy = ReductoParseLegacyConfig
.map_ocr_params(&overrides, &supplied, "parse-legacy")
.map_ocr_params(&overrides, "parse-legacy")
.unwrap();
assert_eq!(
serde_json::to_value(legacy).unwrap(),
json!({
"formatting":{"old":true}, "enhance":null, "extension":false
"enhance":null
})
);
}
@ -558,7 +516,7 @@ mod tests {
serde_json::from_value(json!({"formatting":null,"settings":{},"unknown":true}))
.unwrap();
let params = ReductoParseV3Config
.parse_options(&overrides, "parse-v3")
.map_ocr_params(&overrides, "parse-v3")
.unwrap();
let client = crate::ocr::test_support::ocr_client();
let connection = OcrConnection::default();
@ -586,7 +544,7 @@ mod tests {
})
);
let absent = ReductoParseV3Config
.parse_options(&OcrArguments::default(), "parse-v3")
.map_ocr_params(&crate::call_arguments::CallArguments::default(), "parse-v3")
.unwrap();
assert_eq!(serde_json::to_value(absent).unwrap(), json!({}));
}
@ -603,7 +561,7 @@ mod tests {
let overrides =
serde_json::from_value(json!({"enhance":value,"unknown":true})).unwrap();
let params = ReductoParseLegacyConfig
.parse_options(&overrides, "parse-legacy")
.map_ocr_params(&overrides, "parse-legacy")
.unwrap();
assert_eq!(
serde_json::to_value(build_legacy_body(

View file

@ -154,7 +154,7 @@ impl VertexAIDeepSeekOCRConfig {
request: &LiteLLMOcrRequest,
client: &OcrClient,
) -> Result<reqwest::Request, crate::ocr::Error> {
let params = self.parse_options(&request.optional_params, &request.model)?;
let params = self.map_ocr_params(&request.optional_params, &request.model)?;
let config = VertexConfig::from_sourced_optional_params(
&request.optional_params,
&request.input_sources,
@ -396,17 +396,25 @@ mod tests {
use super::{VertexAIDeepSeekOCRConfig, provider_model};
#[test]
fn inherited_parameter_mapping_only_returns_supplied_optional_params() {
fn unconsumed_options_remain_available_for_body_composition() {
use crate::llms::base_llm::ocr::transformation::BaseOcrConfig;
use serde_json::json;
let non_default = serde_json::from_value(json!({"temperature":0.5})).unwrap();
let supplied = serde_json::from_value(json!({"max_tokens":100,"extension":null})).unwrap();
let arguments =
serde_json::from_value(json!({"temperature":0.5,"extension":null})).unwrap();
assert_eq!(
VertexAIDeepSeekOCRConfig
.map_ocr_params(&non_default, &supplied, "deepseek-ocr")
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(),
supplied
json!({"model":"deepseek-ocr","temperature":0.5,"extension":null})
);
}

View file

@ -2,7 +2,6 @@ use super::common_utils::validate_destination;
use crate::llms::base_llm::ocr::transformation::{BaseOcrConfig, OcrRequestContext};
use crate::llms::mistral::ocr::MistralOcrResponse;
use crate::llms::mistral::ocr::transformation::{MistralOCRConfig, MistralOcrRequest};
use crate::ocr::OcrArguments;
use crate::ocr::OcrClient;
use crate::ocr::document::{inline_remote_document, validate_inline_document};
use crate::ocr::prepare::{credential_env, transform_request_body};
@ -24,15 +23,6 @@ impl BaseOcrConfig for VertexAIOCRConfig {
MistralOCRConfig.get_supported_ocr_params(model)
}
fn map_ocr_params(
&self,
non_default_params: &OcrArguments,
optional_params: &OcrArguments,
model: &str,
) -> Result<OcrArguments, crate::ocr::Error> {
MistralOCRConfig.map_ocr_params(non_default_params, optional_params, model)
}
async fn async_transform_ocr_request(
&self,
model: &str,
@ -65,7 +55,7 @@ impl VertexAIOCRConfig {
request: &LiteLLMOcrRequest,
client: &OcrClient,
) -> Result<reqwest::Request, crate::ocr::Error> {
let params = self.parse_options(&request.optional_params, &request.model)?;
let params = self.map_ocr_params(&request.optional_params, &request.model)?;
let config = VertexConfig::from_sourced_optional_params(
&request.optional_params,
&request.input_sources,
@ -103,7 +93,7 @@ impl VertexAIOCRConfig {
&authentication.headers,
retains_document,
body,
|body| validate_inline_document(&body.document),
|body| validate_inline_document(&crate::ocr::prepare::body_document(body)?),
)
.await
}

View file

@ -1,93 +0,0 @@
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 OcrArguments(Map<String, Value>);
impl OcrArguments {
pub(crate) fn parse<T: DeserializeOwned>(&self) -> Result<T, super::Error> {
super::wire::decode_request_value(Value::Object(self.0.clone()), "optional_params")
}
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()
}
pub(crate) fn compose_body<B: Serialize>(
&self,
body: &B,
consumed: &[&str],
) -> Result<Value, super::Error> {
let Value::Object(fields) =
serde_json::to_value(body).map_err(|_| crate::params::Error::Body)?
else {
return Err(crate::params::Error::Body.into());
};
let overrides = match self.get("extra_body") {
None | Some(Value::Null) => None,
Some(Value::Object(fields)) => Some(fields),
Some(_) => return Err(crate::params::Error::ExtraBody.into()),
};
let extensions = self.iter().filter(|(name, _)| {
!consumed.contains(&name.as_str())
&& name.as_str() != "extra_body"
&& !crate::params::is_control_param(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 OcrArguments {
type Target = Map<String, Value>;
fn deref(&self) -> &Self::Target {
&self.0
}
}
impl From<Map<String, Value>> for OcrArguments {
fn from(values: Map<String, Value>) -> Self {
Self(values)
}
}
impl From<OcrArguments> for Map<String, Value> {
fn from(arguments: OcrArguments) -> Self {
arguments.0
}
}
impl FromIterator<(String, Value)> for OcrArguments {
fn from_iter<T: IntoIterator<Item = (String, Value)>>(iter: T) -> Self {
Self(iter.into_iter().collect())
}
}
impl IntoIterator for OcrArguments {
type Item = (String, Value);
type IntoIter = serde_json::map::IntoIter;
fn into_iter(self) -> Self::IntoIter {
self.0.into_iter()
}
}

View file

@ -88,6 +88,14 @@ pub enum Error {
Headers(#[from] crate::http_utils::HeaderError),
}
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 is_request(&self) -> bool {
matches!(

View file

@ -1,6 +1,4 @@
mod arguments;
mod error;
pub use arguments::OcrArguments;
pub use error::Error;
pub mod client;
pub(crate) mod document;

View file

@ -1,4 +1,4 @@
use serde::{Serialize, de::DeserializeOwned};
use serde::Serialize;
use serde_json::Value;
use super::OcrClient;
@ -12,21 +12,19 @@ pub(crate) async fn transform_request_body<B>(
headers: &[(String, String)],
retains_document: bool,
body: B,
validate: impl Fn(&B) -> Result<(), super::Error>,
validate: impl Fn(&Value) -> Result<(), super::Error>,
) -> Result<reqwest::Request, super::Error>
where
B: Serialize + DeserializeOwned,
B: Serialize,
{
let composed = request.optional_params.compose_body(
let composed = crate::call_arguments::compose_body(
&request.optional_params,
&body,
request.config.get_supported_ocr_params(&request.model),
)?;
let composed = OcrWireBody::<B>::decode(composed, "body")?;
validate(&composed.body)?;
validate(&composed)?;
let (body, headers) = if request.hooks.intercepts_requests() {
let body = serde_json::to_value(composed).map_err(|_| super::Error::RequestField {
path: "body".into(),
})?;
let body = composed;
let retained_fields = request
.optional_params
.keys()
@ -52,9 +50,13 @@ where
retained_fields,
})
.await?;
let body = OcrWireBody::<B>::decode(changed.body, "guardrail.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 {
(composed, headers.to_vec())
};
@ -106,24 +108,19 @@ pub(crate) async fn guardrail_document(
Ok((document, changed.headers))
}
#[derive(Serialize)]
struct OcrWireBody<B> {
#[serde(skip)]
body: B,
#[serde(flatten)]
fields: serde_json::Map<String, Value>,
}
impl<B: Serialize + DeserializeOwned> OcrWireBody<B> {
fn decode(value: Value, prefix: &str) -> Result<Self, super::Error> {
let body: B = super::wire::decode_request_value(value.clone(), prefix)?;
let Value::Object(fields) = value else {
return Err(super::Error::RequestField {
path: prefix.into(),
});
};
Ok(Self { body, fields })
}
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 source = document
.iter()
.filter(|(name, _)| matches!(name.as_str(), "type" | "image_url" | "document_url"))
.map(|(name, value)| (name.clone(), value.clone()))
.collect();
super::wire::decode_request_value(Value::Object(source), "body.document")
}
pub(crate) fn credential_env(name: &str) -> Option<String> {

View file

@ -7,9 +7,9 @@ use serde::{Deserialize, Serialize};
use serde_json::{Map, Value};
use serde_with::serde_as;
use super::OcrArguments;
use super::hooks::{NoopOcrHooks, OcrHooks};
use super::provider_config::{OcrConfigKind, resolve_provider_config};
use crate::call_arguments::CallArguments;
use crate::constants::OCR_HTTP_TIMEOUT_SECS;
use litellm_auth::{InputSource, TokenProviderHandle};
@ -97,7 +97,7 @@ pub struct LiteLLMOcrRequest {
pub connection: OcrConnection,
pub hooks: Arc<dyn OcrHooks>,
pub litellm_call_id: Option<String>,
pub optional_params: OcrArguments,
pub optional_params: CallArguments,
pub input_sources: BTreeMap<String, InputSource>,
pub azure_ad_token_provider: Option<TokenProviderHandle>,
pub(crate) config: OcrConfigKind,
@ -108,7 +108,7 @@ impl LiteLLMOcrRequest {
model: String,
document: OcrDocument,
custom_llm_provider: Option<&str>,
optional_params: OcrArguments,
optional_params: CallArguments,
) -> Result<Self, super::Error> {
let (model, config) = resolve_provider_config(&model, custom_llm_provider)?;

View file

@ -1,8 +1,8 @@
use std::collections::BTreeMap;
use std::time::Duration;
use super::OcrArguments;
use super::types::{LiteLLMOcrRequest, OcrConnection, OcrDocument};
use crate::call_arguments::{ArgumentSpec, CallArguments};
use litellm_auth::InputSource;
use serde::{
Deserialize,
@ -11,6 +11,7 @@ use serde::{
use serde_json::{Map, Value};
const COMMON_OPTION_FIELDS: &[&str] = &["req_format", "extra_body", "max_response_bytes"];
pub const BOUND_FIELDS: &[&str] = &["model", "document", "timeout", "input_sources"];
const AZURE_AUTH_OPTION_FIELDS: &[&str] = &[
"azure_ad_token",
"tenant_id",
@ -31,12 +32,6 @@ const VERTEX_AUTH_OPTION_FIELDS: &[&str] = &[
"vertex_ai_location",
];
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub struct OptionalParamSpec {
pub name: &'static str,
pub secret: bool,
}
#[derive(Debug)]
pub struct DecodedOcrResponse<T> {
pub data: T,
@ -54,7 +49,7 @@ pub struct OcrWireRequest {
pub custom_llm_provider: Option<String>,
pub extra_headers: Option<Map<String, Value>>,
#[serde(default)]
pub optional_params: OcrArguments,
pub optional_params: CallArguments,
#[serde(default)]
pub input_sources: BTreeMap<String, InputSource>,
pub timeout_seconds: Option<f64>,
@ -91,11 +86,11 @@ pub fn consumed_optional_param_names(
pub fn consumed_optional_params(
model: &str,
custom_llm_provider: Option<&str>,
) -> Result<Vec<OptionalParamSpec>, crate::ocr::Error> {
) -> Result<Vec<ArgumentSpec>, crate::ocr::Error> {
consumed_optional_param_names(model, custom_llm_provider).map(|names| {
names
.into_iter()
.map(|name| OptionalParamSpec {
.map(|name| ArgumentSpec {
name,
secret: matches!(
name,
@ -110,17 +105,6 @@ pub fn consumed_optional_params(
})
}
pub fn project_argument(
name: &str,
consumed: &[OptionalParamSpec],
host_fields: &[String],
) -> bool {
consumed.iter().any(|field| field.name == name)
|| (!host_fields.iter().any(|field| field == name)
&& !crate::params::is_control_param(name)
&& !matches!(name, "model" | "document" | "timeout" | "input_sources"))
}
pub fn decode_request(wire: OcrWireRequest) -> Result<LiteLLMOcrRequest, crate::ocr::Error> {
let api_key_source = source_for(&wire.input_sources, "api_key");
let api_base_source = source_for(&wire.input_sources, "api_base");
@ -269,14 +253,14 @@ mod tests {
#[test]
fn core_selects_consumed_values_without_serializing_host_objects() {
let fields = consumed_optional_params("mistral/model", None).unwrap();
let host_fields = vec!["metadata".into(), "callbacks".into(), "id".into()];
assert!(project_argument("future_option", &fields, &host_fields));
assert!(project_argument("extra_body", &fields, &host_fields));
assert!(project_argument("id", &fields, &host_fields));
assert!(!project_argument("metadata", &fields, &host_fields));
assert!(!project_argument("callbacks", &fields, &host_fields));
assert!(!project_argument("api_key", &fields, &host_fields));
assert!(!project_argument("document", &fields, &host_fields));
use crate::call_arguments::should_project;
assert!(should_project("future_option", &fields, BOUND_FIELDS));
assert!(should_project("extra_body", &fields, BOUND_FIELDS));
assert!(should_project("id", &fields, BOUND_FIELDS));
assert!(!should_project("metadata", &fields, BOUND_FIELDS));
assert!(!should_project("callbacks", &fields, BOUND_FIELDS));
assert!(!should_project("api_key", &fields, BOUND_FIELDS));
assert!(!should_project("document", &fields, BOUND_FIELDS));
}
#[test]

View file

@ -90,19 +90,15 @@ pub(crate) fn python_timeout_seconds(py: Python<'_>, timeout: Py<PyAny>) -> PyRe
pub(crate) fn project_optional_fields(
kwargs: &Bound<'_, PyDict>,
fields: &[litellm_core::ocr::wire::OptionalParamSpec],
fields: &[litellm_core::call_arguments::ArgumentSpec],
bound_fields: &[&str],
) -> PyResult<Map<String, Value>> {
let controls: Vec<String> = kwargs
.py()
.import("litellm.types.utils")?
.getattr("all_litellm_params")?
.extract()?;
kwargs
.iter()
.map(|(name, value)| Ok((name.extract::<String>()?, value)))
.filter_map(|entry: PyResult<_>| match entry {
Ok((name, value))
if litellm_core::ocr::wire::project_argument(&name, fields, &controls) =>
if litellm_core::call_arguments::should_project(&name, fields, bound_fields) =>
{
Some(from_py(&value).map(|value| (name, value)))
}

View file

@ -29,12 +29,12 @@ pub(super) struct ProjectedOcrCall {
pub fields: ProjectedOcrFields,
}
struct OcrArguments<'a, 'py> {
struct PythonOcrFields<'a, 'py> {
request: &'a Bound<'py, PyAny>,
kwargs: &'a Bound<'py, PyDict>,
}
impl<'py> OcrArguments<'_, 'py> {
impl<'py> PythonOcrFields<'_, 'py> {
fn lookup(&self, name: &str) -> PyResult<Bound<'py, PyAny>> {
match self.kwargs.get_item(name)? {
Some(value) => Ok(value),
@ -116,7 +116,7 @@ pub(super) fn project_request(
kwargs: &Bound<'_, PyDict>,
) -> PyResult<ProjectedOcrCall> {
let boundary_request = request.clone().unbind();
let arguments = OcrArguments { request, kwargs };
let arguments = PythonOcrFields { request, kwargs };
let model = arguments.model()?;
let custom_llm_provider = arguments.custom_llm_provider()?;
let (wire_document, retained_document) =
@ -124,7 +124,8 @@ pub(super) fn project_request(
let api_key = arguments.api_key()?;
let specs = consumed_optional_params(&model, custom_llm_provider.as_deref())
.map_err(ocr_error_to_pyerr)?;
let optional_params = project_optional_fields(kwargs, &specs)?;
let optional_params =
project_optional_fields(kwargs, &specs, litellm_core::ocr::wire::BOUND_FIELDS)?;
let input_sources = request_input_sources(
kwargs,
optional_params
@ -191,8 +192,8 @@ mod tests {
fn arguments<'a, 'py>(
request: &'a Bound<'py, PyAny>,
kwargs: &'a Bound<'py, PyDict>,
) -> OcrArguments<'a, 'py> {
OcrArguments { request, kwargs }
) -> PythonOcrFields<'a, 'py> {
PythonOcrFields { request, kwargs }
}
fn project_document(

View file

@ -93,6 +93,49 @@ async def test_request_level_custom_pricing_reaches_logging_params_and_bills_the
assert "ocr_cost_per_page" not in ocr_server.requests[0].body
@pytest.mark.asyncio
@pytest.mark.parametrize("asynchronous", [False, True])
async def test_shared_boundary_preserves_provider_fields_and_python_objects(
ocr_server: RecordingServer, asynchronous: bool
) -> None:
marker: Final = object()
observed: Final = []
class Observe(Logging):
def pre_call(self, input, api_key, additional_args):
observed.append(self.model_call_details["litellm_params"]["metadata"]["marker"])
logger: Final = Observe(
model="mistral-ocr-latest",
messages=[],
stream=False,
call_type="aocr" if asynchronous else "ocr",
start_time=datetime.datetime.now(),
litellm_call_id="shared-boundary",
function_id="shared-boundary",
)
arguments: Final = {
"id": "provider-id",
"future_option": {"old": True},
"explicit_null": None,
"metadata": {"marker": marker},
"shared_session": marker,
"litellm_logging_obj": logger,
"extra_body": {"future_option": {"nested": [None, False, 0]}, "metadata": {"provider": True}},
}
if asynchronous:
await call_aocr(ocr_server, **arguments)
else:
call_ocr(ocr_server, **arguments)
body: Final = ocr_server.requests[0].body
assert body["id"] == "provider-id"
assert body["future_option"] == {"nested": [None, False, 0]}
assert "explicit_null" in body and body["explicit_null"] is None
assert body["metadata"] == {"provider": True}
assert "shared_session" not in body
assert observed == [marker] and observed[0] is marker
@pytest.mark.asyncio
async def test_response_replacement_finalized_before_dispatch_in_caller_task(ocr_server: RecordingServer) -> None:
caller: Final = asyncio.current_task()