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 157dc63a337..0d16899fc05 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,11 +1,11 @@ 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::ocr::OcrClient; use crate::ocr::document::{inline_remote_document, validate_inline_document}; use crate::ocr::prepare::{credential_env, transform_request_body}; use crate::ocr::types::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrDocument}; -use crate::params::OpaqueParams; use crate::url_utils::ApiUrl; use litellm_auth_azure::AzureAuthInputs; @@ -25,10 +25,10 @@ impl BaseOcrConfig for AzureAICohereParseConfig { fn map_ocr_params( &self, - non_default_params: &OpaqueParams, - optional_params: &OpaqueParams, + non_default_params: &OcrArguments, + optional_params: &OcrArguments, model: &str, - ) -> Result { + ) -> Result { CohereParseConfig.map_ocr_params(non_default_params, optional_params, model) } @@ -65,11 +65,7 @@ impl AzureAICohereParseConfig { request: &LiteLLMOcrRequest, client: &OcrClient, ) -> Result { - let params = self.map_ocr_params( - &request.optional_params, - &OpaqueParams::default(), - &request.model, - )?; + let params = self.parse_options(&request.optional_params, &request.model)?; let config = AzureAuthInputs { azure_ad_token_provider: request.azure_ad_token_provider.clone(), ..AzureAuthInputs::from_sourced_optional_params( 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 794e0eb2e04..b40b59f096f 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 @@ -6,6 +6,7 @@ use base64::{Engine, engine::general_purpose::STANDARD}; use reqwest::Url; use serde::{Deserialize, Deserializer, Serialize}; use serde_json::{Map, Value}; +use serde_with::serde_as; use tokio::time::Instant; use litellm_auth::{InputSource, Sourced}; @@ -18,6 +19,7 @@ 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; @@ -29,6 +31,7 @@ use crate::ocr::types::{ }; use crate::ocr::wire::DecodedOcrResponse; use crate::params::OpaqueParams; +use crate::serde_compat::{FiniteF64, LaxI64}; use crate::url_utils::ApiUrl; #[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] @@ -127,13 +130,15 @@ struct AzureDocumentIntelligenceAnalyzeResult { pub key_value_pairs: Option>>, } +#[serde_as] #[derive(Clone, Debug, Deserialize)] struct AzureDocumentIntelligencePage { - #[serde(rename = "pageNumber", default, deserialize_with = "optional_i64")] + #[serde(rename = "pageNumber")] + #[serde_as(deserialize_as = "Option")] pub page_number: Option, - #[serde(default, deserialize_with = "optional_f64")] + #[serde_as(deserialize_as = "Option")] pub width: Option, - #[serde(default, deserialize_with = "optional_f64")] + #[serde_as(deserialize_as = "Option")] pub height: Option, pub unit: Option, #[serde(default)] @@ -145,39 +150,6 @@ struct AzureDocumentIntelligenceLine { pub content: Option, } -fn optional_i64<'de, D: Deserializer<'de>>(deserializer: D) -> Result, D::Error> { - match Option::::deserialize(deserializer)? { - None | Some(Value::Null) => Ok(None), - Some(Value::Number(number)) => number - .as_i64() - .map(Some) - .ok_or_else(|| serde::de::Error::custom("expected an integer")), - Some(Value::String(value)) => value - .parse::() - .map(Some) - .map_err(|_| serde::de::Error::custom("expected an integer")), - Some(_) => Err(serde::de::Error::custom("expected an integer")), - } -} - -fn optional_f64<'de, D: Deserializer<'de>>(deserializer: D) -> Result, D::Error> { - match Option::::deserialize(deserializer)? { - None | Some(Value::Null) => Ok(None), - Some(Value::Number(number)) => number - .as_f64() - .filter(|value| value.is_finite()) - .map(Some) - .ok_or_else(|| serde::de::Error::custom("expected a finite number")), - Some(Value::String(value)) => value - .parse::() - .ok() - .filter(|value| value.is_finite()) - .map(Some) - .ok_or_else(|| serde::de::Error::custom("expected a finite number")), - Some(_) => Err(serde::de::Error::custom("expected a number")), - } -} - fn decode_input_params( params: Map, prefix: &str, @@ -330,7 +302,7 @@ fn transform_completed_response( let pages = result .pages .into_iter() - .map(normalize_page) + .map(transform_azure_page) .collect::, _>>()?; let pages_processed = i64::try_from(pages.len()).map_err(|_| crate::ocr::Error::NumericRange("pages"))?; @@ -346,26 +318,16 @@ fn transform_completed_response( }) } -fn normalize_page(page: AzureDocumentIntelligencePage) -> Result { +fn transform_azure_page(page: AzureDocumentIntelligencePage) -> Result { let index = page .page_number .unwrap_or(1) .checked_sub(1) .ok_or(crate::ocr::Error::NumericRange("page.pageNumber"))?; - let scale = if page.unit.as_deref().unwrap_or("inch") == "inch" { - AZURE_DI_DEFAULT_DPI as f64 - } else { - 1.0 - }; - let width = pixel_dimension( + let dimensions = convert_dimensions( page.width.unwrap_or(AZURE_DI_DEFAULT_WIDTH), - scale, - "page.width", - )?; - let height = pixel_dimension( page.height.unwrap_or(AZURE_DI_DEFAULT_HEIGHT), - scale, - "page.height", + page.unit.as_deref().unwrap_or("inch"), )?; let markdown = page .lines @@ -376,18 +338,31 @@ fn normalize_page(page: AzureDocumentIntelligencePage) -> Result Result { + let scale = if unit == "inch" { + AZURE_DI_DEFAULT_DPI as f64 + } else { + 1.0 + }; + Ok(OcrPageDimensions { + width: Some(pixel_dimension(width, scale, "page.width")?), + height: Some(pixel_dimension(height, scale, "page.height")?), + dpi: Some(AZURE_DI_DEFAULT_DPI), + }) +} + fn pixel_dimension(value: f64, scale: f64, field: &'static str) -> Result { let value = value * scale; - if !value.is_finite() || value < i64::MIN as f64 || value > i64::MAX as f64 { + if !value.is_finite() || value < i64::MIN as f64 || value >= -(i64::MIN as f64) { return Err(crate::ocr::Error::NumericRange(field)); } Ok(value.trunc() as i64) @@ -510,10 +485,10 @@ impl BaseOcrConfig for AzureDocumentIntelligenceOCRConfig { fn map_ocr_params( &self, - non_default_params: &OpaqueParams, - optional_params: &OpaqueParams, + non_default_params: &OcrArguments, + optional_params: &OcrArguments, _model: &str, - ) -> Result { + ) -> Result { let mapped = normalize_ocr_params(decode_input_params( non_default_params .iter() @@ -556,7 +531,7 @@ impl BaseOcrConfig for AzureDocumentIntelligenceOCRConfig { ) })) .collect(); - crate::ocr::wire::decode_request_value(Value::Object(fields), "optional_params") + Ok(fields) } async fn async_transform_ocr_request( @@ -617,11 +592,7 @@ impl AzureDocumentIntelligenceOCRConfig { request: &LiteLLMOcrRequest, client: &OcrClient, ) -> Result { - let params = self.map_ocr_params( - &request.optional_params, - &OpaqueParams::default(), - &request.model, - )?; + let params = self.parse_options(&request.optional_params, &request.model)?; let config = AzureAuthInputs { azure_ad_token_provider: request.azure_ad_token_provider.clone(), ..AzureAuthInputs::from_sourced_optional_params( @@ -752,13 +723,31 @@ mod tests { } #[test] - fn mapping_preserves_supplied_options_when_overrides_are_empty() { - let supplied = - serde_json::from_value(json!({"pages":"4", "features":"languages", "extension":true})) - .unwrap(); + fn empty_options_do_not_create_query_fields() { let overrides = serde_json::from_value(json!({"pages":[], "features":null, "req_format":"native"})) .unwrap(); + let mapped = AzureDocumentIntelligenceOCRConfig + .parse_options(&overrides, "model") + .unwrap(); + assert_eq!( + serde_json::to_value(mapped).unwrap(), + json!({ + "req_format":"native" + }) + ); + } + + #[test] + fn mapping_preserves_supplied_options_when_overrides_are_empty() { + let supplied = 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") .unwrap(); @@ -770,6 +759,20 @@ mod tests { ); } + #[test] + fn response_numbers_follow_python_validation_before_dimension_conversion() { + let response = AzureDocumentIntelligenceOCRConfig.decode_and_normalize_response( + "model", + br#"{"status":"succeeded","analyzeResult":{"pages":[{"pageNumber":2.0,"width":" 8.5 ","height":true}]}}"#, + OcrResponseFormat::Litellm, + ).unwrap(); + assert_eq!(response.pages[0].index, 1); + let dimensions = response.pages[0].dimensions.as_ref().unwrap(); + assert_eq!(dimensions.width, Some(816)); + assert_eq!(dimensions.height, Some(96)); + assert!(pixel_dimension(9_223_372_036_854_775_808.0, 1.0, "width").is_err()); + } + #[rstest] #[case(json!([0, 1, 2]), Some("1,2,3"))] #[case(json!([2, 0, 0, 1]), Some("1,2,3"))] 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 4c8a0a8b0ee..b0cd3710e94 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,6 +2,7 @@ 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}; @@ -28,10 +29,10 @@ impl BaseOcrConfig for AzureAIOCRConfig { fn map_ocr_params( &self, - non_default_params: &OpaqueParams, - optional_params: &OpaqueParams, + non_default_params: &OcrArguments, + optional_params: &OcrArguments, model: &str, - ) -> Result { + ) -> Result { MistralOCRConfig.map_ocr_params(non_default_params, optional_params, model) } @@ -67,11 +68,7 @@ impl AzureAIOCRConfig { request: &LiteLLMOcrRequest, client: &OcrClient, ) -> Result { - let params = self.map_ocr_params( - &request.optional_params, - &OpaqueParams::default(), - &request.model, - )?; + let params = self.parse_options(&request.optional_params, &request.model)?; let config = AzureAuthInputs { azure_ad_token_provider: request.azure_ad_token_provider.clone(), ..AzureAuthInputs::from_sourced_optional_params( 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 548041589e6..6af6e50a0f2 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,14 +4,14 @@ use std::sync::Arc; use serde::Serialize; use serde::de::DeserializeOwned; +use crate::ocr::OcrArguments; use crate::ocr::OcrClient; use crate::ocr::hooks::OcrHooks; use crate::ocr::types::{LiteLLMOcrResponse, OcrConnection, OcrDocument, OcrResponseFormat}; -use crate::params::OpaqueParams; pub(crate) trait BaseOcrConfig: Send + Sync + Sized + 'static { type OcrParams: DeserializeOwned + Send + Sync; - type ProviderRequest: Serialize + DeserializeOwned + Send; + type ProviderRequest: Serialize + Send; type ProviderResponse: DeserializeOwned + Send; fn get_supported_ocr_params(&self, _model: &str) -> &'static [&'static str] { @@ -20,14 +20,20 @@ pub(crate) trait BaseOcrConfig: Send + Sync + Sized + 'static { fn map_ocr_params( &self, - _non_default_params: &OpaqueParams, - optional_params: &OpaqueParams, + _non_default_params: &OcrArguments, + optional_params: &OcrArguments, _model: &str, + ) -> Result { + Ok(optional_params.clone()) + } + + fn parse_options( + &self, + arguments: &OcrArguments, + model: &str, ) -> Result { - crate::ocr::wire::decode_request_value( - serde_json::Value::Object(optional_params.clone().into()), - "optional_params", - ) + self.map_ocr_params(arguments, &OcrArguments::default(), model)? + .parse() } fn async_transform_ocr_request( @@ -45,6 +51,22 @@ pub(crate) trait BaseOcrConfig: Send + Sync + Sized + 'static { response: Self::ProviderResponse, ) -> Result; + fn decode_and_normalize_response( + &self, + model: &str, + raw_response: &[u8], + request_format: OcrResponseFormat, + ) -> Result { + let decoded = crate::ocr::wire::decode_response::( + raw_response, + request_format == OcrResponseFormat::Native, + )?; + Ok(LiteLLMOcrResponse { + provider_native_response: decoded.native, + ..self.normalize_response(model, decoded.data)? + }) + } + fn async_transform_ocr_response( &self, model: &str, @@ -58,14 +80,7 @@ pub(crate) trait BaseOcrConfig: Send + Sync + Sized + 'static { ) .await?; crate::ocr::handler::post_call(context.hooks, &bytes).await?; - let decoded = crate::ocr::wire::decode_response::( - &bytes, - context.request_format == OcrResponseFormat::Native, - )?; - Ok(LiteLLMOcrResponse { - provider_native_response: decoded.native, - ..self.normalize_response(model, decoded.data)? - }) + self.decode_and_normalize_response(model, &bytes, context.request_format) } } } 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 80b881c57ae..60c3fc1fa67 100644 --- a/litellm-rust/crates/core/src/llms/cohere/ocr/transformation.rs +++ b/litellm-rust/crates/core/src/llms/cohere/ocr/transformation.rs @@ -5,6 +5,7 @@ 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}; @@ -12,7 +13,6 @@ use crate::ocr::types::{ LiteLLMOcrRequest, LiteLLMOcrResponse, OcrConnection, OcrDocument, OcrPage, OcrPageImage, OcrUsageInfo, }; -use crate::params::OpaqueParams; use crate::url_utils::ApiUrl; #[derive(Clone, Copy, Debug, Default, Deserialize, Serialize)] @@ -27,17 +27,13 @@ pub(crate) enum OutputFormat { pub(crate) struct CohereParams { #[serde(skip_serializing_if = "Option::is_none")] pub output_format: Option, - #[serde(skip_serializing_if = "Option::is_none")] - pub req_format: Option, - #[serde(flatten)] - pub extra_fields: OpaqueParams, } #[derive(Deserialize, Serialize)] pub(crate) struct CohereRequest { pub model: String, pub document: CohereParseDocument, - pub output_format: OutputFormat, + pub output_format: String, } #[derive(Deserialize, Serialize)] @@ -119,25 +115,25 @@ impl BaseOcrConfig for CohereParseConfig { fn map_ocr_params( &self, - non_default_params: &OpaqueParams, - optional_params: &OpaqueParams, + non_default_params: &OcrArguments, + optional_params: &OcrArguments, _model: &str, - ) -> Result { - if let Some(value) = non_default_params - .get("req_format") - .filter(|value| !value.is_null()) - { + ) -> 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)?; } - let fields = optional_params + Ok(optional_params .iter() - .chain(non_default_params.iter().filter(|(name, value)| { - matches!(name.as_str(), "output_format" | "req_format") && !value.is_null() - })) + .chain(overrides.iter()) .map(|(name, value)| (name.clone(), value.clone())) - .collect(); - crate::ocr::wire::decode_request_value(Value::Object(fields), "optional_params") + .collect()) } async fn async_transform_ocr_request( @@ -166,11 +162,7 @@ impl CohereParseConfig { request: &LiteLLMOcrRequest, client: &OcrClient, ) -> Result { - let params = self.map_ocr_params( - &request.optional_params, - &OpaqueParams::default(), - &request.model, - )?; + let params = self.parse_options(&request.optional_params, &request.model)?; let headers = self.validate_environment(&request.connection, &credential_env)?; let url = self.get_complete_url( request @@ -248,7 +240,11 @@ fn build_request(model: &str, image_url: String, params: &CohereParams) -> Coher CohereRequest { model: model.into(), document: CohereParseDocument::ImageUrl { image_url }, - output_format: params.output_format.unwrap_or_default(), + output_format: match params.output_format.unwrap_or_default() { + OutputFormat::Markdown => "markdown", + OutputFormat::Blocks => "blocks", + } + .into(), } } @@ -362,6 +358,45 @@ 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 + })) + .unwrap(); + let overrides = serde_json::from_value(json!({ + "output_format":null, "req_format":null, "ignored":true + })) + .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") + } else { + CohereParseConfig.map_ocr_params(&overrides, &supplied, "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() + ); + } + 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); + } + #[test] fn billed_pages_accept_integral_doubles_and_reject_fractional_counts() { let response = serde_json::from_str::( @@ -422,19 +457,17 @@ mod tests { } #[test] - fn parameter_mapping_merges_supplied_options_and_ignores_null_overrides() { - let supplied = - serde_json::from_value(json!({"output_format":"blocks", "extension":true})).unwrap(); - let overrides = serde_json::from_value( - json!({"output_format":null,"req_format":"native","unknown":true}), + fn provider_options_exclude_response_controls_and_extensions() { + let arguments = serde_json::from_value( + json!({"output_format":"blocks","req_format":"native","unknown":true}), ) .unwrap(); let params = CohereParseConfig - .map_ocr_params(&overrides, &supplied, "parse") + .parse_options(&arguments, "parse") .unwrap(); assert_eq!( serde_json::to_value(¶ms).unwrap(), - json!({"output_format":"blocks","req_format":"native","extension":true}) + json!({"output_format":"blocks"}) ); let document = serde_json::from_value( json!({"type":"image_url","image_url":"https://example.com/a.png","ignored":"field"}), 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 84cc8eaed04..2d652bf275a 100644 --- a/litellm-rust/crates/core/src/llms/mistral/ocr/transformation.rs +++ b/litellm-rust/crates/core/src/llms/mistral/ocr/transformation.rs @@ -3,6 +3,7 @@ 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::{ @@ -79,16 +80,13 @@ impl BaseOcrConfig for MistralOCRConfig { fn map_ocr_params( &self, - non_default_params: &OpaqueParams, - _optional_params: &OpaqueParams, + non_default_params: &OcrArguments, + _optional_params: &OcrArguments, model: &str, - ) -> Result { - let supported = self.get_supported_ocr_params(model); + ) -> Result { Ok(non_default_params - .iter() - .filter(|(name, _)| supported.contains(&name.as_str())) - .map(|(name, value)| (name.clone(), value.clone())) - .collect()) + .select(self.get_supported_ocr_params(model)) + .into()) } async fn async_transform_ocr_request( @@ -117,11 +115,7 @@ impl MistralOCRConfig { request: &LiteLLMOcrRequest, client: &OcrClient, ) -> Result { - let params = self.map_ocr_params( - &request.optional_params, - &OpaqueParams::default(), - &request.model, - )?; + let params = self.parse_options(&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 @@ -285,7 +279,7 @@ mod tests { let input = serde_json::from_value(json!({"pages":null,"extract_header":false,"unknown":true})) .unwrap(); - let supplied = serde_json::from_value(json!({"pages":[1],"extract_footer":true})).unwrap(); + let supplied = serde_json::from_value(json!({"pages":[9],"extension":true})).unwrap(); let params = MistralOCRConfig .map_ocr_params(&input, &supplied, "model") .unwrap(); @@ -295,11 +289,49 @@ mod tests { ); } + #[test] + fn request_transform_uses_already_mapped_params_without_filtering_again() { + let params = serde_json::from_value(json!({"extension":{"nested":null}})).unwrap(); + let body = MistralOCRConfig + .transform_ocr_request("model", document(), ¶ms, &[]) + .unwrap(); + assert_eq!( + serde_json::to_value(body).unwrap()["extension"], + json!({"nested":null}) + ); + } + + #[test] + fn raw_response_transform_keeps_native_payload_separate_from_typed_normalization() { + let raw = br#"{"pages":[{"index":"2","markdown":"text"}],"provider_extension":false}"#; + let response = MistralOCRConfig + .decode_and_normalize_response( + "model", + raw, + crate::ocr::types::OcrResponseFormat::Native, + ) + .unwrap(); + assert_eq!(response.pages[0].index, 2); + let native = response.provider_native_response.unwrap(); + assert_eq!(native["pages"][0]["index"], "2"); + assert_eq!(native["provider_extension"], false); + assert!(response.extra_fields.is_empty()); + assert!( + MistralOCRConfig + .decode_and_normalize_response( + "model", + br#"{"pages":[{"index":0}]}"#, + crate::ocr::types::OcrResponseFormat::Litellm + ) + .is_err() + ); + } + fn mapped_params(value: Value) -> Value { - let params = serde_json::from_value::(value).unwrap(); + let params = serde_json::from_value::(value).unwrap(); serde_json::to_value( MistralOCRConfig - .map_ocr_params(¶ms, &OpaqueParams::default(), "model") + .map_ocr_params(¶ms, &OcrArguments::default(), "model") .unwrap(), ) .unwrap() 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 9c48681deed..a0bb3648a69 100644 --- a/litellm-rust/crates/core/src/llms/reducto/ocr/transformation.rs +++ b/litellm-rust/crates/core/src/llms/reducto/ocr/transformation.rs @@ -5,11 +5,10 @@ use serde_json::{Map, Value, json}; 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, merge_extra_params, -}; +use crate::ocr::prepare::{build_http_request, credential_env, guardrail_document}; use crate::ocr::types::{ LiteLLMOcrRequest, LiteLLMOcrResponse, OcrConnection, OcrDocument, OcrPage, OcrUsageInfo, }; @@ -90,15 +89,15 @@ impl BaseOcrConfig for ReductoParseV3Config { fn map_ocr_params( &self, - non_default_params: &OpaqueParams, - optional_params: &OpaqueParams, + non_default_params: &OcrArguments, + optional_params: &OcrArguments, model: &str, - ) -> Result { - map_ocr_params( + ) -> Result { + Ok(map_ocr_params( non_default_params, optional_params, self.get_supported_ocr_params(model), - ) + )) } #[tracing::instrument( @@ -137,11 +136,7 @@ impl ReductoParseV3Config { request: &LiteLLMOcrRequest, client: &OcrClient, ) -> Result { - let params = self.map_ocr_params( - &request.optional_params, - &OpaqueParams::default(), - &request.model, - )?; + let params = self.parse_options(&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?; @@ -157,10 +152,9 @@ impl ReductoParseV3Config { }, ) .await?; - let extra_params = request + let body = request .optional_params - .without(self.get_supported_ocr_params(&request.model)); - let body = merge_extra_params(&body, extra_params)?; + .compose_body(&body, self.get_supported_ocr_params(&request.model))?; build_http_request(client, request, &url, &headers, &body) } } @@ -179,15 +173,15 @@ impl BaseOcrConfig for ReductoParseLegacyConfig { fn map_ocr_params( &self, - non_default_params: &OpaqueParams, - optional_params: &OpaqueParams, + non_default_params: &OcrArguments, + optional_params: &OcrArguments, model: &str, - ) -> Result { - map_ocr_params( + ) -> Result { + Ok(map_ocr_params( non_default_params, optional_params, self.get_supported_ocr_params(model), - ) + )) } #[tracing::instrument( @@ -223,11 +217,7 @@ impl ReductoParseLegacyConfig { request: &LiteLLMOcrRequest, client: &OcrClient, ) -> Result { - let params = self.map_ocr_params( - &request.optional_params, - &OpaqueParams::default(), - &request.model, - )?; + let params = self.parse_options(&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?; @@ -243,20 +233,19 @@ impl ReductoParseLegacyConfig { }, ) .await?; - let extra_params = request + let body = request .optional_params - .without(self.get_supported_ocr_params(&request.model)); - let body = merge_extra_params(&body, extra_params)?; + .compose_body(&body, self.get_supported_ocr_params(&request.model))?; build_http_request(client, request, &url, &headers, &body) } } -fn map_ocr_params( - non_default_params: &OpaqueParams, - optional_params: &OpaqueParams, +fn map_ocr_params( + non_default_params: &OcrArguments, + optional_params: &OcrArguments, supported_params: &[&str], -) -> Result { - let params = optional_params +) -> OcrArguments { + optional_params .iter() .chain( non_default_params @@ -264,8 +253,7 @@ fn map_ocr_params( .filter(|(name, _)| supported_params.contains(&name.as_str())), ) .map(|(name, value)| (name.clone(), value.clone())) - .collect(); - crate::ocr::wire::decode_request_value(Value::Object(params), "optional_params") + .collect() } fn present_nullable<'de, D: Deserializer<'de>, T: Deserialize<'de>>( @@ -510,6 +498,36 @@ async fn upload_bytes_async( 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(); + let overrides = serde_json::from_value(json!({ + "formatting":null, "enhance":null, "ignored":true + })) + .unwrap(); + let v3 = ReductoParseV3Config + .map_ocr_params(&overrides, &supplied, "parse-v3") + .unwrap(); + assert_eq!( + serde_json::to_value(v3).unwrap(), + json!({ + "formatting":null, "enhance":true, "extension":false + }) + ); + let legacy = ReductoParseLegacyConfig + .map_ocr_params(&overrides, &supplied, "parse-legacy") + .unwrap(); + assert_eq!( + serde_json::to_value(legacy).unwrap(), + json!({ + "formatting":{"old":true}, "enhance":null, "extension":false + }) + ); + } + #[test] fn usage_uses_shared_validation_while_block_page_numbers_are_best_effort() { for usage in [ @@ -535,16 +553,12 @@ mod tests { } #[tokio::test] - async fn v3_mapping_preserves_explicit_null_and_supplied_options() { + async fn v3_options_preserve_explicit_null() { let overrides = serde_json::from_value(json!({"formatting":null,"settings":{},"unknown":true})) .unwrap(); - let supplied = serde_json::from_value( - json!({"formatting":{"old":true},"retrieval":{"mode":"page"},"extension":false}), - ) - .unwrap(); let params = ReductoParseV3Config - .map_ocr_params(&overrides, &supplied, "parse-v3") + .parse_options(&overrides, "parse-v3") .unwrap(); let client = crate::ocr::test_support::ocr_client(); let connection = OcrConnection::default(); @@ -568,15 +582,11 @@ mod tests { assert_eq!( serde_json::to_value(body).unwrap(), json!({ - "input":"reducto://ready.pdf", "formatting":null, "settings":{}, "retrieval":{"mode":"page"}, "extension":false + "input":"reducto://ready.pdf", "formatting":null, "settings":{} }) ); let absent = ReductoParseV3Config - .map_ocr_params( - &OpaqueParams::default(), - &OpaqueParams::default(), - "parse-v3", - ) + .parse_options(&OcrArguments::default(), "parse-v3") .unwrap(); assert_eq!(serde_json::to_value(absent).unwrap(), json!({})); } @@ -592,9 +602,8 @@ mod tests { ] { let overrides = serde_json::from_value(json!({"enhance":value,"unknown":true})).unwrap(); - let supplied = serde_json::from_value(json!({"enhance":{"old":true}})).unwrap(); let params = ReductoParseLegacyConfig - .map_ocr_params(&overrides, &supplied, "parse-legacy") + .parse_options(&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 bcfe722e434..7dfa7acb114 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 @@ -1,6 +1,8 @@ use serde::{Deserialize, Serialize}; use serde_json::{Map, Value}; +use litellm_auth_gcp::{self as vertex, VertexConfig}; + use super::transformation::VertexAIOCRConfig; use crate::llms::base_llm::ocr::transformation::{BaseOcrConfig, OcrRequestContext}; use crate::ocr::OcrClient; @@ -12,7 +14,6 @@ use crate::ocr::types::{ use crate::params::OpaqueParams; use crate::routing_utils::model::{ModelNamespace, ProviderModel, RoutedModel}; use crate::url_utils::ApiUrl; -use litellm_auth_gcp::{self as vertex, VertexConfig}; const DEFAULT_API_BASE: &str = "https://aiplatform.googleapis.com"; const MODEL_NAMESPACE: &str = "deepseek-ai"; @@ -153,11 +154,7 @@ impl VertexAIDeepSeekOCRConfig { request: &LiteLLMOcrRequest, client: &OcrClient, ) -> Result { - let params = self.map_ocr_params( - &OpaqueParams::default(), - &request.optional_params, - &request.model, - )?; + let params = self.parse_options(&request.optional_params, &request.model)?; let config = VertexConfig::from_sourced_optional_params( &request.optional_params, &request.input_sources, @@ -398,6 +395,21 @@ impl VertexAIDeepSeekOCRConfig { mod tests { use super::{VertexAIDeepSeekOCRConfig, provider_model}; + #[test] + fn inherited_parameter_mapping_only_returns_supplied_optional_params() { + 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(); + assert_eq!( + VertexAIDeepSeekOCRConfig + .map_ocr_params(&non_default, &supplied, "deepseek-ocr") + .unwrap(), + supplied + ); + } + #[test] fn config_owns_model_namespace_and_endpoint() { assert_eq!( 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 b33601d5b45..87efdb02bd2 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,6 +2,7 @@ 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}; @@ -25,10 +26,10 @@ impl BaseOcrConfig for VertexAIOCRConfig { fn map_ocr_params( &self, - non_default_params: &OpaqueParams, - optional_params: &OpaqueParams, + non_default_params: &OcrArguments, + optional_params: &OcrArguments, model: &str, - ) -> Result { + ) -> Result { MistralOCRConfig.map_ocr_params(non_default_params, optional_params, model) } @@ -64,11 +65,7 @@ impl VertexAIOCRConfig { request: &LiteLLMOcrRequest, client: &OcrClient, ) -> Result { - let params = self.map_ocr_params( - &request.optional_params, - &OpaqueParams::default(), - &request.model, - )?; + let params = self.parse_options(&request.optional_params, &request.model)?; let config = VertexConfig::from_sourced_optional_params( &request.optional_params, &request.input_sources, diff --git a/litellm-rust/crates/core/src/ocr/arguments.rs b/litellm-rust/crates/core/src/ocr/arguments.rs new file mode 100644 index 00000000000..e8ad6bd3cd3 --- /dev/null +++ b/litellm-rust/crates/core/src/ocr/arguments.rs @@ -0,0 +1,93 @@ +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/mod.rs b/litellm-rust/crates/core/src/ocr/mod.rs index 89cb2165b60..b1e4a52dc62 100644 --- a/litellm-rust/crates/core/src/ocr/mod.rs +++ b/litellm-rust/crates/core/src/ocr/mod.rs @@ -1,4 +1,6 @@ +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 f49978518fe..d15d875bb77 100644 --- a/litellm-rust/crates/core/src/ocr/prepare.rs +++ b/litellm-rust/crates/core/src/ocr/prepare.rs @@ -5,8 +5,6 @@ use super::OcrClient; use super::hooks::OcrDuringCallRequest; use super::types::{LiteLLMOcrRequest, OcrDocument}; -pub(crate) use crate::params::merge_extra_params; - pub(crate) async fn transform_request_body( client: &OcrClient, request: &LiteLLMOcrRequest, @@ -19,10 +17,10 @@ pub(crate) async fn transform_request_body( where B: Serialize + DeserializeOwned, { - let extras = request - .optional_params - .without(request.config.get_supported_ocr_params(&request.model)); - let composed = merge_extra_params(&body, extras)?; + let composed = request.optional_params.compose_body( + &body, + request.config.get_supported_ocr_params(&request.model), + )?; let composed = OcrWireBody::::decode(composed, "body")?; validate(&composed.body)?; let (body, headers) = if request.hooks.intercepts_requests() { @@ -110,10 +108,10 @@ pub(crate) async fn guardrail_document( #[derive(Serialize)] struct OcrWireBody { - #[serde(flatten)] + #[serde(skip)] body: B, #[serde(flatten)] - extra: crate::params::OpaqueParams, + fields: serde_json::Map, } impl OcrWireBody { @@ -124,14 +122,7 @@ impl OcrWireBody { path: prefix.into(), }); }; - let known = serde_json::to_value(&body).map_err(|_| super::Error::RequestField { - path: prefix.into(), - })?; - let extra = fields - .into_iter() - .filter(|(key, _)| known.get(key).is_none()) - .collect(); - Ok(Self { body, extra }) + Ok(Self { body, fields }) } } diff --git a/litellm-rust/crates/core/src/ocr/types.rs b/litellm-rust/crates/core/src/ocr/types.rs index d7a4bd53f0f..b97ff9394f0 100644 --- a/litellm-rust/crates/core/src/ocr/types.rs +++ b/litellm-rust/crates/core/src/ocr/types.rs @@ -7,10 +7,10 @@ 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::constants::OCR_HTTP_TIMEOUT_SECS; -use crate::params::OpaqueParams; use litellm_auth::{InputSource, TokenProviderHandle}; #[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] @@ -97,7 +97,7 @@ pub struct LiteLLMOcrRequest { pub connection: OcrConnection, pub hooks: Arc, pub litellm_call_id: Option, - pub optional_params: OpaqueParams, + pub optional_params: OcrArguments, 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: OpaqueParams, + optional_params: OcrArguments, ) -> 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 e11cf6eeab2..cf22975b5a0 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::params::OpaqueParams; use litellm_auth::InputSource; use serde::{ Deserialize, @@ -54,7 +54,7 @@ pub struct OcrWireRequest { pub custom_llm_provider: Option, pub extra_headers: Option>, #[serde(default)] - pub optional_params: OpaqueParams, + pub optional_params: OcrArguments, #[serde(default)] pub input_sources: BTreeMap, pub timeout_seconds: Option, @@ -70,9 +70,9 @@ pub fn consumed_optional_param_names( ) -> Result, crate::ocr::Error> { use super::provider_config::OcrConfigKind; - let (provider_model, config) = + let (model, config) = super::provider_config::resolve_provider_config(model, custom_llm_provider)?; - let provider_fields = config.get_supported_ocr_params(&provider_model); + let provider_fields = config.get_supported_ocr_params(&model); let auth_fields: &[&str] = match config { OcrConfigKind::AzureAi | OcrConfigKind::AzureDocumentIntelligence @@ -110,6 +110,17 @@ 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"); @@ -255,6 +266,19 @@ pub fn decode_response( mod tests { use super::*; + #[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)); + } + #[test] fn option_projection_is_provider_specific_and_excludes_opaque_fields() { let mistral = consumed_optional_param_names("mistral/model", None).unwrap(); diff --git a/litellm-rust/crates/core/tests/reducto_ocr.rs b/litellm-rust/crates/core/tests/reducto_ocr.rs index f7c2e184397..d7bcf0b8d12 100644 --- a/litellm-rust/crates/core/tests/reducto_ocr.rs +++ b/litellm-rust/crates/core/tests/reducto_ocr.rs @@ -247,6 +247,52 @@ async fn facade_omits_native_response_by_default_and_preserves_auth_priority() { struct RewriteDocument; +struct RewriteHeaders; + +impl OcrHooks for RewriteHeaders { + fn intercepts_requests(&self) -> bool { + true + } + + fn during_call( + &self, + request: OcrDuringCallRequest, + ) -> OcrHookFuture<'_, OcrDuringCallRequest> { + Box::pin(async move { + Ok(OcrDuringCallRequest { + headers: vec![("authorization".into(), "Bearer guarded".into())], + ..request + }) + }) + } +} + +#[rstest] +#[case("reducto/parse-v3")] +#[case("reducto/parse-legacy")] +#[tokio::test] +async fn guardrail_headers_reach_upload_and_parse(#[case] model: &str) { + let (base, seen, server) = mock_server(vec![ + MockResponse::json(json!({"file_id":"reducto://uploaded.pdf"})), + MockResponse::json(json!({"result":{"chunks":[]}})), + ]) + .await; + let mut request = wire_request(model, &base, json!({})); + request.connection.extra_headers = vec![("authorization".into(), "Bearer original".into())]; + request.hooks = Arc::new(RewriteHeaders); + + perform_ocr(request).await.unwrap(); + server.await.unwrap(); + let requests = seen.lock().unwrap(); + assert_eq!(requests.len(), 2); + assert!(requests[0].starts_with("POST /upload ")); + assert!(requests[1].starts_with("POST /parse ")); + for request in requests.iter() { + assert!(request.contains("authorization: Bearer guarded")); + assert!(!request.contains("Bearer original")); + } +} + impl OcrHooks for RewriteDocument { fn intercepts_requests(&self) -> bool { true diff --git a/litellm-rust/crates/python-bridge/src/marshal.rs b/litellm-rust/crates/python-bridge/src/marshal.rs index 2639b0ff915..356e5b6d2dd 100644 --- a/litellm-rust/crates/python-bridge/src/marshal.rs +++ b/litellm-rust/crates/python-bridge/src/marshal.rs @@ -90,7 +90,7 @@ pub(crate) fn python_timeout_seconds(py: Python<'_>, timeout: Py) -> PyRe pub(crate) fn project_optional_fields( kwargs: &Bound<'_, PyDict>, - names: &[&str], + fields: &[litellm_core::ocr::wire::OptionalParamSpec], ) -> PyResult> { let controls: Vec = kwargs .py() @@ -102,13 +102,7 @@ pub(crate) fn project_optional_fields( .map(|(name, value)| Ok((name.extract::()?, value))) .filter_map(|entry: PyResult<_>| match entry { Ok((name, value)) - if names.contains(&name.as_str()) - || (!controls.contains(&name) - && !litellm_core::params::is_control_param(&name) - && !matches!( - name.as_str(), - "model" | "document" | "timeout" | "input_sources" - )) => + if litellm_core::ocr::wire::project_argument(&name, fields, &controls) => { 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 20ee2627060..006f1975fed 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/project.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/project.rs @@ -124,8 +124,7 @@ 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 names = specs.iter().map(|spec| spec.name).collect::>(); - let optional_params = project_optional_fields(kwargs, &names)?; + let optional_params = project_optional_fields(kwargs, &specs)?; let input_sources = request_input_sources( kwargs, optional_params diff --git a/tests/test_litellm_rust/ocr/test_cohere.py b/tests/test_litellm_rust/ocr/test_cohere.py index dd5e3e7d380..639857dd5c8 100644 --- a/tests/test_litellm_rust/ocr/test_cohere.py +++ b/tests/test_litellm_rust/ocr/test_cohere.py @@ -67,6 +67,40 @@ async def test_public_cohere_blocks_and_usage_fallback(recording_server: Recordi assert response.get_provider_native_response() is None +@pytest.mark.asyncio +@pytest.mark.parametrize("model", MODELS) +@pytest.mark.parametrize("asynchronous", [False, True]) +async def test_explicit_body_preserves_provider_native_values( + recording_server: RecordingServer, model: str, asynchronous: bool +) -> None: + recording_server.enqueue(ResponseSpec(body=PAYLOAD)) + document: Final = {**IMAGE, "provider_options": {"future": [None, False, 0]}} + arguments: Final = { + "model": model, + "document": IMAGE, + "api_base": recording_server.base_url, + "api_key": "test-key", + "output_format": "markdown", + "future_option": {"original": True}, + "extra_body": { + "output_format": "future-format", + "document": document, + "future_option": None, + }, + } + if asynchronous: + await litellm.aocr(**arguments) + else: + litellm.ocr(**arguments) + assert not recording_server.requests[0].headers.get("user-agent", "").startswith("python-httpx") + assert recording_server.requests[0].body == { + "model": model.split("/", 1)[1], + "document": document, + "output_format": "future-format", + "future_option": None, + } + + @pytest.mark.asyncio @pytest.mark.parametrize("model", MODELS) @pytest.mark.parametrize( diff --git a/tests/test_litellm_rust/ocr/test_lifecycle.py b/tests/test_litellm_rust/ocr/test_lifecycle.py index ff832dbead0..9875cb1027d 100644 --- a/tests/test_litellm_rust/ocr/test_lifecycle.py +++ b/tests/test_litellm_rust/ocr/test_lifecycle.py @@ -726,7 +726,7 @@ async def test_reducto_lifecycle_retains_upload_parse_and_post_call_boundaries( @pytest.mark.asyncio -async def test_document_intelligence_post_call_observes_submission_and_final_result( +async def test_document_intelligence_post_call_observes_submission_before_polling( ocr_server: RecordingServer, ) -> None: ocr_server.expected_requests = 2 @@ -757,9 +757,8 @@ async def test_document_intelligence_post_call_observes_submission_and_final_res response: Final = await call_aocr( ocr_server, model="azure_ai/doc-intelligence/prebuilt-read", litellm_logging_obj=logger ) - assert [methods for methods, _ in boundaries] == [("POST",), ("POST", "GET")] + assert [methods for methods, _ in boundaries] == [("POST",)] assert json.loads(boundaries[0][1])["status"] == "running" - assert json.loads(boundaries[1][1])["status"] == "succeeded" assert [request.method for request in ocr_server.requests] == ["POST", "GET"] assert ocr_server.requests[1].path == "/operations/1" assert response.pages == []