From ee295aea686ae41a9ccc05c63504c4bd94d9a351 Mon Sep 17 00:00:00 2001 From: yujonglee Date: Fri, 9 Oct 2026 16:29:45 -0700 Subject: [PATCH] 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. --- litellm-rust/Cargo.lock | 4 + .../crates/gateway-inference/src/messages.rs | 1 + .../crates/inference-chat/src/common_utils.rs | 8 +- .../crates/inference-messages/Cargo.toml | 1 + .../inference-messages/src/common_utils.rs | 8 +- .../crates/inference-messages/src/handler.rs | 4 +- .../crates/inference-messages/src/lib.rs | 3 +- .../crates/inference-messages/src/prepare.rs | 17 +- .../crates/inference-messages/src/types.rs | 15 +- .../inference-messages/tests/caching.rs | 2 +- .../inference-messages/tests/messages/main.rs | 1 + .../tests/messages/request.rs | 10 + .../tests/messages/response.rs | 1 + .../tests/messages/secrets.rs | 2 + .../tests/messages/stream.rs | 2 + .../crates/inference-ocr/src/prepare.rs | 9 +- .../inference-ocr/src/provider_config.rs | 5 +- .../inference-transcription/src/prepare.rs | 10 +- litellm-rust/crates/llms/Cargo.toml | 1 + .../crates/llms/src/anthropic/common_utils.rs | 29 +++ .../src/anthropic/messages/transformation.rs | 40 +++- .../src/azure_ai/messages/transformation.rs | 207 ++++++++++-------- litellm-rust/crates/llms/src/base_llm/auth.rs | 11 + .../llms/src/base_llm/messages/context.rs | 3 + .../messages/mid_conversation_system.rs | 201 +++++++++++++++++ .../crates/llms/src/base_llm/messages/mod.rs | 1 + .../src/base_llm/messages/normalization.rs | 128 ++++++++--- .../src/base_llm/messages/transformation.rs | 23 +- .../anthropic_claude3_transformation.rs | 180 ++++++++++++--- .../llms/tests/messages_normalization.rs | 23 +- litellm-rust/crates/python-bridge/Cargo.toml | 1 + .../python-bridge/src/python_settings.rs | 2 +- .../python-bridge/src/routes/messages/host.rs | 153 ++++++++++++- litellm-rust/crates/router/Cargo.toml | 1 + litellm-rust/crates/router/src/deployment.rs | 2 + litellm-rust/crates/router/src/lib.rs | 1 + litellm-rust/crates/router/tests/router.rs | 27 +++ litellm/rust_bridge/model_capabilities.py | 1 + .../rust_bridge/test_model_capabilities.py | 3 + 39 files changed, 901 insertions(+), 240 deletions(-) create mode 100644 litellm-rust/crates/llms/src/base_llm/messages/mid_conversation_system.rs diff --git a/litellm-rust/Cargo.lock b/litellm-rust/Cargo.lock index ec8d10b58de..b0c5b02a245 100644 --- a/litellm-rust/Cargo.lock +++ b/litellm-rust/Cargo.lock @@ -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", ] diff --git a/litellm-rust/crates/gateway-inference/src/messages.rs b/litellm-rust/crates/gateway-inference/src/messages.rs index 8969ad80934..e56308b7ae1 100644 --- a/litellm-rust/crates/gateway-inference/src/messages.rs +++ b/litellm-rust/crates/gateway-inference/src/messages.rs @@ -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, diff --git a/litellm-rust/crates/inference-chat/src/common_utils.rs b/litellm-rust/crates/inference-chat/src/common_utils.rs index 14abd42f6b7..90b75dcddaa 100644 --- a/litellm-rust/crates/inference-chat/src/common_utils.rs +++ b/litellm-rust/crates/inference-chat/src/common_utils.rs @@ -33,13 +33,7 @@ pub(super) fn chat_completions_provider(provider: LlmProviders) -> Option 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, } } diff --git a/litellm-rust/crates/inference-messages/Cargo.toml b/litellm-rust/crates/inference-messages/Cargo.toml index 8dd516732d0..6801c758e54 100644 --- a/litellm-rust/crates/inference-messages/Cargo.toml +++ b/litellm-rust/crates/inference-messages/Cargo.toml @@ -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 diff --git a/litellm-rust/crates/inference-messages/src/common_utils.rs b/litellm-rust/crates/inference-messages/src/common_utils.rs index 145192e672b..5102d063a8a 100644 --- a/litellm-rust/crates/inference-messages/src/common_utils.rs +++ b/litellm-rust/crates/inference-messages/src/common_utils.rs @@ -44,13 +44,7 @@ pub(crate) fn messages_provider(provider: LlmProviders) -> Option 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, } } diff --git a/litellm-rust/crates/inference-messages/src/handler.rs b/litellm-rust/crates/inference-messages/src/handler.rs index ee056a533b0..4e883320303 100644 --- a/litellm-rust/crates/inference-messages/src/handler.rs +++ b/litellm-rust/crates/inference-messages/src/handler.rs @@ -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, ) diff --git a/litellm-rust/crates/inference-messages/src/lib.rs b/litellm-rust/crates/inference-messages/src/lib.rs index cc48a2a39b1..6cf7e34e52a 100644 --- a/litellm-rust/crates/inference-messages/src/lib.rs +++ b/litellm-rust/crates/inference-messages/src/lib.rs @@ -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)] diff --git a/litellm-rust/crates/inference-messages/src/prepare.rs b/litellm-rust/crates/inference-messages/src/prepare.rs index b2a09f3baf9..0c189d0a309 100644 --- a/litellm-rust/crates/inference-messages/src/prepare.rs +++ b/litellm-rust/crates/inference-messages/src/prepare.rs @@ -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, diff --git a/litellm-rust/crates/inference-messages/src/types.rs b/litellm-rust/crates/inference-messages/src/types.rs index 8521345b00e..669f7b283a1 100644 --- a/litellm-rust/crates/inference-messages/src/types.rs +++ b/litellm-rust/crates/inference-messages/src/types.rs @@ -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, pub api_base: Option, pub custom_llm_provider: Option, + pub litellm_params: LitellmParams, pub extra_headers: Option>, pub provider_specific_header: Option, pub timeout: Option, @@ -27,8 +29,16 @@ pub fn messages_body(body: Map) -> Result 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) -> Result { + 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, diff --git a/litellm-rust/crates/inference-messages/tests/caching.rs b/litellm-rust/crates/inference-messages/tests/caching.rs index dbec5b87d81..e9005123119 100644 --- a/litellm-rust/crates/inference-messages/tests/caching.rs +++ b/litellm-rust/crates/inference-messages/tests/caching.rs @@ -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); diff --git a/litellm-rust/crates/inference-messages/tests/messages/main.rs b/litellm-rust/crates/inference-messages/tests/messages/main.rs index 509a8b19847..8f92b31c999 100644 --- a/litellm-rust/crates/inference-messages/tests/messages/main.rs +++ b/litellm-rust/crates/inference-messages/tests/messages/main.rs @@ -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)), diff --git a/litellm-rust/crates/inference-messages/tests/messages/request.rs b/litellm-rust/crates/inference-messages/tests/messages/request.rs index 707e91c2d0d..523b25f81d6 100644 --- a/litellm-rust/crates/inference-messages/tests/messages/request.rs +++ b/litellm-rust/crates/inference-messages/tests/messages/request.rs @@ -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 { diff --git a/litellm-rust/crates/inference-messages/tests/messages/response.rs b/litellm-rust/crates/inference-messages/tests/messages/response.rs index 13b99d45241..b5d623a5582 100644 --- a/litellm-rust/crates/inference-messages/tests/messages/response.rs +++ b/litellm-rust/crates/inference-messages/tests/messages/response.rs @@ -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 diff --git a/litellm-rust/crates/inference-messages/tests/messages/secrets.rs b/litellm-rust/crates/inference-messages/tests/messages/secrets.rs index c565edac9dc..a33fdd130f1 100644 --- a/litellm-rust/crates/inference-messages/tests/messages/secrets.rs +++ b/litellm-rust/crates/inference-messages/tests/messages/secrets.rs @@ -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 }, ) diff --git a/litellm-rust/crates/inference-messages/tests/messages/stream.rs b/litellm-rust/crates/inference-messages/tests/messages/stream.rs index 4dec4de2c02..049a620e454 100644 --- a/litellm-rust/crates/inference-messages/tests/messages/stream.rs +++ b/litellm-rust/crates/inference-messages/tests/messages/stream.rs @@ -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, diff --git a/litellm-rust/crates/inference-ocr/src/prepare.rs b/litellm-rust/crates/inference-ocr/src/prepare.rs index e3b0e819fb2..14120606864 100644 --- a/litellm-rust/crates/inference-ocr/src/prepare.rs +++ b/litellm-rust/crates/inference-ocr/src/prepare.rs @@ -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(|| { diff --git a/litellm-rust/crates/inference-ocr/src/provider_config.rs b/litellm-rust/crates/inference-ocr/src/provider_config.rs index 28cb96238a3..da708d17b86 100644 --- a/litellm-rust/crates/inference-ocr/src/provider_config.rs +++ b/litellm-rust/crates/inference-ocr/src/provider_config.rs @@ -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(), )); diff --git a/litellm-rust/crates/inference-transcription/src/prepare.rs b/litellm-rust/crates/inference-transcription/src/prepare.rs index 91b2d11e43a..15ca7627339 100644 --- a/litellm-rust/crates/inference-transcription/src/prepare.rs +++ b/litellm-rust/crates/inference-transcription/src/prepare.rs @@ -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, } } diff --git a/litellm-rust/crates/llms/Cargo.toml b/litellm-rust/crates/llms/Cargo.toml index ab9c8bc5646..2d3ea85d1fc 100644 --- a/litellm-rust/crates/llms/Cargo.toml +++ b/litellm-rust/crates/llms/Cargo.toml @@ -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" diff --git a/litellm-rust/crates/llms/src/anthropic/common_utils.rs b/litellm-rust/crates/llms/src/anthropic/common_utils.rs index ccf3a6d15f2..335731fa001 100644 --- a/litellm-rust/crates/llms/src/anthropic/common_utils.rs +++ b/litellm-rust/crates/llms/src/anthropic/common_utils.rs @@ -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 { + 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 = 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) -> Vec { 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), } ); diff --git a/litellm-rust/crates/llms/src/anthropic/messages/transformation.rs b/litellm-rust/crates/llms/src/anthropic/messages/transformation.rs index 74ac35752ab..7021a4199df 100644 --- a/litellm-rust/crates/llms/src/anthropic/messages/transformation.rs +++ b/litellm-rust/crates/llms/src/anthropic/messages/transformation.rs @@ -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, ) -> Result { 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, ) -> Result { 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 diff --git a/litellm-rust/crates/llms/src/azure_ai/messages/transformation.rs b/litellm-rust/crates/llms/src/azure_ai/messages/transformation.rs index ffd079b2d04..3e6abf2d14e 100644 --- a/litellm-rust/crates/llms/src/azure_ai/messages/transformation.rs +++ b/litellm-rust/crates/llms/src/azure_ai/messages/transformation.rs @@ -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, ) -> Result { complete_azure_anthropic_url(api_base, env_lookup) @@ -57,20 +57,14 @@ impl BaseMessagesConfig for AzureAnthropicMessagesConfig { request: MessagesRequest, context: &MessagesTransformContext, ) -> Result { - 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, ) -> Result { 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 diff --git a/litellm-rust/crates/llms/src/base_llm/auth.rs b/litellm-rust/crates/llms/src/base_llm/auth.rs index 92a19e3ef2a..de271a269d3 100644 --- a/litellm-rust/crates/llms/src/base_llm/auth.rs +++ b/litellm-rust/crates/llms/src/base_llm/auth.rs @@ -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 }, /// 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, diff --git a/litellm-rust/crates/llms/src/base_llm/messages/context.rs b/litellm-rust/crates/llms/src/base_llm/messages/context.rs index a20a98f56b1..e59a63813b0 100644 --- a/litellm-rust/crates/llms/src/base_llm/messages/context.rs +++ b/litellm-rust/crates/llms/src/base_llm/messages/context.rs @@ -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(), } } diff --git a/litellm-rust/crates/llms/src/base_llm/messages/mid_conversation_system.rs b/litellm-rust/crates/llms/src/base_llm/messages/mid_conversation_system.rs new file mode 100644 index 00000000000..ec1a9863786 --- /dev/null +++ b/litellm-rust/crates/llms/src/base_llm/messages/mid_conversation_system.rs @@ -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) -> Vec { + 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 { + 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, + follower_run: Vec, +) -> Vec { + 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) -> Vec> { + messages + .into_iter() + .fold(Vec::new(), |mut runs: Vec>, 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) -> Vec { + 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) -> Vec { + 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 { + 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"}) + ); + } +} diff --git a/litellm-rust/crates/llms/src/base_llm/messages/mod.rs b/litellm-rust/crates/llms/src/base_llm/messages/mod.rs index ac2112176f9..e9591daa9da 100644 --- a/litellm-rust/crates/llms/src/base_llm/messages/mod.rs +++ b/litellm-rust/crates/llms/src/base_llm/messages/mod.rs @@ -1,4 +1,5 @@ pub mod context; +pub mod mid_conversation_system; pub mod normalization; pub mod streaming; pub mod transformation; diff --git a/litellm-rust/crates/llms/src/base_llm/messages/normalization.rs b/litellm-rust/crates/llms/src/base_llm/messages/normalization.rs index bddd829304e..3ae4145ae32 100644 --- a/litellm-rust/crates/llms/src/base_llm/messages/normalization.rs +++ b/litellm-rust/crates/llms/src/base_llm/messages/normalization.rs @@ -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 { - match content { - MessageContent::Text(text) => vec![ContentBlock::text(text)], - MessageContent::Blocks(blocks) => blocks, - } -} - -fn system_into_blocks(system: Option) -> Vec { - 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, Vec) = 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 = 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 = messages.by_ref().take(leading_count).collect(); + let remaining: Vec = 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 diff --git a/litellm-rust/crates/llms/src/base_llm/messages/transformation.rs b/litellm-rust/crates/llms/src/base_llm/messages/transformation.rs index 2f2d3bcf909..f42adcdcfe6 100644 --- a/litellm-rust/crates/llms/src/base_llm/messages/transformation.rs +++ b/litellm-rust/crates/llms/src/base_llm/messages/transformation.rs @@ -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, ) -> Result; - fn complete_stream_url( - &self, - api_base: Option<&str>, - model: &str, - env_lookup: &dyn Fn(&str) -> Option, - ) -> Result { - 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, ) -> Result; @@ -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, ) -> Result { 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, ) -> Result { Ok(ValidatedEnvironment { diff --git a/litellm-rust/crates/llms/src/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.rs b/litellm-rust/crates/llms/src/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.rs index c5bf7d8fc95..59452e0b693 100644 --- a/litellm-rust/crates/llms/src/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.rs +++ b/litellm-rust/crates/llms/src/bedrock/messages/invoke_transformations/anthropic_claude3_transformation.rs @@ -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, - 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}", ®ion)); 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, ) -> Result { - 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, - ) -> Result { - 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, ) -> Result { 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(), ¶ms, env_lookup), + region: resolve_bedrock_region(model_region.as_deref(), params, env_lookup), service: BEDROCK_SERVICE, - credentials: Box::new(AwsCredentialSource::from_params(¶ms, 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 { 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(); diff --git a/litellm-rust/crates/llms/tests/messages_normalization.rs b/litellm-rust/crates/llms/tests/messages_normalization.rs index 27dba22c662..bd54618aef0 100644 --- a/litellm-rust/crates/llms/tests/messages_normalization.rs +++ b/litellm-rust/crates/llms/tests/messages_normalization.rs @@ -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 = 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 + ); } diff --git a/litellm-rust/crates/python-bridge/Cargo.toml b/litellm-rust/crates/python-bridge/Cargo.toml index 282f5b8544c..6fc05c43146 100644 --- a/litellm-rust/crates/python-bridge/Cargo.toml +++ b/litellm-rust/crates/python-bridge/Cargo.toml @@ -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 diff --git a/litellm-rust/crates/python-bridge/src/python_settings.rs b/litellm-rust/crates/python-bridge/src/python_settings.rs index f0fea5adb06..88900544779 100644 --- a/litellm-rust/crates/python-bridge/src/python_settings.rs +++ b/litellm-rust/crates/python-bridge/src/python_settings.rs @@ -60,7 +60,7 @@ impl PythonSettings { } } -fn missing_module(py: Python<'_>, error: &PyErr, expected: &str) -> PyResult { +pub(crate) fn missing_module(py: Python<'_>, error: &PyErr, expected: &str) -> PyResult { if !error.is_instance_of::(py) { return Ok(false); } diff --git a/litellm-rust/crates/python-bridge/src/routes/messages/host.rs b/litellm-rust/crates/python-bridge/src/routes/messages/host.rs index 841ef3081e3..d34047c735d 100644 --- a/litellm-rust/crates/python-bridge/src/routes/messages/host.rs +++ b/litellm-rust/crates/python-bridge/src/routes/messages/host.rs @@ -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.`, 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>> { + 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>>, + global: impl Fn(&str) -> PyResult>>, +) -> PyResult> { + 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 { 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, + ) { + 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, + ) { + 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, + ) { + assert_projection(kwargs, globals, expected); + } + + fn assert_projection(kwargs: Value, globals: Value, expected: Result) { + Python::initialize(); + Python::attach(|py| { + let dict = |value: &Value| { + to_py(py, value) + .unwrap() + .into_bound(py) + .cast_into::() + .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)] diff --git a/litellm-rust/crates/router/Cargo.toml b/litellm-rust/crates/router/Cargo.toml index 8c01c1e5a99..a356bedcf16 100644 --- a/litellm-rust/crates/router/Cargo.toml +++ b/litellm-rust/crates/router/Cargo.toml @@ -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 diff --git a/litellm-rust/crates/router/src/deployment.rs b/litellm-rust/crates/router/src/deployment.rs index 4a091f8845c..d57284bae5a 100644 --- a/litellm-rust/crates/router/src/deployment.rs +++ b/litellm-rust/crates/router/src/deployment.rs @@ -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, pub api_base: Option, pub custom_llm_provider: Option, + pub litellm_params: LitellmParams, pub timeout: Option, pub shaping: MessagesShaping, } diff --git a/litellm-rust/crates/router/src/lib.rs b/litellm-rust/crates/router/src/lib.rs index da33bfb04bd..fb48213e8a7 100644 --- a/litellm-rust/crates/router/src/lib.rs +++ b/litellm-rust/crates/router/src/lib.rs @@ -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() }, ) diff --git a/litellm-rust/crates/router/tests/router.rs b/litellm-rust/crates/router/tests/router.rs index 97ee544f120..bbaeb806b96 100644 --- a/litellm-rust/crates/router/tests/router.rs +++ b/litellm-rust/crates/router/tests/router.rs @@ -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 { diff --git a/litellm/rust_bridge/model_capabilities.py b/litellm/rust_bridge/model_capabilities.py index eb8496c722f..cbf59dfb97b 100644 --- a/litellm/rust_bridge/model_capabilities.py +++ b/litellm/rust_bridge/model_capabilities.py @@ -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")}, } diff --git a/tests/unit/rust_bridge/test_model_capabilities.py b/tests/unit/rust_bridge/test_model_capabilities.py index 9f02660a98f..93f7dfdc04a 100644 --- a/tests/unit/rust_bridge/test_model_capabilities.py +++ b/tests/unit/rust_bridge/test_model_capabilities.py @@ -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)