diff --git a/litellm-rust/crates/core/src/call_arguments.rs b/litellm-rust/crates/core/src/call_arguments.rs new file mode 100644 index 00000000000..6bd7fb4f8dc --- /dev/null +++ b/litellm-rust/crates/core/src/call_arguments.rs @@ -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); + +impl CallArguments { + pub(crate) fn select(&self, names: &[&str]) -> Map { + 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(arguments: &CallArguments) -> Result { + 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( + arguments: &CallArguments, + body: &B, + consumed: &[&str], +) -> Result { + 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; + + fn deref(&self) -> &Self::Target { + &self.0 + } +} + +impl From> for CallArguments { + fn from(values: Map) -> Self { + Self(values) + } +} + +impl From for Map { + fn from(arguments: CallArguments) -> Self { + arguments.0 + } +} + +impl FromIterator<(String, Value)> for CallArguments { + fn from_iter>(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, + } + let arguments: CallArguments = + serde_json::from_value(json!({"enabled":null,"future":0})).unwrap(); + assert!( + parse_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::(&invalid).err().unwrap().path, + "enabled" + ); + } +} diff --git a/litellm-rust/crates/core/src/lib.rs b/litellm-rust/crates/core/src/lib.rs index 5228160ac9a..f67d0377516 100644 --- a/litellm-rust/crates/core/src/lib.rs +++ b/litellm-rust/crates/core/src/lib.rs @@ -1,4 +1,5 @@ pub mod audio_transcription; +pub mod call_arguments; pub mod call_lifecycle; pub mod chat_completions; pub mod constants; diff --git a/litellm-rust/crates/core/src/llms/azure_ai/ocr/cohere_parse_transformation.rs b/litellm-rust/crates/core/src/llms/azure_ai/ocr/cohere_parse_transformation.rs index 0d16899fc05..97c327b9e69 100644 --- a/litellm-rust/crates/core/src/llms/azure_ai/ocr/cohere_parse_transformation.rs +++ b/litellm-rust/crates/core/src/llms/azure_ai/ocr/cohere_parse_transformation.rs @@ -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 { - 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 { @@ -65,7 +55,7 @@ impl AzureAICohereParseConfig { request: &LiteLLMOcrRequest, client: &OcrClient, ) -> Result { - 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 diff --git a/litellm-rust/crates/core/src/llms/azure_ai/ocr/document_intelligence/transformation.rs b/litellm-rust/crates/core/src/llms/azure_ai/ocr/document_intelligence/transformation.rs index b40b59f096f..afafd5bb194 100644 --- a/litellm-rust/crates/core/src/llms/azure_ai/ocr/document_intelligence/transformation.rs +++ b/litellm-rust/crates/core/src/llms/azure_ai/ocr/document_intelligence/transformation.rs @@ -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, #[serde(skip_serializing_if = "Option::is_none")] pub features: Option, - #[serde(skip_serializing_if = "Option::is_none")] - pub req_format: Option, - #[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 { - 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 { + 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::(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 { - 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] diff --git a/litellm-rust/crates/core/src/llms/azure_ai/ocr/transformation.rs b/litellm-rust/crates/core/src/llms/azure_ai/ocr/transformation.rs index b0cd3710e94..237b77598c3 100644 --- a/litellm-rust/crates/core/src/llms/azure_ai/ocr/transformation.rs +++ b/litellm-rust/crates/core/src/llms/azure_ai/ocr/transformation.rs @@ -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 { - 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 { - 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 } diff --git a/litellm-rust/crates/core/src/llms/base_llm/ocr/transformation.rs b/litellm-rust/crates/core/src/llms/base_llm/ocr/transformation.rs index 6af6e50a0f2..e7cc8eaa481 100644 --- a/litellm-rust/crates/core/src/llms/base_llm/ocr/transformation.rs +++ b/litellm-rust/crates/core/src/llms/base_llm/ocr/transformation.rs @@ -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 { - Ok(optional_params.clone()) - } - - fn parse_options( - &self, - arguments: &OcrArguments, + arguments: &CallArguments, model: &str, ) -> Result { - 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( diff --git a/litellm-rust/crates/core/src/llms/cohere/ocr/mod.rs b/litellm-rust/crates/core/src/llms/cohere/ocr/mod.rs index 477e9afd141..1c04ff00a67 100644 --- a/litellm-rust/crates/core/src/llms/cohere/ocr/mod.rs +++ b/litellm-rust/crates/core/src/llms/cohere/ocr/mod.rs @@ -1,3 +1,3 @@ pub(crate) mod transformation; -pub(crate) use transformation::{CohereParams, CohereResponse, validate_document}; +pub(crate) use transformation::{CohereOptions, CohereResponse, validate_document}; diff --git a/litellm-rust/crates/core/src/llms/cohere/ocr/transformation.rs b/litellm-rust/crates/core/src/llms/cohere/ocr/transformation.rs index 60c3fc1fa67..27b24b8b6db 100644 --- a/litellm-rust/crates/core/src/llms/cohere/ocr/transformation.rs +++ b/litellm-rust/crates/core/src/llms/cohere/ocr/transformation.rs @@ -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, } @@ -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 { 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 { - let overrides: OcrArguments = non_default_params - .select(&["output_format", "req_format"]) - .into_iter() - .filter(|(_, value)| !value.is_null()) - .collect(); - overrides.parse::()?; - if let Some(value) = overrides.get("req_format") { - serde_json::from_value::(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 { @@ -162,7 +128,7 @@ impl CohereParseConfig { request: &LiteLLMOcrRequest, client: &OcrClient, ) -> Result { - 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 { 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(¶ms).unwrap(), @@ -668,10 +653,10 @@ mod tests { Err(crate::ocr::Error::CohereImageOnly) ); } - assert!(serde_json::from_value::(json!({"output_format":"html"})).is_err()); + assert!(serde_json::from_value::(json!({"output_format":"html"})).is_err()); for format in ["markdown", "blocks"] { assert!( - serde_json::from_value::(json!({"output_format":format})).is_ok() + serde_json::from_value::(json!({"output_format":format})).is_ok() ); } let request = CohereParseConfig diff --git a/litellm-rust/crates/core/src/llms/mistral/ocr/transformation.rs b/litellm-rust/crates/core/src/llms/mistral/ocr/transformation.rs index 2d652bf275a..02640af480c 100644 --- a/litellm-rust/crates/core/src/llms/mistral/ocr/transformation.rs +++ b/litellm-rust/crates/core/src/llms/mistral/ocr/transformation.rs @@ -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 { - 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 { - 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::(value).unwrap(); - serde_json::to_value( - MistralOCRConfig - .map_ocr_params(¶ms, &OcrArguments::default(), "model") - .unwrap(), - ) - .unwrap() + let params = serde_json::from_value(value).unwrap(); + serde_json::to_value(MistralOCRConfig.map_ocr_params(¶ms, "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()); diff --git a/litellm-rust/crates/core/src/llms/reducto/ocr/transformation.rs b/litellm-rust/crates/core/src/llms/reducto/ocr/transformation.rs index a0bb3648a69..92c50d06916 100644 --- a/litellm-rust/crates/core/src/llms/reducto/ocr/transformation.rs +++ b/litellm-rust/crates/core/src/llms/reducto/ocr/transformation.rs @@ -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 { - 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 { - 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 { - 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 { - 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>, 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( diff --git a/litellm-rust/crates/core/src/llms/vertex_ai/ocr/deepseek_transformation.rs b/litellm-rust/crates/core/src/llms/vertex_ai/ocr/deepseek_transformation.rs index 7dfa7acb114..21e69e89210 100644 --- a/litellm-rust/crates/core/src/llms/vertex_ai/ocr/deepseek_transformation.rs +++ b/litellm-rust/crates/core/src/llms/vertex_ai/ocr/deepseek_transformation.rs @@ -154,7 +154,7 @@ impl VertexAIDeepSeekOCRConfig { request: &LiteLLMOcrRequest, client: &OcrClient, ) -> Result { - 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}) ); } diff --git a/litellm-rust/crates/core/src/llms/vertex_ai/ocr/transformation.rs b/litellm-rust/crates/core/src/llms/vertex_ai/ocr/transformation.rs index 87efdb02bd2..a32e0eb55da 100644 --- a/litellm-rust/crates/core/src/llms/vertex_ai/ocr/transformation.rs +++ b/litellm-rust/crates/core/src/llms/vertex_ai/ocr/transformation.rs @@ -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 { - 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 { - 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 } diff --git a/litellm-rust/crates/core/src/ocr/arguments.rs b/litellm-rust/crates/core/src/ocr/arguments.rs deleted file mode 100644 index e8ad6bd3cd3..00000000000 --- a/litellm-rust/crates/core/src/ocr/arguments.rs +++ /dev/null @@ -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); - -impl OcrArguments { - pub(crate) fn parse(&self) -> Result { - super::wire::decode_request_value(Value::Object(self.0.clone()), "optional_params") - } - - pub(crate) fn select(&self, names: &[&str]) -> Map { - self.iter() - .filter(|(name, _)| names.contains(&name.as_str())) - .map(|(name, value)| (name.clone(), value.clone())) - .collect() - } - - pub(crate) fn compose_body( - &self, - body: &B, - consumed: &[&str], - ) -> Result { - 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; - - fn deref(&self) -> &Self::Target { - &self.0 - } -} - -impl From> for OcrArguments { - fn from(values: Map) -> Self { - Self(values) - } -} - -impl From for Map { - fn from(arguments: OcrArguments) -> Self { - arguments.0 - } -} - -impl FromIterator<(String, Value)> for OcrArguments { - fn from_iter>(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() - } -} diff --git a/litellm-rust/crates/core/src/ocr/error.rs b/litellm-rust/crates/core/src/ocr/error.rs index a8be0a2d207..19d916804ed 100644 --- a/litellm-rust/crates/core/src/ocr/error.rs +++ b/litellm-rust/crates/core/src/ocr/error.rs @@ -88,6 +88,14 @@ pub enum Error { Headers(#[from] crate::http_utils::HeaderError), } +impl From 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!( diff --git a/litellm-rust/crates/core/src/ocr/mod.rs b/litellm-rust/crates/core/src/ocr/mod.rs index b1e4a52dc62..89cb2165b60 100644 --- a/litellm-rust/crates/core/src/ocr/mod.rs +++ b/litellm-rust/crates/core/src/ocr/mod.rs @@ -1,6 +1,4 @@ -mod arguments; mod error; -pub use arguments::OcrArguments; pub use error::Error; pub mod client; pub(crate) mod document; diff --git a/litellm-rust/crates/core/src/ocr/prepare.rs b/litellm-rust/crates/core/src/ocr/prepare.rs index d15d875bb77..e7625a2dc92 100644 --- a/litellm-rust/crates/core/src/ocr/prepare.rs +++ b/litellm-rust/crates/core/src/ocr/prepare.rs @@ -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( headers: &[(String, String)], retains_document: bool, body: B, - validate: impl Fn(&B) -> Result<(), super::Error>, + validate: impl Fn(&Value) -> Result<(), super::Error>, ) -> Result 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::::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::::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 { - #[serde(skip)] - body: B, - #[serde(flatten)] - fields: serde_json::Map, -} - -impl OcrWireBody { - fn decode(value: Value, prefix: &str) -> Result { - 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 { + 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 { diff --git a/litellm-rust/crates/core/src/ocr/types.rs b/litellm-rust/crates/core/src/ocr/types.rs index b97ff9394f0..f617bb7095a 100644 --- a/litellm-rust/crates/core/src/ocr/types.rs +++ b/litellm-rust/crates/core/src/ocr/types.rs @@ -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, pub litellm_call_id: Option, - pub optional_params: OcrArguments, + pub optional_params: CallArguments, pub input_sources: BTreeMap, pub azure_ad_token_provider: Option, 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 { let (model, config) = resolve_provider_config(&model, custom_llm_provider)?; diff --git a/litellm-rust/crates/core/src/ocr/wire.rs b/litellm-rust/crates/core/src/ocr/wire.rs index cf22975b5a0..996e5d3ecf4 100644 --- a/litellm-rust/crates/core/src/ocr/wire.rs +++ b/litellm-rust/crates/core/src/ocr/wire.rs @@ -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 { pub data: T, @@ -54,7 +49,7 @@ pub struct OcrWireRequest { pub custom_llm_provider: Option, pub extra_headers: Option>, #[serde(default)] - pub optional_params: OcrArguments, + pub optional_params: CallArguments, #[serde(default)] pub input_sources: BTreeMap, pub timeout_seconds: Option, @@ -91,11 +86,11 @@ pub fn consumed_optional_param_names( pub fn consumed_optional_params( model: &str, custom_llm_provider: Option<&str>, -) -> Result, crate::ocr::Error> { +) -> Result, 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 { 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] diff --git a/litellm-rust/crates/python-bridge/src/marshal.rs b/litellm-rust/crates/python-bridge/src/marshal.rs index 356e5b6d2dd..e9efd1c61ec 100644 --- a/litellm-rust/crates/python-bridge/src/marshal.rs +++ b/litellm-rust/crates/python-bridge/src/marshal.rs @@ -90,19 +90,15 @@ pub(crate) fn python_timeout_seconds(py: Python<'_>, timeout: Py) -> 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> { - let controls: Vec = kwargs - .py() - .import("litellm.types.utils")? - .getattr("all_litellm_params")? - .extract()?; kwargs .iter() .map(|(name, value)| Ok((name.extract::()?, 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))) } diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/project.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/project.rs index 006f1975fed..6b07fa02068 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/project.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/project.rs @@ -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> { match self.kwargs.get_item(name)? { Some(value) => Ok(value), @@ -116,7 +116,7 @@ pub(super) fn project_request( kwargs: &Bound<'_, PyDict>, ) -> PyResult { 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( diff --git a/tests/test_litellm_rust/ocr/test_lifecycle.py b/tests/test_litellm_rust/ocr/test_lifecycle.py index 9875cb1027d..9b96d820155 100644 --- a/tests/test_litellm_rust/ocr/test_lifecycle.py +++ b/tests/test_litellm_rust/ocr/test_lifecycle.py @@ -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()