refactor(rust): give Messages configs typed litellm params (#45611)

* refactor(rust): give Messages configs typed litellm params

* feat(rust): port _normalize_system_role_messages for Messages hosts

mid_conversation_system mirrors Python's module function for function: a
leading run of system turns is hoisted into the top-level system, a later
turn stays in place when the model map flags supports_mid_conversation_system
and otherwise becomes a user turn where it was, never between a tool_use and
its tool_result, and billing blocks are stripped from the hoisted system. The
capability joins MessagesModelCapabilities and the bridge projection. Azure AI
uses it in place of fold_system_role_messages, which hoisted every turn.

* fix(rust): keep the Bedrock runtime endpoint behind a blank api_base

* refactor(rust): import LitellmParams from its crate instead of a re-export

* refactor(rust): template the Messages request decode errors

Both decode failures go through ErrorDetail::invalid with the subject and the
serde error as source instead of a preformatted sentence.

* refactor(rust_bridge): read the module globals the param specs name

The Messages host no longer keeps its own list of litellm globals; a spec
whose spellings the call leaves out names the global to read, so the next
family that has one needs no host change.

* refactor(rust): give each route lookup Python's implicit None

ProviderConfigManager's per-route methods list only the providers the route
serves and fall through to None; the Rust lookups now do the same, so a new
LlmProviders variant touches only the route that serves it.
This commit is contained in:
yujonglee 2026-10-09 16:29:45 -07:00 • committed by GitHub
parent 3d2abf9a44
commit ee295aea68
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
39 changed files with 901 additions and 240 deletions

View file

@ -4031,6 +4031,7 @@ dependencies = [
"litellm-inference-testing",
"litellm-llms",
"litellm-llms-types",
"litellm-router-types",
"litellm-secrets",
"litellm-tracing",
"reqwest 0.12.28",
@ -4151,6 +4152,7 @@ dependencies = [
"litellm-http",
"litellm-llms-types",
"litellm-python-compat",
"litellm-router-types",
"litellm-secrets",
"reqwest 0.12.28",
"rstest",
@ -4222,6 +4224,7 @@ dependencies = [
"litellm-inference-transcription",
"litellm-llms",
"litellm-llms-types",
"litellm-router-types",
"litellm-secrets",
"litellm-secrets-aws",
"litellm-secrets-types",
@ -4270,6 +4273,7 @@ dependencies = [
"litellm-config",
"litellm-inference",
"litellm-inference-messages",
"litellm-router-types",
"rstest",
]

View file

@ -79,6 +79,7 @@ fn project(
api_key: deployment.api_key.clone(),
api_base: deployment.api_base.clone(),
custom_llm_provider: deployment.custom_llm_provider.clone(),
litellm_params: deployment.litellm_params.clone(),
extra_headers: None,
provider_specific_header: anthropic_api_headers(headers),
timeout: deployment.timeout,

View file

@ -33,13 +33,7 @@ pub(super) fn chat_completions_provider(provider: LlmProviders) -> Option<ChatPr
LlmProviders::Anthropic => Some(ChatProvider::Anthropic),
LlmProviders::Bedrock => Some(ChatProvider::Bedrock),
LlmProviders::OpenaiLike => Some(ChatProvider::OpenaiLike),
LlmProviders::AwsTextract
| LlmProviders::AzureAi
| LlmProviders::Cohere
| LlmProviders::Mistral
| LlmProviders::Openai
| LlmProviders::Reducto
| LlmProviders::VertexAi => None,
_ => None,
}
}

View file

@ -17,6 +17,7 @@ litellm-http.workspace = true
litellm-inference.workspace = true
litellm-llms.workspace = true
litellm-llms-types.workspace = true
litellm-router-types.workspace = true
litellm-secrets.workspace = true
litellm-tracing.workspace = true
reqwest.workspace = true

View file

@ -44,13 +44,7 @@ pub(crate) fn messages_provider(provider: LlmProviders) -> Option<MessagesProvid
LlmProviders::Anthropic => Some(MessagesProvider::Anthropic),
LlmProviders::AzureAi => Some(MessagesProvider::AzureAi),
LlmProviders::Bedrock => Some(MessagesProvider::Bedrock),
LlmProviders::AwsTextract
| LlmProviders::Cohere
| LlmProviders::Mistral
| LlmProviders::Openai
| LlmProviders::OpenaiLike
| LlmProviders::Reducto
| LlmProviders::VertexAi => None,
_ => None,
}
}

View file

@ -70,7 +70,9 @@ impl MessagesRoute {
WireRequest {
url,
headers: authenticated.headers,
body: serde_json::to_value(&body).map_err(serialize_failure)?,
body: provider
.config()
.wire_body(serde_json::to_value(&body).map_err(serialize_failure)?),
},
request_context,
)

View file

@ -15,7 +15,8 @@ use std::sync::Arc;
pub use litellm_inference::RouteError as Error;
pub use types::{
MessagesCall, MessagesCallResponse, MessagesSettings, MessagesShaping, messages_body,
MessagesCall, MessagesCallResponse, MessagesSettings, MessagesShaping, litellm_params,
messages_body,
};
#[derive(Clone)]

View file

@ -69,6 +69,7 @@ fn prepare_provider_request(
body,
api_key,
api_base,
litellm_params,
extra_headers,
provider_specific_header,
timeout,
@ -98,6 +99,7 @@ fn prepare_provider_request(
forwarded,
api_key.as_deref(),
&transformed.model,
&litellm_params,
&env_lookup,
)?;
let environment = ValidatedEnvironment {
@ -108,11 +110,13 @@ fn prepare_provider_request(
auth: validated.auth,
};
let url = if transformed.params.stream == Some(true) {
config.complete_stream_url(api_base.as_deref(), &transformed.model, &env_lookup)?
} else {
config.get_complete_url(api_base.as_deref(), &transformed.model, &env_lookup)?
};
let url = config.get_complete_url(
api_base.as_deref(),
&transformed.model,
&litellm_params,
transformed.params.stream == Some(true),
&env_lookup,
)?;
Ok(ProviderMessagesRequest {
provider,
@ -227,6 +231,7 @@ mod tests {
api_key: None,
api_base: None,
custom_llm_provider: Some("anthropic".into()),
litellm_params: Default::default(),
extra_headers: None,
provider_specific_header: None,
timeout: None,
@ -253,6 +258,7 @@ mod tests {
api_key: Some("sk-test".into()),
api_base: Some("https://anthropic.test".into()),
custom_llm_provider: Some("anthropic".into()),
litellm_params: Default::default(),
extra_headers: None,
provider_specific_header: None,
timeout: None,
@ -360,6 +366,7 @@ mod tests {
api_key: Some("sk-test".into()),
api_base: Some("https://resource.services.ai.azure.com".into()),
custom_llm_provider: custom_llm_provider.map(Into::into),
litellm_params: Default::default(),
extra_headers: Some(Map::from_iter([("x-priority".into(), json!("extra"))])),
provider_specific_header: Some(configured),
timeout: None,

View file

@ -2,11 +2,12 @@ use std::time::Duration;
use bytes::Bytes;
use litellm_host::call::CallOutput;
use litellm_llms::base_llm::messages::context::MessagesModelCapabilities;
use litellm_llms::{ErrorDetail, base_llm::messages::context::MessagesModelCapabilities};
use litellm_llms_types::{
formats::messages::{MessagesRequest, MessagesResponse},
headers::ProviderSpecificHeaders,
};
use litellm_router_types::LitellmParams;
use serde::{Deserialize, Serialize};
use serde_json::{Map, Value};
@ -17,6 +18,7 @@ pub struct MessagesCall {
pub api_key: Option<String>,
pub api_base: Option<String>,
pub custom_llm_provider: Option<String>,
pub litellm_params: LitellmParams,
pub extra_headers: Option<Map<String, Value>>,
pub provider_specific_header: Option<ProviderSpecificHeaders>,
pub timeout: Option<Duration>,
@ -27,8 +29,16 @@ pub fn messages_body(body: Map<String, Value>) -> Result<MessagesRequest, Error>
serde_json::from_value(Value::Object(body)).map_err(invalid_request)
}
/// The caller's litellm params, projected by a host from the keys [`LitellmParams::fields`]
/// names. A key present with a value of the wrong type is a request error, as it is for
/// Python's `GenericLiteLLMParams(**kwargs)`.
pub fn litellm_params(fields: Map<String, Value>) -> Result<LitellmParams, Error> {
serde_json::from_value(Value::Object(fields))
.map_err(|err| Error::InvalidRequest(ErrorDetail::invalid("litellm params", err)))
}
pub(super) fn invalid_request(err: serde_json::Error) -> Error {
Error::InvalidRequest(format!("invalid Anthropic messages request: {err}").into())
Error::InvalidRequest(ErrorDetail::invalid("Anthropic messages request", err))
}
pub type MessagesCallResponse =
@ -126,6 +136,7 @@ mod tests {
supports_output_config: true,
supports_sampling_params: false,
supports_speed: true,
supports_mid_conversation_system: false,
effort_tiers: SupportedEffortTiers {
minimal: false,
low: true,

View file

@ -206,7 +206,7 @@ async fn messages_cache_identity_follows_resolved_configuration_and_request_call
let cache = ScopedCache::new(cache.clone(), CacheScope::Shared);
support::messages_route(secrets.clone()).with_cache(cache).execute(MessagesCall {
body: serde_json::from_value(json!({"model":"anthropic/cache-test-model","messages":[{"role":"user","content":"hello"}],"max_tokens":32})).unwrap(),
api_key:None,api_base:None,custom_llm_provider:None,extra_headers:None,provider_specific_header:None,timeout:None,shaping:Default::default(),
api_key:None,api_base:None,custom_llm_provider:None,litellm_params:Default::default(),extra_headers:None,provider_specific_header:None,timeout:None,shaping:Default::default(),
}, &hooks, None).await.unwrap();
}
assert_eq!(hooks.calls.load(Ordering::SeqCst), 4);

View file

@ -81,6 +81,7 @@ fn call() -> MessagesCall {
api_key: None,
api_base: None,
custom_llm_provider: Some("anthropic".into()),
litellm_params: Default::default(),
extra_headers: None,
provider_specific_header: None,
timeout: Some(Duration::from_secs(5)),

View file

@ -58,6 +58,7 @@ async fn credentials_become_exactly_one_auth_header(
run_message(MessagesCall {
custom_llm_provider: Some(provider.into()),
litellm_params: Default::default(),
api_key: api_key.map(Into::into),
api_base: Some(upstream.uri()),
extra_headers: headers(extra_headers.iter().copied()),
@ -85,6 +86,7 @@ async fn a_call_without_credentials_fails_before_sending(
let error = run(MessagesCall {
custom_llm_provider: Some(provider.into()),
litellm_params: Default::default(),
api_base: Some(upstream.uri()),
..call
})
@ -124,6 +126,7 @@ async fn each_provider_posts_to_its_messages_endpoint(
run_message(MessagesCall {
custom_llm_provider: provider.map(Into::into),
litellm_params: Default::default(),
api_key: Some("sk".into()),
api_base: Some(format!("{}{base_suffix}", upstream.uri())),
..with_model(call, model)
@ -154,6 +157,7 @@ async fn unsupported_providers_are_rejected_before_sending(
) {
let error = run(MessagesCall {
custom_llm_provider: provider.map(Into::into),
litellm_params: Default::default(),
api_key: Some("sk".into()),
api_base: Some(UNREACHABLE_BASE.into()),
..with_model(call, model)
@ -206,6 +210,7 @@ async fn cache_scope_removal_is_selected_by_the_provider(
run_message(MessagesCall {
custom_llm_provider: Some(provider.into()),
litellm_params: Default::default(),
api_key: Some("sk-azure".into()),
api_base: Some(upstream.uri()),
body: body(json!({
@ -349,6 +354,7 @@ async fn caller_protocol_headers_win_over_the_defaults(call: MessagesCall, #[cas
run_message(MessagesCall {
custom_llm_provider: Some(provider.into()),
litellm_params: Default::default(),
api_key: Some("sk".into()),
api_base: Some(upstream.uri()),
extra_headers: headers([
@ -400,6 +406,7 @@ async fn unsupported_params_are_dropped_under_drop_params_and_rejected_without_i
},
body: call.body.clone(),
custom_llm_provider: call.custom_llm_provider.clone(),
litellm_params: Default::default(),
extra_headers: None,
provider_specific_header: None,
timeout: call.timeout,
@ -577,6 +584,7 @@ async fn metadata_is_reduced_to_the_user_id(call: MessagesCall, #[case] provider
run_message(with_fields(
MessagesCall {
custom_llm_provider: Some(provider.into()),
litellm_params: Default::default(),
api_key: Some("sk".into()),
api_base: Some(upstream.uri()),
..call
@ -638,6 +646,7 @@ async fn system_message_folding_is_selected_by_the_provider(
run_message(with_fields(
MessagesCall {
custom_llm_provider: Some(provider.into()),
litellm_params: Default::default(),
api_key: Some("sk-azure".into()),
api_base: Some(upstream.uri()),
shaping: MessagesShaping {
@ -702,6 +711,7 @@ async fn provider_validation_runs_before_caller_parameter_removal(
let result = run(with_fields(
MessagesCall {
custom_llm_provider: Some(provider.into()),
litellm_params: Default::default(),
api_key: Some("sk-test".into()),
api_base: Some(upstream.uri()),
shaping: MessagesShaping {

View file

@ -104,6 +104,7 @@ async fn the_provider_message_is_returned(call: MessagesCall, #[case] provider:
let message = run_message(MessagesCall {
custom_llm_provider: Some(provider.into()),
litellm_params: Default::default(),
api_key: Some("sk".into()),
api_base: Some(upstream.uri()),
..call

View file

@ -38,6 +38,7 @@ async fn the_credential_and_base_come_from_the_secret_source(
secrets.clone(),
MessagesCall {
custom_llm_provider: Some(provider.into()),
litellm_params: Default::default(),
..call
},
)
@ -189,6 +190,7 @@ async fn azure_without_a_base_anywhere_fails_before_sending(call: MessagesCall)
Arc::new(RecordingSecrets::new([("AZURE_API_KEY", "sk-azure")])),
MessagesCall {
custom_llm_provider: Some("azure_ai".into()),
litellm_params: Default::default(),
..call
},
)

View file

@ -346,6 +346,7 @@ async fn the_sdk_returns_stream_headers_and_every_sse_byte(
.execute(
MessagesCall {
custom_llm_provider: Some(provider.into()),
litellm_params: Default::default(),
..streaming(call, upstream.uri())
},
&(),
@ -470,6 +471,7 @@ async fn a_host_on_anthropic_sse_is_relayed_byte_for_byte(call: MessagesCall) {
let host = RecordingStreamHost::new(
MessagesCall {
custom_llm_provider: Some("azure_ai".into()),
litellm_params: Default::default(),
..streaming(call, upstream.uri())
},
usize::MAX,

View file

@ -21,14 +21,7 @@ pub(crate) fn prepare_request(
Some("MISTRAL_AZURE_API_BASE"),
),
LlmProviders::AzureAi => (None, Some("AZURE_AI_API_BASE")),
LlmProviders::Anthropic
| LlmProviders::AwsTextract
| LlmProviders::Bedrock
| LlmProviders::Cohere
| LlmProviders::Openai
| LlmProviders::OpenaiLike
| LlmProviders::Reducto
| LlmProviders::VertexAi => (None, None),
_ => (None, None),
};
let secret = |name: &str| secrets.truthy(name);
let dynamic_api_key = credentials.dynamic_api_key.or_else(|| {

View file

@ -225,10 +225,7 @@ pub(crate) fn resolve_provider_config(
OcrConfigKind::VertexDeepSeek
}
LlmProviders::VertexAi => OcrConfigKind::VertexAi,
LlmProviders::Anthropic
| LlmProviders::Bedrock
| LlmProviders::Openai
| LlmProviders::OpenaiLike => {
_ => {
return Err(Error::InvalidProvider(
provider.custom_llm_provider.to_string(),
));

View file

@ -17,15 +17,7 @@ use litellm_inference::provider::resolve_llm_provider;
fn provider_config(provider: LlmProviders) -> Option<&'static dyn BaseAudioTranscriptionConfig> {
match provider {
LlmProviders::Bedrock => Some(&BEDROCK_AUDIO_TRANSCRIPTION_CONFIG),
LlmProviders::Anthropic
| LlmProviders::AwsTextract
| LlmProviders::AzureAi
| LlmProviders::Cohere
| LlmProviders::Mistral
| LlmProviders::Openai
| LlmProviders::OpenaiLike
| LlmProviders::Reducto
| LlmProviders::VertexAi => None,
_ => None,
}
}

View file

@ -20,6 +20,7 @@ litellm-framer.workspace = true
litellm-http.workspace = true
litellm-secrets.workspace = true
litellm-python-compat.workspace = true
litellm-router-types.workspace = true
base64.workspace = true
bytes.workspace = true
data-url = "0.3.2"

View file

@ -7,6 +7,7 @@ use litellm_http::request::{
use litellm_llms_types::{
formats::messages::{
ContentBlock, ContentBlockType, EffortLevel, Message, MessageContent, MessagesTool,
SystemPrompt,
},
providers::anthropic::{AnthropicBeta, BetaSet},
recognized::Recognized,
@ -385,6 +386,33 @@ pub fn is_encrypted_reasoning_block(block: &ContentBlock) -> bool {
field.is_some_and(|value| value.starts_with(ENCRYPTED_REASONING_SIGNATURE_PREFIX))
}
const BILLING_HEADER_PREFIX: &str = "x-anthropic-billing-header:";
fn is_billing_header_block(block: &ContentBlock) -> bool {
block.block_type == Some(ContentBlockType::Text)
&& block
.text
.as_deref()
.is_some_and(|text| text.starts_with(BILLING_HEADER_PREFIX))
}
/// Python's `AnthropicMessagesConfig._filter_billing_headers_from_system`: the Claude Code
/// attribution blocks the first-party API reads, dropped for hosts that reject them. `None`
/// when nothing else was in the system prompt.
pub fn filter_billing_headers_from_system(system: SystemPrompt) -> Option<SystemPrompt> {
match system {
SystemPrompt::Text(text) if text.starts_with(BILLING_HEADER_PREFIX) => None,
SystemPrompt::Text(text) => Some(SystemPrompt::Text(text)),
SystemPrompt::Blocks(blocks) => {
let kept: Vec<ContentBlock> = blocks
.into_iter()
.filter(|block| !is_billing_header_block(block))
.collect();
(!kept.is_empty()).then_some(SystemPrompt::Blocks(kept))
}
}
}
pub fn strip_encrypted_reasoning_blocks(messages: Vec<Message>) -> Vec<Message> {
retain_blocks(messages, |block| !is_encrypted_reasoning_block(block))
}
@ -1858,6 +1886,7 @@ mod tests {
supports_output_config: false,
supports_sampling_params: true,
supports_speed: false,
supports_mid_conversation_system: false,
effort_tiers: tiers(false, false, false, false, false, false),
}
);

View file

@ -3,9 +3,10 @@ use litellm_llms_types::{
formats::messages::{
ContextEdit, ContextManagement, Message, MessagesOptionalParams, MessagesRequest, Speed,
},
providers::anthropic::{AnthropicBeta, BetaSet},
providers::anthropic::{AnthropicBeta, BetaProvider, BetaSet},
recognized::Recognized,
};
use litellm_router_types::LitellmParams;
use serde_json::{Map, Value, json};
use super::{handler::shape_anthropic_messages_request, thinking::translate_thinking};
@ -47,6 +48,8 @@ impl BaseMessagesConfig for AnthropicMessagesConfig {
&self,
api_base: Option<&str>,
_model: &str,
_litellm_params: &LitellmParams,
_stream: bool,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> Result<String, Error> {
Ok(complete_anthropic_url(api_base, env_lookup))
@ -76,6 +79,7 @@ impl BaseMessagesConfig for AnthropicMessagesConfig {
headers: Headers,
api_key: Option<&str>,
_model: &str,
_litellm_params: &LitellmParams,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> Result<ValidatedEnvironment, Error> {
let headers = match optionally_handle_anthropic_oauth(headers, api_key) {
@ -150,6 +154,15 @@ pub(crate) fn update_headers_with_anthropic_beta(
merge_beta_headers(headers, feature_betas(request))
}
/// The betas a request needs on a host other than the first-party API: the Anthropic set
/// filtered and renamed through the host's `BetaProvider` policy.
pub fn provider_feature_betas(request: &MessagesRequest, provider: BetaProvider) -> BetaSet {
feature_betas(request)
.iter()
.filter_map(|beta| beta.on(provider))
.collect()
}
fn feature_betas(request: &MessagesRequest) -> BetaSet {
let params = &request.params;
let tools = params.tools.as_deref();
@ -815,7 +828,13 @@ mod tests {
#[case] expected: &str,
) {
assert_eq!(
ANTHROPIC_MESSAGES_CONFIG.get_complete_url(api_base, "claude", &env(vars)),
ANTHROPIC_MESSAGES_CONFIG.get_complete_url(
api_base,
"claude",
&LitellmParams::default(),
false,
&env(vars),
),
Ok(expected.to_string())
);
}
@ -833,6 +852,7 @@ mod tests {
headers(forwarded),
api_key,
"claude",
&LitellmParams::default(),
&env(vars),
)
}
@ -1108,8 +1128,20 @@ mod tests {
requested.borrow_mut().push(name.to_string());
None
};
let _ = ANTHROPIC_MESSAGES_CONFIG.validate_environment(Vec::new(), None, "claude", &record);
let _ = ANTHROPIC_MESSAGES_CONFIG.get_complete_url(None, "claude", &record);
let _ = ANTHROPIC_MESSAGES_CONFIG.validate_environment(
Vec::new(),
None,
"claude",
&LitellmParams::default(),
&record,
);
let _ = ANTHROPIC_MESSAGES_CONFIG.get_complete_url(
None,
"claude",
&LitellmParams::default(),
false,
&record,
);
let requested = requested.into_inner();
assert!(!requested.is_empty());
let undeclared: Vec<&String> = requested

View file

@ -1,9 +1,7 @@
use litellm_auth::{CredentialPlacement, SecretValue};
use litellm_http::request::{has_bearer_auth, has_header};
use litellm_llms_types::formats::messages::{
CacheControl, ContentBlock, Message, MessageContent, MessagesOptionalParams, MessagesRequest,
SystemPrompt,
};
use litellm_llms_types::formats::messages::MessagesRequest;
use litellm_router_types::LitellmParams;
use crate::{
Error,
@ -20,7 +18,7 @@ use crate::{
auth::{AuthScheme, Headers, ValidatedEnvironment},
messages::{
context::MessagesTransformContext,
normalization::fold_system_role_messages,
normalization::{normalize_system_role_messages, strip_cache_control_scope},
transformation::{BaseMessagesConfig, MESSAGES_PATH_SUFFIX},
},
},
@ -47,6 +45,8 @@ impl BaseMessagesConfig for AzureAnthropicMessagesConfig {
&self,
api_base: Option<&str>,
_model: &str,
_litellm_params: &LitellmParams,
_stream: bool,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> Result<String, Error> {
complete_azure_anthropic_url(api_base, env_lookup)
@ -57,20 +57,14 @@ impl BaseMessagesConfig for AzureAnthropicMessagesConfig {
request: MessagesRequest,
context: &MessagesTransformContext,
) -> Result<MessagesRequest, Error> {
let request = fold_system_role_messages(request);
transform_messages_request(
MessagesRequest {
messages: request
.messages
.into_iter()
.map(strip_scope_from_message)
.collect(),
params: MessagesOptionalParams {
system: request.params.system.map(strip_scope_from_system),
..request.params
},
..request
},
strip_cache_control_scope(normalize_system_role_messages(
request,
context
.thinking
.capabilities
.supports_mid_conversation_system,
)),
context,
)
}
@ -86,6 +80,7 @@ impl BaseMessagesConfig for AzureAnthropicMessagesConfig {
headers: Headers,
api_key: Option<&str>,
_model: &str,
_litellm_params: &LitellmParams,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> Result<ValidatedEnvironment, Error> {
if has_header(&headers, API_KEY_PLACEMENT.header_name()) || has_bearer_auth(&headers) {
@ -129,37 +124,6 @@ pub fn complete_azure_anthropic_url(
Ok(format!("{with_anthropic}{MESSAGES_PATH_SUFFIX}"))
}
fn strip_scope_from_block(block: ContentBlock) -> ContentBlock {
ContentBlock {
cache_control: block.cache_control.map(|cache_control| CacheControl {
scope: None,
..cache_control
}),
..block
}
}
fn strip_scope_from_system(system: SystemPrompt) -> SystemPrompt {
match system {
SystemPrompt::Blocks(blocks) => {
SystemPrompt::Blocks(blocks.into_iter().map(strip_scope_from_block).collect())
}
text => text,
}
}
fn strip_scope_from_message(message: Message) -> Message {
Message {
content: match message.content {
MessageContent::Blocks(blocks) => {
MessageContent::Blocks(blocks.into_iter().map(strip_scope_from_block).collect())
}
text => text,
},
..message
}
}
#[cfg(test)]
mod tests {
use litellm_llms_types::formats::messages::MessagesResponse;
@ -169,7 +133,9 @@ mod tests {
use litellm_auth::CredentialPlacement;
use super::*;
use crate::base_llm::messages::context::MessagesModelCapabilities;
use crate::base_llm::messages::{
context::MessagesModelCapabilities, mid_conversation_system::CONVERTED_SYSTEM_NOTE,
};
fn request_from(value: serde_json::Value) -> MessagesRequest {
serde_json::from_value(value).expect("valid request")
@ -268,6 +234,7 @@ mod tests {
.collect(),
api_key,
"claude",
&LitellmParams::default(),
&|_| None,
)
.unwrap()
@ -427,62 +394,101 @@ mod tests {
assert_eq!(transformed, body);
}
#[test]
fn transform_request_folds_system_role_message_into_top_level_system() {
let request = request_from(json!({
"model": "claude-sonnet-4-5",
"max_tokens": 256,
"system": [{"type": "text", "text": "base system"}],
"messages": [
{"role": "user", "content": "fix the bug"},
{"role": "system", "content": "Available agent types: claude"}
]
}));
let transformed = to_value(
AZURE_ANTHROPIC_MESSAGES_CONFIG
.transform_anthropic_messages_request(request, &MessagesTransformContext::default())
.expect("request transforms"),
);
assert_eq!(
transformed["messages"],
json!([{"role": "user", "content": "fix the bug"}])
);
assert_eq!(
transformed["system"],
json!([
{"type": "text", "text": "base system"},
{"type": "text", "text": "Available agent types: claude"}
])
);
fn context_with(supports_mid_conversation_system: bool) -> MessagesTransformContext {
MessagesTransformContext::new(
MessagesModelCapabilities {
supports_mid_conversation_system,
..MessagesModelCapabilities::default()
},
false,
)
}
#[test]
fn transform_request_folds_system_role_when_no_top_level_system() {
let request = request_from(json!({
"model": "claude-sonnet-4-5",
"max_tokens": 256,
#[rstest]
#[case::leading_turn_joins_the_top_level_system(
json!({
"system": [{"type": "text", "text": "base system"}],
"messages": [
{"role": "user", "content": [{"type": "text", "text": "hi"}]},
{"role": "system", "content": [{"type": "text", "text": "sys block"}]}
{"role": "system", "content": "Available agent types: claude"},
{"role": "user", "content": "fix the bug"}
]
}));
}),
false,
json!([{"type": "text", "text": "base system"}, {"type": "text", "text": "Available agent types: claude"}]),
json!([{"role": "user", "content": "fix the bug"}]),
)]
#[case::leading_turn_becomes_the_system_when_there_is_none(
json!({"messages": [
{"role": "system", "content": [{"type": "text", "text": "sys block"}]},
{"role": "user", "content": [{"type": "text", "text": "hi"}]}
]}),
false,
json!([{"type": "text", "text": "sys block"}]),
json!([{"role": "user", "content": [{"type": "text", "text": "hi"}]}]),
)]
#[case::billing_block_in_a_leading_turn_is_dropped(
json!({"messages": [
{"role": "system", "content": [{"type": "text", "text": "x-anthropic-billing-header: cc_version=1"}, {"type": "text", "text": "keep"}]},
{"role": "user", "content": "hi"}
]}),
false,
json!([{"type": "text", "text": "keep"}]),
json!([{"role": "user", "content": "hi"}]),
)]
#[case::later_turn_is_converted_in_place_for_a_model_without_the_capability(
json!({
"system": "be terse",
"messages": [
{"role": "user", "content": "fix the bug"},
{"role": "system", "content": "reminder"}
]
}),
false,
json!("be terse"),
json!([
{"role": "user", "content": "fix the bug"},
{"role": "user", "content": [
{"type": "text", "text": CONVERTED_SYSTEM_NOTE},
{"type": "text", "text": "reminder"}
]}
]),
)]
#[case::later_turn_stays_for_a_model_with_the_capability(
json!({
"system": "be terse",
"messages": [
{"role": "user", "content": "fix the bug"},
{"role": "system", "content": "reminder"}
]
}),
true,
json!("be terse"),
json!([
{"role": "user", "content": "fix the bug"},
{"role": "system", "content": "reminder"}
]),
)]
fn transform_request_normalizes_system_role_messages_like_python(
#[case] body: serde_json::Value,
#[case] supports_mid_conversation_system: bool,
#[case] system: serde_json::Value,
#[case] messages: serde_json::Value,
) {
let mut body = body;
body["model"] = json!("claude-sonnet-4-5");
body["max_tokens"] = json!(256);
let transformed = to_value(
AZURE_ANTHROPIC_MESSAGES_CONFIG
.transform_anthropic_messages_request(request, &MessagesTransformContext::default())
.transform_anthropic_messages_request(
request_from(body),
&context_with(supports_mid_conversation_system),
)
.expect("request transforms"),
);
assert_eq!(
transformed["messages"],
json!([{"role": "user", "content": [{"type": "text", "text": "hi"}]}])
);
assert_eq!(
transformed["system"],
json!([{"type": "text", "text": "sys block"}])
);
assert_eq!(transformed["system"], system);
assert_eq!(transformed["messages"], messages);
}
#[test]
@ -607,9 +613,16 @@ mod tests {
Vec::new(),
None,
"claude",
&LitellmParams::default(),
&record,
);
let _ = AZURE_ANTHROPIC_MESSAGES_CONFIG.get_complete_url(
None,
"claude",
&LitellmParams::default(),
false,
&record,
);
let _ = AZURE_ANTHROPIC_MESSAGES_CONFIG.get_complete_url(None, "claude", &record);
let requested = requested.into_inner();
assert!(!requested.is_empty());
let undeclared: Vec<&String> = requested

View file

@ -7,6 +7,7 @@
use litellm_auth::{AuthServices, CredentialPlacement, SecretValue, TokenProviderHandle};
use litellm_auth_aws::{AwsCredentialSource, SigV4Signer};
use litellm_auth_gcp::VertexConfig;
use litellm_http::request::with_header;
pub type Headers = Vec<(String, String)>;
@ -24,6 +25,9 @@ pub enum AuthScheme {
/// A bearer acquired when the request is sent, from a token source such as a cloud SDK
/// or a caller-supplied callable.
Token { provider: TokenProviderHandle },
/// A Google access token for the Vertex AI project the config names, acquired when the
/// request is sent through the shared GCP token cache.
GcpAccessToken { config: Box<VertexConfig> },
/// AWS SigV4 over the bytes that go on the wire, so the handler signs after the body is
/// serialized.
AwsSigV4 {
@ -73,6 +77,13 @@ pub async fn resolve_auth(
signer: None,
})
}
AuthScheme::GcpAccessToken { config } => {
let token = services.gcp.access_token(&config, env_lookup).await?;
Ok(Authenticated {
headers: with_credential(headers, CredentialPlacement::Bearer, &token),
signer: None,
})
}
AuthScheme::AwsSigV4 {
region,
service,

View file

@ -40,6 +40,8 @@ pub struct MessagesModelCapabilities {
#[serde(default)]
pub supports_speed: bool,
#[serde(default)]
pub supports_mid_conversation_system: bool,
#[serde(default)]
pub effort_tiers: SupportedEffortTiers,
}
@ -57,6 +59,7 @@ impl Default for MessagesModelCapabilities {
supports_output_config: false,
supports_sampling_params: true,
supports_speed: false,
supports_mid_conversation_system: false,
effort_tiers: SupportedEffortTiers::default(),
}
}

View file

@ -0,0 +1,201 @@
use litellm_llms_types::formats::messages::{
ContentBlock, ContentBlockType, Message, MessageContent, SystemPrompt,
};
use serde_json::Map;
/// Python's `CONVERTED_SYSTEM_NOTE`.
pub const CONVERTED_SYSTEM_NOTE: &str = "Operator note (not from the user): the following was originally a mid-conversation system-role reminder.";
const SYSTEM_ROLE: &str = "system";
const USER_ROLE: &str = "user";
pub fn as_system_content_blocks(value: Option<SystemPrompt>) -> Vec<ContentBlock> {
match value {
None => Vec::new(),
Some(SystemPrompt::Text(text)) => vec![ContentBlock::text(text)],
Some(SystemPrompt::Blocks(blocks)) => blocks,
}
}
pub fn message_content_blocks(content: MessageContent) -> Vec<ContentBlock> {
match content {
MessageContent::Text(text) => vec![ContentBlock::text(text)],
MessageContent::Blocks(blocks) => blocks,
}
}
pub fn is_system_role_message(message: &Message) -> bool {
message.role == SYSTEM_ROLE
}
pub fn system_role_message_as_user(message: Message) -> Message {
Message {
role: USER_ROLE.into(),
content: MessageContent::Blocks(
[ContentBlock::text(CONVERTED_SYSTEM_NOTE)]
.into_iter()
.chain(message_content_blocks(message.content))
.collect(),
),
extra: Map::new(),
}
}
pub fn opens_with_tool_results(message: &Message) -> bool {
message.role == USER_ROLE
&& matches!(
&message.content,
MessageContent::Blocks(blocks)
if blocks.first().and_then(|block| block.block_type.as_ref()) == Some(&ContentBlockType::ToolResult)
)
}
fn system_run_placed_after_tool_results(
system_run: Vec<Message>,
follower_run: Vec<Message>,
) -> Vec<Message> {
let mut follower_run = follower_run.into_iter();
match follower_run.next() {
Some(first) if opens_with_tool_results(&first) => [first]
.into_iter()
.chain(system_run)
.chain(follower_run)
.collect(),
first => system_run
.into_iter()
.chain(first)
.chain(follower_run)
.collect(),
}
}
fn runs(messages: Vec<Message>) -> Vec<Vec<Message>> {
messages
.into_iter()
.fold(Vec::new(), |mut runs: Vec<Vec<Message>>, message| {
match runs.last_mut() {
Some(run)
if is_system_role_message(&run[0]) == is_system_role_message(&message) =>
{
run.push(message);
}
_ => runs.push(vec![message]),
}
runs
})
}
fn system_turns_after_tool_results(messages: Vec<Message>) -> Vec<Message> {
let mut runs = runs(messages).into_iter();
let mut placed = Vec::new();
while let Some(run) = runs.next() {
if !is_system_role_message(&run[0]) {
placed.extend(run);
continue;
}
let follower_run = runs.next().unwrap_or_default();
placed.extend(system_run_placed_after_tool_results(run, follower_run));
}
placed
}
pub fn convert_mid_conversation_system_turns(messages: Vec<Message>) -> Vec<Message> {
system_turns_after_tool_results(messages)
.into_iter()
.map(|message| {
if is_system_role_message(&message) {
system_role_message_as_user(message)
} else {
message
}
})
.collect()
}
#[cfg(test)]
mod tests {
use rstest::rstest;
use serde_json::{Value, json};
use super::*;
fn messages(value: Value) -> Vec<Message> {
serde_json::from_value(value).unwrap()
}
fn roles(messages: &[Message]) -> Vec<&str> {
messages
.iter()
.map(|message| message.role.as_str())
.collect()
}
#[rstest]
#[case::no_system_turns(
json!([{"role": "user", "content": "a"}, {"role": "assistant", "content": "b"}]),
vec!["user", "assistant"],
)]
#[case::a_turn_between_messages_becomes_a_user_turn_in_place(
json!([
{"role": "user", "content": "a"},
{"role": "assistant", "content": "b"},
{"role": "system", "content": "reminder"},
{"role": "user", "content": "c"}
]),
vec!["user", "assistant", "user", "user"],
)]
#[case::a_run_between_a_tool_use_and_its_result_moves_after_the_result(
json!([
{"role": "user", "content": "a"},
{"role": "assistant", "content": [{"type": "tool_use", "id": "t", "name": "n", "input": {}}]},
{"role": "system", "content": "reminder"},
{"role": "user", "content": [{"type": "tool_result", "tool_use_id": "t", "content": "ok"}, {"type": "text", "text": "next"}]},
{"role": "user", "content": "d"}
]),
vec!["user", "assistant", "user", "user", "user"],
)]
fn turns_are_converted_in_place_and_never_split_a_tool_call_from_its_result(
#[case] input: Value,
#[case] expected_roles: Vec<&str>,
) {
let converted = convert_mid_conversation_system_turns(messages(input));
assert_eq!(roles(&converted), expected_roles);
assert!(
converted
.iter()
.all(|message| !is_system_role_message(message))
);
}
#[rstest]
fn a_converted_turn_carries_the_note_then_the_original_content_only() {
let converted = convert_mid_conversation_system_turns(messages(json!([
{"role": "user", "content": "a"},
{"role": "system", "content": "reminder", "name": "ignored"}
])));
assert_eq!(
serde_json::to_value(&converted[1]).unwrap(),
json!({"role": "user", "content": [
{"type": "text", "text": CONVERTED_SYSTEM_NOTE},
{"type": "text", "text": "reminder"}
]})
);
}
#[rstest]
fn the_tool_result_turn_keeps_its_place_before_the_moved_run() {
let converted = convert_mid_conversation_system_turns(messages(json!([
{"role": "assistant", "content": [{"type": "tool_use", "id": "t", "name": "n", "input": {}}]},
{"role": "system", "content": "reminder"},
{"role": "user", "content": [{"type": "tool_result", "tool_use_id": "t", "content": "ok"}]}
])));
assert!(opens_with_tool_results(&converted[1]));
assert_eq!(
serde_json::to_value(&converted[2]).unwrap()["content"][1],
json!({"type": "text", "text": "reminder"})
);
}
}

View file

@ -1,4 +1,5 @@
pub mod context;
pub mod mid_conversation_system;
pub mod normalization;
pub mod streaming;
pub mod transformation;

View file

@ -1,47 +1,103 @@
use litellm_llms_types::formats::messages::{
ContentBlock, Message, MessageContent, MessagesOptionalParams, MessagesRequest, SystemPrompt,
CacheControl, ContentBlock, Message, MessageContent, MessagesOptionalParams, MessagesRequest,
SystemPrompt,
};
const SYSTEM_ROLE: &str = "system";
use crate::{
anthropic::common_utils::filter_billing_headers_from_system,
base_llm::messages::mid_conversation_system::{
as_system_content_blocks, convert_mid_conversation_system_turns, is_system_role_message,
message_content_blocks,
},
};
fn content_into_blocks(content: MessageContent) -> Vec<ContentBlock> {
match content {
MessageContent::Text(text) => vec![ContentBlock::text(text)],
MessageContent::Blocks(blocks) => blocks,
}
}
fn system_into_blocks(system: Option<SystemPrompt>) -> Vec<ContentBlock> {
match system {
None => Vec::new(),
Some(SystemPrompt::Text(text)) => vec![ContentBlock::text(text)],
Some(SystemPrompt::Blocks(blocks)) => blocks,
}
}
pub fn fold_system_role_messages(request: MessagesRequest) -> MessagesRequest {
if !request.messages.iter().any(|msg| msg.role == SYSTEM_ROLE) {
return request;
}
let (system_messages, chat_messages): (Vec<Message>, Vec<Message>) = request
/// Python's `_normalize_system_role_messages`: a leading run of `role: "system"` entries is
/// hoisted into the top-level `system`, a later one stays in place when the model accepts it
/// and otherwise becomes a user turn where it was, and billing blocks are stripped from the
/// top-level `system` either way.
pub fn normalize_system_role_messages(
request: MessagesRequest,
supports_mid_conversation_system: bool,
) -> MessagesRequest {
let leading_count = request
.messages
.into_iter()
.partition(|msg| msg.role == SYSTEM_ROLE);
let folded_system: Vec<ContentBlock> = system_into_blocks(request.params.system)
.into_iter()
.chain(
system_messages
.iter()
.take_while(|message| is_system_role_message(message))
.count();
let mut messages = request.messages.into_iter();
let hoisted: Vec<Message> = messages.by_ref().take(leading_count).collect();
let remaining: Vec<Message> = messages.collect();
let remaining = if supports_mid_conversation_system {
remaining
} else {
convert_mid_conversation_system_turns(remaining)
};
let system = if hoisted.is_empty() {
request.params.system
} else {
Some(SystemPrompt::Blocks(
as_system_content_blocks(request.params.system)
.into_iter()
.flat_map(|msg| content_into_blocks(msg.content)),
)
.collect();
.chain(
hoisted
.into_iter()
.flat_map(|message| message_content_blocks(message.content)),
)
.collect(),
))
};
MessagesRequest {
messages: chat_messages,
messages: remaining,
params: MessagesOptionalParams {
system: (!folded_system.is_empty()).then_some(SystemPrompt::Blocks(folded_system)),
system: system.and_then(filter_billing_headers_from_system),
..request.params
},
..request
}
}
fn strip_scope_from_block(block: ContentBlock) -> ContentBlock {
ContentBlock {
cache_control: block.cache_control.map(|cache_control| CacheControl {
scope: None,
..cache_control
}),
..block
}
}
fn strip_scope_from_system(system: SystemPrompt) -> SystemPrompt {
match system {
SystemPrompt::Blocks(blocks) => {
SystemPrompt::Blocks(blocks.into_iter().map(strip_scope_from_block).collect())
}
text => text,
}
}
fn strip_scope_from_message(message: Message) -> Message {
Message {
content: match message.content {
MessageContent::Blocks(blocks) => {
MessageContent::Blocks(blocks.into_iter().map(strip_scope_from_block).collect())
}
text => text,
},
..message
}
}
/// Python's `_remove_scope_from_cache_control`: hosts other than the first-party API reject
/// `cache_control.scope`, so it is dropped from the system prompt and every message block.
pub fn strip_cache_control_scope(request: MessagesRequest) -> MessagesRequest {
MessagesRequest {
messages: request
.messages
.into_iter()
.map(strip_scope_from_message)
.collect(),
params: MessagesOptionalParams {
system: request.params.system.map(strip_scope_from_system),
..request.params
},
..request

View file

@ -1,4 +1,6 @@
use litellm_llms_types::formats::messages::{MessagesRequest, MessagesResponse};
use litellm_router_types::LitellmParams;
use serde_json::Value;
use super::context::MessagesTransformContext;
@ -20,18 +22,11 @@ pub trait BaseMessagesConfig: Sync {
&self,
api_base: Option<&str>,
model: &str,
litellm_params: &LitellmParams,
stream: bool,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> Result<String, Error>;
fn complete_stream_url(
&self,
api_base: Option<&str>,
model: &str,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> Result<String, Error> {
self.get_complete_url(api_base, model, env_lookup)
}
fn transform_anthropic_messages_request(
&self,
request: MessagesRequest,
@ -58,6 +53,7 @@ pub trait BaseMessagesConfig: Sync {
headers: Headers,
api_key: Option<&str>,
model: &str,
litellm_params: &LitellmParams,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> Result<ValidatedEnvironment, Error>;
@ -75,6 +71,12 @@ pub trait BaseMessagesConfig: Sync {
fn request_headers(&self, headers: Headers, _request: &MessagesRequest) -> Headers {
headers
}
/// The JSON that goes on the wire, for a host whose body differs from the typed request:
/// Python's configs `pop("model")` when the model is addressed by the URL.
fn wire_body(&self, body: Value) -> Value {
body
}
}
#[cfg(test)]
@ -94,6 +96,8 @@ mod tests {
&self,
_api_base: Option<&str>,
_model: &str,
_litellm_params: &LitellmParams,
_stream: bool,
_env_lookup: &dyn Fn(&str) -> Option<String>,
) -> Result<String, Error> {
Ok(String::new())
@ -104,6 +108,7 @@ mod tests {
headers: Headers,
_api_key: Option<&str>,
_model: &str,
_litellm_params: &LitellmParams,
_env_lookup: &dyn Fn(&str) -> Option<String>,
) -> Result<ValidatedEnvironment, Error> {
Ok(ValidatedEnvironment {

View file

@ -19,6 +19,7 @@ use litellm_llms_types::formats::messages::{
MessagesRequest,
streaming::{MessagesStreamEvent, MessagesStreamUsage},
};
use litellm_router_types::LitellmParams;
use serde_json::{Map, Value};
use crate::{
@ -74,16 +75,26 @@ fn bearer_token(
fn invoke_url(
api_base: Option<&str>,
model: &str,
aws: &AwsParams,
stream: bool,
env_lookup: &dyn Fn(&str) -> Option<String>,
path: &str,
) -> String {
let (model_id, model_region) =
bedrock_model_id_and_region(model.strip_prefix(INVOKE_MODEL_PREFIX).unwrap_or(model));
let region = resolve_bedrock_region(model_region.as_deref(), &AwsParams::default(), env_lookup);
let endpoint = api_base
.map(str::trim)
.filter(|value| !value.is_empty())
.map(str::to_string)
let region = resolve_bedrock_region(model_region.as_deref(), aws, env_lookup);
let path = if stream {
INVOKE_STREAM_PATH
} else {
INVOKE_PATH
};
let configured = |value: Option<&str>| {
value
.map(str::trim)
.filter(|value| !value.is_empty())
.map(str::to_string)
};
let endpoint = configured(api_base)
.or_else(|| configured(aws.aws_bedrock_runtime_endpoint.as_deref()))
.or_else(|| env_lookup(AWS_BEDROCK_RUNTIME_ENDPOINT))
.unwrap_or_else(|| BEDROCK_RUNTIME_ENDPOINT_TEMPLATE.replace("{region}", &region));
format!("{}/model/{model_id}/{path}", endpoint.trim_end_matches('/'))
@ -102,18 +113,17 @@ impl BaseMessagesConfig for AmazonAnthropicClaudeMessagesConfig {
&self,
api_base: Option<&str>,
model: &str,
litellm_params: &LitellmParams,
stream: bool,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> Result<String, Error> {
Ok(invoke_url(api_base, model, env_lookup, INVOKE_PATH))
}
fn complete_stream_url(
&self,
api_base: Option<&str>,
model: &str,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> Result<String, Error> {
Ok(invoke_url(api_base, model, env_lookup, INVOKE_STREAM_PATH))
Ok(invoke_url(
api_base,
model,
&litellm_params.aws,
stream,
env_lookup,
))
}
fn transform_anthropic_messages_request(
@ -137,6 +147,7 @@ impl BaseMessagesConfig for AmazonAnthropicClaudeMessagesConfig {
headers: Headers,
api_key: Option<&str>,
model: &str,
litellm_params: &LitellmParams,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> Result<ValidatedEnvironment, Error> {
if let Some(token) = bearer_token(api_key, env_lookup) {
@ -150,13 +161,13 @@ impl BaseMessagesConfig for AmazonAnthropicClaudeMessagesConfig {
}
let (_, model_region) =
bedrock_model_id_and_region(model.strip_prefix(INVOKE_MODEL_PREFIX).unwrap_or(model));
let params = AwsParams::default();
let params = &litellm_params.aws;
Ok(ValidatedEnvironment {
headers,
auth: AuthScheme::AwsSigV4 {
region: resolve_bedrock_region(model_region.as_deref(), &params, env_lookup),
region: resolve_bedrock_region(model_region.as_deref(), params, env_lookup),
service: BEDROCK_SERVICE,
credentials: Box::new(AwsCredentialSource::from_params(&params, env_lookup)),
credentials: Box::new(AwsCredentialSource::from_params(params, env_lookup)),
},
})
}
@ -503,19 +514,129 @@ mod tests {
assert_eq!(sse, expected);
}
#[test]
fn config_uses_the_streaming_url_only_for_streams() {
let env = |_: &str| -> Option<String> { None };
let config = AmazonAnthropicClaudeMessagesConfig;
fn in_region(region: &str) -> LitellmParams {
LitellmParams {
aws: AwsParams {
aws_region_name: Some(region.into()),
..AwsParams::default()
},
..LitellmParams::default()
}
}
fn with_runtime_endpoint(endpoint: &str) -> LitellmParams {
LitellmParams {
aws: AwsParams {
aws_bedrock_runtime_endpoint: Some(endpoint.into()),
..in_region("us-west-2").aws
},
..LitellmParams::default()
}
}
#[rstest]
#[case::invoke_in_the_params_region(
None,
in_region("us-west-2"),
false,
None,
"https://bedrock-runtime.us-west-2.amazonaws.com/model/anthropic.claude-3/invoke"
)]
#[case::stream_in_the_params_region(
None,
in_region("us-west-2"),
true,
None,
"https://bedrock-runtime.us-west-2.amazonaws.com/model/anthropic.claude-3/invoke-with-response-stream"
)]
#[case::api_base_outranks_the_params_endpoint(
Some("https://base.test/"),
with_runtime_endpoint("https://params.test"),
false,
None,
"https://base.test/model/anthropic.claude-3/invoke"
)]
#[case::a_blank_api_base_does_not_hide_the_params_endpoint(
Some(" "),
with_runtime_endpoint("https://params.test"),
false,
Some("https://env.test"),
"https://params.test/model/anthropic.claude-3/invoke"
)]
#[case::params_endpoint_outranks_the_environment(
None,
with_runtime_endpoint("https://params.test"),
false,
Some("https://env.test"),
"https://params.test/model/anthropic.claude-3/invoke"
)]
#[case::environment_outranks_the_region_template(
None,
in_region("us-west-2"),
false,
Some("https://env.test"),
"https://env.test/model/anthropic.claude-3/invoke"
)]
fn url_follows_python_endpoint_precedence_and_the_stream_path(
#[case] api_base: Option<&str>,
#[case] litellm_params: LitellmParams,
#[case] stream: bool,
#[case] env_endpoint: Option<&str>,
#[case] expected: &str,
) {
let env = |name: &str| {
(name == AWS_BEDROCK_RUNTIME_ENDPOINT)
.then_some(env_endpoint)
.flatten()
.map(str::to_string)
};
assert_eq!(
config
.get_complete_url(None, "anthropic.claude-3", &env)
AmazonAnthropicClaudeMessagesConfig
.get_complete_url(
api_base,
"anthropic.claude-3",
&litellm_params,
stream,
&env
)
.unwrap(),
config
.complete_stream_url(None, "anthropic.claude-3", &env)
.unwrap()
.replace(INVOKE_STREAM_PATH, INVOKE_PATH)
expected
);
}
#[test]
fn sigv4_scope_and_credentials_come_from_the_litellm_params() {
let litellm_params = LitellmParams {
aws: AwsParams {
aws_access_key_id: Some("AKIAPARAMS".into()),
aws_secret_access_key: Some("params-secret".into()),
..in_region("eu-central-1").aws
},
..LitellmParams::default()
};
let validated = AmazonAnthropicClaudeMessagesConfig
.validate_environment(
Vec::new(),
None,
"anthropic.claude-3",
&litellm_params,
&|_| None,
)
.unwrap();
let AuthScheme::AwsSigV4 {
region,
credentials,
..
} = validated.auth
else {
panic!("expected SigV4, got {:?}", validated.auth);
};
let AwsCredentialSource::HostSupplied(credentials) = *credentials else {
panic!("expected the params' static keys, got {credentials:?}");
};
assert_eq!(
(region.as_str(), credentials.access_key_id()),
("eu-central-1", "AKIAPARAMS")
);
}
@ -538,6 +659,7 @@ mod tests {
vec![("authorization".into(), "Bearer forwarded".into())],
api_key,
"anthropic.claude-3",
&LitellmParams::default(),
&env,
)
.unwrap();

View file

@ -1,4 +1,4 @@
use litellm_llms::base_llm::messages::normalization::fold_system_role_messages;
use litellm_llms::base_llm::messages::normalization::normalize_system_role_messages;
use litellm_llms_types::formats::messages::MessagesRequest;
use rstest::rstest;
use serde_json::{Value, json};
@ -7,42 +7,47 @@ use serde_json::{Value, json};
#[case::text_system(json!("existing"), json!([{"type": "text", "text": "existing"}]))]
#[case::block_system(json!([{ "type": "future", "payload": 7 }]), json!([{ "type": "future", "payload": 7 }]))]
#[case::no_system(Value::Null, json!([]))]
fn folding_preserves_block_fields_order_and_unrelated_request_fields(
fn hoisting_preserves_block_fields_order_and_unrelated_request_fields(
#[case] system: Value,
#[case] initial_blocks: Value,
) {
let cache_control = json!({"type": "ephemeral", "scope": "global", "future": true});
let folded_block = json!({"type": "text", "text": "second", "cache_control": cache_control});
let leading_block = json!({"type": "text", "text": "second", "cache_control": cache_control});
let user = json!({"role": "user", "content": "hello", "future_message": 42});
let later_turn = json!({"role": "system", "content": "mid-conversation reminder"});
let request: MessagesRequest = serde_json::from_value(json!({
"model": "test-model",
"max_tokens": 64,
"system": system,
"messages": [
{"role": "system", "content": "first"},
{"role": "system", "content": [leading_block]},
user,
{"role": "system", "content": [folded_block]}
later_turn
],
"future_request": {"nested": true}
}))
.unwrap();
let folded = fold_system_role_messages(request);
let normalized = normalize_system_role_messages(request, true);
let expected_blocks: Vec<Value> = initial_blocks
.as_array()
.unwrap()
.iter()
.cloned()
.chain([json!({"type": "text", "text": "first"}), folded_block])
.chain([json!({"type": "text", "text": "first"}), leading_block])
.collect();
assert_eq!(
serde_json::to_value(&folded).unwrap(),
serde_json::to_value(&normalized).unwrap(),
json!({
"model": "test-model",
"max_tokens": 64,
"system": expected_blocks,
"messages": [user],
"messages": [user, later_turn],
"future_request": {"nested": true}
})
);
assert_eq!(fold_system_role_messages(folded.clone()), folded);
assert_eq!(
normalize_system_role_messages(normalized.clone(), true),
normalized
);
}

View file

@ -46,6 +46,7 @@ litellm-llms.workspace = true
litellm-secrets = { workspace = true, features = ["aws", "azure", "google", "hashicorp", "cyberark"] }
litellm-secrets-types.workspace = true
litellm-llms-types.workspace = true
litellm-router-types.workspace = true
litellm-host-python.workspace = true
litellm-token-counter = { path = "../token-counter", default-features = false }
pyo3.workspace = true

View file

@ -60,7 +60,7 @@ impl PythonSettings {
}
}
fn missing_module(py: Python<'_>, error: &PyErr, expected: &str) -> PyResult<bool> {
pub(crate) fn missing_module(py: Python<'_>, error: &PyErr, expected: &str) -> PyResult<bool> {
if !error.is_instance_of::<PyModuleNotFoundError>(py) {
return Ok(false);
}

View file

@ -5,11 +5,12 @@ use bytes::Bytes;
use litellm_host_python::{InvokeError, PythonBinding, from_py, present, to_py};
use litellm_http::transport::Error as TransportError;
use litellm_inference_messages::{
Error, MessagesCall, MessagesSettings, MessagesShaping, messages_body,
Error, MessagesCall, MessagesSettings, MessagesShaping, litellm_params, messages_body,
route::{Messages, MessagesStreamHead},
};
use litellm_llms::base_llm::messages::context::MessagesModelCapabilities;
use litellm_llms_types::headers::ProviderSpecificHeaders;
use litellm_router_types::LitellmParams;
use pyo3::{
exceptions::{PyException, PyValueError},
gc::{PyTraverseError, PyVisit},
@ -21,6 +22,7 @@ use serde_json::{Map, Value};
use crate::{
errors::{RustUpstreamError, route_error_to_pyerr},
marshal::{optional_timeout, project_optional_fields, public_response, python_timeout_seconds},
python_settings::missing_module,
};
const ROUTE_HOST_MODULE: &str = "litellm.rust_bridge.messages.route_host";
@ -63,6 +65,33 @@ fn merge_headers(
(!merged.is_empty()).then_some(merged)
}
/// Reads `litellm.<name>`, the module global a param spec names, the way
/// `VertexBase.get_vertex_ai_project` reads `litellm.vertex_project`.
fn module_global<'py>(py: Python<'py>, name: &str) -> PyResult<Option<Bound<'py, PyAny>>> {
let module = match py.import("litellm") {
Ok(module) => module,
Err(error) if missing_module(py, &error, "litellm")? => return Ok(None),
Err(error) => return Err(error),
};
let value = module.getattr(name)?;
Ok((!value.is_none()).then_some(value))
}
/// The caller's litellm params, read from the kwargs by the names the typed params declare,
/// so a key the configs do not read is never converted; a spec whose spellings the call all
/// leaves out falls back to the module global it names, which stays the host's concern.
fn project_litellm_params<'py>(
argument: impl Fn(&str) -> PyResult<Option<Bound<'py, PyAny>>>,
global: impl Fn(&str) -> PyResult<Option<Bound<'py, PyAny>>>,
) -> PyResult<Result<LitellmParams, Error>> {
let fields = project_optional_fields(LitellmParams::fields(), &argument)?;
let globals = LitellmParams::specs()
.filter(|spec| !spec.wire.iter().any(|name| fields.contains_key(*name)))
.filter_map(|spec| spec.module_global);
let folded = project_optional_fields(globals, &global)?;
Ok(litellm_params(fields.into_iter().chain(folded).collect()))
}
fn native_error(py: Python<'_>, error: Error) -> PyResult<PyErr> {
match error {
Error::Transport(TransportError::Http { status, body }) => {
@ -132,15 +161,19 @@ impl MessagesPythonHost {
let api_base = string("api_base")?;
let extra_headers = self.merged_headers(py, arguments)?;
let provider_specific_header = self.provider_specific_header(py, arguments)?;
Ok(messages_body(body).map(|body| MessagesCall {
body,
api_key,
api_base,
extra_headers,
provider_specific_header,
custom_llm_provider,
timeout: optional_timeout(timeout),
shaping,
let litellm_params = project_litellm_params(argument, |name| module_global(py, name))?;
Ok(messages_body(body).and_then(|body| {
Ok(MessagesCall {
body,
api_key,
api_base,
extra_headers,
provider_specific_header,
custom_llm_provider,
litellm_params: litellm_params?,
timeout: optional_timeout(timeout),
shaping,
})
}))
}
@ -351,6 +384,106 @@ mod tests {
);
}
#[rstest]
#[case::model_and_credentials_are_projected(
json!({"model": "m", "api_key": "k", "api_base": "b", "api_version": "v", "timeout": 5, "messages": []}),
Ok(json!({"model": "m", "api_key": "k", "api_base": "b", "api_version": "v"})),
)]
#[case::aws_keys_are_projected(
json!({"model": "m", "aws_region_name": "eu-central-1", "aws_access_key_id": "AKIA", "messages": [{"role": "user"}]}),
Ok(json!({"model": "m", "aws_region_name": "eu-central-1", "aws_access_key_id": "AKIA"})),
)]
#[case::vertex_keys_are_projected_in_both_spellings(
json!({"model": "m", "vertex_project": "p", "vertex_ai_location": "us-east5", "vertex_credentials": {"type": "service_account"}}),
Ok(json!({"model": "m", "vertex_project": "p", "vertex_ai_location": "us-east5", "vertex_credentials": {"type": "service_account"}})),
)]
#[case::an_explicit_none_is_absent(json!({"model": "m", "aws_region_name": null}), Ok(json!({"model": "m"})))]
#[case::a_wrong_type_is_a_request_error(json!({"model": "m", "aws_region_name": 7}), Err(()))]
#[case::a_missing_model_is_a_request_error(json!({"aws_region_name": "eu-central-1"}), Err(()))]
fn litellm_params_are_projected_by_their_declared_names(
#[case] kwargs: Value,
#[case] expected: Result<Value, ()>,
) {
assert_projection(kwargs, json!({}), expected);
}
#[rstest]
#[case::global_fills_a_missing_vertex_project(
json!({"model": "m"}),
json!({"vertex_project": "from-global", "vertex_location": "us-east5"}),
Ok(json!({"model": "m", "vertex_project": "from-global", "vertex_location": "us-east5"})),
)]
#[case::the_call_wins_over_the_global(
json!({"model": "m", "vertex_project": "from-call"}),
json!({"vertex_project": "from-global"}),
Ok(json!({"model": "m", "vertex_project": "from-call"})),
)]
#[case::an_explicit_none_in_the_call_still_falls_back(
json!({"model": "m", "vertex_project": null}),
json!({"vertex_project": "from-global"}),
Ok(json!({"model": "m", "vertex_project": "from-global"})),
)]
#[case::a_global_of_the_wrong_type_is_a_request_error(
json!({"model": "m"}),
json!({"vertex_location": 5}),
Err(()),
)]
fn module_globals_fill_the_litellm_params_the_call_leaves_out(
#[case] kwargs: Value,
#[case] globals: Value,
#[case] expected: Result<Value, ()>,
) {
assert_projection(kwargs, globals, expected);
}
#[rstest]
#[case::a_global_without_a_spec_is_never_read(
json!({"model": "m"}),
json!({"aws_region_name": "eu-central-1", "api_key": "k"}),
Ok(json!({"model": "m"})),
)]
#[case::only_the_spec_named_global_is_read(
json!({"model": "m"}),
json!({"vertex_ai_project": "legacy-global", "vertex_project": "from-global"}),
Ok(json!({"model": "m", "vertex_project": "from-global"})),
)]
fn only_the_globals_the_specs_name_are_consulted(
#[case] kwargs: Value,
#[case] globals: Value,
#[case] expected: Result<Value, ()>,
) {
assert_projection(kwargs, globals, expected);
}
fn assert_projection(kwargs: Value, globals: Value, expected: Result<Value, ()>) {
Python::initialize();
Python::attach(|py| {
let dict = |value: &Value| {
to_py(py, value)
.unwrap()
.into_bound(py)
.cast_into::<PyDict>()
.unwrap()
};
let kwargs = dict(&kwargs);
let globals = dict(&globals);
let bound = PyDict::new(py);
let projected = project_litellm_params(
|name| present(&kwargs, &bound, name),
|name| present(&globals, &bound, name),
)
.unwrap();
match (projected, expected) {
(Ok(projected), Ok(fields)) => assert_eq!(
projected,
litellm_params(serde_json::from_value(fields).unwrap()).unwrap()
),
(Err(error), Err(())) => assert!(matches!(error, Error::InvalidRequest(_))),
(projected, expected) => panic!("got {projected:?}, expected {expected:?}"),
}
});
}
#[rstest]
#[case::rejected_request(Error::InvalidRequest("does not support top_k=5".into()), true)]
#[case::missing_field(Error::MissingField("max_tokens"), true)]

View file

@ -9,6 +9,7 @@ repository.workspace = true
litellm-config.workspace = true
litellm-inference.workspace = true
litellm-inference-messages.workspace = true
litellm-router-types.workspace = true
[dev-dependencies]
rstest.workspace = true

View file

@ -1,6 +1,7 @@
use std::time::Duration;
use litellm_inference_messages::MessagesShaping;
use litellm_router_types::LitellmParams;
#[derive(Clone, Debug, Default)]
pub struct Deployment {
@ -8,6 +9,7 @@ pub struct Deployment {
pub api_key: Option<String>,
pub api_base: Option<String>,
pub custom_llm_provider: Option<String>,
pub litellm_params: LitellmParams,
pub timeout: Option<Duration>,
pub shaping: MessagesShaping,
}

View file

@ -25,6 +25,7 @@ impl Router {
.map(|value| value.expose().to_string()),
api_base: model.litellm_params.api_base.clone(),
custom_llm_provider: model.litellm_params.custom_llm_provider.clone(),
litellm_params: model.litellm_params.clone(),
..Deployment::default()
},
)

View file

@ -29,6 +29,32 @@ fn configuration_preserves_deployment_parameters(#[case] parameters: &str) {
assert_eq!(deployment.custom_llm_provider, params.custom_llm_provider);
assert_eq!(deployment.timeout, Deployment::default().timeout);
assert_eq!(deployment.shaping, Deployment::default().shaping);
assert_eq!(deployment.litellm_params, *params);
}
#[rstest]
fn configuration_types_the_provider_params_of_a_deployment() {
let config = Config::from_yaml(
"model_list:
- model_name: public-model
litellm_params:
model: bedrock/anthropic.claude-3
aws_region_name: eu-central-1
aws_bedrock_runtime_endpoint: https://runtime.example
azure_ad_token: ignored-by-every-group",
)
.unwrap();
let router = Router::from_model_list(&config.model_list);
let params = &router.get("public-model").unwrap().litellm_params;
assert_eq!(
(
params.aws.aws_region_name.as_deref(),
params.aws.aws_bedrock_runtime_endpoint.as_deref(),
params.aws.aws_access_key_id.as_deref(),
),
(Some("eu-central-1"), Some("https://runtime.example"), None)
);
}
#[rstest]
@ -72,6 +98,7 @@ fn programmatic_deployments_preserve_overrides_and_last_entry_wins() {
api_key: Some("test-key".into()),
api_base: Some("https://provider.example/v1".into()),
custom_llm_provider: Some("test-provider".into()),
litellm_params: Default::default(),
timeout: Some(Duration::from_secs(7)),
shaping: MessagesShaping {
settings: MessagesSettings {

View file

@ -31,5 +31,6 @@ def anthropic_model_capabilities(model: str, custom_llm_provider: str | None) ->
"supports_output_config": supports("supports_output_config"),
"supports_sampling_params": AnthropicModelInfo._supports_sampling_params(resolved_model), # pyright: ignore[reportPrivateUsage] # same gate the handler applies
"supports_speed": AnthropicConfig._model_supports_speed_param(resolved_model, provider), # pyright: ignore[reportPrivateUsage] # same gate the handler applies
"supports_mid_conversation_system": supports("supports_mid_conversation_system"),
"effort_tiers": {level: tier(level) for level in ("minimal", "low", "medium", "high", "xhigh", "max")},
}

View file

@ -32,12 +32,14 @@ def test_capabilities_come_from_the_model_map_under_the_callers_provider(monkeyp
supports_output_config=True,
supports_xhigh_reasoning_effort=True,
supports_sampling_params=False,
supports_mid_conversation_system=True,
)
capabilities: Final = anthropic_model_capabilities("anthropic/claude-test-adaptive", None)
assert capabilities["supports_adaptive_thinking"]
assert capabilities["supports_output_config"]
assert capabilities["supports_mid_conversation_system"]
assert not capabilities["supports_legacy_thinking"]
assert not capabilities["supports_sampling_params"]
assert capabilities["effort_tiers"] == {
@ -56,6 +58,7 @@ def test_unmapped_model_keeps_sampling_params_and_no_reasoning_features() -> Non
assert capabilities["supports_sampling_params"]
assert not capabilities["supports_reasoning"]
assert not capabilities["supports_adaptive_thinking"]
assert not capabilities["supports_mid_conversation_system"]
assert capabilities["effort_tiers"] == dict.fromkeys(("minimal", "low", "medium", "high", "xhigh", "max"), False)