diff --git a/litellm-rust/crates/core/src/messages/route.rs b/litellm-rust/crates/core/src/messages/route.rs index 8cd3eaf3aa3..127c3a89253 100644 --- a/litellm-rust/crates/core/src/messages/route.rs +++ b/litellm-rust/crates/core/src/messages/route.rs @@ -12,7 +12,7 @@ use litellm_host::{ machine::{HostChannel, MachineFault, RouteMachine}, route::Route, }; -use litellm_secrets::source::SecretSource; +use litellm_secrets::source::{SecretSource, resolve_on_demand}; use litellm_types::{ llms::anthropic_messages::anthropic_response::AnthropicMessagesResponse, utils::ProviderSpecificHeaders, @@ -139,23 +139,24 @@ async fn execute( ) -> Result { let MessagesOpResult::Request(call) = host.route(MessagesOp::ProjectRequest).await?; let stream = call.streams(); - 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(), - )?; + let request = resolve_on_demand(secrets.as_ref(), |secrets| { + 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(), + }, + resolve_provider(&call.model, call.custom_llm_provider.as_deref())?, + secrets, + ) + }) + .await?; 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 ce48752864a..bde2ca581c6 100644 --- a/litellm-rust/crates/core/src/messages/tests.rs +++ b/litellm-rust/crates/core/src/messages/tests.rs @@ -110,18 +110,54 @@ async fn route_reads_the_provider_credential_and_base_from_the_secret_source() { .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::>() + *secrets.requested.lock().unwrap(), + [ + "ANTHROPIC_API_KEY", + "ANTHROPIC_API_BASE", + "ANTHROPIC_BASE_URL" + ] ); } +#[tokio::test] +async fn route_reads_no_secret_when_the_caller_supplies_the_key_and_base() { + 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 failing = Arc::new(RecordingSecrets::new(Vec::new(), true)); + + let output = litellm_host::run::run( + messages_machine(failing.clone()), + &LocalMessagesHost::new(MessagesCall { + api_key: Some("sk-caller".into()), + api_base: Some(format!("http://{addr}")), + ..secrets_call() + }), + ) + .await + .expect("messages request succeeds without touching the secret manager"); + + assert!(matches!(output, MessagesOutput::Message(_))); + assert!( + server + .await + .expect("server task completes") + .to_ascii_lowercase() + .contains("x-api-key: sk-caller") + ); + assert!(failing.requested.lock().unwrap().is_empty()); +} + #[tokio::test] async fn route_surfaces_a_secret_manager_failure_before_the_call() { let Err(error) = litellm_host::run::run( 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 59280c04a70..1f725fc8d08 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 @@ -20,7 +20,6 @@ use crate::{ }; 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"; @@ -97,15 +96,6 @@ 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, @@ -866,26 +856,4 @@ 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 c409f7f687e..77431301660 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 @@ -76,10 +76,6 @@ 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() } @@ -600,26 +596,4 @@ 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 8db14687214..55becdd27e4 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,8 +61,6 @@ 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") } @@ -119,10 +117,6 @@ mod tests { } impl BaseAnthropicMessagesConfig for StubConfig { - fn secret_names(&self) -> &'static [&'static str] { - &[] - } - fn get_complete_url( &self, _api_base: Option<&str>, @@ -154,10 +148,6 @@ 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/secrets/Cargo.toml b/litellm-rust/crates/secrets/Cargo.toml index 17acc01682b..68337f11106 100644 --- a/litellm-rust/crates/secrets/Cargo.toml +++ b/litellm-rust/crates/secrets/Cargo.toml @@ -34,6 +34,7 @@ tokio = { workspace = true, features = ["fs"] } [dev-dependencies] rstest.workspace = true +tokio = { workspace = true, features = ["macros", "rt"] } wiremock = "0.6.5" tempfile = "3" aws-sdk-kms = "1.120.0" diff --git a/litellm-rust/crates/secrets/src/source.rs b/litellm-rust/crates/secrets/src/source.rs index a1615bad055..7d0ef9ea2fb 100644 --- a/litellm-rust/crates/secrets/src/source.rs +++ b/litellm-rust/crates/secrets/src/source.rs @@ -1,4 +1,4 @@ -use std::{collections::HashMap, sync::Arc}; +use std::{cell::OnceCell, collections::HashMap, sync::Arc}; use futures_util::future::{BoxFuture, try_join_all}; use litellm_core_utils::settings::Lookup; @@ -71,3 +71,147 @@ impl Lookup for SecretSnapshot { .map(|value| value.expose().to_owned()) } } + +struct OnDemand<'a> { + fetched: &'a [(String, Option)], + first_missing: OnceCell, +} + +impl Lookup for OnDemand<'_> { + fn get(&self, name: &str) -> Option { + match self.fetched.iter().find(|(fetched, _)| fetched == name) { + Some((_, value)) => value.as_ref().map(|value| value.expose().to_owned()), + None => { + let _ = self.first_missing.set(name.to_owned()); + None + } + } + } +} + +pub async fn resolve_on_demand( + source: &dyn SecretSource, + attempt: impl Fn(&dyn Lookup) -> Result, +) -> Result +where + E: From, +{ + let mut fetched: Vec<(String, Option)> = Vec::new(); + loop { + let lookup = OnDemand { + fetched: &fetched, + first_missing: OnceCell::new(), + }; + let outcome = attempt(&lookup); + let Some(name) = lookup.first_missing.into_inner() else { + return outcome; + }; + let value = source.get_secret_str(&name).await?; + fetched.push((name, value)); + } +} + +#[cfg(test)] +mod tests { + use std::sync::Mutex; + + use rstest::rstest; + + use super::*; + + struct RecordingSource { + values: &'static [(&'static str, &'static str)], + failing: Option<&'static str>, + reads: Mutex>, + } + + impl RecordingSource { + fn new(values: &'static [(&'static str, &'static str)]) -> Self { + Self { + values, + failing: None, + reads: Mutex::new(Vec::new()), + } + } + + fn reads(&self) -> Vec { + self.reads.lock().unwrap().clone() + } + } + + impl SecretSource for RecordingSource { + fn get_secret_str<'a>( + &'a self, + name: &'a str, + ) -> BoxFuture<'a, Result, Error>> { + Box::pin(async move { + self.reads.lock().unwrap().push(name.to_owned()); + if self.failing == Some(name) { + return Err(Error::ManagedSecretMissing); + } + Ok(self + .values + .iter() + .find(|(key, _)| *key == name) + .map(|(_, value)| SecretValue::new(*value))) + }) + } + } + + fn key_then_token(lookup: &dyn Lookup) -> Result { + lookup + .get("KEY") + .or_else(|| lookup.get("TOKEN")) + .ok_or(Error::MissingEnvironment) + } + + #[rstest] + #[case::first_name_found(&[("KEY", "k"), ("TOKEN", "t")], Ok("k"), &["KEY"])] + #[case::falls_through_to_the_second_name(&[("TOKEN", "t")], Ok("t"), &["KEY", "TOKEN"])] + #[case::nothing_found(&[], Err(()), &["KEY", "TOKEN"])] + #[tokio::test] + async fn reads_only_the_names_the_attempt_asks_for_in_order( + #[case] values: &'static [(&'static str, &'static str)], + #[case] expected: Result<&str, ()>, + #[case] expected_reads: &[&str], + ) { + let source = RecordingSource::new(values); + let outcome = resolve_on_demand(&source, key_then_token).await; + assert_eq!( + (outcome.as_deref().map_err(|_| ()), source.reads()), + ( + expected, + expected_reads.iter().map(ToString::to_string).collect() + ) + ); + } + + #[tokio::test] + async fn an_attempt_that_needs_no_secret_reads_none() { + let source = RecordingSource { + failing: Some("KEY"), + ..RecordingSource::new(&[]) + }; + let outcome: Result<&str, Error> = resolve_on_demand(&source, |_| Ok("given")).await; + assert_eq!( + (outcome.ok(), source.reads()), + (Some("given"), Vec::::new()) + ); + } + + #[tokio::test] + async fn a_failed_read_ends_the_resolution() { + let source = RecordingSource { + failing: Some("KEY"), + ..RecordingSource::new(&[("TOKEN", "t")]) + }; + let outcome = resolve_on_demand(&source, key_then_token).await; + assert_eq!( + ( + matches!(outcome, Err(Error::ManagedSecretMissing)), + source.reads() + ), + (true, vec!["KEY".to_string()]) + ); + } +}