mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
fix(rust): read Messages secrets only when the request needs them
The route resolved every declared provider secret before preparing the call, so a caller that passed api_key and api_base still hit the secret manager for all four names, paying its latency and failing the call if a read errored. Python reads each secret lazily through `api_key or get_secret_str(...)`. resolve_on_demand in litellm-secrets now runs the pure prepare step against the secrets fetched so far and fetches only the first name it asks for that is not yet known, repeating until nothing is missing. Reads happen in the same order and with the same short circuits as Python, and the per-config secret_names list is gone since nothing needs to be declared up front
This commit is contained in:
parent
8edd279027
commit
2f421ac770
7 changed files with 209 additions and 95 deletions
|
|
@ -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<MessagesOutput, Error> {
|
||||
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"));
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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::<Vec<_>>()
|
||||
*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(
|
||||
|
|
|
|||
|
|
@ -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::<String>::new());
|
||||
let record = |name: &str| -> Option<String> {
|
||||
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());
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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::<String>::new());
|
||||
let record = |name: &str| -> Option<String> {
|
||||
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());
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -61,8 +61,6 @@ pub trait BaseAnthropicMessagesConfig: Sync {
|
|||
env_lookup: &dyn Fn(&str) -> Option<String>,
|
||||
) -> Result<String, Error>;
|
||||
|
||||
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>,
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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<SecretValue>)],
|
||||
first_missing: OnceCell<String>,
|
||||
}
|
||||
|
||||
impl Lookup for OnDemand<'_> {
|
||||
fn get(&self, name: &str) -> Option<String> {
|
||||
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<T, E>(
|
||||
source: &dyn SecretSource,
|
||||
attempt: impl Fn(&dyn Lookup) -> Result<T, E>,
|
||||
) -> Result<T, E>
|
||||
where
|
||||
E: From<Error>,
|
||||
{
|
||||
let mut fetched: Vec<(String, Option<SecretValue>)> = 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<Vec<String>>,
|
||||
}
|
||||
|
||||
impl RecordingSource {
|
||||
fn new(values: &'static [(&'static str, &'static str)]) -> Self {
|
||||
Self {
|
||||
values,
|
||||
failing: None,
|
||||
reads: Mutex::new(Vec::new()),
|
||||
}
|
||||
}
|
||||
|
||||
fn reads(&self) -> Vec<String> {
|
||||
self.reads.lock().unwrap().clone()
|
||||
}
|
||||
}
|
||||
|
||||
impl SecretSource for RecordingSource {
|
||||
fn get_secret_str<'a>(
|
||||
&'a self,
|
||||
name: &'a str,
|
||||
) -> BoxFuture<'a, Result<Option<SecretValue>, 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<String, Error> {
|
||||
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::<String>::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()])
|
||||
);
|
||||
}
|
||||
}
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue