From 09509f702e395e28cbd70939315cd67b7d39995e Mon Sep 17 00:00:00 2001 From: Yujong Lee Date: Thu, 24 Sep 2026 10:45:33 -0700 Subject: [PATCH] fix(rust): resolve Messages credentials through the secret source and scope headers by resolved provider The native Messages route read ANTHROPIC_API_KEY, ANTHROPIC_AUTH_TOKEN and the base URL straight from the process environment, so a key or base held in a configured secret manager was never found. Each provider config now declares its secret names and the route resolves them through the same SecretSource the OCR route uses, with the Python bridge passing in litellm's configured manager provider_specific_header entries were scoped by the explicit custom_llm_provider only, falling back to anthropic, so an azure_ai/ model lost its azure_ai scoped headers. Scoping now happens in the route after the provider is resolved from the model, as Python's handler does The Azure config now adds the same anthropic-beta feature headers Python's Azure route adds, and the metadata allowlist, reasoning auto summary and history sanitizers move from the core route into the llms crate, mirroring their home in Python's messages handler --- .../src/get_provider_specific_headers.rs | 93 ++++ litellm-rust/crates/core-utils/src/lib.rs | 1 + .../crates/core/src/messages/error.rs | 28 ++ litellm-rust/crates/core/src/messages/mod.rs | 7 +- .../crates/core/src/messages/prepare.rs | 459 +++++++----------- .../crates/core/src/messages/route.rs | 50 +- .../crates/core/src/messages/tests.rs | 134 ++++- .../crates/core/src/messages/types.rs | 2 + .../messages/handler.rs | 270 +++++++++++ .../messages/headers.rs | 16 +- .../experimental_pass_through/messages/mod.rs | 1 + .../messages/transformation.rs | 36 +- .../anthropic/messages_transformation.rs | 83 +++- .../anthropic_messages/transformation.rs | 10 + .../python-bridge/src/routes/messages/host.rs | 144 +----- .../python-bridge/src/routes/messages/mod.rs | 3 +- litellm-rust/crates/types/src/utils.rs | 15 + .../rust_bridge/messages/test_secrets.py | 111 +++++ 18 files changed, 1036 insertions(+), 427 deletions(-) create mode 100644 litellm-rust/crates/core-utils/src/get_provider_specific_headers.rs create mode 100644 litellm-rust/crates/llms/src/anthropic/experimental_pass_through/messages/handler.rs create mode 100644 tests/test_litellm/rust_bridge/messages/test_secrets.py diff --git a/litellm-rust/crates/core-utils/src/get_provider_specific_headers.rs b/litellm-rust/crates/core-utils/src/get_provider_specific_headers.rs new file mode 100644 index 00000000000..bfcd448e2d8 --- /dev/null +++ b/litellm-rust/crates/core-utils/src/get_provider_specific_headers.rs @@ -0,0 +1,93 @@ +use litellm_types::utils::{ProviderSpecificHeader, ProviderSpecificHeaders}; +use serde_json::{Map, Value}; + +pub fn get_provider_specific_headers( + provider_specific_header: Option<&ProviderSpecificHeaders>, + custom_llm_provider: &str, +) -> Map { + let entries: &[ProviderSpecificHeader] = match provider_specific_header { + None => &[], + Some(ProviderSpecificHeaders::One(entry)) => std::slice::from_ref(entry), + Some(ProviderSpecificHeaders::Many(entries)) => entries, + }; + entries + .iter() + .filter(|entry| { + entry + .custom_llm_provider + .split(',') + .any(|scoped| scoped.trim() == custom_llm_provider) + }) + .flat_map(|entry| entry.extra_headers.clone()) + .collect() +} + +#[cfg(test)] +mod tests { + use rstest::rstest; + use serde_json::json; + + use super::*; + + #[rstest] + #[case::single_entry_for_the_provider( + json!({"custom_llm_provider": "anthropic", "extra_headers": {"Authorization": "Bearer t", "Custom-Header": "v"}}), + json!({"Authorization": "Bearer t", "Custom-Header": "v"}), + )] + #[case::single_entry_for_another_provider( + json!({"custom_llm_provider": "openai", "extra_headers": {"Authorization": "Bearer t"}}), + json!({}), + )] + #[case::provider_in_a_comma_separated_scope( + json!({"custom_llm_provider": "bedrock,anthropic,vertex_ai", "extra_headers": {"anthropic-beta": "context-1m-2025-08-07"}}), + json!({"anthropic-beta": "context-1m-2025-08-07"}), + )] + #[case::provider_missing_from_a_comma_separated_scope( + json!({"custom_llm_provider": "bedrock,vertex_ai", "extra_headers": {"anthropic-beta": "test"}}), + json!({}), + )] + #[case::scope_with_spaces( + json!({"custom_llm_provider": "bedrock, anthropic , vertex_ai", "extra_headers": {"anthropic-beta": "test"}}), + json!({"anthropic-beta": "test"}), + )] + #[case::scope_names_must_match_exactly( + json!({"custom_llm_provider": "anthropic_text", "extra_headers": {"anthropic-beta": "test"}}), + json!({}), + )] + #[case::entries_scope_independently( + json!([ + {"custom_llm_provider": "anthropic,bedrock,vertex_ai", "extra_headers": {"anthropic-beta": "context-1m-2025-08-07"}}, + {"custom_llm_provider": "bedrock", "extra_headers": {"x-bedrock-only": "no"}}, + {"custom_llm_provider": "anthropic", "extra_headers": {"authorization": "Bearer sk-ant-oat01-fake-token"}} + ]), + json!({"anthropic-beta": "context-1m-2025-08-07", "authorization": "Bearer sk-ant-oat01-fake-token"}), + )] + #[case::later_entries_win( + json!([ + {"custom_llm_provider": "anthropic", "extra_headers": {"x-scoped": "first"}}, + {"custom_llm_provider": "anthropic", "extra_headers": {"x-scoped": "second"}} + ]), + json!({"x-scoped": "second"}), + )] + #[case::empty_list(json!([]), json!({}))] + #[case::entry_without_scope(json!({"extra_headers": {"x-scoped": "yes"}}), json!({}))] + #[case::entry_without_headers(json!({"custom_llm_provider": "anthropic"}), json!({}))] + fn provider_specific_headers_match_the_scoped_provider( + #[case] configured: Value, + #[case] expected: Value, + ) { + let configured: ProviderSpecificHeaders = serde_json::from_value(configured).unwrap(); + assert_eq!( + Value::Object(get_provider_specific_headers( + Some(&configured), + "anthropic" + )), + expected + ); + } + + #[test] + fn no_configured_headers_match_nothing() { + assert_eq!(get_provider_specific_headers(None, "anthropic"), Map::new()); + } +} diff --git a/litellm-rust/crates/core-utils/src/lib.rs b/litellm-rust/crates/core-utils/src/lib.rs index 31d967f9048..a937f55654e 100644 --- a/litellm-rust/crates/core-utils/src/lib.rs +++ b/litellm-rust/crates/core-utils/src/lib.rs @@ -3,6 +3,7 @@ pub mod core_helpers; pub mod dot_notation_indexing; pub mod exception_mapping_utils; pub mod get_llm_provider_logic; +pub mod get_provider_specific_headers; pub mod params; pub mod prompt_templates; pub mod secret_redaction; diff --git a/litellm-rust/crates/core/src/messages/error.rs b/litellm-rust/crates/core/src/messages/error.rs index 51fb764032c..2a9723beb38 100644 --- a/litellm-rust/crates/core/src/messages/error.rs +++ b/litellm-rust/crates/core/src/messages/error.rs @@ -1,3 +1,5 @@ +use std::sync::Arc; + use litellm_llms::base_llm::chat::transformation::Error as LlmError; #[derive(Clone, Debug, PartialEq, Eq, thiserror::Error)] @@ -18,8 +20,34 @@ pub enum Error { Transport(#[from] litellm_http::transport::Error), #[error(transparent)] Headers(#[from] litellm_http::request::HeaderError), + #[error(transparent)] + Secret(#[from] SecretError), } +#[derive(Clone, Debug, thiserror::Error)] +#[error(transparent)] +pub struct SecretError(Arc); + +impl SecretError { + pub fn source_error(&self) -> &litellm_secrets::Error { + &self.0 + } +} + +impl From for Error { + fn from(error: litellm_secrets::Error) -> Self { + Self::Secret(SecretError(Arc::new(error))) + } +} + +impl PartialEq for SecretError { + fn eq(&self, other: &Self) -> bool { + Arc::ptr_eq(&self.0, &other.0) + } +} + +impl Eq for SecretError {} + impl From for Error { fn from(error: LlmError) -> Self { match error { diff --git a/litellm-rust/crates/core/src/messages/mod.rs b/litellm-rust/crates/core/src/messages/mod.rs index f588208d731..8795d4f8507 100644 --- a/litellm-rust/crates/core/src/messages/mod.rs +++ b/litellm-rust/crates/core/src/messages/mod.rs @@ -12,6 +12,9 @@ mod common_utils; mod handler; mod prepare; pub mod route; +use std::sync::Arc; + +use litellm_secrets::source::EnvironmentSecrets; use litellm_types::llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse; use route::{LocalMessagesHost, MessagesCall, MessagesOutput, messages_machine}; use serde_json::Value; @@ -31,10 +34,12 @@ pub async fn messages(request: MessagesRequest<'_>) -> Result Ok(*message), MessagesOutput::Streamed => Err(Error::Unsupported( "streamed responses need a streaming host", diff --git a/litellm-rust/crates/core/src/messages/prepare.rs b/litellm-rust/crates/core/src/messages/prepare.rs index c56181db59b..dc4b3562e3f 100644 --- a/litellm-rust/crates/core/src/messages/prepare.rs +++ b/litellm-rust/crates/core/src/messages/prepare.rs @@ -1,57 +1,79 @@ use litellm_core_utils::{ dot_notation_indexing::delete_nested_value, get_llm_provider_logic::{CustomLlmProvider, get_custom_llm_provider}, + get_provider_specific_headers::get_provider_specific_headers, + settings::Lookup, }; use litellm_llms::{ - anthropic::common_utils::{ - flatten_unencrypted_web_search_results, sanitize_tool_use_ids, strip_empty_content_blocks, - strip_provider_specific_fields, + anthropic::experimental_pass_through::messages::handler::shape_anthropic_messages_request, + base_llm::anthropic_messages::transformation::{ + BaseAnthropicMessagesConfig, MessagesTransformContext, }, - base_llm::anthropic_messages::transformation::MessagesTransformContext, }; -use litellm_types::llms::anthropic_messages::anthropic_request::{ - AnthropicMessage, AnthropicMessagesRequest, -}; -use serde_json::{Map, Value, json}; +use litellm_types::llms::anthropic_messages::anthropic_request::AnthropicMessagesRequest; +use serde_json::{Map, Value}; use super::{ Error, common_utils::{messages_provider_config, string_headers}, }; -use crate::messages::types::{MessagesRequest, MessagesShaping, ProviderMessagesRequest}; +use crate::messages::types::{MessagesRequest, ProviderMessagesRequest}; -pub(super) fn prepare_provider_request( - request: MessagesRequest<'_>, -) -> Result { - let provider_info = get_custom_llm_provider(request.model, request.custom_llm_provider) +pub(super) struct ResolvedProvider<'a> { + pub(super) model: &'a str, + pub(super) provider: &'a str, + pub(super) config: &'static dyn BaseAnthropicMessagesConfig, +} + +pub(super) fn resolve_provider<'a>( + model: &'a str, + custom_llm_provider: Option<&'a str>, +) -> Result, Error> { + let CustomLlmProvider { + model, + custom_llm_provider: provider, + } = get_custom_llm_provider(model, custom_llm_provider) .or_else(|| { - request - .custom_llm_provider - .map(|provider| CustomLlmProvider { - model: request.model, - custom_llm_provider: provider, - }) + custom_llm_provider.map(|provider| CustomLlmProvider { + model, + custom_llm_provider: provider, + }) }) .ok_or_else(|| { Error::InvalidProvider( "unable to resolve custom_llm_provider for messages request".to_string(), ) })?; - let model = provider_info.model.to_string(); - let provider = provider_info.custom_llm_provider; - let config = messages_provider_config(provider) .ok_or_else(|| Error::InvalidProvider(provider.to_string()))?; - let env_lookup = |key: &str| std::env::var(key).ok(); + Ok(ResolvedProvider { + model, + provider, + config, + }) +} + +pub(super) fn prepare_provider_request( + request: MessagesRequest<'_>, + resolved: ResolvedProvider<'_>, + secrets: &dyn Lookup, +) -> Result { + let ResolvedProvider { + model, + provider, + config, + } = resolved; + let model = model.to_string(); + let env_lookup = |key: &str| secrets.get(key); let typed_request: AnthropicMessagesRequest = serde_json::from_value(request.body).map_err(invalid_request)?; - let sanitized = sanitize_request( + let sanitized = shape_anthropic_messages_request( AnthropicMessagesRequest { model: model.clone(), ..typed_request }, - &request.shaping, + request.shaping.reasoning_auto_summary, )?; let trimmed = without_additional_drop_params(sanitized, &request.shaping.additional_drop_params)?; @@ -60,7 +82,15 @@ pub(super) fn prepare_provider_request( &MessagesTransformContext::new(request.shaping.capabilities, request.shaping.drop_params), )?; - let forwarded = string_headers(request.extra_headers)?; + let scoped = get_provider_specific_headers(request.provider_specific_header.as_ref(), provider); + let forwarded = string_headers(Some( + request + .extra_headers + .into_iter() + .flatten() + .chain(scoped) + .collect(), + ))?; let authenticated = config.authenticate(forwarded, request.api_key, &env_lookup)?; let headers = config.request_headers( with_default_headers(authenticated, config.default_headers()), @@ -115,59 +145,6 @@ fn without_additional_drop_params( serde_json::from_value(Value::Object(merged)).map_err(invalid_request) } -fn sanitize_request( - request: AnthropicMessagesRequest, - shaping: &MessagesShaping, -) -> Result { - Ok(AnthropicMessagesRequest { - messages: sanitize_messages(request.messages), - metadata: request - .metadata - .as_ref() - .map(allowed_metadata) - .transpose()?, - thinking: with_reasoning_auto_summary(request.thinking, shaping.reasoning_auto_summary), - ..request - }) -} - -fn sanitize_messages(messages: Vec) -> Vec { - strip_provider_specific_fields(flatten_unencrypted_web_search_results( - sanitize_tool_use_ids(strip_empty_content_blocks(messages)), - )) -} - -fn allowed_metadata(metadata: &Value) -> Result { - let Value::Object(fields) = metadata else { - return Err(Error::InvalidRequest(format!( - "metadata must be an object, got {metadata}" - ))); - }; - match fields.get("user_id") { - None | Some(Value::Null) => Ok(json!({})), - Some(Value::String(user_id)) => Ok(json!({"user_id": user_id})), - Some(other) => Err(Error::InvalidRequest(format!( - "metadata.user_id must be a string, got {other}" - ))), - } -} - -fn with_reasoning_auto_summary(thinking: Option, enabled: bool) -> Option { - let Some(Value::Object(thinking)) = thinking else { - return thinking; - }; - if !enabled || thinking.get("type").and_then(Value::as_str) == Some("disabled") { - return Some(Value::Object(thinking)); - } - Some(Value::Object( - thinking - .into_iter() - .filter(|(key, _)| key != "display") - .chain([("display".to_string(), json!("summarized"))]) - .collect(), - )) -} - fn with_default_headers( headers: Vec<(String, String)>, defaults: &[(&str, &str)], @@ -186,194 +163,105 @@ fn with_default_headers( #[cfg(test)] mod tests { + use litellm_types::utils::ProviderSpecificHeaders; use rstest::{fixture, rstest}; + use serde_json::json; use super::*; - - fn messages(value: Value) -> Vec { - serde_json::from_value(value).unwrap() - } - - fn request(body: Value) -> AnthropicMessagesRequest { - serde_json::from_value(body).unwrap() - } + use crate::messages::types::MessagesShaping; #[fixture] fn shaping() -> MessagesShaping { MessagesShaping::default() } + fn prepare(request: MessagesRequest<'_>) -> Result { + prepare_with_secrets(request, &|_: &str| None) + } + + fn prepare_with_secrets( + request: MessagesRequest<'_>, + secrets: &dyn Lookup, + ) -> Result { + let resolved = resolve_provider(request.model, request.custom_llm_provider)?; + prepare_provider_request(request, resolved, secrets) + } + + #[rstest] + #[case::api_key( + &[("ANTHROPIC_API_KEY", "sk-secret")], + &[("x-api-key", "sk-secret")], + "https://api.anthropic.com/v1/messages" + )] + #[case::auth_token( + &[("ANTHROPIC_AUTH_TOKEN", "token")], + &[("authorization", "Bearer token")], + "https://api.anthropic.com/v1/messages" + )] + #[case::api_base( + &[("ANTHROPIC_API_KEY", "sk-secret"), ("ANTHROPIC_API_BASE", "https://gateway.test")], + &[("x-api-key", "sk-secret")], + "https://gateway.test/v1/messages" + )] + #[case::sdk_base_url( + &[("ANTHROPIC_API_KEY", "sk-secret"), ("ANTHROPIC_BASE_URL", "https://sdk.test")], + &[("x-api-key", "sk-secret")], + "https://sdk.test/v1/messages" + )] + fn credentials_and_base_come_from_the_resolved_secrets( + shaping: MessagesShaping, + #[case] secrets: &[(&str, &str)], + #[case] expected_auth: &[(&str, &str)], + #[case] expected_url: &str, + ) { + let lookup = |name: &str| { + secrets + .iter() + .find(|(key, _)| *key == name) + .map(|(_, value)| value.to_string()) + }; + let prepared = prepare_with_secrets( + MessagesRequest { + model: "claude-test", + body: json!({"model": "claude-test", "messages": [{"role": "user", "content": "hi"}], "max_tokens": 16}), + api_key: None, + api_base: None, + custom_llm_provider: Some("anthropic"), + extra_headers: None, + provider_specific_header: None, + timeout: None, + shaping, + }, + &lookup, + ) + .unwrap(); + let auth: Vec<(&str, &str)> = prepared + .upstream_headers + .iter() + .filter(|(name, _)| matches!(name.as_str(), "x-api-key" | "authorization")) + .map(|(name, value)| (name.as_str(), value.as_str())) + .collect(); + assert_eq!( + (auth.as_slice(), prepared.url.as_str()), + (expected_auth, expected_url) + ); + } + fn prepared_body(body: Value, shaping: MessagesShaping) -> Result { - prepare_provider_request(MessagesRequest { + prepare(MessagesRequest { model: "anthropic/claude-test", body, api_key: Some("sk-test"), api_base: Some("https://anthropic.test"), custom_llm_provider: Some("anthropic"), extra_headers: None, + provider_specific_header: None, timeout: None, shaping, }) .map(|prepared| prepared.body) } - #[rstest] - #[case::empty_text_next_to_a_tool_use( - json!([{"role": "assistant", "content": [ - {"type": "text", "text": " "}, - {"type": "tool_use", "id": "t", "name": "B", "input": {}} - ]}]), - json!([{"role": "assistant", "content": [ - {"type": "tool_use", "id": "t", "name": "B", "input": {}} - ]}]), - )] - #[case::cross_provider_tool_ids( - json!([ - {"role": "assistant", "content": [{"type": "tool_use", "id": "functions.Bash:0", "name": "Bash", "input": {}}]}, - {"role": "user", "content": [{"type": "tool_result", "tool_use_id": "functions.Bash:0", "content": "ok"}]} - ]), - json!([ - {"role": "assistant", "content": [{"type": "tool_use", "id": "functions_Bash_0", "name": "Bash", "input": {}}]}, - {"role": "user", "content": [{"type": "tool_result", "tool_use_id": "functions_Bash_0", "content": "ok"}]} - ]), - )] - #[case::replayed_unencrypted_web_search_results( - json!([ - {"role": "user", "content": "latest litellm version?"}, - {"role": "assistant", "content": [ - {"type": "server_tool_use", "id": "srvtoolu_1", "name": "web_search", "input": {"query": "latest litellm version"}}, - {"type": "web_search_tool_result", "tool_use_id": "srvtoolu_1", "content": [{ - "type": "web_search_result", - "url": "https://github.com/BerriAI/litellm/releases", - "title": "Releases", - "page_age": null, - "encrypted_content": "", - "snippet": "Latest release v1.95.0" - }]} - ]}, - {"role": "user", "content": "which version?"} - ]), - json!([ - {"role": "user", "content": "latest litellm version?"}, - {"role": "assistant", "content": [{ - "type": "text", - "text": "Web search results for 'latest litellm version':\n\nTitle: Releases\nURL: https://github.com/BerriAI/litellm/releases\nSnippet: Latest release v1.95.0" - }]}, - {"role": "user", "content": "which version?"} - ]), - )] - #[case::replayed_provider_specific_fields( - json!([ - {"role": "assistant", "content": [{ - "type": "tool_use", "id": "toolu_01", "name": "get_weather", "input": {"city": "Paris"}, - "provider_specific_fields": {"signature": "sig_abc"} - }]}, - {"role": "user", "content": [{"type": "tool_result", "tool_use_id": "toolu_01", "content": "Sunny"}]} - ]), - json!([ - {"role": "assistant", "content": [{"type": "tool_use", "id": "toolu_01", "name": "get_weather", "input": {"city": "Paris"}}]}, - {"role": "user", "content": [{"type": "tool_result", "tool_use_id": "toolu_01", "content": "Sunny"}]} - ]), - )] - #[case::ids_are_normalized_before_web_search_results_flatten( - json!([ - {"role": "user", "content": "run it"}, - {"role": "assistant", "content": [ - {"type": "thinking", "thinking": "", "signature": "sig"}, - {"type": "text", "text": ""}, - {"type": "tool_use", "id": "functions.Bash:0", "name": "Bash", "input": {}, "provider_specific_fields": {"x": 1}}, - {"type": "server_tool_use", "id": "srv.1", "name": "web_search", "input": {"query": "q"}, "provider_specific_fields": {"x": 2}}, - {"type": "web_search_tool_result", "tool_use_id": "srv.1", "provider_specific_fields": {"x": 3}, "content": [ - {"type": "web_search_result", "url": "u", "title": "", "encrypted_content": "", "provider_specific_fields": {"x": 4}} - ]} - ]}, - {"role": "user", "content": [{"type": "tool_result", "tool_use_id": "functions.Bash:0", "content": "ok"}]}, - {"role": "assistant", "content": [{"type": "text", "text": " "}]} - ]), - json!([ - {"role": "user", "content": "run it"}, - {"role": "assistant", "content": [ - {"type": "tool_use", "id": "functions_Bash_0", "name": "Bash", "input": {}}, - {"type": "server_tool_use", "id": "srv_1", "name": "web_search", "input": {"query": "q"}}, - {"type": "text", "text": "Web search results:\n\nURL: u"} - ]}, - {"role": "user", "content": [{"type": "tool_result", "tool_use_id": "functions_Bash_0", "content": "ok"}]} - ]), - )] - fn sanitize_messages_cleans_replayed_history(#[case] history: Value, #[case] expected: Value) { - assert_eq!( - serde_json::to_value(sanitize_messages(messages(history))).unwrap(), - expected - ); - } - - #[rstest] - #[case::keeps_only_user_id(json!({"user_id": "u-1", "trace_id": "internal"}), Ok(json!({"user_id": "u-1"})))] - #[case::null_user_id(json!({"user_id": null, "trace_id": "internal"}), Ok(json!({})))] - #[case::no_user_id(json!({"trace_id": "internal"}), Ok(json!({})))] - #[case::empty(json!({}), Ok(json!({})))] - #[case::numeric_user_id( - json!({"user_id": 123}), - Err(Error::InvalidRequest("metadata.user_id must be a string, got 123".to_string())), - )] - #[case::boolean_user_id( - json!({"user_id": true}), - Err(Error::InvalidRequest("metadata.user_id must be a string, got true".to_string())), - )] - #[case::not_an_object( - json!(["u-1"]), - Err(Error::InvalidRequest(r#"metadata must be an object, got ["u-1"]"#.to_string())), - )] - fn allowed_metadata_passes_only_a_string_user_id( - #[case] metadata: Value, - #[case] expected: Result, - ) { - assert_eq!(allowed_metadata(&metadata), expected); - } - - #[rstest] - #[case::adaptive( - Some(json!({"type": "adaptive", "budget_tokens": 5000})), - true, - Some(json!({"type": "adaptive", "budget_tokens": 5000, "display": "summarized"})), - )] - #[case::enabled( - Some(json!({"type": "enabled", "budget_tokens": 10000})), - true, - Some(json!({"type": "enabled", "budget_tokens": 10000, "display": "summarized"})), - )] - #[case::no_type(Some(json!({})), true, Some(json!({"display": "summarized"})))] - #[case::display_omitted_is_overridden( - Some(json!({"type": "enabled", "budget_tokens": 10000, "display": "omitted"})), - true, - Some(json!({"type": "enabled", "budget_tokens": 10000, "display": "summarized"})), - )] - #[case::display_summarized_is_kept( - Some(json!({"type": "enabled", "display": "summarized"})), - true, - Some(json!({"type": "enabled", "display": "summarized"})), - )] - #[case::disabled_thinking(Some(json!({"type": "disabled"})), true, Some(json!({"type": "disabled"})))] - #[case::flag_off( - Some(json!({"type": "enabled", "budget_tokens": 10000})), - false, - Some(json!({"type": "enabled", "budget_tokens": 10000})), - )] - #[case::flag_off_keeps_callers_display( - Some(json!({"type": "enabled", "display": "omitted"})), - false, - Some(json!({"type": "enabled", "display": "omitted"})), - )] - #[case::no_thinking(None, true, None)] - #[case::non_object_thinking(Some(json!("enabled")), true, Some(json!("enabled")))] - fn reasoning_auto_summary_marks_active_thinking_as_summarized( - #[case] thinking: Option, - #[case] enabled: bool, - #[case] expected: Option, - ) { - assert_eq!(with_reasoning_auto_summary(thinking, enabled), expected); - } - #[rstest] #[case::nothing_forwarded( &[], @@ -403,39 +291,6 @@ mod tests { ); } - #[rstest] - fn sanitize_request_shapes_messages_metadata_and_thinking(shaping: MessagesShaping) { - let sanitized = sanitize_request( - request(json!({ - "model": "m", - "messages": [{"role": "assistant", "content": [ - {"type": "text", "text": ""}, - {"type": "tool_use", "id": "functions.Bash:0", "name": "Bash", "input": {}} - ]}], - "metadata": {"user_id": "u", "trace_id": "t"}, - "thinking": {"type": "enabled", "budget_tokens": 1024}, - "safeguards": [{"type": "dangerous_tool_use"}] - })), - &MessagesShaping { - reasoning_auto_summary: true, - ..shaping - }, - ) - .unwrap(); - assert_eq!( - serde_json::to_value(sanitized).unwrap(), - json!({ - "model": "m", - "messages": [{"role": "assistant", "content": [ - {"type": "tool_use", "id": "functions_Bash_0", "name": "Bash", "input": {}} - ]}], - "metadata": {"user_id": "u"}, - "thinking": {"type": "enabled", "budget_tokens": 1024, "display": "summarized"}, - "safeguards": [{"type": "dangerous_tool_use"}] - }) - ); - } - #[rstest] #[case::top_level_and_nested_paths( json!({ @@ -498,6 +353,54 @@ mod tests { ); } + #[rstest] + #[case::model_prefix_picks_the_provider( + "azure_ai/claude-test", + None, + &[("x-priority", "extra"), ("x-scoped", "azure_ai")] + )] + #[case::explicit_provider( + "claude-test", + Some("anthropic"), + &[("x-priority", "scoped"), ("x-scoped", "anthropic")] + )] + #[case::provider_prefix_on_an_anthropic_model( + "anthropic/claude-test", + None, + &[("x-priority", "scoped"), ("x-scoped", "anthropic")] + )] + fn provider_specific_headers_follow_the_resolved_provider( + shaping: MessagesShaping, + #[case] model: &str, + #[case] custom_llm_provider: Option<&str>, + #[case] expected: &[(&str, &str)], + ) { + let configured: ProviderSpecificHeaders = serde_json::from_value(json!([ + {"custom_llm_provider": "azure_ai", "extra_headers": {"x-scoped": "azure_ai"}}, + {"custom_llm_provider": "anthropic", "extra_headers": {"x-scoped": "anthropic", "x-priority": "scoped"}} + ])) + .unwrap(); + let prepared = prepare(MessagesRequest { + model, + body: json!({"model": model, "messages": [{"role": "user", "content": "hi"}], "max_tokens": 16}), + api_key: Some("sk-test"), + api_base: Some("https://resource.services.ai.azure.com"), + custom_llm_provider, + extra_headers: Some(serde_json::from_value(json!({"x-priority": "extra"})).unwrap()), + provider_specific_header: Some(configured), + timeout: None, + shaping, + }) + .unwrap(); + let caller_headers: Vec<(&str, &str)> = prepared + .upstream_headers + .iter() + .filter(|(name, _)| matches!(name.as_str(), "x-priority" | "x-scoped")) + .map(|(name, value)| (name.as_str(), value.as_str())) + .collect(); + assert_eq!(caller_headers, expected); + } + #[rstest] fn prepared_body_carries_the_provider_stripped_model(shaping: MessagesShaping) { assert_eq!( diff --git a/litellm-rust/crates/core/src/messages/route.rs b/litellm-rust/crates/core/src/messages/route.rs index 59c8e413e17..8cd3eaf3aa3 100644 --- a/litellm-rust/crates/core/src/messages/route.rs +++ b/litellm-rust/crates/core/src/messages/route.rs @@ -1,4 +1,7 @@ -use std::{sync::Mutex, time::Duration}; +use std::{ + sync::{Arc, Mutex}, + time::Duration, +}; use bytes::Bytes; use litellm_auth::SecretValue; @@ -9,14 +12,18 @@ use litellm_host::{ machine::{HostChannel, MachineFault, RouteMachine}, route::Route, }; -use litellm_types::llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse; +use litellm_secrets::source::SecretSource; +use litellm_types::{ + llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse, + utils::ProviderSpecificHeaders, +}; use serde_json::{Map, Value}; use super::{ Error, common_utils::messages_provider_config, handler::{decode_response, network, provider_error, send}, - prepare::prepare_provider_request, + prepare::{prepare_provider_request, resolve_provider}, types::{MessagesRequest, MessagesShaping}, }; use crate::constants::ANTHROPIC_MESSAGES_PROVIDER; @@ -38,6 +45,7 @@ pub struct MessagesCall { pub api_base: Option, pub custom_llm_provider: Option, pub extra_headers: Option>, + pub provider_specific_header: Option, pub timeout: Option, pub shaping: MessagesShaping, } @@ -121,23 +129,33 @@ impl Host for LocalMessagesHost { } } -pub fn messages_machine() -> MessagesMachine { - RouteMachine::new(|host| Box::pin(execute(host))) +pub fn messages_machine(secrets: Arc) -> MessagesMachine { + RouteMachine::new(move |host| Box::pin(execute(host, secrets.clone()))) } -async fn execute(host: MessagesHost) -> Result { +async fn execute( + host: MessagesHost, + secrets: Arc, +) -> Result { let MessagesOpResult::Request(call) = host.route(MessagesOp::ProjectRequest).await?; let stream = call.streams(); - let request = prepare_provider_request(MessagesRequest { - model: &call.model, - body: Value::Object(call.body.clone()), - api_key: call.api_key.as_deref(), - api_base: call.api_base.as_deref(), - custom_llm_provider: call.custom_llm_provider.as_deref(), - extra_headers: call.extra_headers.clone(), - timeout: call.timeout, - shaping: call.shaping.clone(), - })?; + let resolved = resolve_provider(&call.model, call.custom_llm_provider.as_deref())?; + let secrets = secrets.resolve(resolved.config.secret_names()).await?; + let request = prepare_provider_request( + MessagesRequest { + model: &call.model, + body: Value::Object(call.body.clone()), + api_key: call.api_key.as_deref(), + api_base: call.api_base.as_deref(), + custom_llm_provider: call.custom_llm_provider.as_deref(), + extra_headers: call.extra_headers.clone(), + provider_specific_header: call.provider_specific_header.clone(), + timeout: call.timeout, + shaping: call.shaping.clone(), + }, + resolved, + secrets.as_ref(), + )?; if stream && request.provider != ANTHROPIC_MESSAGES_PROVIDER { return Err(Error::Unsupported("streaming messages for this provider")); } diff --git a/litellm-rust/crates/core/src/messages/tests.rs b/litellm-rust/crates/core/src/messages/tests.rs index 59b3d03ef9d..ce48752864a 100644 --- a/litellm-rust/crates/core/src/messages/tests.rs +++ b/litellm-rust/crates/core/src/messages/tests.rs @@ -1,6 +1,8 @@ -use std::time::Duration; +use std::{sync::Arc, time::Duration}; +use futures_util::future::BoxFuture; use litellm_http::request::{has_bearer_auth, has_header}; +use litellm_secrets::{SecretValue, source::SecretSource}; use serde_json::{Map, Value, json}; use tokio::{ io::{AsyncReadExt, AsyncWriteExt}, @@ -11,9 +13,131 @@ use super::{ Error, common_utils::{messages_provider_config, string_headers, truncate_error_body}, messages, + route::{LocalMessagesHost, MessagesCall, MessagesOutput, messages_machine}, }; use crate::messages::types::{MessagesRequest, MessagesShaping}; +struct RecordingSecrets { + values: Vec<(&'static str, String)>, + fails: bool, + requested: std::sync::Mutex>, +} + +impl RecordingSecrets { + fn new(values: Vec<(&'static str, String)>, fails: bool) -> Self { + Self { + values, + fails, + requested: std::sync::Mutex::new(Vec::new()), + } + } +} + +impl SecretSource for RecordingSecrets { + fn get_secret_str<'a>( + &'a self, + name: &'a str, + ) -> BoxFuture<'a, Result, litellm_secrets::Error>> { + Box::pin(async move { + self.requested.lock().unwrap().push(name.to_string()); + if self.fails { + return Err(litellm_secrets::Error::ManagedSecretMissing); + } + Ok(self + .values + .iter() + .find(|(key, _)| *key == name) + .map(|(_, value)| SecretValue::new(value.clone()))) + }) + } +} + +fn secrets_call() -> MessagesCall { + let Value::Object(body) = json!({ + "model": "claude-sonnet-4-5", + "max_tokens": 16, + "messages": [{"role": "user", "content": "hi"}] + }) else { + unreachable!("literal object") + }; + MessagesCall { + model: "claude-sonnet-4-5".into(), + body, + api_key: None, + api_base: None, + custom_llm_provider: Some("anthropic".into()), + extra_headers: None, + provider_specific_header: None, + timeout: Some(Duration::from_secs(5)), + shaping: MessagesShaping::default(), + } +} + +#[tokio::test] +async fn route_reads_the_provider_credential_and_base_from_the_secret_source() { + let listener = TcpListener::bind("127.0.0.1:0").await.expect("binds"); + let addr = listener.local_addr().expect("addr"); + let server = tokio::spawn(async move { + let (mut socket, _) = listener.accept().await.expect("accepts request"); + let request = read_http_request(&mut socket).await; + let response_body = r#"{"id":"msg_1","type":"message","role":"assistant","content":[],"model":"claude-sonnet-4-5","stop_reason":"end_turn","usage":{"input_tokens":1,"output_tokens":1}}"#; + socket + .write_all(write_response(response_body).as_bytes()) + .await + .expect("writes response"); + request + }); + let secrets = Arc::new(RecordingSecrets::new( + vec![ + ("ANTHROPIC_API_KEY", "sk-from-manager".to_string()), + ("ANTHROPIC_BASE_URL", format!("http://{addr}")), + ], + false, + )); + + let output = litellm_host::run::run( + messages_machine(secrets.clone()), + &LocalMessagesHost::new(secrets_call()), + ) + .await + .expect("messages request succeeds"); + + assert!(matches!(output, MessagesOutput::Message(_))); + let request = server.await.expect("server task completes"); + assert!( + request + .to_ascii_lowercase() + .contains("x-api-key: sk-from-manager"), + "{request}" + ); + let requested = secrets.requested.lock().unwrap().clone(); + assert_eq!( + requested, + messages_provider_config("anthropic") + .unwrap() + .secret_names() + .iter() + .map(ToString::to_string) + .collect::>() + ); +} + +#[tokio::test] +async fn route_surfaces_a_secret_manager_failure_before_the_call() { + let Err(error) = litellm_host::run::run( + messages_machine(Arc::new(RecordingSecrets::new(Vec::new(), true))), + &LocalMessagesHost::new(secrets_call()), + ) + .await + else { + panic!("a secret manager failure fails the call"); + }; + assert!( + matches!(&error, Error::Secret(source) if matches!(source.source_error(), litellm_secrets::Error::ManagedSecretMissing)), + "{error:?}" + ); +} + async fn read_http_request(socket: &mut TcpStream) -> String { let mut request = Vec::new(); let mut buffer = [0_u8; 1024]; @@ -158,6 +282,7 @@ async fn messages_round_trip_builds_azure_request_and_passes_response_through() api_base: Some(&format!("http://{addr}")), custom_llm_provider: Some("azure_ai"), extra_headers: None, + provider_specific_header: None, timeout: Some(Duration::from_secs(5)), shaping: MessagesShaping::default(), }) @@ -215,6 +340,7 @@ async fn messages_round_trip_builds_native_anthropic_request() { api_base: Some(&format!("http://{addr}")), custom_llm_provider: Some("anthropic"), extra_headers: None, + provider_specific_header: None, timeout: Some(Duration::from_secs(5)), shaping: MessagesShaping::default(), }) @@ -269,6 +395,7 @@ async fn messages_does_not_duplicate_auth_when_x_api_key_supplied() { api_base: Some(&format!("http://{addr}")), custom_llm_provider: Some("azure_ai"), extra_headers: Some(headers), + provider_specific_header: None, timeout: Some(Duration::from_secs(5)), shaping: MessagesShaping::default(), }) @@ -324,6 +451,7 @@ async fn messages_forwards_entra_id_bearer_without_requiring_api_key() { api_base: Some(&format!("http://{addr}")), custom_llm_provider: Some("azure_ai"), extra_headers: Some(headers), + provider_specific_header: None, timeout: Some(Duration::from_secs(5)), shaping: MessagesShaping::default(), }) @@ -349,6 +477,7 @@ async fn messages_requires_auth_when_no_key_and_no_header() { api_base: Some("http://127.0.0.1:1"), custom_llm_provider: Some("azure_ai"), extra_headers: None, + provider_specific_header: None, timeout: Some(Duration::from_millis(50)), shaping: MessagesShaping::default(), }) @@ -388,6 +517,7 @@ async fn messages_ignores_malformed_authorization_and_uses_api_key() { api_base: Some(&format!("http://{addr}")), custom_llm_provider: Some("azure_ai"), extra_headers: Some(headers), + provider_specific_header: None, timeout: Some(Duration::from_secs(5)), shaping: MessagesShaping::default(), }) @@ -430,6 +560,7 @@ async fn messages_maps_provider_error_status_to_http_error() { api_base: Some(&format!("http://{addr}")), custom_llm_provider: Some("azure_ai"), extra_headers: None, + provider_specific_header: None, timeout: Some(Duration::from_secs(5)), shaping: MessagesShaping::default(), }) @@ -451,6 +582,7 @@ async fn messages_rejects_unsupported_provider() { api_base: Some("http://127.0.0.1:1"), custom_llm_provider: Some("openai"), extra_headers: None, + provider_specific_header: None, timeout: Some(Duration::from_millis(50)), shaping: MessagesShaping::default(), }) diff --git a/litellm-rust/crates/core/src/messages/types.rs b/litellm-rust/crates/core/src/messages/types.rs index f97ba1cf08a..4a5dd2926e0 100644 --- a/litellm-rust/crates/core/src/messages/types.rs +++ b/litellm-rust/crates/core/src/messages/types.rs @@ -4,6 +4,7 @@ use litellm_llms::{ anthropic::common_utils::AnthropicModelCapabilities, base_llm::anthropic_messages::transformation::BaseAnthropicMessagesConfig, }; +use litellm_types::utils::ProviderSpecificHeaders; use serde::{Deserialize, Serialize}; use serde_json::{Map, Value}; @@ -26,6 +27,7 @@ pub struct MessagesRequest<'a> { pub api_base: Option<&'a str>, pub custom_llm_provider: Option<&'a str>, pub extra_headers: Option>, + pub provider_specific_header: Option, pub timeout: Option, pub shaping: MessagesShaping, } diff --git a/litellm-rust/crates/llms/src/anthropic/experimental_pass_through/messages/handler.rs b/litellm-rust/crates/llms/src/anthropic/experimental_pass_through/messages/handler.rs new file mode 100644 index 00000000000..0e2ab97956a --- /dev/null +++ b/litellm-rust/crates/llms/src/anthropic/experimental_pass_through/messages/handler.rs @@ -0,0 +1,270 @@ +use litellm_types::llms::anthropic_messages::anthropic_request::{ + AnthropicMessage, AnthropicMessagesRequest, +}; +use serde_json::{Value, json}; + +use crate::{ + anthropic::common_utils::{ + flatten_unencrypted_web_search_results, sanitize_tool_use_ids, strip_empty_content_blocks, + strip_provider_specific_fields, + }, + base_llm::chat::transformation::Error, +}; + +pub fn shape_anthropic_messages_request( + request: AnthropicMessagesRequest, + reasoning_auto_summary: bool, +) -> Result { + Ok(AnthropicMessagesRequest { + messages: sanitize_anthropic_messages(request.messages), + metadata: request + .metadata + .as_ref() + .map(validate_anthropic_api_metadata) + .transpose()?, + thinking: with_reasoning_auto_summary(request.thinking, reasoning_auto_summary), + ..request + }) +} + +fn sanitize_anthropic_messages(messages: Vec) -> Vec { + strip_provider_specific_fields(flatten_unencrypted_web_search_results( + sanitize_tool_use_ids(strip_empty_content_blocks(messages)), + )) +} + +fn validate_anthropic_api_metadata(metadata: &Value) -> Result { + let Value::Object(fields) = metadata else { + return Err(Error::InvalidRequest(format!( + "metadata must be an object, got {metadata}" + ))); + }; + match fields.get("user_id") { + None | Some(Value::Null) => Ok(json!({})), + Some(Value::String(user_id)) => Ok(json!({"user_id": user_id})), + Some(other) => Err(Error::InvalidRequest(format!( + "metadata.user_id must be a string, got {other}" + ))), + } +} + +fn with_reasoning_auto_summary(thinking: Option, enabled: bool) -> Option { + let Some(Value::Object(thinking)) = thinking else { + return thinking; + }; + if !enabled || thinking.get("type").and_then(Value::as_str) == Some("disabled") { + return Some(Value::Object(thinking)); + } + Some(Value::Object( + thinking + .into_iter() + .filter(|(key, _)| key != "display") + .chain([("display".to_string(), json!("summarized"))]) + .collect(), + )) +} + +#[cfg(test)] +mod tests { + use rstest::rstest; + + use super::*; + + fn messages(value: Value) -> Vec { + serde_json::from_value(value).unwrap() + } + + fn request(body: Value) -> AnthropicMessagesRequest { + serde_json::from_value(body).unwrap() + } + + #[rstest] + #[case::empty_text_next_to_a_tool_use( + json!([{"role": "assistant", "content": [ + {"type": "text", "text": " "}, + {"type": "tool_use", "id": "t", "name": "B", "input": {}} + ]}]), + json!([{"role": "assistant", "content": [ + {"type": "tool_use", "id": "t", "name": "B", "input": {}} + ]}]), + )] + #[case::cross_provider_tool_ids( + json!([ + {"role": "assistant", "content": [{"type": "tool_use", "id": "functions.Bash:0", "name": "Bash", "input": {}}]}, + {"role": "user", "content": [{"type": "tool_result", "tool_use_id": "functions.Bash:0", "content": "ok"}]} + ]), + json!([ + {"role": "assistant", "content": [{"type": "tool_use", "id": "functions_Bash_0", "name": "Bash", "input": {}}]}, + {"role": "user", "content": [{"type": "tool_result", "tool_use_id": "functions_Bash_0", "content": "ok"}]} + ]), + )] + #[case::replayed_unencrypted_web_search_results( + json!([ + {"role": "user", "content": "latest litellm version?"}, + {"role": "assistant", "content": [ + {"type": "server_tool_use", "id": "srvtoolu_1", "name": "web_search", "input": {"query": "latest litellm version"}}, + {"type": "web_search_tool_result", "tool_use_id": "srvtoolu_1", "content": [{ + "type": "web_search_result", + "url": "https://github.com/BerriAI/litellm/releases", + "title": "Releases", + "page_age": null, + "encrypted_content": "", + "snippet": "Latest release v1.95.0" + }]} + ]}, + {"role": "user", "content": "which version?"} + ]), + json!([ + {"role": "user", "content": "latest litellm version?"}, + {"role": "assistant", "content": [{ + "type": "text", + "text": "Web search results for 'latest litellm version':\n\nTitle: Releases\nURL: https://github.com/BerriAI/litellm/releases\nSnippet: Latest release v1.95.0" + }]}, + {"role": "user", "content": "which version?"} + ]), + )] + #[case::replayed_provider_specific_fields( + json!([ + {"role": "assistant", "content": [{ + "type": "tool_use", "id": "toolu_01", "name": "get_weather", "input": {"city": "Paris"}, + "provider_specific_fields": {"signature": "sig_abc"} + }]}, + {"role": "user", "content": [{"type": "tool_result", "tool_use_id": "toolu_01", "content": "Sunny"}]} + ]), + json!([ + {"role": "assistant", "content": [{"type": "tool_use", "id": "toolu_01", "name": "get_weather", "input": {"city": "Paris"}}]}, + {"role": "user", "content": [{"type": "tool_result", "tool_use_id": "toolu_01", "content": "Sunny"}]} + ]), + )] + #[case::ids_are_normalized_before_web_search_results_flatten( + json!([ + {"role": "user", "content": "run it"}, + {"role": "assistant", "content": [ + {"type": "thinking", "thinking": "", "signature": "sig"}, + {"type": "text", "text": ""}, + {"type": "tool_use", "id": "functions.Bash:0", "name": "Bash", "input": {}, "provider_specific_fields": {"x": 1}}, + {"type": "server_tool_use", "id": "srv.1", "name": "web_search", "input": {"query": "q"}, "provider_specific_fields": {"x": 2}}, + {"type": "web_search_tool_result", "tool_use_id": "srv.1", "provider_specific_fields": {"x": 3}, "content": [ + {"type": "web_search_result", "url": "u", "title": "", "encrypted_content": "", "provider_specific_fields": {"x": 4}} + ]} + ]}, + {"role": "user", "content": [{"type": "tool_result", "tool_use_id": "functions.Bash:0", "content": "ok"}]}, + {"role": "assistant", "content": [{"type": "text", "text": " "}]} + ]), + json!([ + {"role": "user", "content": "run it"}, + {"role": "assistant", "content": [ + {"type": "tool_use", "id": "functions_Bash_0", "name": "Bash", "input": {}}, + {"type": "server_tool_use", "id": "srv_1", "name": "web_search", "input": {"query": "q"}}, + {"type": "text", "text": "Web search results:\n\nURL: u"} + ]}, + {"role": "user", "content": [{"type": "tool_result", "tool_use_id": "functions_Bash_0", "content": "ok"}]} + ]), + )] + fn sanitize_anthropic_messages_cleans_replayed_history( + #[case] history: Value, + #[case] expected: Value, + ) { + assert_eq!( + serde_json::to_value(sanitize_anthropic_messages(messages(history))).unwrap(), + expected + ); + } + + #[rstest] + #[case::keeps_only_user_id(json!({"user_id": "u-1", "trace_id": "internal"}), Ok(json!({"user_id": "u-1"})))] + #[case::null_user_id(json!({"user_id": null, "trace_id": "internal"}), Ok(json!({})))] + #[case::no_user_id(json!({"trace_id": "internal"}), Ok(json!({})))] + #[case::empty(json!({}), Ok(json!({})))] + #[case::numeric_user_id( + json!({"user_id": 123}), + Err(Error::InvalidRequest("metadata.user_id must be a string, got 123".to_string())), + )] + #[case::boolean_user_id( + json!({"user_id": true}), + Err(Error::InvalidRequest("metadata.user_id must be a string, got true".to_string())), + )] + #[case::not_an_object( + json!(["u-1"]), + Err(Error::InvalidRequest(r#"metadata must be an object, got ["u-1"]"#.to_string())), + )] + fn validate_anthropic_api_metadata_passes_only_a_string_user_id( + #[case] metadata: Value, + #[case] expected: Result, + ) { + assert_eq!(validate_anthropic_api_metadata(&metadata), expected); + } + + #[rstest] + #[case::adaptive( + Some(json!({"type": "adaptive", "budget_tokens": 5000})), + true, + Some(json!({"type": "adaptive", "budget_tokens": 5000, "display": "summarized"})), + )] + #[case::enabled( + Some(json!({"type": "enabled", "budget_tokens": 10000})), + true, + Some(json!({"type": "enabled", "budget_tokens": 10000, "display": "summarized"})), + )] + #[case::no_type(Some(json!({})), true, Some(json!({"display": "summarized"})))] + #[case::display_omitted_is_overridden( + Some(json!({"type": "enabled", "budget_tokens": 10000, "display": "omitted"})), + true, + Some(json!({"type": "enabled", "budget_tokens": 10000, "display": "summarized"})), + )] + #[case::display_summarized_is_kept( + Some(json!({"type": "enabled", "display": "summarized"})), + true, + Some(json!({"type": "enabled", "display": "summarized"})), + )] + #[case::disabled_thinking(Some(json!({"type": "disabled"})), true, Some(json!({"type": "disabled"})))] + #[case::flag_off( + Some(json!({"type": "enabled", "budget_tokens": 10000})), + false, + Some(json!({"type": "enabled", "budget_tokens": 10000})), + )] + #[case::flag_off_keeps_callers_display( + Some(json!({"type": "enabled", "display": "omitted"})), + false, + Some(json!({"type": "enabled", "display": "omitted"})), + )] + #[case::no_thinking(None, true, None)] + #[case::non_object_thinking(Some(json!("enabled")), true, Some(json!("enabled")))] + fn reasoning_auto_summary_marks_active_thinking_as_summarized( + #[case] thinking: Option, + #[case] enabled: bool, + #[case] expected: Option, + ) { + assert_eq!(with_reasoning_auto_summary(thinking, enabled), expected); + } + + #[test] + fn shaping_cleans_messages_metadata_and_thinking() { + let sanitized = shape_anthropic_messages_request( + request(json!({ + "model": "m", + "messages": [{"role": "assistant", "content": [ + {"type": "text", "text": ""}, + {"type": "tool_use", "id": "functions.Bash:0", "name": "Bash", "input": {}} + ]}], + "metadata": {"user_id": "u", "trace_id": "t"}, + "thinking": {"type": "enabled", "budget_tokens": 1024}, + "safeguards": [{"type": "dangerous_tool_use"}] + })), + true, + ) + .unwrap(); + assert_eq!( + serde_json::to_value(sanitized).unwrap(), + json!({ + "model": "m", + "messages": [{"role": "assistant", "content": [ + {"type": "tool_use", "id": "functions_Bash_0", "name": "Bash", "input": {}} + ]}], + "metadata": {"user_id": "u"}, + "thinking": {"type": "enabled", "budget_tokens": 1024, "display": "summarized"}, + "safeguards": [{"type": "dangerous_tool_use"}] + }) + ); + } +} diff --git a/litellm-rust/crates/llms/src/anthropic/experimental_pass_through/messages/headers.rs b/litellm-rust/crates/llms/src/anthropic/experimental_pass_through/messages/headers.rs index 815fcc48c6f..8d48d7a0f5c 100644 --- a/litellm-rust/crates/llms/src/anthropic/experimental_pass_through/messages/headers.rs +++ b/litellm-rust/crates/llms/src/anthropic/experimental_pass_through/messages/headers.rs @@ -1,12 +1,16 @@ use litellm_types::llms::anthropic_messages::anthropic_request::AnthropicMessagesRequest; use serde_json::Value; -use crate::anthropic::{ - ANTHROPIC_OAUTH_TOKEN_PREFIX, - common_utils::{ - ANTHROPIC_OAUTH_BETA_HEADER, beta, has_advisor_tool, is_anthropic_oauth_key, - is_tool_search_used, join_beta_values, requires_native_compaction_beta, split_beta_values, +use crate::{ + anthropic::{ + ANTHROPIC_OAUTH_TOKEN_PREFIX, + common_utils::{ + ANTHROPIC_OAUTH_BETA_HEADER, beta, has_advisor_tool, is_anthropic_oauth_key, + is_tool_search_used, join_beta_values, requires_native_compaction_beta, + split_beta_values, + }, }, + base_llm::anthropic_messages::transformation::Headers, }; const ANTHROPIC_API_KEY_ENV: &str = "ANTHROPIC_API_KEY"; @@ -16,8 +20,6 @@ const AUTHORIZATION: &str = "authorization"; const API_KEY_HEADER: &str = "x-api-key"; const DIRECT_BROWSER_ACCESS_HEADER: &str = "anthropic-dangerous-direct-browser-access"; -pub type Headers = Vec<(String, String)>; - fn header_value<'a>(headers: &'a [(String, String)], name: &str) -> Option<&'a str> { headers .iter() diff --git a/litellm-rust/crates/llms/src/anthropic/experimental_pass_through/messages/mod.rs b/litellm-rust/crates/llms/src/anthropic/experimental_pass_through/messages/mod.rs index 26cc601b06e..5adf5fda16f 100644 --- a/litellm-rust/crates/llms/src/anthropic/experimental_pass_through/messages/mod.rs +++ b/litellm-rust/crates/llms/src/anthropic/experimental_pass_through/messages/mod.rs @@ -1,3 +1,4 @@ +pub mod handler; pub mod headers; pub mod streaming_iterator; pub mod thinking; diff --git a/litellm-rust/crates/llms/src/anthropic/experimental_pass_through/messages/transformation.rs b/litellm-rust/crates/llms/src/anthropic/experimental_pass_through/messages/transformation.rs index d7fb2fc40b6..59280c04a70 100644 --- a/litellm-rust/crates/llms/src/anthropic/experimental_pass_through/messages/transformation.rs +++ b/litellm-rust/crates/llms/src/anthropic/experimental_pass_through/messages/transformation.rs @@ -3,7 +3,7 @@ use litellm_types::llms::anthropic_messages::anthropic_request::AnthropicMessage use serde_json::{Map, Value, json}; use super::{ - headers::{Headers, authenticate, with_feature_betas}, + headers::{authenticate, with_feature_betas}, thinking::{ThinkingBudgets, ThinkingContext, translate_thinking}, }; use crate::{ @@ -13,13 +13,14 @@ use crate::{ }, base_llm::{ anthropic_messages::transformation::{ - BaseAnthropicMessagesConfig, MessagesTransformContext, + BaseAnthropicMessagesConfig, Headers, MessagesTransformContext, }, chat::transformation::Error, }, }; const ANTHROPIC_API_KEY_ENV: &str = "ANTHROPIC_API_KEY"; +const ANTHROPIC_AUTH_TOKEN_ENV: &str = "ANTHROPIC_AUTH_TOKEN"; const ANTHROPIC_API_BASE_ENV: &str = "ANTHROPIC_API_BASE"; const ANTHROPIC_BASE_URL_ENV: &str = "ANTHROPIC_BASE_URL"; const DEFAULT_ANTHROPIC_API_BASE: &str = "https://api.anthropic.com"; @@ -96,6 +97,15 @@ impl BaseAnthropicMessagesConfig for AnthropicMessagesConfig { resolve_anthropic_api_key(api_key, env_lookup).map_err(Error::from) } + fn secret_names(&self) -> &'static [&'static str] { + &[ + ANTHROPIC_API_KEY_ENV, + ANTHROPIC_AUTH_TOKEN_ENV, + ANTHROPIC_API_BASE_ENV, + ANTHROPIC_BASE_URL_ENV, + ] + } + fn authenticate( &self, headers: Headers, @@ -856,4 +866,26 @@ mod tests { ] ); } + + #[test] + fn secret_names_cover_every_credential_and_base_lookup() { + let requested = std::cell::RefCell::new(Vec::::new()); + let record = |name: &str| -> Option { + requested.borrow_mut().push(name.to_string()); + None + }; + let _ = ANTHROPIC_MESSAGES_CONFIG.authenticate(Vec::new(), None, &record); + let _ = ANTHROPIC_MESSAGES_CONFIG.get_complete_url(None, "claude", &record); + let requested = requested.into_inner(); + assert!(!requested.is_empty()); + let undeclared: Vec<&String> = requested + .iter() + .filter(|name| { + !ANTHROPIC_MESSAGES_CONFIG + .secret_names() + .contains(&name.as_str()) + }) + .collect(); + assert_eq!(undeclared, Vec::<&String>::new()); + } } diff --git a/litellm-rust/crates/llms/src/azure_ai/anthropic/messages_transformation.rs b/litellm-rust/crates/llms/src/azure_ai/anthropic/messages_transformation.rs index aa8de60f24e..c409f7f687e 100644 --- a/litellm-rust/crates/llms/src/azure_ai/anthropic/messages_transformation.rs +++ b/litellm-rust/crates/llms/src/azure_ai/anthropic/messages_transformation.rs @@ -11,7 +11,7 @@ use crate::{ }, base_llm::{ anthropic_messages::transformation::{ - BaseAnthropicMessagesConfig, MessagesAuthStrategy, MessagesTransformContext, + BaseAnthropicMessagesConfig, Headers, MessagesAuthStrategy, MessagesTransformContext, }, chat::transformation::Error, }, @@ -76,6 +76,10 @@ impl BaseAnthropicMessagesConfig for AzureAnthropicMessagesConfig { resolve_azure_api_key(api_key, env_lookup) } + fn secret_names(&self) -> &'static [&'static str] { + &[AZURE_API_KEY_ENV, AZURE_API_BASE_ENV] + } + fn auth_strategy(&self) -> MessagesAuthStrategy { self.anthropic.auth_strategy() } @@ -87,6 +91,10 @@ impl BaseAnthropicMessagesConfig for AzureAnthropicMessagesConfig { fn default_headers(&self) -> &'static [(&'static str, &'static str)] { self.anthropic.default_headers() } + + fn request_headers(&self, headers: Headers, request: &AnthropicMessagesRequest) -> Headers { + self.anthropic.request_headers(headers, request) + } } pub fn resolve_azure_api_key( @@ -520,6 +528,57 @@ mod tests { assert!(err.is_data()); } + #[rstest::rstest] + #[case::compact_context_management_edit( + json!({"context_management": {"edits": [{"type": "compact_20260112"}]}}), + &[], + &[("x-api-key", "k"), ("anthropic-beta", "compact-2026-01-12")] + )] + #[case::forwarded_beta_merged_with_structured_output( + json!({"output_config": {"format": {"type": "json_schema"}}}), + &[("anthropic-beta", "web-search-2025-03-05")], + &[("x-api-key", "k"), ("anthropic-beta", "structured-outputs-2025-11-13,web-search-2025-03-05")] + )] + #[case::no_feature_needs_a_beta(json!({}), &[], &[("x-api-key", "k")])] + fn request_headers_carry_the_anthropic_feature_betas( + #[case] fields: serde_json::Value, + #[case] forwarded: &[(&str, &str)], + #[case] expected: &[(&str, &str)], + ) { + let pairs = |pairs: &[(&str, &str)]| -> Vec<(String, String)> { + pairs + .iter() + .map(|(name, value)| (name.to_string(), value.to_string())) + .collect() + }; + let serde_json::Value::Object(fields) = fields else { + panic!("case fields are an object") + }; + let request = request_from(serde_json::Value::Object( + [ + ("model".to_string(), json!("claude-sonnet")), + ("max_tokens".to_string(), json!(16)), + ( + "messages".to_string(), + json!([{"role": "user", "content": "hi"}]), + ), + ] + .into_iter() + .chain(fields) + .collect(), + )); + assert_eq!( + AZURE_ANTHROPIC_MESSAGES_CONFIG.request_headers( + pairs(&[("x-api-key", "k")]) + .into_iter() + .chain(pairs(forwarded)) + .collect(), + &request + ), + pairs(expected) + ); + } + #[test] fn transform_response_passes_through() { let response: AnthropicMessagesResponse = serde_json::from_value(json!({ @@ -541,4 +600,26 @@ mod tests { assert_eq!(value["stop_sequence"], json!(null)); assert_eq!(value["content"][0]["text"], json!("hello")); } + + #[test] + fn secret_names_cover_every_credential_and_base_lookup() { + let requested = std::cell::RefCell::new(Vec::::new()); + let record = |name: &str| -> Option { + requested.borrow_mut().push(name.to_string()); + None + }; + let _ = AZURE_ANTHROPIC_MESSAGES_CONFIG.authenticate(Vec::new(), None, &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 + .iter() + .filter(|name| { + !AZURE_ANTHROPIC_MESSAGES_CONFIG + .secret_names() + .contains(&name.as_str()) + }) + .collect(); + assert_eq!(undeclared, Vec::<&String>::new()); + } } diff --git a/litellm-rust/crates/llms/src/base_llm/anthropic_messages/transformation.rs b/litellm-rust/crates/llms/src/base_llm/anthropic_messages/transformation.rs index 55becdd27e4..8db14687214 100644 --- a/litellm-rust/crates/llms/src/base_llm/anthropic_messages/transformation.rs +++ b/litellm-rust/crates/llms/src/base_llm/anthropic_messages/transformation.rs @@ -61,6 +61,8 @@ pub trait BaseAnthropicMessagesConfig: Sync { env_lookup: &dyn Fn(&str) -> Option, ) -> Result; + fn secret_names(&self) -> &'static [&'static str]; + fn auth_strategy(&self) -> MessagesAuthStrategy { MessagesAuthStrategy::Header("x-api-key") } @@ -117,6 +119,10 @@ mod tests { } impl BaseAnthropicMessagesConfig for StubConfig { + fn secret_names(&self) -> &'static [&'static str] { + &[] + } + fn get_complete_url( &self, _api_base: Option<&str>, @@ -148,6 +154,10 @@ mod tests { struct DefaultsConfig; impl BaseAnthropicMessagesConfig for DefaultsConfig { + fn secret_names(&self) -> &'static [&'static str] { + &[] + } + fn get_complete_url( &self, _api_base: Option<&str>, 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 51122b4e0f6..9d97094aeda 100644 --- a/litellm-rust/crates/python-bridge/src/routes/messages/host.rs +++ b/litellm-rust/crates/python-bridge/src/routes/messages/host.rs @@ -6,13 +6,13 @@ use litellm_core::messages::{ }; use litellm_host_python::{InvokeError, RouteHost, from_py, lookup, to_py}; use litellm_http::transport::Error as TransportError; +use litellm_types::utils::ProviderSpecificHeaders; use pyo3::{ exceptions::{PyException, PyValueError}, gc::{PyTraverseError, PyVisit}, prelude::*, types::{PyBytes, PyDict}, }; -use serde::Deserialize; use serde_json::{Map, Value}; use crate::{ @@ -48,49 +48,14 @@ const BODY_FIELDS: [&str; 22] = [ "safeguards", ]; -#[derive(Deserialize)] -struct ProviderSpecificHeader { - #[serde(default)] - custom_llm_provider: String, - #[serde(default)] - extra_headers: Map, -} - -#[derive(Deserialize)] -#[serde(untagged)] -enum ProviderSpecificHeaders { - One(ProviderSpecificHeader), - Many(Vec), -} - -impl ProviderSpecificHeaders { - fn matching(self, provider: &str) -> impl Iterator { - let entries = match self { - Self::One(entry) => vec![entry], - Self::Many(entries) => entries, - }; - entries - .into_iter() - .filter(move |entry| { - entry - .custom_llm_provider - .split(',') - .any(|scoped| scoped.trim() == provider) - }) - .flat_map(|entry| entry.extra_headers) - } -} - fn merge_headers( forwarded: Option>, extra_headers: Option>, - scoped: impl IntoIterator, ) -> Option> { let merged: Map = forwarded .into_iter() .flatten() .chain(extra_headers.into_iter().flatten()) - .chain(scoped) .collect(); (!merged.is_empty()).then_some(merged) } @@ -162,6 +127,7 @@ impl MessagesRouteHost { api_key: string("api_key")?, api_base: string("api_base")?, extra_headers: self.merged_headers(py, arguments)?, + provider_specific_header: self.provider_specific_header(py, arguments)?, custom_llm_provider, timeout: optional_timeout(timeout), shaping, @@ -180,20 +146,23 @@ impl MessagesRouteHost { .map(|value| from_py(&value)) .transpose() }; - let provider = self.provider(py); - let scoped = lookup(arguments, request, "provider_specific_header")? - .filter(|value| !value.is_none()) - .map(|value| from_py::(&value)) - .transpose()? - .into_iter() - .flat_map(|headers| headers.matching(&provider)); Ok(merge_headers( mapping("headers")?, mapping("extra_headers")?, - scoped, )) } + fn provider_specific_header( + &self, + py: Python<'_>, + arguments: &Bound<'_, PyDict>, + ) -> PyResult> { + lookup(arguments, self.request.bind(py), "provider_specific_header")? + .filter(|value| !value.is_none()) + .map(|value| from_py(&value)) + .transpose() + } + fn shaping( &self, py: Python<'_>, @@ -273,6 +242,11 @@ impl RouteHost for MessagesRouteHost { } fn classify(&self, py: Python<'_>, error: Error) -> PyResult { + if let Error::Secret(source) = &error + && let Some(original) = crate::secrets::python_error(py, source.source_error()) + { + return Ok(original); + } Ok(self.map_failure(py, native_error(py, error)?)) } @@ -299,95 +273,25 @@ mod tests { } #[rstest] - #[case::single_entry_for_the_provider( - json!({"custom_llm_provider": "anthropic", "extra_headers": {"Authorization": "Bearer t", "Custom-Header": "v"}}), - json!({"Authorization": "Bearer t", "Custom-Header": "v"}), - )] - #[case::single_entry_for_another_provider( - json!({"custom_llm_provider": "openai", "extra_headers": {"Authorization": "Bearer t"}}), - json!({}), - )] - #[case::provider_in_a_comma_separated_scope( - json!({"custom_llm_provider": "bedrock,anthropic,vertex_ai", "extra_headers": {"anthropic-beta": "context-1m-2025-08-07"}}), - json!({"anthropic-beta": "context-1m-2025-08-07"}), - )] - #[case::provider_missing_from_a_comma_separated_scope( - json!({"custom_llm_provider": "bedrock,vertex_ai", "extra_headers": {"anthropic-beta": "test"}}), - json!({}), - )] - #[case::scope_with_spaces( - json!({"custom_llm_provider": "bedrock, anthropic , vertex_ai", "extra_headers": {"anthropic-beta": "test"}}), - json!({"anthropic-beta": "test"}), - )] - #[case::scope_names_must_match_exactly( - json!({"custom_llm_provider": "anthropic_text", "extra_headers": {"anthropic-beta": "test"}}), - json!({}), - )] - #[case::entries_scope_independently( - json!([ - {"custom_llm_provider": "anthropic,bedrock,vertex_ai", "extra_headers": {"anthropic-beta": "context-1m-2025-08-07"}}, - {"custom_llm_provider": "bedrock", "extra_headers": {"x-bedrock-only": "no"}}, - {"custom_llm_provider": "anthropic", "extra_headers": {"authorization": "Bearer sk-ant-oat01-fake-token"}} - ]), - json!({"anthropic-beta": "context-1m-2025-08-07", "authorization": "Bearer sk-ant-oat01-fake-token"}), - )] - #[case::later_entries_win( - json!([ - {"custom_llm_provider": "anthropic", "extra_headers": {"x-scoped": "first"}}, - {"custom_llm_provider": "anthropic", "extra_headers": {"x-scoped": "second"}} - ]), - json!({"x-scoped": "second"}), - )] - #[case::empty_list(json!([]), json!({}))] - #[case::entry_without_scope(json!({"extra_headers": {"x-scoped": "yes"}}), json!({}))] - #[case::entry_without_headers(json!({"custom_llm_provider": "anthropic"}), json!({}))] - fn provider_specific_headers_match_the_scoped_provider( - #[case] configured: Value, - #[case] expected: Value, - ) { - let headers: ProviderSpecificHeaders = serde_json::from_value(configured).unwrap(); - assert_eq!( - headers - .matching("anthropic") - .collect::>(), - map(expected) - ); - } - - #[rstest] - #[case::scoped_over_extra_over_forwarded( + #[case::extra_over_forwarded( Some(json!({"X-Priority": "forwarded", "X-Forwarded-Only": "keep"})), Some(json!({"X-Priority": "extra", "X-Extra-Only": "also-keep"})), - json!({"X-Priority": "provider", "X-Provider-Only": "keep-this-too"}), - Some(json!({ - "X-Priority": "provider", - "X-Forwarded-Only": "keep", - "X-Extra-Only": "also-keep", - "X-Provider-Only": "keep-this-too" - })), - )] - #[case::extra_over_forwarded( - Some(json!({"X-Priority": "forwarded"})), - Some(json!({"X-Priority": "extra"})), - json!({}), - Some(json!({"X-Priority": "extra"})), + Some(json!({"X-Priority": "extra", "X-Forwarded-Only": "keep", "X-Extra-Only": "also-keep"})), )] + #[case::only_forwarded(Some(json!({"X-Forwarded": "yes"})), None, Some(json!({"X-Forwarded": "yes"})))] #[case::only_extra_headers( None, Some(json!({"X-Custom-Header": "from-kwargs", "X-Auth-Token": "token123"})), - json!({}), Some(json!({"X-Custom-Header": "from-kwargs", "X-Auth-Token": "token123"})), )] - #[case::only_scoped(None, None, json!({"x-scoped": "yes"}), Some(json!({"x-scoped": "yes"})))] - #[case::nothing(None, Some(json!({})), json!({}), None)] - fn headers_merge_forwarded_then_extra_then_scoped( + #[case::nothing(None, Some(json!({})), None)] + fn headers_merge_forwarded_then_extra( #[case] forwarded: Option, #[case] extra_headers: Option, - #[case] scoped: Value, #[case] expected: Option, ) { assert_eq!( - merge_headers(forwarded.map(map), extra_headers.map(map), map(scoped)), + merge_headers(forwarded.map(map), extra_headers.map(map)), expected.map(map) ); } diff --git a/litellm-rust/crates/python-bridge/src/routes/messages/mod.rs b/litellm-rust/crates/python-bridge/src/routes/messages/mod.rs index 804d883e9ae..fd474e6b2d4 100644 --- a/litellm-rust/crates/python-bridge/src/routes/messages/mod.rs +++ b/litellm-rust/crates/python-bridge/src/routes/messages/mod.rs @@ -39,11 +39,12 @@ fn run_messages( "the Rust Messages route does not serve this provider", )); } + let secrets = crate::secrets::source(py)?; run_legacy_call( py, SURFACE, PublicCall::capture(&request, &args, &kwargs)?, - crate::logger::LoggedMachine::new(messages_machine()), + crate::logger::LoggedMachine::new(messages_machine(secrets)), MessagesRouteHost::new(request.unbind()), asynchronous, ) diff --git a/litellm-rust/crates/types/src/utils.rs b/litellm-rust/crates/types/src/utils.rs index 7f0c18f9f2c..5ca56ec9e49 100644 --- a/litellm-rust/crates/types/src/utils.rs +++ b/litellm-rust/crates/types/src/utils.rs @@ -3,6 +3,21 @@ use serde_json::{Map, Value}; use crate::llms::openai::{ChatCompletionThinkingBlock, ChatCompletionToolCallChunk}; +#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)] +pub struct ProviderSpecificHeader { + #[serde(default)] + pub custom_llm_provider: String, + #[serde(default)] + pub extra_headers: Map, +} + +#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] +#[serde(untagged)] +pub enum ProviderSpecificHeaders { + One(ProviderSpecificHeader), + Many(Vec), +} + /// OpenAI `usage`, including the `prompt_tokens_details` split LiteLLM's Python /// path reports so cost tracking sees the same numbers on either path. #[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)] diff --git a/tests/test_litellm/rust_bridge/messages/test_secrets.py b/tests/test_litellm/rust_bridge/messages/test_secrets.py new file mode 100644 index 00000000000..cf37ed0830b --- /dev/null +++ b/tests/test_litellm/rust_bridge/messages/test_secrets.py @@ -0,0 +1,111 @@ +from __future__ import annotations + +from collections.abc import Awaitable, Mapping +from dataclasses import replace +from types import MappingProxyType +from typing import Final, Protocol, cast # noqa: TID251 # narrows the parametrized path to its protocol + +import httpx +import pytest + +import litellm +from litellm.integrations.custom_secret_manager import CustomSecretManager +from litellm.llms.anthropic.experimental_pass_through.messages.handler import anthropic_messages +from litellm.rust_bridge import settings +from litellm.rust_bridge.messages.entrypoints import NATIVE_AMESSAGES, NATIVE_MESSAGES, LiteLLMMessagesRequest +from litellm.types.secret_managers.main import KeyManagementSettings, KeyManagementSystem +from tests.test_litellm_rust.support.recording_server import ResponseSpec, recording_service +from tests.test_litellm_rust.support.requests import MESSAGES, MESSAGES_MODEL, MESSAGES_RESPONSE + +pytest.importorskip("litellm.rust_bridge._native") + +pytestmark = pytest.mark.usefixtures("local_model_cost_map") + + +class Messages(Protocol): + def __call__(self) -> Awaitable[object]: ... + + +class _ManagedSecrets(CustomSecretManager): + def __init__(self, values: Mapping[str, str]) -> None: + super().__init__(secret_manager_name="rust_bridge_messages_test") + self.values: Final = values + + async def async_read_secret( + self, + secret_name: str, + optional_params: dict[str, object] | None = None, + timeout: float | httpx.Timeout | None = None, + ) -> str | None: + raise AssertionError("get_secret reads custom managers synchronously") + + def sync_read_secret( + self, + secret_name: str, + optional_params: dict[str, object] | None = None, + timeout: float | httpx.Timeout | None = None, + ) -> str | None: + return self.values.get(secret_name) + + +def _native_request() -> LiteLLMMessagesRequest: + return LiteLLMMessagesRequest( + model=MESSAGES_MODEL, + messages=MESSAGES, + max_tokens=8, + stream=None, + api_key=None, + api_base=None, + custom_llm_provider=None, + kwargs=MappingProxyType({}), + ) + + +def _public_kwargs() -> dict[str, object]: + return {"model": MESSAGES_MODEL, "messages": [dict(message) for message in MESSAGES], "max_tokens": 8} + + +async def _python_messages() -> object: + return await anthropic_messages(**_public_kwargs()) + + +async def _rust_messages() -> object: + route: Final = NATIVE_MESSAGES.load() + assert route is not None + return route(_native_request(), (), _public_kwargs()) + + +async def _rust_amessages() -> object: + route: Final = NATIVE_AMESSAGES.load() + assert route is not None + return await route(_native_request(), (), _public_kwargs()) + + +@pytest.fixture( + params=(_python_messages, _rust_messages, _rust_amessages), ids=("python-async", "rust-sync", "rust-async") +) +def messages(request: pytest.FixtureRequest) -> Messages: + return cast(Messages, request.param) + + +async def test_secret_manager_supplies_the_anthropic_key_and_base( + monkeypatch: pytest.MonkeyPatch, messages: Messages +) -> None: + for name in ("ANTHROPIC_API_KEY", "ANTHROPIC_AUTH_TOKEN", "ANTHROPIC_API_BASE", "ANTHROPIC_BASE_URL"): + monkeypatch.delenv(name, raising=False) + with recording_service() as server: + server.default_response = ResponseSpec(body=MESSAGES_RESPONSE) + monkeypatch.setattr( + litellm, + "secret_manager_client", + _ManagedSecrets({"ANTHROPIC_API_KEY": "vault-key", "ANTHROPIC_BASE_URL": server.base_url}), + ) + monkeypatch.setattr(litellm, "_key_management_system", KeyManagementSystem.CUSTOM) + monkeypatch.setattr(litellm, "_key_management_settings", KeyManagementSettings(access_mode="read_only")) + configured: Final = settings.secret_manager + monkeypatch.setattr(settings, "secret_manager", lambda: replace(configured(), native=True)) + + await messages() + + assert len(server.requests) == 1 + assert server.requests[0].headers["x-api-key"] == "vault-key"