diff --git a/litellm-rust/crates/core/src/audio_transcription/tests.rs b/litellm-rust/crates/core/src/audio_transcription/tests.rs index 6693ba2a3cf..8fe185a2c72 100644 --- a/litellm-rust/crates/core/src/audio_transcription/tests.rs +++ b/litellm-rust/crates/core/src/audio_transcription/tests.rs @@ -22,6 +22,8 @@ async fn bedrock_request_is_signed_and_contains_audio() { assert!(request.contains("authorization: AWS4-HMAC-SHA256")); assert!(request.contains("x-amz-date:")); assert!(request.contains("\"bytes\":\"AQI=\"")); + assert!(request.contains("\"future_option\":{\"nested\":[null,false,0]}")); + assert!(!request.contains("\"aws_secret_access_key\"")); assert!(request.contains("Transcribe the audio. Respond with only the transcript.")); let response = b"HTTP/1.1 200 OK\r\nContent-Type: application/json\r\nContent-Length: 53\r\nConnection: close\r\n\r\n{\"output\":{\"message\":{\"content\":[{\"text\":\"hello\"}]}}}"; stream.write_all(response).expect("response"); @@ -31,6 +33,10 @@ async fn bedrock_request_is_signed_and_contains_audio() { ("aws_access_key_id".to_string(), json!("access-key")), ("aws_secret_access_key".to_string(), json!("secret-key")), ("aws_region_name".to_string(), json!("us-east-1")), + ( + "future_option".to_string(), + json!({"nested":[null,false,0]}), + ), ]); let api_base = format!("http://{address}"); let response = audio_transcription(AudioTranscriptionRequest { diff --git a/litellm-rust/crates/core/src/audio_transcription/transformation.rs b/litellm-rust/crates/core/src/audio_transcription/transformation.rs index b15064880fd..6fe13445acc 100644 --- a/litellm-rust/crates/core/src/audio_transcription/transformation.rs +++ b/litellm-rust/crates/core/src/audio_transcription/transformation.rs @@ -18,7 +18,7 @@ pub trait AudioTranscriptionProviderConfig: Sync { #[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)] fn map_transcription_params(&self, params: &OpaqueParams) -> OpaqueParams { - params.retain_supported(self.supported_transcription_params()) + params.provider_params() } fn transform_transcription_request( diff --git a/litellm-rust/crates/core/src/chat_completions/prepare.rs b/litellm-rust/crates/core/src/chat_completions/prepare.rs index e965a0fe5d7..bd0ff5f4571 100644 --- a/litellm-rust/crates/core/src/chat_completions/prepare.rs +++ b/litellm-rust/crates/core/src/chat_completions/prepare.rs @@ -47,6 +47,7 @@ pub(super) fn resolve_request( "chat completions requires at least one message".to_string(), )); } + request.optional_params.clone().into_provider_body()?; if let Some(reason) = config.unsupported_reason(&messages, &request.optional_params) { return Err(Error::Unsupported(reason.0)); } diff --git a/litellm-rust/crates/core/src/chat_completions/transformation.rs b/litellm-rust/crates/core/src/chat_completions/transformation.rs index eb1a54886d4..b6e9766d7a5 100644 --- a/litellm-rust/crates/core/src/chat_completions/transformation.rs +++ b/litellm-rust/crates/core/src/chat_completions/transformation.rs @@ -18,11 +18,6 @@ pub enum ChatCompletionsAuth { /// Why a request cannot be served by the Rust path. /// -/// The core declines rather than guessing: the host turns this into a -/// transparent fallback to the Python implementation, which covers the full -/// surface. Acceptance is an allowlist, so a parameter or message shape the -/// core has never seen declines by construction instead of being translated -/// wrong. #[derive(Clone, Copy, Debug, PartialEq, Eq)] pub struct Unsupported(pub &'static str); @@ -77,12 +72,7 @@ pub trait ChatCompletionsProviderConfig: Sync { messages: &[ChatMessage], optional_params: &OpaqueParams, ) -> Option { - unsupported_param( - self.supported_openai_params(), - self.config_params(), - optional_params, - ) - .or_else(|| messages.iter().find_map(unsupported_message)) + unsupported_param(optional_params).or_else(|| messages.iter().find_map(unsupported_message)) } fn transform_request( @@ -99,26 +89,31 @@ pub trait ChatCompletionsProviderConfig: Sync { ) -> Result; } -pub fn unsupported_param( - supported: &'static [(&'static str, &'static str)], - config: &'static [&'static str], - optional_params: &OpaqueParams, -) -> Option { - if optional_params +pub fn unsupported_param(optional_params: &OpaqueParams) -> Option { + let effective = match optional_params.clone().into_provider_body() { + Ok(fields) => fields, + Err(_) => return None, + }; + if effective .get(STREAM_PARAM) .and_then(Value::as_bool) .unwrap_or(false) { return Some(Unsupported("streaming")); } - optional_params + effective .keys() .any(|key| { - key != STREAM_PARAM - && !supported - .iter() - .any(|(_, provider_name)| *provider_name == key) - && !config.contains(&key.as_str()) + matches!( + key.as_str(), + "tools" + | "tool_choice" + | "toolConfig" + | "thinking" + | "top_k" + | "topK" + | "_parallel_tool_use_config" + ) }) .then_some(Unsupported("unrecognized request parameter")) } diff --git a/litellm-rust/crates/core/src/chat_completions/types.rs b/litellm-rust/crates/core/src/chat_completions/types.rs index 93afcc9e524..8c0586b55df 100644 --- a/litellm-rust/crates/core/src/chat_completions/types.rs +++ b/litellm-rust/crates/core/src/chat_completions/types.rs @@ -73,7 +73,7 @@ pub struct ChatMessage { #[serde(default, skip_serializing_if = "Option::is_none")] pub name: Option, #[serde(flatten)] - pub extra: Map, + pub extra: crate::params::OpaqueParams, } /// OpenAI `usage`, including the `prompt_tokens_details` split LiteLLM's Python 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 ec65bac3774..e708c81ef9c 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 @@ -30,11 +30,16 @@ mod types { } #[derive(Clone, Debug, Serialize, Deserialize)] + #[serde(untagged)] pub(crate) enum DocumentIntelligenceRequest { - #[serde(rename = "urlSource")] - UrlSource(String), - #[serde(rename = "base64Source")] - Base64Source(String), + UrlSource { + #[serde(rename = "urlSource")] + url_source: String, + }, + Base64Source { + #[serde(rename = "base64Source")] + base64_source: String, + }, } #[derive(Clone, Debug, PartialEq)] @@ -385,11 +390,14 @@ mod mapping { return Err(OcrRequestError::MissingDocumentUrl); } Ok(if let Some(document) = InlineDocument::parse(source)? { - DocumentIntelligenceRequest::Base64Source( - STANDARD.encode(document.decode(crate::constants::OCR_INLINE_MAX_BYTES)?), - ) + DocumentIntelligenceRequest::Base64Source { + base64_source: STANDARD + .encode(document.decode(crate::constants::OCR_INLINE_MAX_BYTES)?), + } } else { - DocumentIntelligenceRequest::UrlSource(source.to_string()) + DocumentIntelligenceRequest::UrlSource { + url_source: source.to_string(), + } }) } 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 3f69bb539c1..c02d5863543 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 @@ -13,8 +13,10 @@ pub(crate) trait BaseOcrConfig: Send + Sync + Sized + 'static { fn get_supported_ocr_params(&self, model: &str) -> &'static [&'static str]; - fn map_ocr_params(&self, model: &str, params: &OpaqueParams) -> OpaqueParams { - params.retain_supported(self.get_supported_ocr_params(model)) + fn map_ocr_params(&self, _model: &str, params: &OpaqueParams) -> OpaqueParams { + params + .provider_params() + .without(&["extra_body", "model", "document"]) } fn prepare_request( 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 37442885149..46f002b92be 100644 --- a/litellm-rust/crates/core/src/llms/mistral/ocr/transformation.rs +++ b/litellm-rust/crates/core/src/llms/mistral/ocr/transformation.rs @@ -127,10 +127,10 @@ mod mapping_tests { } #[rstest] - fn map_ocr_params_drops_unknown_params() { + fn map_ocr_params_preserves_unknown_params() { let mapped = mapped_params(json!({"extract_header":true,"unsupported_param":"value"})); assert_eq!(mapped["extract_header"], true); - assert!(mapped.get("unsupported_param").is_none()); + assert_eq!(mapped["unsupported_param"], "value"); } #[rstest] 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 5f7747168b5..348fbb9d178 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 @@ -37,6 +37,8 @@ mod types { pub(crate) struct DeepSeekOcrMessage { pub role: UserRole, pub content: Vec, + #[serde(default, flatten)] + pub extra: crate::params::OpaqueParams, } #[derive(Clone, Debug, Serialize, Deserialize)] @@ -125,6 +127,7 @@ mod mapping { messages: vec![DeepSeekOcrMessage { role: UserRole::User, content: vec![content], + extra: Default::default(), }], params: params.clone(), }) diff --git a/litellm-rust/crates/core/src/messages/prepare.rs b/litellm-rust/crates/core/src/messages/prepare.rs index 3c2f6b9d372..5895ae5ea5d 100644 --- a/litellm-rust/crates/core/src/messages/prepare.rs +++ b/litellm-rust/crates/core/src/messages/prepare.rs @@ -33,7 +33,11 @@ pub(super) fn prepare_provider_request( let headers = validate_environment(config, request.extra_headers, request.api_key, &env_lookup)?; - let typed_request = serde_json::from_value(request.body).map_err(|err| { + let params: crate::params::OpaqueParams = serde_json::from_value(request.body) + .map_err(|_| Error::InvalidRequest("messages body must be an object".into()))?; + let mut fields = params.into_provider_body()?; + fields.insert("model".into(), Value::String(model.clone())); + let typed_request = serde_json::from_value(Value::Object(fields)).map_err(|err| { Error::InvalidRequest(format!("invalid Anthropic messages request: {err}")) })?; let transformed = config.transform_request(typed_request)?; diff --git a/litellm-rust/crates/core/src/messages/tests.rs b/litellm-rust/crates/core/src/messages/tests.rs index df9f7051011..af1047e1df3 100644 --- a/litellm-rust/crates/core/src/messages/tests.rs +++ b/litellm-rust/crates/core/src/messages/tests.rs @@ -199,7 +199,9 @@ async fn messages_round_trip_builds_native_anthropic_request() { body: json!({ "model": "claude-sonnet-4-5", "max_tokens": 1024, - "messages": [{"role": "user", "content": "hi"}] + "messages": [{"role": "user", "content": "hi", "future_message": null}], + "future_option": {"nested": [null, false, 0]}, + "extra_body": {"max_tokens": 2048, "model": "wrong"} }), api_key: Some("sk-ant"), api_base: Some(&format!("http://{addr}")), @@ -214,7 +216,16 @@ async fn messages_round_trip_builds_native_anthropic_request() { assert_eq!(response.stop_reason.as_deref(), Some("end_turn")); let request = server.await.expect("server task completes"); - let (head, _) = request.split_once("\r\n\r\n").expect("has body"); + let (head, body) = request.split_once("\r\n\r\n").expect("has body"); + let body: Value = serde_json::from_str(body).unwrap(); + assert_eq!(body["future_option"], json!({"nested":[null,false,0]})); + assert_eq!( + body["messages"][0].get("future_message"), + Some(&Value::Null) + ); + assert_eq!(body["max_tokens"], 2048); + assert_eq!(body["model"], "claude-sonnet-4-5"); + assert!(body.get("extra_body").is_none()); assert!(head.starts_with("POST /v1/messages "), "{head}"); let head_lower = head.to_ascii_lowercase(); assert!(head_lower.contains("x-api-key: sk-ant"), "{head}"); diff --git a/litellm-rust/crates/core/src/messages/types.rs b/litellm-rust/crates/core/src/messages/types.rs index b9f807c29fd..a2a2cdb669c 100644 --- a/litellm-rust/crates/core/src/messages/types.rs +++ b/litellm-rust/crates/core/src/messages/types.rs @@ -44,7 +44,7 @@ pub struct ContentBlock { #[serde(skip_serializing_if = "Option::is_none")] pub cache_control: Option, #[serde(flatten)] - pub extra: Map, + pub extra: crate::params::OpaqueParams, } #[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)] @@ -56,7 +56,7 @@ pub struct CacheControl { #[serde(skip_serializing_if = "Option::is_none")] pub scope: Option, #[serde(flatten)] - pub extra: Map, + pub extra: crate::params::OpaqueParams, } #[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] @@ -64,7 +64,7 @@ pub struct AnthropicMessage { pub role: String, pub content: MessageContent, #[serde(flatten)] - pub extra: Map, + pub extra: crate::params::OpaqueParams, } #[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] @@ -110,7 +110,7 @@ pub struct AnthropicMessagesRequest { #[serde(skip_serializing_if = "Option::is_none")] pub inference_geo: Option, #[serde(flatten)] - pub extra: Map, + pub extra: crate::params::OpaqueParams, } #[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] @@ -130,5 +130,5 @@ pub struct AnthropicMessagesResponse { #[serde(skip_serializing_if = "Option::is_none")] pub container: Option, #[serde(flatten)] - pub extra: Map, + pub extra: crate::params::OpaqueParams, } diff --git a/litellm-rust/crates/core/src/ocr/prepare.rs b/litellm-rust/crates/core/src/ocr/prepare.rs index 2ba1e460652..524f8f258a1 100644 --- a/litellm-rust/crates/core/src/ocr/prepare.rs +++ b/litellm-rust/crates/core/src/ocr/prepare.rs @@ -1,61 +1,22 @@ -use serde::{Deserialize, Serialize, de::DeserializeOwned}; -use serde_json::{Map, Value}; +use serde::{Serialize, de::DeserializeOwned}; +use serde_json::Value; use super::OcrClient; use super::error::{OcrError, OcrRequestError}; use super::hooks::OcrDuringCallRequest; use super::types::{LiteLLMOcrRequest, OcrDocument}; -#[derive(Debug, Deserialize)] -pub(crate) struct ParsedProviderParams { - #[serde(flatten)] - pub known: T, - #[serde(default, flatten)] - pub extra_params: Map, -} +pub(crate) use crate::params::{ParsedProviderParams, merge_extra_params}; pub(crate) fn _prepare_ocr_request( request: &LiteLLMOcrRequest, ) -> Result, OcrRequestError> { super::wire::decode_request_value( - Value::Object(request.optional_params.clone().into()), + Value::Object(request.optional_params.provider_params().into()), "optional_params", ) } -pub(crate) fn merge_extra_params( - body: &B, - extra_params: Map, -) -> Result { - let Value::Object(fields) = - serde_json::to_value(body).map_err(|_| OcrRequestError::RequestField { - path: "body".into(), - })? - else { - return Err(OcrRequestError::RequestField { - path: "body".into(), - }); - }; - let extra_body = extra_params - .get("extra_body") - .and_then(Value::as_object) - .cloned() - .unwrap_or_default() - .into_iter() - .collect::>(); - Ok(Value::Object( - fields - .into_iter() - .chain( - extra_params - .into_iter() - .filter(|(name, _)| name != "extra_body"), - ) - .chain(extra_body) - .collect(), - )) -} - pub(crate) async fn transform_request_body( client: &OcrClient, request: &LiteLLMOcrRequest, @@ -63,13 +24,19 @@ pub(crate) async fn transform_request_body( headers: &[(String, String)], retains_document: bool, body: B, - validate: impl FnOnce(&B) -> Result<(), OcrRequestError>, + validate: impl Fn(&B) -> Result<(), OcrRequestError>, ) -> Result 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 = OcrWireBody::::decode(composed, "body")?; + validate(&composed.body)?; let (body, headers) = if request.hooks.intercepts_requests() { - let body = serde_json::to_value(body).map_err(|_| OcrRequestError::RequestField { + let body = serde_json::to_value(composed).map_err(|_| OcrRequestError::RequestField { path: "body".into(), })?; let retained_fields = request @@ -78,6 +45,13 @@ where .filter(|name| body.get(*name).is_some()) .cloned() .chain(retains_document.then(|| "document".to_string())) + .filter(|name| { + request + .optional_params + .get("extra_body") + .and_then(Value::as_object) + .is_none_or(|overrides| !overrides.contains_key(name)) + }) .collect(); let changed = request .hooks @@ -90,17 +64,11 @@ where retained_fields, }) .await?; - let body = OcrWireBody::::decode(changed.body)?; + let body = OcrWireBody::::decode(changed.body, "guardrail.body")?; validate(&body.body)?; (body, changed.headers) } else { - ( - OcrWireBody { - body, - extra: Map::new(), - }, - headers.to_vec(), - ) + (composed, headers.to_vec()) }; build_http_request(client, request, url, &headers, &body) } @@ -155,19 +123,19 @@ struct OcrWireBody { #[serde(flatten)] body: B, #[serde(flatten)] - extra: Map, + extra: crate::params::OpaqueParams, } impl OcrWireBody { - fn decode(value: Value) -> Result { - let body: B = super::wire::decode_request_value(value.clone(), "guardrail.body")?; + 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(OcrRequestError::RequestField { - path: "guardrail.body".into(), + path: prefix.into(), }); }; let known = serde_json::to_value(&body).map_err(|_| OcrRequestError::RequestField { - path: "guardrail.body".into(), + path: prefix.into(), })?; let extra = fields .into_iter() @@ -182,6 +150,7 @@ pub(crate) fn credential_env(name: &str) -> Option { } #[cfg(test)] mod tests { + use serde::Deserialize; use serde_json::json; use super::*; diff --git a/litellm-rust/crates/core/src/params.rs b/litellm-rust/crates/core/src/params.rs index d5d89a78eb8..4793abc154f 100644 --- a/litellm-rust/crates/core/src/params.rs +++ b/litellm-rust/crates/core/src/params.rs @@ -1,4 +1,4 @@ -use std::ops::Deref; +use std::ops::{Deref, DerefMut}; use serde::{Deserialize, Serialize}; use serde_json::{Map, Value}; @@ -7,19 +7,125 @@ use serde_json::{Map, Value}; #[serde(transparent)] pub struct OpaqueParams(Map); +#[derive(Debug, Deserialize)] +pub(crate) struct ParsedProviderParams { + #[serde(flatten)] + pub known: T, + #[serde(default, flatten)] + pub extra_params: OpaqueParams, +} + +pub fn is_control_param(name: &str) -> bool { + matches!( + name, + "api_key" + | "api_base" + | "custom_llm_provider" + | "extra_headers" + | "timeout" + | "timeout_seconds" + | "request_timeout" + | "max_retries" + | "req_format" + | "max_response_bytes" + | "litellm_call_id" + | "litellm_logging_obj" + | "litellm_metadata" + | "proxy_server_request" + | "callbacks" + | "success_callback" + | "failure_callback" + | "guardrails" + | "azure_ad_token" + | "azure_ad_token_provider" + | "tenant_id" + | "client_id" + | "client_secret" + | "azure_scope" + | "azure_authority_host" + | "azure_credential" + | "azure_federated_token_file" + | "enable_azure_ad_token_refresh" + | "vertex_credentials" + | "vertex_ai_credentials" + | "vertex_project" + | "vertex_ai_project" + | "vertex_location" + | "vertex_ai_location" + | "aws_access_key_id" + | "aws_secret_access_key" + | "aws_session_token" + | "aws_region_name" + | "aws_session_name" + | "aws_profile_name" + | "aws_role_name" + | "aws_web_identity_token" + | "aws_sts_endpoint" + | "aws_external_id" + | "aws_bedrock_runtime_endpoint" + ) +} + impl OpaqueParams { pub fn into_inner(self) -> Map { self.0 } - pub fn retain_supported(&self, supported: &[&str]) -> Self { - Self( - self.iter() - .filter(|(name, _)| supported.contains(&name.as_str())) - .map(|(name, value)| (name.clone(), value.clone())) - .collect(), - ) + pub fn without(&self, names: &[&str]) -> Self { + self.iter() + .filter(|(name, _)| !names.contains(&name.as_str())) + .map(|(name, value)| (name.clone(), value.clone())) + .collect() } + + pub fn provider_params(&self) -> Self { + self.iter() + .filter(|(name, _)| !is_control_param(name)) + .map(|(name, value)| (name.clone(), value.clone())) + .collect() + } + + pub fn into_provider_body(self) -> Result, crate::Error> { + let mut fields = self.0; + let overrides = match fields.remove("extra_body") { + None | Some(Value::Null) => Map::new(), + Some(Value::Object(fields)) => fields, + Some(_) => { + return Err(crate::Error::InvalidRequest( + "extra_body must be an object".into(), + )); + } + }; + Ok(fields + .into_iter() + .chain(overrides) + .filter(|(name, _)| name != "extra_body" && !is_control_param(name)) + .collect()) + } +} + +pub(crate) fn merge_extra_params( + body: &B, + extra_params: OpaqueParams, +) -> Result { + let Value::Object(fields) = serde_json::to_value(body) + .map_err(|_| crate::Error::InvalidRequest("body must be a JSON object".into()))? + else { + return Err(crate::Error::InvalidRequest( + "body must be a JSON object".into(), + )); + }; + Ok(Value::Object( + fields + .into_iter() + .chain( + extra_params + .into_provider_body()? + .into_iter() + .filter(|(name, _)| name != "model"), + ) + .collect(), + )) } impl Deref for OpaqueParams { @@ -30,6 +136,12 @@ impl Deref for OpaqueParams { } } +impl DerefMut for OpaqueParams { + fn deref_mut(&mut self) -> &mut Self::Target { + &mut self.0 + } +} + impl From> for OpaqueParams { fn from(value: Map) -> Self { Self(value) @@ -64,15 +176,55 @@ mod tests { use super::*; #[test] - fn supported_keys_preserve_opaque_values() { + fn extras_merge_shallowly_and_preserve_values_without_leaking_controls() { + let extras: OpaqueParams = serde_json::from_value(json!({ + "future": {"nested": [false, 0, null]}, + "explicit_null": null, + "azure_ad_token": "secret", + "req_format": "native", + "extra_body": { + "future": {"replacement": true}, + "temperature": 0.5, + "model": "override", + "aws_secret_access_key": "secret" + } + })) + .unwrap(); + let body = + merge_extra_params(&json!({"model":"resolved", "temperature":0.1}), extras).unwrap(); + assert_eq!( + body, + json!({ + "model":"resolved", "temperature":0.5, + "future":{"replacement":true}, "explicit_null":null + }) + ); + } + + #[test] + fn invalid_extra_body_is_rejected_and_null_is_empty() { + for value in [json!(false), json!([]), json!("value"), json!(1)] { + let params: OpaqueParams = serde_json::from_value(json!({"extra_body":value})).unwrap(); + assert!(params.into_provider_body().is_err()); + } + let params: OpaqueParams = + serde_json::from_value(json!({"extra_body":null,"future":null})).unwrap(); + assert_eq!( + Value::Object(params.into_provider_body().unwrap()), + json!({"future":null}) + ); + } + + #[test] + fn provider_params_preserve_opaque_values() { let params: OpaqueParams = serde_json::from_value(json!({ "object": {"future": [1, null]}, "null": null, - "unsupported": true + "azure_ad_token": "secret" })) .unwrap(); - let retained = params.retain_supported(&["object", "null"]); + let retained = params.provider_params(); assert_eq!( serde_json::to_value(retained).unwrap(), diff --git a/litellm-rust/crates/core/src/providers/anthropic/chat_completions/tests.rs b/litellm-rust/crates/core/src/providers/anthropic/chat_completions/tests.rs index 6ac704a3567..83af55e87dc 100644 --- a/litellm-rust/crates/core/src/providers/anthropic/chat_completions/tests.rs +++ b/litellm-rust/crates/core/src/providers/anthropic/chat_completions/tests.rs @@ -29,6 +29,31 @@ fn reason(msgs: Value, opts: Value) -> Option { ANTHROPIC_CHAT_COMPLETIONS_CONFIG.unsupported_reason(&messages(msgs), ¶ms(opts)) } +#[test] +fn forwards_provider_extensions_and_applies_explicit_overrides() { + let opts = json!({"metadata":{"user_id":"u1"}, "future":{"nested":[null,false,0]}, + "temperature":0.1, "extra_body":{"temperature":0.7, "model":"wrong"}}); + let msgs = json!([{"role":"user","content":"hi"}]); + assert_eq!(reason(msgs.clone(), opts.clone()), None); + let body = transform("resolved", msgs, opts); + assert_eq!(body["future"], json!({"nested":[null,false,0]})); + assert_eq!(body["metadata"], json!({"user_id":"u1"})); + assert_eq!(body["temperature"], 0.7); + assert_eq!(body["model"], "resolved"); + assert!(body.get("extra_body").is_none()); +} + +#[test] +fn overrides_cannot_hide_unsupported_streaming() { + assert_eq!( + reason( + json!([{"role":"user","content":"hi"}]), + json!({"stream":false,"extra_body":{"stream":true}}) + ), + Some(Unsupported("streaming")) + ); +} + #[test] fn builds_the_messages_body_python_builds() { let body = transform( @@ -160,14 +185,11 @@ fn accepts_an_explicit_stream_false() { } #[test] -fn declines_any_param_outside_the_allowlist() { +fn declines_params_requiring_unsupported_behavior() { for param in [ json!({"tools": []}), json!({"tool_choice": {"type": "auto"}}), json!({"thinking": {"type": "enabled"}}), - json!({"system": "injected"}), - json!({"metadata": {"user_id": "u1"}}), - json!({"output_config": {"effort": "high"}}), ] { assert_eq!( reason(json!([{"role": "user", "content": "hi"}]), param.clone()), diff --git a/litellm-rust/crates/core/src/providers/anthropic/chat_completions/transformation.rs b/litellm-rust/crates/core/src/providers/anthropic/chat_completions/transformation.rs index 8c674c868a1..a473605fd62 100644 --- a/litellm-rust/crates/core/src/providers/anthropic/chat_completions/transformation.rs +++ b/litellm-rust/crates/core/src/providers/anthropic/chat_completions/transformation.rs @@ -127,7 +127,7 @@ impl ChatCompletionsProviderConfig for AnthropicChatCompletionsConfig { messages: &[ChatMessage], optional_params: &OpaqueParams, ) -> Option { - unsupported_param(self.supported_openai_params(), &[], optional_params) + unsupported_param(optional_params) .or_else(|| messages.iter().find_map(unsupported_message)) // Anthropic rejects a request whose first turn is not a user turn. // Python only repairs that under `litellm.modify_params`, which the @@ -145,11 +145,10 @@ impl ChatCompletionsProviderConfig for AnthropicChatCompletionsConfig { optional_params: OpaqueParams, ) -> Result { Ok(ProviderChatRequestData { - body: anthropic_body( - model, - &build_conversation(&messages), - optional_params.into_inner(), - ), + body: crate::params::merge_extra_params( + &anthropic_body(model, &build_conversation(&messages), Map::new()), + optional_params.without(&["stream"]), + )?, }) } diff --git a/litellm-rust/crates/core/src/providers/azure_ai/messages/transformation.rs b/litellm-rust/crates/core/src/providers/azure_ai/messages/transformation.rs index 8edda1901e5..175f3d20d94 100644 --- a/litellm-rust/crates/core/src/providers/azure_ai/messages/transformation.rs +++ b/litellm-rust/crates/core/src/providers/azure_ai/messages/transformation.rs @@ -86,7 +86,7 @@ fn text_content_block(text: String) -> ContentBlock { ]); ContentBlock { cache_control: None, - extra, + extra: extra.into(), } } diff --git a/litellm-rust/crates/core/src/providers/bedrock/audio_transcription.rs b/litellm-rust/crates/core/src/providers/bedrock/audio_transcription.rs index 3b078ddf037..e31f645db2c 100644 --- a/litellm-rust/crates/core/src/providers/bedrock/audio_transcription.rs +++ b/litellm-rust/crates/core/src/providers/bedrock/audio_transcription.rs @@ -70,17 +70,20 @@ impl AudioTranscriptionProviderConfig for BedrockAudioTranscriptionConfig { inference_config.insert("temperature".to_string(), temperature.clone()); } Ok(AudioTranscriptionRequestData { - body: json!({ - "messages": [{ - "role": "user", - "content": [ - {"audio": {"format": format, "source": {"bytes": data}}}, - {"text": instruction} - ] - }], - "system": [{"text": "You are a transcription assistant."}], - "inferenceConfig": inference_config, - }), + body: crate::params::merge_extra_params( + &json!({ + "messages": [{ + "role": "user", + "content": [ + {"audio": {"format": format, "source": {"bytes": data}}}, + {"text": instruction} + ] + }], + "system": [{"text": "You are a transcription assistant."}], + "inferenceConfig": inference_config, + }), + optional_params.without(SUPPORTED_PARAMS), + )?, }) } @@ -178,7 +181,8 @@ mod tests { ] }], "system": [{"text": "You are a transcription assistant."}], - "inferenceConfig": {"maxTokens": 4096, "temperature": 0} + "inferenceConfig": {"maxTokens": 4096, "temperature": 0}, + "timestamp_granularities": ["word"] }) ); } diff --git a/litellm-rust/crates/core/src/providers/bedrock/chat_completions/tests.rs b/litellm-rust/crates/core/src/providers/bedrock/chat_completions/tests.rs index 05f8be1f11a..69c05b60468 100644 --- a/litellm-rust/crates/core/src/providers/bedrock/chat_completions/tests.rs +++ b/litellm-rust/crates/core/src/providers/bedrock/chat_completions/tests.rs @@ -35,6 +35,24 @@ fn reason(msgs: Value, opts: Value) -> Option { BEDROCK_CHAT_COMPLETIONS_CONFIG.unsupported_reason(&messages(msgs), ¶ms(opts)) } +#[test] +fn forwards_native_extension_locations_without_guessing_inference_keys() { + let opts = json!({"maxTokens":64, "future":null, + "additionalModelRequestFields":{"top_k":40,"new_option":[false,0]}, + "extra_body":{"requestMetadata":{"key":"value"}}, "aws_secret_access_key":"secret"}); + let msgs = json!([{"role":"user","content":"hi"}]); + assert_eq!(reason(msgs.clone(), opts.clone()), None); + let body = transform(msgs, opts); + assert_eq!(body["inferenceConfig"], json!({"maxTokens":64})); + assert_eq!( + body["additionalModelRequestFields"], + json!({"top_k":40,"new_option":[false,0]}) + ); + assert_eq!(body.get("future"), Some(&Value::Null)); + assert_eq!(body["requestMetadata"], json!({"key":"value"})); + assert!(body.get("aws_secret_access_key").is_none()); +} + #[test] fn builds_the_converse_body_python_builds() { let body = transform( @@ -123,13 +141,11 @@ fn declines_top_k_because_python_routes_it_by_base_model() { } #[test] -fn declines_tools_and_other_params_outside_the_allowlist() { +fn declines_params_requiring_unsupported_behavior() { for param in [ json!({"tools": []}), json!({"tool_choice": {"auto": {}}}), json!({"thinking": {"type": "enabled"}}), - json!({"requestMetadata": {"k": "v"}}), - json!({"outputConfig": {}}), json!({"_parallel_tool_use_config": {}}), ] { assert_eq!( diff --git a/litellm-rust/crates/core/src/providers/bedrock/chat_completions/transformation.rs b/litellm-rust/crates/core/src/providers/bedrock/chat_completions/transformation.rs index 92a784ccf5a..5c805dcb41a 100644 --- a/litellm-rust/crates/core/src/providers/bedrock/chat_completions/transformation.rs +++ b/litellm-rust/crates/core/src/providers/bedrock/chat_completions/transformation.rs @@ -177,36 +177,32 @@ impl ChatCompletionsProviderConfig for BedrockChatCompletionsConfig { messages: &[ChatMessage], optional_params: &OpaqueParams, ) -> Option { - unsupported_param( - self.supported_openai_params(), - CONFIG_PARAMS, - optional_params, - ) - .or_else(|| messages.iter().find_map(unsupported_message)) - // Python's Converse translation drops blank text blocks instead of - // substituting the placeholder the shared conversation builder - // applies, so decline blank text rather than diverge. - .or_else(|| { - messages - .iter() - .any(has_blank_text) - .then_some(Unsupported("blank message text")) - }) - // Converse has no assistant prefill: Python inserts a continue turn - // when a conversation opens or closes on an assistant message, and - // only under `litellm.modify_params`, which the core cannot see. - // Declining both ends also keeps the shared builder's final - // assistant right-strip (an Anthropic rule) unreachable here. - .or_else(|| { - let conversation = build_conversation(messages); - let ends_on_assistant = conversation - .turns - .last() - .is_some_and(|turn| turn.role == TurnRole::Assistant); - (!conversation.opens_on_user_turn() || ends_on_assistant).then_some(Unsupported( - "conversation does not run user turn to user turn", - )) - }) + unsupported_param(optional_params) + .or_else(|| messages.iter().find_map(unsupported_message)) + // Python's Converse translation drops blank text blocks instead of + // substituting the placeholder the shared conversation builder + // applies, so decline blank text rather than diverge. + .or_else(|| { + messages + .iter() + .any(has_blank_text) + .then_some(Unsupported("blank message text")) + }) + // Converse has no assistant prefill: Python inserts a continue turn + // when a conversation opens or closes on an assistant message, and + // only under `litellm.modify_params`, which the core cannot see. + // Declining both ends also keeps the shared builder's final + // assistant right-strip (an Anthropic rule) unreachable here. + .or_else(|| { + let conversation = build_conversation(messages); + let ends_on_assistant = conversation + .turns + .last() + .is_some_and(|turn| turn.role == TurnRole::Assistant); + (!conversation.opens_on_user_turn() || ends_on_assistant).then_some(Unsupported( + "conversation does not run user turn to user turn", + )) + }) } fn transform_request( @@ -216,7 +212,16 @@ impl ChatCompletionsProviderConfig for BedrockChatCompletionsConfig { optional_params: OpaqueParams, ) -> Result { Ok(ProviderChatRequestData { - body: converse_body(&build_conversation(&messages), &optional_params), + body: crate::params::merge_extra_params( + &converse_body(&build_conversation(&messages), &optional_params), + optional_params.without(&[ + "maxTokens", + "temperature", + "topP", + "stopSequences", + "stream", + ]), + )?, }) } diff --git a/litellm-rust/crates/core/src/providers/openai/responses/transformation.rs b/litellm-rust/crates/core/src/providers/openai/responses/transformation.rs index be86bb90311..a35f98602de 100644 --- a/litellm-rust/crates/core/src/providers/openai/responses/transformation.rs +++ b/litellm-rust/crates/core/src/providers/openai/responses/transformation.rs @@ -16,8 +16,10 @@ impl ResponsesWebSocketProviderConfig for OpenAIResponsesWsConfig { event: &ResponsesWsEvent, model: &str, ) -> Result { + let mut event = event.clone(); + event.data = event.data.into_provider_body()?.into(); Ok(ResponsesWsTransformResult::passthrough(enforce_model( - event, model, + &event, model, ))) } @@ -34,6 +36,24 @@ impl ResponsesWebSocketProviderConfig for OpenAIResponsesWsConfig { mod tests { use super::*; + #[test] + fn request_extensions_round_trip_with_overrides_and_resolved_model() { + let event = serde_json::from_value(serde_json::json!({ + "type":"response.create", "future":{"nested":[null,false,0]}, + "extra_body":{"model":"wrong","provider_option":null} + })) + .unwrap(); + let result = OPENAI_RESPONSES_WS_CONFIG + .transform_ws_request(&event, "resolved") + .unwrap(); + assert_eq!( + serde_json::to_value(&result.events[0]).unwrap(), + serde_json::json!({ + "type":"response.create", "model":"resolved", "future":{"nested":[null,false,0]}, "provider_option":null + }) + ); + } + #[test] fn openai_config_is_native_and_enforces_model() { let event: ResponsesWsEvent = diff --git a/litellm-rust/crates/core/src/responses/types.rs b/litellm-rust/crates/core/src/responses/types.rs index 4942309992e..00ac8c11daa 100644 --- a/litellm-rust/crates/core/src/responses/types.rs +++ b/litellm-rust/crates/core/src/responses/types.rs @@ -1,5 +1,5 @@ use serde::{Deserialize, Deserializer, Serialize, Serializer}; -use serde_json::{Map, Value}; +use serde_json::Value; #[derive(Clone, Debug, PartialEq, Eq)] pub enum ResponsesWsEventType { @@ -58,7 +58,7 @@ pub struct ResponsesWsEvent { #[serde(rename = "type")] pub event_type: ResponsesWsEventType, #[serde(flatten)] - pub data: Map, + pub data: crate::params::OpaqueParams, } impl ResponsesWsEvent { diff --git a/litellm-rust/crates/core/tests/azure_document_intelligence_ocr.rs b/litellm-rust/crates/core/tests/azure_document_intelligence_ocr.rs index 62d0c657a1a..5837076f877 100644 --- a/litellm-rust/crates/core/tests/azure_document_intelligence_ocr.rs +++ b/litellm-rust/crates/core/tests/azure_document_intelligence_ocr.rs @@ -21,7 +21,7 @@ async fn facade_maps_pages_features_and_url_document() { let mut request = wire_request( "azure_ai/doc-intelligence/prebuilt-read", &base, - json!({"pages":[2,0,0,1],"features":["keyValuePairs","languages"]}), + json!({"pages":[2,0,0,1],"features":["keyValuePairs","languages"], "future_option": {"nested":null}, "extra_body":{"provider_option":false}}), ); request.document = serde_json::from_value(json!({ "type":"document_url", @@ -42,7 +42,7 @@ async fn facade_maps_pages_features_and_url_document() { let body: Value = serde_json::from_str(request.split_once("\r\n\r\n").unwrap().1).unwrap(); assert_eq!( body, - json!({"urlSource":"https://example.com/document.pdf"}) + json!({"urlSource":"https://example.com/document.pdf", "future_option":{"nested":null}, "provider_option":false}) ); } diff --git a/litellm-rust/crates/core/tests/ocr.rs b/litellm-rust/crates/core/tests/ocr.rs index 0ce2e4c1c39..39223beaabc 100644 --- a/litellm-rust/crates/core/tests/ocr.rs +++ b/litellm-rust/crates/core/tests/ocr.rs @@ -59,7 +59,7 @@ async fn facade_executes_direct_mistral_once() { let result = perform_ocr(wire_request( "mistral/model", &base, - json!({"pages":"0,2-4","extract_header":true,"unknown":"ignored"}), + json!({"pages":"0,2-4","extract_header":true,"unknown":{"nested":[null,false,0]}}), )) .await .unwrap(); @@ -81,7 +81,8 @@ async fn facade_executes_direct_mistral_once() { "model":"model", "document":{"type":"document_url","document_url":"data:application/pdf;base64,YWJj"}, "pages":"0,2-4", - "extract_header":true + "extract_header":true, + "unknown":{"nested":[null,false,0]} }) ); } @@ -132,6 +133,60 @@ struct RecordingHooks { block: bool, } +struct ExtensionHooks; + +impl OcrHooks for ExtensionHooks { + fn intercepts_requests(&self) -> bool { + true + } + + fn during_call( + &self, + mut request: OcrDuringCallRequest, + ) -> OcrHookFuture<'_, OcrDuringCallRequest> { + Box::pin(async move { + assert_eq!(request.body["pages"], json!([2])); + assert_eq!(request.body.get("future"), Some(&Value::Null)); + assert!( + !request + .retained_fields + .iter() + .any(|field| field == "pages" || field == "document") + ); + request.body.as_object_mut().unwrap().remove("future"); + request.body["hook_option"] = json!({"nested":[null,false,0]}); + Ok(request) + }) + } +} + +#[tokio::test] +async fn composed_extensions_reach_hooks_and_removed_fields_stay_removed() { + let (base, seen, server) = mock_server(vec![MockResponse::json(json!({"pages":[]}))]).await; + let request = super::LiteLLMOcrRequest { + hooks: Arc::new(ExtensionHooks), + ..wire_request( + "mistral/model", + &base, + json!({ + "pages":[0], "future":null, "extra_body":{"pages":[2], + "document":{"type":"document_url","document_url":"data:application/pdf;base64,eHl6"}} + }), + ) + }; + perform_ocr(request).await.unwrap(); + server.await.unwrap(); + let requests = seen.lock().unwrap(); + let body: Value = serde_json::from_str(requests[0].split_once("\r\n\r\n").unwrap().1).unwrap(); + assert_eq!(body["pages"], json!([2])); + assert_eq!( + body["document"]["document_url"], + "data:application/pdf;base64,eHl6" + ); + assert_eq!(body["hook_option"], json!({"nested":[null,false,0]})); + assert!(body.get("future").is_none()); +} + impl OcrHooks for RecordingHooks { fn intercepts_requests(&self) -> bool { true diff --git a/litellm-rust/crates/core/tests/vertex_ai_deepseek_ocr.rs b/litellm-rust/crates/core/tests/vertex_ai_deepseek_ocr.rs index 676799eb2fe..8cefe086b45 100644 --- a/litellm-rust/crates/core/tests/vertex_ai_deepseek_ocr.rs +++ b/litellm-rust/crates/core/tests/vertex_ai_deepseek_ocr.rs @@ -45,7 +45,9 @@ async fn facade_executes_vertex_deepseek_at_the_openai_endpoint() { let body = request_body(&requests[0]); assert_eq!(body["model"], "deepseek-ai/deepseek-ocr-maas"); assert_eq!(body["temperature"], 0.1); - assert!(body.get("future_ocr_option").is_none()); + assert_eq!(body["future_ocr_option"], true); + assert_eq!(body["provider_option"], "value"); + assert!(body.get("vertex_project").is_none()); assert!(body.get("extra_body").is_none()); assert_eq!( body["messages"][0]["content"][0], diff --git a/litellm-rust/crates/core/tests/vertex_ai_ocr.rs b/litellm-rust/crates/core/tests/vertex_ai_ocr.rs index adbf6de7434..64ceb9942f1 100644 --- a/litellm-rust/crates/core/tests/vertex_ai_ocr.rs +++ b/litellm-rust/crates/core/tests/vertex_ai_ocr.rs @@ -110,7 +110,7 @@ async fn configs_build_complete_requests_and_share_mistral_normalization() { "include_image_base64": true, "vertex_project": "project-1", "vertex_location": "us-central1", - "unknown": "ignored" + "unknown": "preserved" }); let direct = wire_request( "mistral/mistral-ocr-maas", @@ -143,7 +143,8 @@ async fn configs_build_complete_requests_and_share_mistral_normalization() { "model": "mistral-ocr-maas", "document": {"type": "document_url", "document_url": "data:application/pdf;base64,YWJj"}, "pages": [0, 2], - "include_image_base64": true + "include_image_base64": true, + "unknown": "preserved" }) ); } diff --git a/litellm-rust/crates/python-bridge/src/marshal.rs b/litellm-rust/crates/python-bridge/src/marshal.rs index 5f7633a64a0..68652e1b41a 100644 --- a/litellm-rust/crates/python-bridge/src/marshal.rs +++ b/litellm-rust/crates/python-bridge/src/marshal.rs @@ -92,11 +92,27 @@ pub(crate) fn project_optional_fields( kwargs: &Bound<'_, PyDict>, names: &[&str], ) -> PyResult> { - names + let controls: Vec = kwargs + .py() + .import("litellm.types.utils")? + .getattr("all_litellm_params")? + .extract()?; + kwargs .iter() - .filter_map(|name| match kwargs.get_item(name) { - Ok(Some(value)) => Some(from_py(&value).map(|value| ((*name).to_string(), value))), - Ok(None) => None, + .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" + )) => + { + Some(from_py(&value).map(|value| (name, value))) + } + Ok(_) => None, Err(error) => Some(Err(error)), }) .collect() 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 28ccf466e21..42aa9ecf37e 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/project.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/project.rs @@ -128,9 +128,9 @@ pub(super) fn project_request( let optional_params = project_optional_fields(kwargs, &names)?; let input_sources = request_input_sources( kwargs, - names - .iter() - .copied() + optional_params + .keys() + .map(String::as_str) .chain(["api_key", "api_base", "extra_headers"]), )?; let azure_ad_token_provider = kwargs diff --git a/tests/test_litellm_rust/ocr/test_cohere.py b/tests/test_litellm_rust/ocr/test_cohere.py index 2a35dc62bd1..dd5e3e7d380 100644 --- a/tests/test_litellm_rust/ocr/test_cohere.py +++ b/tests/test_litellm_rust/ocr/test_cohere.py @@ -40,7 +40,9 @@ async def test_public_cohere_request_and_normalization( request: Final = recording_server.requests[0] assert request.path == ("/providers/cohere/v2/parse" if model.startswith("azure_ai/") else "/v2/parse") assert request.headers["authorization"] == "Bearer test-key" - assert request.body == {"model": model.split("/", 1)[1], "document": IMAGE, "output_format": "markdown"} + assert request.body == { + "model": model.split("/", 1)[1], "document": IMAGE, "output_format": "markdown", "unrecognized": True + } assert [page.index for page in response.pages] == [4, 1] assert response.pages[0].markdown == "receipt" assert response.pages[0].images[0].bbox == BOX diff --git a/tests/test_litellm_rust/ocr/test_lifecycle.py b/tests/test_litellm_rust/ocr/test_lifecycle.py index dfcd63d3019..ff832dbead0 100644 --- a/tests/test_litellm_rust/ocr/test_lifecycle.py +++ b/tests/test_litellm_rust/ocr/test_lifecycle.py @@ -529,7 +529,7 @@ async def test_retained_argument_aliases_and_body_roots_survive_envelope_replace ) -> None: pages: Final = [0] document: Final = {"type": "document_url", "document_url": "data:application/pdf;base64,YWJj"} - opaque: Final = object() + opaque: Final = {"nested": [None, False, 0]} observed: Final = [] class Observe(Logging): diff --git a/tests/test_litellm_rust/ocr/test_requests.py b/tests/test_litellm_rust/ocr/test_requests.py index 4f4b39fa6c6..c600b80845a 100644 --- a/tests/test_litellm_rust/ocr/test_requests.py +++ b/tests/test_litellm_rust/ocr/test_requests.py @@ -63,6 +63,38 @@ def test_native_ocr_sends_model_and_document_to_mistral_ocr_path(ocr_server: Rec assert ocr_server.requests[0].body == {"model": "mistral-ocr-latest", "document": OCR_DOCUMENT} +@pytest.mark.asyncio +@pytest.mark.parametrize("asynchronous", [False, True]) +async def test_provider_extensions_and_overrides_survive_native_lifecycle( + ocr_server: RecordingServer, asynchronous: bool +) -> None: + options: Final = { + "pages": [0], + "future_option": {"nested": [None, False, 0]}, + "null_option": None, + "extra_body": {"pages": [2], "override_option": True, "model": "wrong"}, + } + if asynchronous: + await call_native_aocr(ocr_server, **options) + else: + call_native_ocr(ocr_server, **options) + assert_native_request(ocr_server) + assert ocr_server.requests[0].body == { + "model": "mistral-ocr-latest", + "document": OCR_DOCUMENT, + "pages": [2], + "future_option": {"nested": [None, False, 0]}, + "null_option": None, + "override_option": True, + } + + +def test_non_json_provider_extension_fails_before_http(ocr_server: RecordingServer) -> None: + ocr_server.expected_requests = 0 + with pytest.raises(litellm.APIConnectionError, match="unsupported type object"): + call_native_ocr(ocr_server, future_option=object()) + + def test_native_ocr_prepares_file_document_like_python(ocr_server: RecordingServer) -> None: response: Final = call_native_ocr( ocr_server, @@ -479,7 +511,7 @@ async def test_native_ocr_inherits_named_credentials_without_overwriting_argumen from litellm.models.credentials import CredentialItem pages: Final = [0] - opaque: Final = object() + opaque: Final = {"nested": [None, False, 0]} monkeypatch.setenv("MISTRAL_API_KEY", "environment-key") monkeypatch.setattr( litellm, @@ -518,6 +550,7 @@ async def test_native_ocr_inherits_named_credentials_without_overwriting_argumen assert ocr_server.requests[0].body["pages"] == [0, 2] + @pytest.mark.parametrize("source", ["sdk", "proxy"]) @pytest.mark.parametrize( "filename,mime", [("scan.PNG", "image/png"), ("document.pdf", "application/pdf"), ("note.txt", "text/plain")] diff --git a/tests/test_litellm_rust/test_ocr.py b/tests/test_litellm_rust/test_ocr.py index e0e06d685b8..f1f0a16c35b 100644 --- a/tests/test_litellm_rust/test_ocr.py +++ b/tests/test_litellm_rust/test_ocr.py @@ -122,14 +122,14 @@ def test_native_lifecycle_core_encodes_python_file_input( document={"type": "file", "file": file_input, "mime_type": mime_type}, api_key="test-key", api_base=f"http://127.0.0.1:{server.server_port}", - opaque_extension=object(), + opaque_extension={"nested": [None, False, 0]}, ) assert response.pages[0].markdown == "native OCR response" assert requests[0]["body"]["document"] == { "type": expected_type, expected_field: expected_uri, } - assert "opaque_extension" not in requests[0]["body"] + assert requests[0]["body"]["opaque_extension"] == {"nested": [None, False, 0]} @pytest.mark.parametrize("asynchronous", [False, True])