This commit is contained in:
Yujong Lee 2026-09-15 09:16:44 -07:00
parent ff3185f385
commit cc88a9479e
32 changed files with 519 additions and 193 deletions

View file

@ -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 {

View file

@ -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(

View file

@ -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));
}

View file

@ -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> {
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<ChatCompletionsResponse, Error>;
}
pub fn unsupported_param(
supported: &'static [(&'static str, &'static str)],
config: &'static [&'static str],
optional_params: &OpaqueParams,
) -> Option<Unsupported> {
if optional_params
pub fn unsupported_param(optional_params: &OpaqueParams) -> Option<Unsupported> {
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"))
}

View file

@ -73,7 +73,7 @@ pub struct ChatMessage {
#[serde(default, skip_serializing_if = "Option::is_none")]
pub name: Option<String>,
#[serde(flatten)]
pub extra: Map<String, Value>,
pub extra: crate::params::OpaqueParams,
}
/// OpenAI `usage`, including the `prompt_tokens_details` split LiteLLM's Python

View file

@ -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(),
}
})
}

View file

@ -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(

View file

@ -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]

View file

@ -37,6 +37,8 @@ mod types {
pub(crate) struct DeepSeekOcrMessage {
pub role: UserRole,
pub content: Vec<crate::ocr::types::OcrDocument>,
#[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(),
})

View file

@ -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)?;

View file

@ -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}");

View file

@ -44,7 +44,7 @@ pub struct ContentBlock {
#[serde(skip_serializing_if = "Option::is_none")]
pub cache_control: Option<CacheControl>,
#[serde(flatten)]
pub extra: Map<String, Value>,
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<String>,
#[serde(flatten)]
pub extra: Map<String, Value>,
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<String, Value>,
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<String>,
#[serde(flatten)]
pub extra: Map<String, Value>,
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<Value>,
#[serde(flatten)]
pub extra: Map<String, Value>,
pub extra: crate::params::OpaqueParams,
}

View file

@ -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<T> {
#[serde(flatten)]
pub known: T,
#[serde(default, flatten)]
pub extra_params: Map<String, Value>,
}
pub(crate) use crate::params::{ParsedProviderParams, merge_extra_params};
pub(crate) fn _prepare_ocr_request<T: DeserializeOwned>(
request: &LiteLLMOcrRequest,
) -> Result<ParsedProviderParams<T>, OcrRequestError> {
super::wire::decode_request_value(
Value::Object(request.optional_params.clone().into()),
Value::Object(request.optional_params.provider_params().into()),
"optional_params",
)
}
pub(crate) fn merge_extra_params<B: Serialize>(
body: &B,
extra_params: Map<String, Value>,
) -> Result<Value, OcrRequestError> {
let Value::Object(fields) =
serde_json::to_value(body).map_err(|_| OcrRequestError::RequestField {
path: "body".into(),
})?
else {
return Err(OcrRequestError::RequestField {
path: "body".into(),
});
};
let extra_body = extra_params
.get("extra_body")
.and_then(Value::as_object)
.cloned()
.unwrap_or_default()
.into_iter()
.collect::<Map<String, Value>>();
Ok(Value::Object(
fields
.into_iter()
.chain(
extra_params
.into_iter()
.filter(|(name, _)| name != "extra_body"),
)
.chain(extra_body)
.collect(),
))
}
pub(crate) async fn transform_request_body<B>(
client: &OcrClient,
request: &LiteLLMOcrRequest,
@ -63,13 +24,19 @@ pub(crate) async fn transform_request_body<B>(
headers: &[(String, String)],
retains_document: bool,
body: B,
validate: impl FnOnce(&B) -> Result<(), OcrRequestError>,
validate: impl Fn(&B) -> Result<(), OcrRequestError>,
) -> Result<reqwest::Request, OcrError>
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::<B>::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::<B>::decode(changed.body)?;
let body = OcrWireBody::<B>::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<B> {
#[serde(flatten)]
body: B,
#[serde(flatten)]
extra: Map<String, Value>,
extra: crate::params::OpaqueParams,
}
impl<B: Serialize + DeserializeOwned> OcrWireBody<B> {
fn decode(value: Value) -> Result<Self, OcrRequestError> {
let body: B = super::wire::decode_request_value(value.clone(), "guardrail.body")?;
fn decode(value: Value, prefix: &str) -> Result<Self, OcrRequestError> {
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<String> {
}
#[cfg(test)]
mod tests {
use serde::Deserialize;
use serde_json::json;
use super::*;

View file

@ -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<String, Value>);
#[derive(Debug, Deserialize)]
pub(crate) struct ParsedProviderParams<T> {
#[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<String, Value> {
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<Map<String, Value>, 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<B: Serialize>(
body: &B,
extra_params: OpaqueParams,
) -> Result<Value, crate::Error> {
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<Map<String, Value>> for OpaqueParams {
fn from(value: Map<String, Value>) -> 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(),

View file

@ -29,6 +29,31 @@ fn reason(msgs: Value, opts: Value) -> Option<Unsupported> {
ANTHROPIC_CHAT_COMPLETIONS_CONFIG.unsupported_reason(&messages(msgs), &params(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()),

View file

@ -127,7 +127,7 @@ impl ChatCompletionsProviderConfig for AnthropicChatCompletionsConfig {
messages: &[ChatMessage],
optional_params: &OpaqueParams,
) -> Option<Unsupported> {
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<ProviderChatRequestData, Error> {
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"]),
)?,
})
}

View file

@ -86,7 +86,7 @@ fn text_content_block(text: String) -> ContentBlock {
]);
ContentBlock {
cache_control: None,
extra,
extra: extra.into(),
}
}

View file

@ -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"]
})
);
}

View file

@ -35,6 +35,24 @@ fn reason(msgs: Value, opts: Value) -> Option<Unsupported> {
BEDROCK_CHAT_COMPLETIONS_CONFIG.unsupported_reason(&messages(msgs), &params(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!(

View file

@ -177,36 +177,32 @@ impl ChatCompletionsProviderConfig for BedrockChatCompletionsConfig {
messages: &[ChatMessage],
optional_params: &OpaqueParams,
) -> Option<Unsupported> {
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<ProviderChatRequestData, Error> {
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",
]),
)?,
})
}

View file

@ -16,8 +16,10 @@ impl ResponsesWebSocketProviderConfig for OpenAIResponsesWsConfig {
event: &ResponsesWsEvent,
model: &str,
) -> Result<ResponsesWsTransformResult, Error> {
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 =

View file

@ -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<String, Value>,
pub data: crate::params::OpaqueParams,
}
impl ResponsesWsEvent {

View file

@ -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})
);
}

View file

@ -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

View file

@ -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],

View file

@ -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"
})
);
}

View file

@ -92,11 +92,27 @@ pub(crate) fn project_optional_fields(
kwargs: &Bound<'_, PyDict>,
names: &[&str],
) -> PyResult<Map<String, Value>> {
names
let controls: Vec<String> = 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::<String>()?, 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()

View file

@ -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

View file

@ -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

View file

@ -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):

View file

@ -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")]

View file

@ -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])