mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-07 02:59:05 +00:00
wip
This commit is contained in:
parent
ff3185f385
commit
cc88a9479e
32 changed files with 519 additions and 193 deletions
|
|
@ -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 {
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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));
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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"))
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(),
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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(),
|
||||
})
|
||||
|
|
|
|||
|
|
@ -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)?;
|
||||
|
|
|
|||
|
|
@ -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}");
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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::*;
|
||||
|
|
|
|||
|
|
@ -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(),
|
||||
|
|
|
|||
|
|
@ -29,6 +29,31 @@ fn reason(msgs: Value, opts: Value) -> Option<Unsupported> {
|
|||
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()),
|
||||
|
|
|
|||
|
|
@ -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"]),
|
||||
)?,
|
||||
})
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -86,7 +86,7 @@ fn text_content_block(text: String) -> ContentBlock {
|
|||
]);
|
||||
ContentBlock {
|
||||
cache_control: None,
|
||||
extra,
|
||||
extra: extra.into(),
|
||||
}
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
})
|
||||
);
|
||||
}
|
||||
|
|
|
|||
|
|
@ -35,6 +35,24 @@ fn reason(msgs: Value, opts: Value) -> Option<Unsupported> {
|
|||
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!(
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
]),
|
||||
)?,
|
||||
})
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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 =
|
||||
|
|
|
|||
|
|
@ -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 {
|
||||
|
|
|
|||
|
|
@ -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})
|
||||
);
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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],
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
})
|
||||
);
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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")]
|
||||
|
|
|
|||
|
|
@ -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])
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue