From d6ffc554ecf103c73eb47f9a3343bc325980c658 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Tue, 22 Sep 2026 09:02:15 -0700 Subject: [PATCH] feat(rust): align secret manager operation contexts (#42480) * refactor(rust): simplify Python secret callbacks * feat(rust): align secret manager operation contexts * test(rust): port secret name validation cases * fix(rust): restore secret manager CI checks * fix(rust): validate Vault contexts and preserve rotation timeouts --------- Co-authored-by: Yujong Lee --- litellm-rust/Cargo.lock | 1 + litellm-rust/crates/python-bridge/Cargo.toml | 3 +- .../python-bridge/src/routes/ocr/host.rs | 2 +- .../python-bridge/src/secrets/callback.rs | 225 +++++----- .../crates/python-bridge/src/secrets/error.rs | 27 ++ .../crates/python-bridge/src/secrets/mod.rs | 3 + litellm-rust/crates/secrets-aws/src/error.rs | 4 + .../crates/secrets-aws/src/secret_manager.rs | 235 +++++++++-- litellm-rust/crates/secrets-aws/tests/kms.rs | 71 +++- .../secrets-aws/tests/secret_manager.rs | 202 ++++++++- .../crates/secrets-azure/src/key_vault.rs | 5 +- .../crates/secrets-azure/tests/key_vault.rs | 90 ++-- .../crates/secrets-azure/tests/live.rs | 8 +- .../secrets-cyberark/src/secret_manager.rs | 154 +++++-- .../secrets-cyberark/tests/secret_manager.rs | 105 +++-- .../crates/secrets-google/tests/kms.rs | 57 ++- .../secrets-google/tests/secret_manager.rs | 103 +++-- .../crates/secrets-hashicorp/src/error.rs | 4 + .../secrets-hashicorp/src/secret_manager.rs | 211 ++++++++-- .../secrets-hashicorp/tests/secret_manager.rs | 392 ++++++++++++++++-- .../secrets-types/src/base_secret_manager.rs | 32 +- .../crates/secrets-types/src/context.rs | 65 +++ litellm-rust/crates/secrets-types/src/lib.rs | 5 + .../crates/secrets-types/tests/config.rs | 56 ++- .../crates/secrets-types/tests/context.rs | 110 +++++ .../crates/secrets-types/tests/rotation.rs | 97 ++++- litellm-rust/crates/secrets/src/handler.rs | 18 +- .../crates/secrets/tests/resolution.rs | 40 +- 28 files changed, 1839 insertions(+), 486 deletions(-) create mode 100644 litellm-rust/crates/python-bridge/src/secrets/error.rs create mode 100644 litellm-rust/crates/secrets-types/src/context.rs create mode 100644 litellm-rust/crates/secrets-types/tests/context.rs diff --git a/litellm-rust/Cargo.lock b/litellm-rust/Cargo.lock index 37384cbfa53..76004c5d978 100644 --- a/litellm-rust/Cargo.lock +++ b/litellm-rust/Cargo.lock @@ -3065,6 +3065,7 @@ dependencies = [ "serde_json", "serde_with", "sha2 0.10.9", + "thiserror 2.0.19", "tokio", "tokio-tungstenite", "url", diff --git a/litellm-rust/crates/python-bridge/Cargo.toml b/litellm-rust/crates/python-bridge/Cargo.toml index 7846beef28a..dd0975219e4 100644 --- a/litellm-rust/crates/python-bridge/Cargo.toml +++ b/litellm-rust/crates/python-bridge/Cargo.toml @@ -52,8 +52,9 @@ pyo3-async-runtimes.workspace = true reqwest.workspace = true redis = { version = "1.7.0", features = ["tls-rustls"] } serde_json.workspace = true -url.workspace = true +thiserror.workspace = true tokio = { workspace = true, features = ["rt", "sync"] } +url.workspace = true [dev-dependencies] litellm-secrets-aws.workspace = true diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/host.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/host.rs index 325377e5285..8bf99cd355f 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/host.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/host.rs @@ -131,7 +131,7 @@ impl RouteHost for OcrRouteHost { fn classify(&self, py: Python<'_>, error: Error) -> PyResult { if let Error::Secret(source) = &error - && let Some(original) = crate::secrets::callback::python_error(py, source) + && let Some(original) = crate::secrets::python_error(py, source) { return Ok(original); } diff --git a/litellm-rust/crates/python-bridge/src/secrets/callback.rs b/litellm-rust/crates/python-bridge/src/secrets/callback.rs index 2b98081acef..a8e5850fe3f 100644 --- a/litellm-rust/crates/python-bridge/src/secrets/callback.rs +++ b/litellm-rust/crates/python-bridge/src/secrets/callback.rs @@ -1,45 +1,20 @@ -use std::{fmt, future::Future, pin::Pin}; +use std::{future::Future, pin::Pin}; use litellm_core_utils::settings::Lookup; use litellm_secrets::{ Error, ExternalSecretManager, KeyManagementSettings, KeyManagementSystem, Secret, SecretValue, }; -use pyo3::{exceptions::PyBaseException, prelude::*, types::PyDict}; +use pyo3::{prelude::*, types::PyDict}; + +use super::error::external_error; const HANDLER_MODULE: &str = "litellm.secret_managers.secret_manager_handler"; -struct PythonSecretError(Py); - -impl fmt::Debug for PythonSecretError { - fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { - formatter.write_str("PythonSecretError") - } -} - -impl fmt::Display for PythonSecretError { - fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { - formatter.write_str("Python secret manager failed") - } -} - -impl std::error::Error for PythonSecretError {} - -pub(crate) fn python_error(py: Python<'_>, error: &Error) -> Option { - let Error::ExternalManager(source) = error else { - return None; - }; - source - .downcast_ref::() - .map(|error| PyErr::from_value(error.0.clone_ref(py).into_bound(py).into_any())) -} - /// A secret manager whose reads execute in Python: a custom manager, a legacy compatible /// client, or a manually assigned SDK client. pub(crate) struct PythonSecretManager { client: Py, system: Option, - /// The `key_manager` name Python's handler dispatches on. - key_manager: &'static str, settings: Option>, } @@ -52,7 +27,6 @@ impl PythonSecretManager { Self { client, system, - key_manager: system.map_or("local", python_name), settings, } } @@ -78,14 +52,12 @@ impl PythonSecretManager { } let kwargs = PyDict::new(py); kwargs.set_item("client", client)?; - kwargs.set_item("key_manager", self.key_manager)?; + kwargs.set_item("key_manager", self.system.map_or("local", python_name))?; kwargs.set_item("secret_name", name)?; - kwargs.set_item( - "key_management_settings", - self.settings - .as_ref() - .map_or_else(|| py.None(), |settings| settings.clone_ref(py)), - )?; + match &self.settings { + Some(settings) => kwargs.set_item("key_management_settings", settings.bind(py))?, + None => kwargs.set_item("key_management_settings", py.None())?, + } py.import(HANDLER_MODULE)? .getattr("get_secret_from_manager")? .call((), Some(&kwargs))? @@ -123,9 +95,7 @@ impl ExternalSecretManager for PythonSecretManager { Python::attach(|py| { self.read(py, name) .map(|value| value.map(SecretValue::new).map(Secret::String)) - .map_err(|error| { - Error::ExternalManager(Box::new(PythonSecretError(error.into_value(py)))) - }) + .map_err(|error| external_error(py, error)) }) }) } @@ -140,19 +110,27 @@ mod tests { SecretManagerState, SecretResolver, }; use pyo3::{prelude::*, types::PyDict}; + use rstest::rstest; - use super::{HANDLER_MODULE, PythonSecretManager, python_error, python_name}; + use super::{HANDLER_MODULE, PythonSecretManager, python_name}; + use crate::secrets::python_error; + #[rstest] + #[case::value_error("ValueError", None)] + #[case::value_error_with_fallback("ValueError", Some("environment-key"))] + #[case::cancelled("asyncio.CancelledError", None)] + #[case::cancelled_with_fallback("asyncio.CancelledError", Some("environment-key"))] #[tokio::test] - async fn callback_failures_preserve_python_exceptions_even_with_environment_fallback() { + async fn callback_failures_preserve_python_exceptions_even_with_environment_fallback( + #[case] failure_type: &str, + #[case] fallback: Option<&'static str>, + ) { Python::initialize(); - for failure_type in ["ValueError", "asyncio.CancelledError"] { - for fallback in [None, Some("environment-key")] { - let (reader, locals) = Python::attach(|py| { - let locals = PyDict::new(py); - locals.set_item("failure_type", failure_type).unwrap(); - py.run( - c" + let (reader, locals) = Python::attach(|py| { + let locals = PyDict::new(py); + locals.set_item("failure_type", failure_type).unwrap(); + py.run( + c" import asyncio failure = eval(failure_type)('secret manager failed') cause = RuntimeError('original cause') @@ -164,48 +142,46 @@ class Manager: raise failure manager = Manager() ", - Some(&locals), - Some(&locals), - ) - .unwrap(); - let reader = PythonSecretManager::new( - locals.get_item("manager").unwrap().unwrap().unbind(), - None, - None, - ); - (reader, locals.unbind()) - }); - let resolver = SecretResolver::new( - Arc::new(SecretManagerState::new( - SecretManager::External(Arc::new(reader)), - KeyManagementSettings::default(), - )), - Arc::new(move |_: &str| fallback.map(str::to_owned)), - OidcResolver::default(), - ) - .with_failure_policy(FailurePolicy::EnvironmentFallback); - let error = resolver.get_secret("API_KEY", None).await.unwrap_err(); - Python::attach(|py| { - let original = python_error(py, &error).unwrap(); - let locals = locals.bind(py); - assert!( - original - .value(py) - .is(locals.get_item("failure").unwrap().unwrap()) - ); - for (attribute, name) in [("__cause__", "cause"), ("__context__", "context")] { - assert!( - original - .value(py) - .getattr(attribute) - .unwrap() - .is(locals.get_item(name).unwrap().unwrap()) - ); - } - assert!(original.traceback(py).is_some()); - }); + Some(&locals), + Some(&locals), + ) + .unwrap(); + let reader = PythonSecretManager::new( + locals.get_item("manager").unwrap().unwrap().unbind(), + None, + None, + ); + (reader, locals.unbind()) + }); + let resolver = SecretResolver::new( + Arc::new(SecretManagerState::new( + SecretManager::External(Arc::new(reader)), + KeyManagementSettings::default(), + )), + Arc::new(move |_: &str| fallback.map(str::to_owned)), + OidcResolver::default(), + ) + .with_failure_policy(FailurePolicy::EnvironmentFallback); + let error = resolver.get_secret("API_KEY", None).await.unwrap_err(); + Python::attach(|py| { + let original = python_error(py, &error).unwrap(); + let locals = locals.bind(py); + assert!( + original + .value(py) + .is(locals.get_item("failure").unwrap().unwrap()) + ); + for (attribute, name) in [("__cause__", "cause"), ("__context__", "context")] { + assert!( + original + .value(py) + .getattr(attribute) + .unwrap() + .is(locals.get_item(name).unwrap().unwrap()) + ); } - } + assert!(original.traceback(py).is_some()); + }); } /// Installs a fake `get_secret_from_manager` that records its kwargs, runs `body`, and @@ -245,47 +221,57 @@ for name in installed: .unwrap(); } - #[test] - fn python_names_round_trip_through_serde() { - for system in [ - KeyManagementSystem::GoogleKms, - KeyManagementSystem::AzureKeyVault, - KeyManagementSystem::AwsSecretManager, - KeyManagementSystem::GoogleSecretManager, - KeyManagementSystem::HashicorpVault, - KeyManagementSystem::Cyberark, - KeyManagementSystem::Local, - KeyManagementSystem::AwsKms, - KeyManagementSystem::Custom, - ] { - assert_eq!( - serde_json::to_value(system).unwrap(), - serde_json::Value::String(python_name(system).to_owned()) - ); - } + #[rstest] + #[case::google_kms(KeyManagementSystem::GoogleKms)] + #[case::azure_key_vault(KeyManagementSystem::AzureKeyVault)] + #[case::aws_secret_manager(KeyManagementSystem::AwsSecretManager)] + #[case::google_secret_manager(KeyManagementSystem::GoogleSecretManager)] + #[case::hashicorp_vault(KeyManagementSystem::HashicorpVault)] + #[case::cyberark(KeyManagementSystem::Cyberark)] + #[case::local(KeyManagementSystem::Local)] + #[case::aws_kms(KeyManagementSystem::AwsKms)] + #[case::custom(KeyManagementSystem::Custom)] + fn python_names_round_trip_through_serde(#[case] system: KeyManagementSystem) { + assert_eq!( + serde_json::to_value(system).unwrap(), + serde_json::Value::String(python_name(system).to_owned()) + ); } - #[test] - fn custom_readers_without_a_system_are_called_directly() { + #[rstest] + #[case::legacy(None, false)] + #[case::custom(Some(KeyManagementSystem::Custom), true)] + fn direct_readers_receive_compatible_kwargs( + #[case] system: Option, + #[case] expects_optional_params: bool, + ) { Python::initialize(); Python::attach(|py| { let locals = PyDict::new(py); py.run( c" +class Settings: + def model_dump(self): + return {'scope': 'custom'} class Manager: def __init__(self): self.names = [] + self.optional_params = [] def sync_read_secret(self, secret_name, optional_params=None, timeout=None): self.names.append(secret_name) + self.optional_params.append(optional_params) return 'direct-' + secret_name manager = Manager() +settings = Settings() ", Some(&locals), Some(&locals), ) .unwrap(); let manager = locals.get_item("manager").unwrap().unwrap(); - let reader = PythonSecretManager::new(manager.clone().unbind(), None, None); + let settings = expects_optional_params + .then(|| locals.get_item("settings").unwrap().unwrap().unbind()); + let reader = PythonSecretManager::new(manager.clone().unbind(), system, settings); assert_eq!( reader.read(py, "API_KEY").unwrap().as_deref(), Some("direct-API_KEY") @@ -298,6 +284,23 @@ manager = Manager() .unwrap(), ["API_KEY"] ); + let optional_params = manager + .getattr("optional_params") + .unwrap() + .get_item(0) + .unwrap(); + if expects_optional_params { + assert_eq!( + optional_params + .get_item("scope") + .unwrap() + .extract::() + .unwrap(), + "custom" + ); + } else { + assert!(optional_params.is_none()); + } }); } diff --git a/litellm-rust/crates/python-bridge/src/secrets/error.rs b/litellm-rust/crates/python-bridge/src/secrets/error.rs new file mode 100644 index 00000000000..1bf9351ee64 --- /dev/null +++ b/litellm-rust/crates/python-bridge/src/secrets/error.rs @@ -0,0 +1,27 @@ +use std::fmt; + +use litellm_secrets::Error; +use pyo3::{exceptions::PyBaseException, prelude::*}; + +#[derive(thiserror::Error)] +#[error("Python secret manager failed")] +struct PythonSecretError(Py); + +impl fmt::Debug for PythonSecretError { + fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result { + formatter.write_str("PythonSecretError") + } +} + +pub(super) fn external_error(py: Python<'_>, error: PyErr) -> Error { + Error::ExternalManager(Box::new(PythonSecretError(error.into_value(py)))) +} + +pub(crate) fn python_error(py: Python<'_>, error: &Error) -> Option { + let Error::ExternalManager(source) = error else { + return None; + }; + source + .downcast_ref::() + .map(|error| PyErr::from_value(error.0.clone_ref(py).into_bound(py).into_any())) +} diff --git a/litellm-rust/crates/python-bridge/src/secrets/mod.rs b/litellm-rust/crates/python-bridge/src/secrets/mod.rs index f6ca57b08d1..c753015aa9e 100644 --- a/litellm-rust/crates/python-bridge/src/secrets/mod.rs +++ b/litellm-rust/crates/python-bridge/src/secrets/mod.rs @@ -1,3 +1,6 @@ pub(crate) mod callback; pub(crate) mod config; +mod error; pub(crate) mod resolved; + +pub(crate) use error::python_error; diff --git a/litellm-rust/crates/secrets-aws/src/error.rs b/litellm-rust/crates/secrets-aws/src/error.rs index 23595397a13..cc9e6b69786 100644 --- a/litellm-rust/crates/secrets-aws/src/error.rs +++ b/litellm-rust/crates/secrets-aws/src/error.rs @@ -6,6 +6,10 @@ pub enum Error { Auth(#[from] #[redact] litellm_auth_aws::Error), #[error("AWS region is not configured")] MissingRegion, + #[error("AWS Secrets Manager received a non-AWS operation context")] + InvalidOperationContext, + #[error("AWS Secrets Manager was constructed without context-aware configuration")] + OperationContextUnavailable, #[error("KMS response has no plaintext")] MissingPlaintext, #[error("AWS request timed out")] diff --git a/litellm-rust/crates/secrets-aws/src/secret_manager.rs b/litellm-rust/crates/secrets-aws/src/secret_manager.rs index 493cb1d2e8f..76c61534861 100644 --- a/litellm-rust/crates/secrets-aws/src/secret_manager.rs +++ b/litellm-rust/crates/secrets-aws/src/secret_manager.rs @@ -16,7 +16,8 @@ use litellm_auth_aws::constants::{ }; use litellm_core_utils::settings::Lookup; use litellm_secrets_types::{ - BaseSecretManager, KeyManagementSettings, Secret, SecretValue, async_rotate_secret, + AwsOperationContext, BaseSecretManager, KeyManagementSettings, Secret, SecretOperationContext, + SecretValue, SecretWriteContext, async_rotate_secret, }; use serde_json::Value; @@ -25,9 +26,17 @@ use crate::{Error, auth}; #[derive(Clone)] pub struct AwsSecretsManagerV2 { client: Client, + context_client_factory: Option>, write_settings: AwsSecretWriteSettings, } +#[derive(Clone)] +struct ContextClientFactory { + settings: KeyManagementSettings, + environment: Arc, + endpoint_url: Option, +} + #[derive(Clone, Debug, Default)] pub struct AwsSecretWriteSettings { pub kms_key_id: Option, @@ -55,6 +64,19 @@ impl AwsSecretsManagerV2 { pub fn new(client: Client, write_settings: AwsSecretWriteSettings) -> Self { Self { client, + context_client_factory: None, + write_settings, + } + } + + fn with_context_client_factory( + client: Client, + write_settings: AwsSecretWriteSettings, + context_client_factory: ContextClientFactory, + ) -> Self { + Self { + client, + context_client_factory: Some(Box::new(context_client_factory)), write_settings, } } @@ -67,19 +89,18 @@ impl AwsSecretsManagerV2 { if use_aws_secret_manager != Some(true) { return Ok(None); } - let builder = aws_sdk_secretsmanager::Config::builder() - .behavior_version(BehaviorVersion::latest()) - .region(Region::new(auth::region(&settings, environment.as_ref())?)) - .credentials_provider(auth::Credentials::new(&settings, environment.clone())); - let config = match environment.get(AWS_BEDROCK_RUNTIME_ENDPOINT) { - Some(url) => builder - .endpoint_url(url.replace("bedrock-runtime", "secretsmanager")) - .build(), - None => builder.build(), + let context_client_factory = ContextClientFactory { + settings: settings.clone(), + environment: environment.clone(), + endpoint_url: environment + .get(AWS_BEDROCK_RUNTIME_ENDPOINT) + .map(|url| url.replace("bedrock-runtime", "secretsmanager")), }; - Ok(Some(Self::new( - Client::from_conf(config), + let client = context_client_factory.client(&AwsOperationContext::default())?; + Ok(Some(Self::with_context_client_factory( + client, (&settings).into(), + context_client_factory, ))) } @@ -118,7 +139,14 @@ impl AwsSecretsManagerV2 { } pub async fn async_read_secret(&self, name: &str) -> Result, Error> { - match self.client.get_secret_value().secret_id(name).send().await { + Self::async_read_secret_with_client(&self.client, name).await + } + + async fn async_read_secret_with_client( + client: &Client, + name: &str, + ) -> Result, Error> { + match client.get_secret_value().secret_id(name).send().await { Ok(response) => response .secret_string .map(SecretValue::new) @@ -149,8 +177,19 @@ impl AwsSecretsManagerV2 { value: &SecretValue, description: Option<&str>, ) -> Result { - let response = self - .client + self.async_write_secret_with_client_and_tags(&self.client, name, value, description, None) + .await + } + + async fn async_write_secret_with_client_and_tags( + &self, + client: &Client, + name: &str, + value: &SecretValue, + description: Option<&str>, + tags: Option<&BTreeMap>, + ) -> Result { + let response = client .create_secret() .name(name) .secret_string(value.expose()) @@ -161,7 +200,7 @@ impl AwsSecretsManagerV2 { .clone() .filter(|v| !v.is_empty()), ) - .set_tags(self.write_settings.tags.as_ref().map(|tags| { + .set_tags(tags.or(self.write_settings.tags.as_ref()).map(|tags| { tags.iter() .map(|(key, value)| Tag::builder().key(key).value(value).build()) .collect() @@ -171,7 +210,10 @@ impl AwsSecretsManagerV2 { .map_err(|error| Error::Create(Box::new(error)))?; if let Some(regions) = &self.write_settings.replica_regions && !regions.is_empty() - && self.async_replicate_secret(name, regions).await.is_err() + && self + .async_replicate_secret_with_client(client, name, regions) + .await + .is_err() { tracing::warn!("secret created but replication failed"); } @@ -182,11 +224,21 @@ impl AwsSecretsManagerV2 { &self, name: &str, regions: &[String], + ) -> Result, Error> { + self.async_replicate_secret_with_client(&self.client, name, regions) + .await + } + + async fn async_replicate_secret_with_client( + &self, + client: &Client, + name: &str, + regions: &[String], ) -> Result, Error> { if regions.is_empty() { return Ok(None); } - self.client + client .replicate_secret_to_regions() .secret_id(name) .set_add_replica_regions(Some( @@ -206,7 +258,17 @@ impl AwsSecretsManagerV2 { name: &str, value: &SecretValue, ) -> Result { - self.client + self.async_put_secret_value_with_client(&self.client, name, value) + .await + } + + async fn async_put_secret_value_with_client( + &self, + client: &Client, + name: &str, + value: &SecretValue, + ) -> Result { + client .put_secret_value() .secret_id(name) .secret_string(value.expose()) @@ -218,12 +280,22 @@ impl AwsSecretsManagerV2 { pub async fn async_delete_secret( &self, name: &str, - recovery_window_in_days: i64, + recovery_window_in_days: Option, ) -> Result { - self.client + self.async_delete_secret_with_client(&self.client, name, recovery_window_in_days) + .await + } + + async fn async_delete_secret_with_client( + &self, + client: &Client, + name: &str, + recovery_window_in_days: Option, + ) -> Result { + client .delete_secret() .secret_id(name) - .recovery_window_in_days(recovery_window_in_days) + .set_recovery_window_in_days(recovery_window_in_days.map(i64::from)) .send() .await .map_err(|error| Error::Delete(Box::new(error))) @@ -234,17 +306,105 @@ impl AwsSecretsManagerV2 { current_name: &str, new_name: &str, value: &SecretValue, + ) -> Result { + self.async_rotate_secret_with_context( + current_name, + new_name, + value, + &SecretOperationContext::default(), + ) + .await + } + + pub async fn async_rotate_secret_with_context( + &self, + current_name: &str, + new_name: &str, + value: &SecretValue, + context: &SecretOperationContext, ) -> Result { if current_name == new_name { + let client = self.client_for_context(context)?; return self - .async_put_secret_value(current_name, value) + .async_put_secret_value_with_client(&client, current_name, value) .await .map(RotationResponse::Updated); } - async_rotate_secret(self, current_name, new_name, value) + async_rotate_secret(self, current_name, new_name, value, context) .await .map(RotationResponse::Created) } + + fn client_for_context(&self, context: &SecretOperationContext) -> Result { + match context { + SecretOperationContext::Default => Ok(self.client.clone()), + SecretOperationContext::Aws(context) if context == &AwsOperationContext::default() => { + Ok(self.client.clone()) + } + SecretOperationContext::Aws(context) => self + .context_client_factory + .as_ref() + .ok_or(Error::OperationContextUnavailable)? + .client(context), + _ => Err(Error::InvalidOperationContext), + } + } +} + +impl ContextClientFactory { + fn client(&self, context: &AwsOperationContext) -> Result { + let settings = KeyManagementSettings { + aws_region_name: context + .region_name + .clone() + .or_else(|| self.settings.aws_region_name.clone()), + aws_role_name: context + .role_name + .clone() + .or_else(|| self.settings.aws_role_name.clone()), + aws_session_name: context + .session_name + .clone() + .or_else(|| self.settings.aws_session_name.clone()), + aws_external_id: context + .external_id + .clone() + .or_else(|| self.settings.aws_external_id.clone()), + aws_profile_name: context + .profile_name + .clone() + .or_else(|| self.settings.aws_profile_name.clone()), + aws_web_identity_token: context + .web_identity_token + .clone() + .or_else(|| self.settings.aws_web_identity_token.clone()), + aws_sts_endpoint: context + .sts_endpoint + .clone() + .or_else(|| self.settings.aws_sts_endpoint.clone()), + ..self.settings.clone() + }; + let builder = aws_sdk_secretsmanager::Config::builder() + .behavior_version(BehaviorVersion::latest()) + .region(Region::new(auth::region( + &settings, + self.environment.as_ref(), + )?)) + .credentials_provider(auth::Credentials::new(&settings, self.environment.clone())); + let builder = match context.timeout { + Some(timeout) => builder.timeout_config( + aws_sdk_secretsmanager::config::timeout::TimeoutConfig::builder() + .operation_timeout(timeout) + .build(), + ), + None => builder, + }; + let config = match &self.endpoint_url { + Some(endpoint_url) => builder.endpoint_url(endpoint_url.clone()).build(), + None => builder.build(), + }; + Ok(Client::from_conf(config)) + } } impl BaseSecretManager for AwsSecretsManagerV2 { @@ -252,25 +412,40 @@ impl BaseSecretManager for AwsSecretsManagerV2 { type WriteResponse = CreateSecretOutput; type DeleteResponse = DeleteSecretOutput; - async fn async_read_secret(&self, name: &str) -> Result, Error> { - self.async_read_secret(name).await + async fn async_read_secret( + &self, + name: &str, + context: &SecretOperationContext, + ) -> Result, Error> { + let client = self.client_for_context(context)?; + Self::async_read_secret_with_client(&client, name).await } async fn async_write_secret( &self, name: &str, value: &SecretValue, - description: Option<&str>, + context: &SecretWriteContext, ) -> Result { - self.async_write_secret(name, value, description).await + let client = self.client_for_context(&context.operation)?; + self.async_write_secret_with_client_and_tags( + &client, + name, + value, + context.description.as_deref(), + (!context.tags.is_empty()).then_some(&context.tags), + ) + .await } async fn async_delete_secret( &self, name: &str, - recovery_window_in_days: i64, + recovery_window_in_days: Option, + context: &SecretOperationContext, ) -> Result { - self.async_delete_secret(name, recovery_window_in_days) + let client = self.client_for_context(context)?; + self.async_delete_secret_with_client(&client, name, recovery_window_in_days) .await } } diff --git a/litellm-rust/crates/secrets-aws/tests/kms.rs b/litellm-rust/crates/secrets-aws/tests/kms.rs index 39a50297551..687cca104e0 100644 --- a/litellm-rust/crates/secrets-aws/tests/kms.rs +++ b/litellm-rust/crates/secrets-aws/tests/kms.rs @@ -3,12 +3,15 @@ use aws_sdk_kms::{ config::{BehaviorVersion, Credentials, Region, retry::RetryConfig}, }; use base64::{Engine, engine::general_purpose::STANDARD}; -use litellm_secrets_aws::AwsKms; +use litellm_secrets_aws::{AwsKms, Error, load_aws_kms}; +use litellm_secrets_types::KeyManagementSettings; +use rstest::rstest; use wiremock::{ Mock, MockServer, ResponseTemplate, matchers::{body_json, header}, }; +#[rstest] #[tokio::test] async fn kms_decrypt_calls_the_sdk_without_applying_lookup_policy() { let server = MockServer::start().await; @@ -40,20 +43,56 @@ async fn kms_decrypt_calls_the_sdk_without_applying_lookup_policy() { ); } -#[test] -fn disabled_kms_loader_does_not_require_environment_configuration() { - use litellm_secrets_aws::load_aws_kms; - use litellm_secrets_types::KeyManagementSettings; +#[rstest] +#[case::unset(None)] +#[case::disabled(Some(false))] +fn disabled_kms_loader_does_not_require_environment_configuration(#[case] enabled: Option) { use std::sync::Arc; - for enabled in [None, Some(false)] { - assert!( - load_aws_kms( - enabled, - &KeyManagementSettings::default(), - Arc::new(|_: &str| None) - ) - .unwrap() - .is_none() - ); - } + assert!( + load_aws_kms( + enabled, + &KeyManagementSettings::default(), + Arc::new(|_: &str| None) + ) + .unwrap() + .is_none() + ); +} + +#[rstest] +#[case::settings(Some("configured-region"), None)] +#[case::environment(None, Some("environment-region"))] +fn enabled_kms_loader_accepts_either_region_source( + #[case] configured_region: Option<&'static str>, + #[case] environment_region: Option<&'static str>, +) { + use std::sync::Arc; + let settings = KeyManagementSettings { + aws_region_name: configured_region.map(str::to_owned), + ..KeyManagementSettings::default() + }; + let environment = Arc::new(move |name: &str| { + (name == "AWS_REGION_NAME") + .then(|| environment_region.map(str::to_owned)) + .flatten() + }); + + assert!( + load_aws_kms(Some(true), &settings, environment) + .unwrap() + .is_some() + ); +} + +#[rstest] +fn enabled_kms_loader_rejects_missing_region() { + use std::sync::Arc; + assert!(matches!( + load_aws_kms( + Some(true), + &KeyManagementSettings::default(), + Arc::new(|_: &str| None), + ), + Err(Error::MissingRegion) + )); } diff --git a/litellm-rust/crates/secrets-aws/tests/secret_manager.rs b/litellm-rust/crates/secrets-aws/tests/secret_manager.rs index a410767cb5a..7dbc8b61374 100644 --- a/litellm-rust/crates/secrets-aws/tests/secret_manager.rs +++ b/litellm-rust/crates/secrets-aws/tests/secret_manager.rs @@ -1,11 +1,21 @@ -use std::sync::atomic::{AtomicUsize, Ordering}; +use std::{ + sync::{ + Arc, + atomic::{AtomicUsize, Ordering}, + }, + time::Duration, +}; use aws_sdk_secretsmanager::{ Client, config::{BehaviorVersion, Credentials, Region, retry::RetryConfig}, }; use litellm_secrets_aws::{AwsSecretsManagerV2, Error, RotationResponse}; -use litellm_secrets_types::{KeyManagementSettings, SecretValue}; +use litellm_secrets_types::{ + AwsOperationContext, BaseSecretManager, KeyManagementSettings, SecretOperationContext, + SecretValue, SecretWriteContext, +}; +use rstest::{fixture, rstest}; use serde_json::json; use wiremock::{ Mock, MockServer, ResponseTemplate, @@ -25,12 +35,39 @@ fn manager(server: &MockServer, settings: KeyManagementSettings) -> AwsSecretsMa AwsSecretsManagerV2::new(client, (&settings).into()) } -#[rstest::rstest] +fn loaded_manager(server: &MockServer) -> AwsSecretsManagerV2 { + let endpoint_url = server.uri(); + let environment: Arc = + Arc::new(move |name: &str| match name { + "AWS_BEDROCK_RUNTIME_ENDPOINT" => Some(endpoint_url.clone()), + "AWS_ACCESS_KEY_ID" => Some("test".into()), + "AWS_SECRET_ACCESS_KEY" => Some("test".into()), + _ => None, + }); + AwsSecretsManagerV2::load_aws_secret_manager( + Some(true), + KeyManagementSettings { + aws_region_name: Some("us-east-1".into()), + ..Default::default() + }, + environment, + ) + .unwrap() + .unwrap() +} + +#[fixture] +fn default_settings() -> KeyManagementSettings { + KeyManagementSettings::default() +} + +#[rstest] #[case::string_value("KEY", Some("value"))] #[case::missing_value("missing", None)] #[case::non_string_value("BOOL", None)] #[tokio::test] async fn primary_lookup_preserves_read_semantics( + default_settings: KeyManagementSettings, #[case] name: &str, #[case] expected: Option<&str>, ) { @@ -45,7 +82,7 @@ async fn primary_lookup_preserves_read_semantics( .expect(1) .mount(&server) .await; - let manager = manager(&server, KeyManagementSettings::default()); + let manager = manager(&server, default_settings); assert_eq!( manager .read_secret_for_resolver(name, Some("primary"), &|_: &str| None) @@ -57,16 +94,19 @@ async fn primary_lookup_preserves_read_semantics( ); } -#[rstest::rstest] +#[rstest] #[case::access_key("AWS_ACCESS_KEY_ID")] #[case::secret_access_key("AWS_SECRET_ACCESS_KEY")] #[case::region_name("AWS_REGION_NAME")] #[case::region("AWS_REGION")] #[case::bedrock_endpoint("AWS_BEDROCK_RUNTIME_ENDPOINT")] #[tokio::test] -async fn bootstrap_keys_bypass_primary_lookup(#[case] name: &str) { +async fn bootstrap_keys_bypass_primary_lookup( + default_settings: KeyManagementSettings, + #[case] name: &str, +) { let server = MockServer::start().await; - let manager = manager(&server, KeyManagementSettings::default()); + let manager = manager(&server, default_settings); assert_eq!( manager .read_secret_for_resolver(name, Some("primary"), &|_: &str| Some("bootstrap".into())) @@ -79,8 +119,11 @@ async fn bootstrap_keys_bypass_primary_lookup(#[case] name: &str) { ); } +#[rstest] #[tokio::test] -async fn failed_read_returns_none_but_invalid_primary_json_is_an_error() { +async fn failed_read_returns_none_but_invalid_primary_json_is_an_error( + default_settings: KeyManagementSettings, +) { let server = MockServer::start().await; Mock::given(body_partial_json(json!({"SecretId":"missing"}))) .respond_with( @@ -92,7 +135,11 @@ async fn failed_read_returns_none_but_invalid_primary_json_is_an_error() { .respond_with(ResponseTemplate::new(200).set_body_json(json!({"SecretString":"not-json"}))) .mount(&server) .await; - let manager = manager(&server, KeyManagementSettings::default()); + Mock::given(body_partial_json(json!({"SecretId":"no-string"}))) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({"Name":"no-string"}))) + .mount(&server) + .await; + let manager = manager(&server, default_settings); assert!( manager .async_read_secret("missing") @@ -106,10 +153,17 @@ async fn failed_read_returns_none_but_invalid_primary_json_is_an_error() { .await, Err(Error::PrimarySecret) )); + assert!(matches!( + manager.async_read_secret("no-string").await, + Err(Error::MissingString) + )); } +#[rstest] #[tokio::test] -async fn same_name_rotation_uses_put_and_returns_its_response() { +async fn same_name_rotation_uses_put_and_returns_its_response( + default_settings: KeyManagementSettings, +) { let server = MockServer::start().await; Mock::given(header("x-amz-target", "secretsmanager.PutSecretValue")) .and(body_partial_json( @@ -121,7 +175,7 @@ async fn same_name_rotation_uses_put_and_returns_its_response() { .expect(1) .mount(&server) .await; - let response = manager(&server, KeyManagementSettings::default()) + let response = manager(&server, default_settings) .async_rotate_secret("key", "key", &SecretValue::new("replacement")) .await .unwrap(); @@ -132,8 +186,11 @@ async fn same_name_rotation_uses_put_and_returns_its_response() { assert_eq!(server.received_requests().await.unwrap().len(), 1); } +#[rstest] #[tokio::test] -async fn renamed_rotation_reads_creates_verifies_then_deletes() { +async fn renamed_rotation_reads_creates_verifies_then_deletes( + default_settings: KeyManagementSettings, +) { let server = MockServer::start().await; let step = AtomicUsize::new(0); Mock::given(wiremock::matchers::method("POST")) @@ -176,7 +233,7 @@ async fn renamed_rotation_reads_creates_verifies_then_deletes() { .mount(&server) .await; assert!(matches!( - manager(&server, KeyManagementSettings::default()) + manager(&server, default_settings) .async_rotate_secret("old", "new", &SecretValue::new("replacement")) .await .unwrap(), @@ -184,6 +241,7 @@ async fn renamed_rotation_reads_creates_verifies_then_deletes() { )); } +#[rstest] #[tokio::test] async fn creation_passes_tags_and_kms_and_survives_replication_failure() { let server = MockServer::start().await; @@ -230,6 +288,112 @@ async fn creation_passes_tags_and_kms_and_survives_replication_failure() { ); } +#[rstest] +#[tokio::test] +async fn trait_write_uses_typed_write_context(default_settings: KeyManagementSettings) { + let server = MockServer::start().await; + Mock::given(header("x-amz-target", "secretsmanager.CreateSecret")) + .and(body_partial_json(json!({ + "Name": "key", + "SecretString": "value", + "Description": "created by caller", + "Tags": [{"Key": "stage", "Value": "test"}], + }))) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({"Name": "key"}))) + .expect(1) + .mount(&server) + .await; + let context = SecretWriteContext { + description: Some("created by caller".into()), + tags: std::collections::BTreeMap::from([("stage".into(), "test".into())]), + ..Default::default() + }; + let response = BaseSecretManager::async_write_secret( + &manager(&server, default_settings), + "key", + &SecretValue::new("value"), + &context, + ) + .await + .unwrap(); + assert_eq!(response.name(), Some("key")); +} + +#[rstest] +#[tokio::test] +async fn trait_delete_accepts_an_unspecified_recovery_window( + default_settings: KeyManagementSettings, +) { + let server = MockServer::start().await; + Mock::given(header("x-amz-target", "secretsmanager.DeleteSecret")) + .and(body_partial_json(json!({"SecretId": "key"}))) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({"Name": "key"}))) + .expect(1) + .mount(&server) + .await; + let response = BaseSecretManager::async_delete_secret( + &manager(&server, default_settings), + "key", + None, + &SecretOperationContext::default(), + ) + .await + .unwrap(); + assert_eq!(response.name(), Some("key")); +} + +#[rstest] +#[tokio::test] +async fn trait_read_uses_the_aws_region_from_its_operation_context() { + let server = MockServer::start().await; + Mock::given(header("x-amz-target", "secretsmanager.GetSecretValue")) + .respond_with(|request: &wiremock::Request| { + let authorization = request + .headers + .get("authorization") + .unwrap() + .to_str() + .unwrap(); + assert!(authorization.contains("/us-west-2/secretsmanager/aws4_request")); + ResponseTemplate::new(200).set_body_json(json!({"SecretString": "value"})) + }) + .expect(1) + .mount(&server) + .await; + let context = SecretOperationContext::Aws(AwsOperationContext { + region_name: Some("us-west-2".into()), + ..Default::default() + }); + let value = BaseSecretManager::async_read_secret(&loaded_manager(&server), "key", &context) + .await + .unwrap(); + assert_eq!(value.unwrap().expose(), "value"); +} + +#[rstest] +#[tokio::test] +async fn trait_read_applies_the_aws_operation_timeout() { + let server = MockServer::start().await; + Mock::given(header("x-amz-target", "secretsmanager.GetSecretValue")) + .respond_with( + ResponseTemplate::new(200) + .set_delay(Duration::from_secs(1)) + .set_body_json(json!({"SecretString": "late"})), + ) + .expect(1) + .mount(&server) + .await; + let context = SecretOperationContext::Aws(AwsOperationContext { + timeout: Some(Duration::from_millis(30)), + ..Default::default() + }); + assert!(matches!( + BaseSecretManager::async_read_secret(&loaded_manager(&server), "key", &context).await, + Err(Error::Timeout) + )); +} + +#[rstest] #[tokio::test] async fn credential_failures_are_not_swallowed_as_missing_secrets() { use aws_credential_types::provider::{ProvideCredentials, error::CredentialsError, future}; @@ -260,9 +424,9 @@ async fn credential_failures_are_not_swallowed_as_missing_secrets() { assert!(server.received_requests().await.unwrap().is_empty()); } +#[rstest] #[tokio::test] async fn read_timeout_is_an_error_and_cannot_be_mistaken_for_missing() { - use std::time::Duration; let server = MockServer::start().await; Mock::given(wiremock::matchers::method("POST")) .respond_with( @@ -291,12 +455,16 @@ async fn read_timeout_is_an_error_and_cannot_be_mistaken_for_missing() { )); } -#[rstest::rstest] +#[rstest] #[case::denied(400, "AccessDeniedException")] #[case::throttled(400, "ThrottlingException")] #[case::unavailable(503, "ServiceUnavailableException")] #[tokio::test] -async fn service_failures_remain_errors(#[case] status: u16, #[case] code: &str) { +async fn service_failures_remain_errors( + default_settings: KeyManagementSettings, + #[case] status: u16, + #[case] code: &str, +) { let server = MockServer::start().await; Mock::given(header("x-amz-target", "secretsmanager.GetSecretValue")) .respond_with(ResponseTemplate::new(status).set_body_json(json!({"__type":code}))) @@ -304,7 +472,7 @@ async fn service_failures_remain_errors(#[case] status: u16, #[case] code: &str) .mount(&server) .await; assert!(matches!( - manager(&server, KeyManagementSettings::default()) + manager(&server, default_settings) .async_read_secret("key") .await, Err(Error::Read(_)) diff --git a/litellm-rust/crates/secrets-azure/src/key_vault.rs b/litellm-rust/crates/secrets-azure/src/key_vault.rs index e12289b83f5..81ad0331ca7 100644 --- a/litellm-rust/crates/secrets-azure/src/key_vault.rs +++ b/litellm-rust/crates/secrets-azure/src/key_vault.rs @@ -76,10 +76,7 @@ impl AzureKeyVault { .unwrap_or_default() } - pub async fn get_secret_from_azure_key_vault( - &self, - name: &str, - ) -> Result, Error> { + pub async fn get_secret(&self, name: &str) -> Result, Error> { let token = self .auth .get_azure_ad_token(&self.inputs, &|key| self.environment.get(key)) diff --git a/litellm-rust/crates/secrets-azure/tests/key_vault.rs b/litellm-rust/crates/secrets-azure/tests/key_vault.rs index cf9102d0b45..3b199f81264 100644 --- a/litellm-rust/crates/secrets-azure/tests/key_vault.rs +++ b/litellm-rust/crates/secrets-azure/tests/key_vault.rs @@ -2,6 +2,7 @@ use std::sync::Arc; use litellm_secrets_azure::{AzureKeyVault, Error}; use litellm_secrets_types::{Secret, SecretValue}; +use rstest::{fixture, rstest}; use serde::Deserialize; use wiremock::{ Mock, MockServer, ResponseTemplate, @@ -17,6 +18,7 @@ fn manager(server: &MockServer) -> AzureKeyVault { .unwrap() } +#[rstest] #[tokio::test] async fn reads_secret_with_bearer_token_and_api_version() { let server = MockServer::start().await; @@ -32,7 +34,7 @@ async fn reads_secret_with_bearer_token_and_api_version() { .await; let secret = manager(&server) - .get_secret_from_azure_key_vault("OPENAI-API-KEY") + .get_secret("OPENAI-API-KEY") .await .unwrap() .unwrap(); @@ -40,6 +42,7 @@ async fn reads_secret_with_bearer_token_and_api_version() { assert_eq!(secret, Secret::String(SecretValue::new("s3cret"))); } +#[rstest] #[tokio::test] async fn percent_encodes_secret_name_path_segment() { let server = MockServer::start().await; @@ -53,7 +56,7 @@ async fn percent_encodes_secret_name_path_segment() { .await; let secret = manager(&server) - .get_secret_from_azure_key_vault("name/with spaces") + .get_secret("name/with spaces") .await .unwrap() .unwrap(); @@ -61,9 +64,12 @@ async fn percent_encodes_secret_name_path_segment() { assert_eq!(secret.as_str(), Some("value")); } -#[rstest::rstest] +#[rstest] #[case::not_found(404, None)] +#[case::unauthorized(401, Some(401))] #[case::forbidden(403, Some(403))] +#[case::throttled(429, Some(429))] +#[case::server_error(500, Some(500))] #[tokio::test] async fn handles_statuses(#[case] status: u16, #[case] expected_status: Option) { let server = MockServer::start().await; @@ -73,9 +79,7 @@ async fn handles_statuses(#[case] status: u16, #[case] expected_status: Option assert_eq!(result.unwrap(), None), @@ -83,6 +87,7 @@ async fn handles_statuses(#[case] status: u16, #[case] expected_status: Option, + #[case] missing_environment: bool, +) { + let result = AzureKeyVault::new(Arc::new(move |name: &str| { + (name == "AZURE_KEY_VAULT_URI") + .then(|| uri.map(str::to_owned)) + .flatten() + })); + + if missing_environment { + assert!(matches!( + result, + Err(Error::MissingEnvironment("AZURE_KEY_VAULT_URI")) + )); + } else { + assert!(matches!(result, Err(Error::VaultUri))); + } } -#[rstest::rstest] -#[case("https://myvault.vault.azure.net", "https://vault.azure.net/.default")] -#[case( +#[rstest] +#[case::public_cloud("https://myvault.vault.azure.net", "https://vault.azure.net/.default")] +#[case::government_cloud( "https://v.vault.usgovcloudapi.net/", "https://vault.usgovcloudapi.net/.default" )] -#[case("http://localhost:8080", "https://localhost/.default")] -#[test] +#[case::local("http://localhost:8080", "https://localhost/.default")] fn derives_scope_from_vault_host(#[case] uri: &str, #[case] expected: &str) { let manager = AzureKeyVault::with_client( reqwest::Client::new(), @@ -139,6 +146,7 @@ fn derives_scope_from_vault_host(#[case] uri: &str, #[case] expected: &str) { assert_eq!(manager.scope(), expected); } +#[rstest] #[tokio::test] async fn missing_credentials_do_not_request_vault() { let server = MockServer::start().await; @@ -150,7 +158,7 @@ async fn missing_credentials_do_not_request_vault() { assert!( manager_without_credentials(&server) - .get_secret_from_azure_key_vault("NAME") + .get_secret("NAME") .await .is_err() ); @@ -192,11 +200,15 @@ struct FixtureExpected { error: Option, } +#[fixture] +fn parity_fixture() -> Fixture { + serde_json::from_str(include_str!("fixtures/key_vault_parity.json")).unwrap() +} + +#[rstest] #[tokio::test] -async fn parity_fixture_matches_python_backend_contract() { - let fixture: Fixture = - serde_json::from_str(include_str!("fixtures/key_vault_parity.json")).unwrap(); - for case in fixture.cases { +async fn parity_fixture_matches_python_backend_contract(parity_fixture: Fixture) { + for case in parity_fixture.cases { let server = MockServer::start().await; Mock::given(path(format!("/secrets/{}", case.secret_name))) .respond_with( @@ -205,9 +217,7 @@ async fn parity_fixture_matches_python_backend_contract() { .expect(1) .mount(&server) .await; - let result = manager(&server) - .get_secret_from_azure_key_vault(&case.secret_name) - .await; + let result = manager(&server).get_secret(&case.secret_name).await; if case.expected.missing == Some(true) { assert_eq!(result.unwrap(), None); } else if case.expected.error == Some(true) { diff --git a/litellm-rust/crates/secrets-azure/tests/live.rs b/litellm-rust/crates/secrets-azure/tests/live.rs index a062ba95070..18306382613 100644 --- a/litellm-rust/crates/secrets-azure/tests/live.rs +++ b/litellm-rust/crates/secrets-azure/tests/live.rs @@ -3,18 +3,16 @@ use std::sync::Arc; use litellm_core_utils::settings::ProcessEnvironment; use litellm_secrets_azure::AzureKeyVault; use litellm_secrets_types::Secret; +use rstest::rstest; +#[rstest] #[tokio::test] #[ignore] async fn reads_a_real_secret() { let environment = Arc::new(ProcessEnvironment); let manager = AzureKeyVault::new(environment).unwrap(); let name = std::env::var("AZURE_KEY_VAULT_LIVE_SECRET_NAME").unwrap(); - let secret = manager - .get_secret_from_azure_key_vault(&name) - .await - .unwrap() - .unwrap(); + let secret = manager.get_secret(&name).await.unwrap().unwrap(); assert!(matches!(&secret, Secret::String(_))); let host = std::env::var("AZURE_KEY_VAULT_URI") .unwrap() diff --git a/litellm-rust/crates/secrets-cyberark/src/secret_manager.rs b/litellm-rust/crates/secrets-cyberark/src/secret_manager.rs index 9d6eaaf1c4e..80e6fb12f9e 100644 --- a/litellm-rust/crates/secrets-cyberark/src/secret_manager.rs +++ b/litellm-rust/crates/secrets-cyberark/src/secret_manager.rs @@ -2,7 +2,10 @@ use std::{fs, sync::Arc, time::Duration}; use base64::{Engine, engine::general_purpose::STANDARD}; use litellm_core_utils::settings::Lookup; -use litellm_secrets_types::{BaseSecretManager, SecretValue, validate_secret_name}; +use litellm_secrets_types::{ + BaseSecretManager, SecretOperationContext, SecretValue, SecretWriteContext, + validate_secret_name, +}; use moka::future::Cache; use percent_encoding::{AsciiSet, NON_ALPHANUMERIC, utf8_percent_encode}; @@ -140,7 +143,7 @@ impl CyberArkSecretManager { .map_err(|_| Error::Endpoint) } - async fn authenticate(&self) -> Result { + async fn authenticate(&self, context: &SecretOperationContext) -> Result { if let Some(token) = self.token.get(&()).await { return Ok(token); } @@ -155,12 +158,12 @@ impl CyberArkSecretManager { self.account, self.username )) .map_err(|_| Error::Endpoint)?; - let response = self - .client - .post(url) - .body(self.api_key.expose().to_owned()) - .send() - .await?; + let response = with_timeout( + self.client.post(url).body(self.api_key.expose().to_owned()), + context, + ) + .send() + .await?; if !response.status().is_success() { return Err(Error::AuthStatus(response.status().as_u16())); } @@ -169,23 +172,37 @@ impl CyberArkSecretManager { Ok(token) } - async fn authorization_header(&self) -> Result { + async fn authorization_header( + &self, + context: &SecretOperationContext, + ) -> Result { Ok(format!( "Token token=\"{}\"", - self.authenticate().await?.expose() + self.authenticate(context).await?.expose() )) } pub async fn async_read_secret(&self, name: &str) -> Result, Error> { + self.async_read_secret_with_context(name, &SecretOperationContext::default()) + .await + } + + pub async fn async_read_secret_with_context( + &self, + name: &str, + context: &SecretOperationContext, + ) -> Result, Error> { if let Some(value) = self.secrets.get(name).await { return Ok(Some(value)); } - let response = self - .client - .get(self.secret_url(name)?) - .header("Authorization", self.authorization_header().await?) - .send() - .await?; + let response = with_timeout( + self.client + .get(self.secret_url(name)?) + .header("Authorization", self.authorization_header(context).await?), + context, + ) + .send() + .await?; if response.status() == reqwest::StatusCode::NOT_FOUND { return Ok(None); } @@ -198,20 +215,38 @@ impl CyberArkSecretManager { } pub async fn async_write_secret( + &self, + name: &str, + value: &SecretValue, + description: Option<&str>, + ) -> Result<(), Error> { + self.async_write_secret_with_context( + name, + value, + description, + &SecretOperationContext::default(), + ) + .await + } + + pub async fn async_write_secret_with_context( &self, name: &str, value: &SecretValue, _description: Option<&str>, + context: &SecretOperationContext, ) -> Result<(), Error> { validate_secret_name(name)?; - self.ensure_variable_exists(name).await; - let response = self - .client - .post(self.secret_url(name)?) - .header("Authorization", self.authorization_header().await?) - .body(value.expose().to_owned()) - .send() - .await?; + self.ensure_variable_exists(name, context).await; + let response = with_timeout( + self.client + .post(self.secret_url(name)?) + .header("Authorization", self.authorization_header(context).await?) + .body(value.expose().to_owned()), + context, + ) + .send() + .await?; if !response.status().is_success() { return Err(Error::Status(response.status().as_u16())); } @@ -219,7 +254,7 @@ impl CyberArkSecretManager { Ok(()) } - async fn ensure_variable_exists(&self, name: &str) { + async fn ensure_variable_exists(&self, name: &str, context: &SecretOperationContext) { let policy_url = self .endpoint .join(&format!("policies/{}/policy/root", self.account)); @@ -227,7 +262,7 @@ impl CyberArkSecretManager { tracing::warn!("Could not build CyberArk policy endpoint"); return; }; - let Ok(authorization) = self.authorization_header().await else { + let Ok(authorization) = self.authorization_header(context).await else { tracing::warn!("Could not authenticate while ensuring CyberArk variable exists"); return; }; @@ -235,14 +270,16 @@ impl CyberArkSecretManager { "- !variable {}\n", serde_json::to_string(name).expect("serializing a string cannot fail") ); - let response = self - .client - .post(policy_url) - .header("Authorization", authorization) - .header("Content-Type", "application/x-yaml") - .body(body) - .send() - .await; + let response = with_timeout( + self.client + .post(policy_url) + .header("Authorization", authorization) + .header("Content-Type", "application/x-yaml") + .body(body), + context, + ) + .send() + .await; match response { Ok(response) if response.status().is_success() => {} Ok(response) @@ -271,7 +308,21 @@ impl CyberArkSecretManager { pub async fn async_delete_secret( &self, name: &str, - _recovery_window_in_days: i64, + recovery_window_in_days: Option, + ) -> Result { + self.async_delete_secret_with_context( + name, + recovery_window_in_days, + &SecretOperationContext::default(), + ) + .await + } + + pub async fn async_delete_secret_with_context( + &self, + name: &str, + _recovery_window_in_days: Option, + _context: &SecretOperationContext, ) -> Result { tracing::warn!( "CyberArk Conjur does not support direct secret deletion. Secrets must be removed through policy updates." @@ -286,29 +337,50 @@ impl BaseSecretManager for CyberArkSecretManager { type WriteResponse = (); type DeleteResponse = DeleteOutcome; - async fn async_read_secret(&self, name: &str) -> Result, Error> { - self.async_read_secret(name).await + async fn async_read_secret( + &self, + name: &str, + context: &SecretOperationContext, + ) -> Result, Error> { + self.async_read_secret_with_context(name, context).await } async fn async_write_secret( &self, name: &str, value: &SecretValue, - description: Option<&str>, + context: &SecretWriteContext, ) -> Result<(), Error> { - self.async_write_secret(name, value, description).await + self.async_write_secret_with_context( + name, + value, + context.description.as_deref(), + &context.operation, + ) + .await } async fn async_delete_secret( &self, name: &str, - recovery_window_in_days: i64, + recovery_window_in_days: Option, + context: &SecretOperationContext, ) -> Result { - self.async_delete_secret(name, recovery_window_in_days) + self.async_delete_secret_with_context(name, recovery_window_in_days, context) .await } } +fn with_timeout( + request: reqwest::RequestBuilder, + context: &SecretOperationContext, +) -> reqwest::RequestBuilder { + match context.timeout() { + Some(timeout) => request.timeout(timeout), + None => request, + } +} + fn normalize_endpoint(mut endpoint: reqwest::Url) -> reqwest::Url { if !endpoint.path().ends_with('/') { endpoint.set_path(&format!("{}/", endpoint.path())); diff --git a/litellm-rust/crates/secrets-cyberark/tests/secret_manager.rs b/litellm-rust/crates/secrets-cyberark/tests/secret_manager.rs index fd7198b70fb..59964d90d50 100644 --- a/litellm-rust/crates/secrets-cyberark/tests/secret_manager.rs +++ b/litellm-rust/crates/secrets-cyberark/tests/secret_manager.rs @@ -2,7 +2,10 @@ use std::{sync::Arc, time::Duration}; use base64::{Engine, engine::general_purpose::STANDARD}; use litellm_secrets_cyberark::{CyberArkSecretManager, DeleteOutcome, Error}; -use litellm_secrets_types::SecretValue; +use litellm_secrets_types::{ + BaseSecretManager, CyberarkOperationContext, SecretOperationContext, SecretValue, +}; +use rstest::{fixture, rstest}; use serde::Deserialize; use wiremock::{ Match, Mock, MockServer, Request, ResponseTemplate, @@ -40,7 +43,8 @@ impl Match for RawPath { } } -fn fixture() -> ParityFixture { +#[fixture] +fn parity_fixture() -> ParityFixture { serde_json::from_str(include_str!("fixtures/parity.json")).unwrap() } @@ -65,6 +69,7 @@ async fn mount_auth(server: &MockServer, expected: u64) { .await; } +#[rstest] #[tokio::test] async fn successful_reads_cache_auth_secret_and_redact_values() { let server = MockServer::start().await; @@ -89,6 +94,7 @@ async fn successful_reads_cache_auth_secret_and_redact_values() { } } +#[rstest] #[tokio::test] async fn concurrent_reads_share_authentication_request() { let server = MockServer::start().await; @@ -122,7 +128,7 @@ async fn concurrent_reads_share_authentication_request() { assert_eq!(second.unwrap().unwrap().expose(), "value"); } -#[rstest::rstest] +#[rstest] #[case::not_found(404)] #[case::unauthorized(401)] #[case::forbidden(403)] @@ -162,6 +168,7 @@ async fn failed_reads_are_not_cached(#[case] status: u16) { } } +#[rstest] #[tokio::test] async fn failed_authentication_is_not_cached_and_does_not_read_secret() { let server = MockServer::start().await; @@ -199,6 +206,29 @@ async fn failed_authentication_is_not_cached_and_does_not_read_secret() { ); } +#[rstest] +#[tokio::test] +async fn trait_read_applies_cyberark_operation_timeout_to_authentication() { + let server = MockServer::start().await; + Mock::given(path("/authn/acct/admin/authenticate")) + .respond_with( + ResponseTemplate::new(200) + .set_body_string(TOKEN_JSON) + .set_delay(Duration::from_millis(50)), + ) + .mount(&server) + .await; + let manager = manager(&server, Duration::from_secs(60)); + let context = SecretOperationContext::Cyberark(CyberarkOperationContext { + timeout: Some(Duration::from_millis(10)), + }); + + let result = BaseSecretManager::async_read_secret(&manager, "key", &context).await; + + assert!(matches!(result, Err(Error::Http(error)) if error.is_timeout())); +} + +#[rstest] #[tokio::test] async fn expired_tokens_and_secrets_are_fetched_again() { let server = MockServer::start().await; @@ -215,13 +245,14 @@ async fn expired_tokens_and_secrets_are_fetched_again() { } } -#[rstest::rstest] +#[rstest] +#[case::plain("OPENAI_API_KEY")] +#[case::path("team/app/key")] +#[case::punctuation("a b+c.d-e_f~g")] +#[case::quote("needs \"quote\"")] #[tokio::test] -async fn secret_names_use_python_quote_encoding( - #[values("OPENAI_API_KEY", "team/app/key", "a b+c.d-e_f~g", "needs \"quote\"")] name: &str, -) { - let fixture = fixture(); - let secret = fixture +async fn secret_names_use_python_quote_encoding(parity_fixture: ParityFixture, #[case] name: &str) { + let secret = parity_fixture .secrets .iter() .find(|secret| secret.name == name) @@ -244,11 +275,11 @@ async fn secret_names_use_python_quote_encoding( ); } -#[rstest::rstest] -#[case(201)] -#[case(409)] -#[case(422)] -#[case(500)] +#[rstest] +#[case::created(201)] +#[case::already_exists(409)] +#[case::unprocessable(422)] +#[case::server_error(500)] #[tokio::test] async fn writes_tolerate_policy_status_and_cache_value(#[case] policy_status: u16) { let server = MockServer::start().await; @@ -282,6 +313,7 @@ async fn writes_tolerate_policy_status_and_cache_value(#[case] policy_status: u1 ); } +#[rstest] #[tokio::test] async fn failed_value_write_is_not_cached() { let server = MockServer::start().await; @@ -319,13 +351,17 @@ async fn failed_value_write_is_not_cached() { ); } +#[rstest] +#[case::parent("../etc")] +#[case::embedded_parent("team/../etc")] +#[case::control("key\n")] #[tokio::test] -async fn unsafe_names_fail_before_http_calls() { +async fn unsafe_names_fail_before_http_calls(#[case] name: &str) { let server = MockServer::start().await; let manager = manager(&server, Duration::from_secs(60)); assert!(matches!( manager - .async_write_secret("../etc", &SecretValue::new("v"), None) + .async_write_secret(name, &SecretValue::new("v"), None) .await, Err(Error::Operation( litellm_secrets_types::Error::UnsafeSecretName @@ -333,6 +369,7 @@ async fn unsafe_names_fail_before_http_calls() { )); } +#[rstest] #[tokio::test] async fn delete_invalidates_cache_and_reports_not_supported() { let server = MockServer::start().await; @@ -353,7 +390,7 @@ async fn delete_invalidates_cache_and_reports_not_supported() { "v" ); assert_eq!( - manager.async_delete_secret("key", 7).await.unwrap(), + manager.async_delete_secret("key", Some(7)).await.unwrap(), DeleteOutcome::NotSupported ); assert_eq!( @@ -367,7 +404,7 @@ async fn delete_invalidates_cache_and_reports_not_supported() { ); } -#[test] +#[rstest] fn new_validates_credentials_before_license_and_configuration() { let empty: Arc = Arc::new(|_: &str| None); @@ -413,6 +450,7 @@ fn new_validates_credentials_before_license_and_configuration() { )); } +#[rstest] #[tokio::test] async fn new_reads_environment_defaults_end_to_end() { let server = MockServer::start().await; @@ -446,7 +484,7 @@ async fn new_reads_environment_defaults_end_to_end() { ); } -#[test] +#[rstest] fn new_reports_missing_client_certificate_files() { assert!(matches!( CyberArkSecretManager::new( @@ -461,6 +499,7 @@ fn new_reports_missing_client_certificate_files() { )); } +#[rstest] #[tokio::test] async fn trailing_slash_endpoint_preserves_base_path() { let server = MockServer::start().await; @@ -494,23 +533,25 @@ async fn trailing_slash_endpoint_preserves_base_path() { ); } -#[test] -fn parity_fixture_matches_authentication_contract() { - let fixture = fixture(); - assert_eq!(fixture.endpoint, "http://conjur.test:8080"); - assert_eq!(fixture.account, "acct"); - assert_eq!(fixture.username, "admin"); - assert_eq!(fixture.api_key, "k3y"); - assert_eq!(fixture.authenticate_path, "/authn/acct/admin/authenticate"); - assert_eq!(fixture.token_json, TOKEN_JSON); +#[rstest] +fn parity_fixture_matches_authentication_contract(parity_fixture: ParityFixture) { + assert_eq!(parity_fixture.endpoint, "http://conjur.test:8080"); + assert_eq!(parity_fixture.account, "acct"); + assert_eq!(parity_fixture.username, "admin"); + assert_eq!(parity_fixture.api_key, "k3y"); assert_eq!( - fixture.authorization_header, + parity_fixture.authenticate_path, + "/authn/acct/admin/authenticate" + ); + assert_eq!(parity_fixture.token_json, TOKEN_JSON); + assert_eq!( + parity_fixture.authorization_header, format!("Token token=\"{}\"", STANDARD.encode(TOKEN_JSON)) ); - assert_eq!(fixture.policy_path, "/policies/acct/policy/root"); - assert_eq!(fixture.secrets.len(), 4); + assert_eq!(parity_fixture.policy_path, "/policies/acct/policy/root"); + assert_eq!(parity_fixture.secrets.len(), 4); assert_eq!( - fixture.secrets[1].policy_body, + parity_fixture.secrets[1].policy_body, "- !variable \"team/app/key\"\n" ); } diff --git a/litellm-rust/crates/secrets-google/tests/kms.rs b/litellm-rust/crates/secrets-google/tests/kms.rs index 667ecd268c8..9739a46b10c 100644 --- a/litellm-rust/crates/secrets-google/tests/kms.rs +++ b/litellm-rust/crates/secrets-google/tests/kms.rs @@ -1,11 +1,14 @@ use base64::{Engine, engine::general_purpose::STANDARD}; use google_cloud_kms_v1::client::KeyManagementService; -use litellm_secrets_google::GoogleKms; +use litellm_core_utils::settings::Lookup; +use litellm_secrets_google::{Error, GoogleKms, kms::validate_environment}; +use rstest::rstest; use wiremock::{ Mock, MockServer, ResponseTemplate, matchers::{body_json, path}, }; +#[rstest] #[tokio::test] async fn google_kms_decrypts_using_the_configured_resource() { let server = MockServer::start().await; @@ -35,15 +38,49 @@ async fn google_kms_decrypts_using_the_configured_resource() { ); } +#[rstest] +#[case::unset(None)] +#[case::disabled(Some(false))] #[tokio::test] -async fn disabled_google_kms_loader_does_not_require_environment_configuration() { +async fn disabled_google_kms_loader_does_not_require_environment_configuration( + #[case] enabled: Option, +) { use std::sync::Arc; - for enabled in [None, Some(false)] { - assert!( - litellm_secrets_google::load_google_kms(enabled, Arc::new(|_: &str| None)) - .await - .unwrap() - .is_none() - ); - } + assert!( + litellm_secrets_google::load_google_kms(enabled, Arc::new(|_: &str| None)) + .await + .unwrap() + .is_none() + ); +} + +#[rstest] +#[case::credentials_missing(None, None, "GOOGLE_APPLICATION_CREDENTIALS")] +#[case::resource_missing(Some("credentials"), None, "GOOGLE_KMS_RESOURCE_NAME")] +fn enabled_google_kms_requires_all_environment_values( + #[case] credentials: Option<&str>, + #[case] resource: Option<&str>, + #[case] missing: &'static str, +) { + let environment = move |name: &str| match name { + "GOOGLE_APPLICATION_CREDENTIALS" => credentials.map(str::to_owned), + "GOOGLE_KMS_RESOURCE_NAME" => resource.map(str::to_owned), + _ => None, + }; + + assert!(matches!( + validate_environment(&environment as &dyn Lookup), + Err(Error::MissingEnvironment(name)) if name == missing + )); +} + +#[rstest] +fn complete_google_kms_environment_is_valid() { + let environment = |name: &str| match name { + "GOOGLE_APPLICATION_CREDENTIALS" => Some("credentials".to_owned()), + "GOOGLE_KMS_RESOURCE_NAME" => Some("resource".to_owned()), + _ => None, + }; + + assert!(validate_environment(&environment as &dyn Lookup).is_ok()); } diff --git a/litellm-rust/crates/secrets-google/tests/secret_manager.rs b/litellm-rust/crates/secrets-google/tests/secret_manager.rs index b3b1d29e62c..867dcf935c1 100644 --- a/litellm-rust/crates/secrets-google/tests/secret_manager.rs +++ b/litellm-rust/crates/secrets-google/tests/secret_manager.rs @@ -2,6 +2,7 @@ use std::{sync::Arc, time::Duration}; use base64::{Engine, engine::general_purpose::STANDARD}; use litellm_secrets_google::{Error, GoogleSecretManager}; +use rstest::{fixture, rstest}; use wiremock::{ Mock, MockServer, ResponseTemplate, @@ -20,11 +21,17 @@ fn manager(server: &MockServer, always_read: bool, ttl: Duration) -> GoogleSecre .unwrap() } -#[rstest::rstest] +#[fixture] +fn default_ttl() -> Duration { + Duration::from_secs(60) +} + +#[rstest] #[case::nonempty("private-value")] #[case::empty("")] #[tokio::test] async fn successful_reads_use_auth_latest_version_and_cache_including_empty_values( + default_ttl: Duration, #[case] value: &str, ) { let server = MockServer::start().await; @@ -39,7 +46,7 @@ async fn successful_reads_use_auth_latest_version_and_cache_including_empty_valu .expect(1) .mount(&server) .await; - let manager = manager(&server, false, Duration::from_secs(60)); + let manager = manager(&server, false, default_ttl); for _ in 0..2 { assert_eq!( manager @@ -54,21 +61,44 @@ async fn successful_reads_use_auth_latest_version_and_cache_including_empty_valu } } -#[rstest::rstest] -#[case::not_found(404, serde_json::json!({}))] -#[case::unauthorized(401, serde_json::json!({}))] -#[case::forbidden(403, serde_json::json!({}))] -#[case::throttled(429, serde_json::json!({}))] -#[case::unavailable(503, serde_json::json!({}))] -#[case::missing_payload(200, serde_json::json!({"payload":{}}))] -#[case::invalid_base64(200, serde_json::json!({"payload":{"data":"%%%"}}))] +enum ExpectedReadFailure { + Missing, + Status(u16), + MissingPayload, + Base64, + Utf8, +} + +#[rstest] +#[case::not_found(404, serde_json::json!({}), ExpectedReadFailure::Missing)] +#[case::unauthorized(401, serde_json::json!({}), ExpectedReadFailure::Status(401))] +#[case::forbidden(403, serde_json::json!({}), ExpectedReadFailure::Status(403))] +#[case::throttled(429, serde_json::json!({}), ExpectedReadFailure::Status(429))] +#[case::unavailable(503, serde_json::json!({}), ExpectedReadFailure::Status(503))] +#[case::missing_payload( + 200, + serde_json::json!({"payload":{}}), + ExpectedReadFailure::MissingPayload +)] +#[case::invalid_base64( + 200, + serde_json::json!({"payload":{"data":"%%%"}}), + ExpectedReadFailure::Base64 +)] +#[case::invalid_utf8( + 200, + serde_json::json!({"payload":{"data":STANDARD.encode([0xff])}}), + ExpectedReadFailure::Utf8 +)] #[tokio::test] async fn failed_or_missing_reads_are_not_cached( + default_ttl: Duration, #[case] status: u16, #[case] body: serde_json::Value, + #[case] expected: ExpectedReadFailure, ) { let server = MockServer::start().await; - let manager = manager(&server, false, Duration::from_secs(60)); + let manager = manager(&server, false, default_ttl); let failing = Mock::given(path( "/v1/projects/project/secrets/key/versions/latest:access", )) @@ -77,13 +107,16 @@ async fn failed_or_missing_reads_are_not_cached( .mount_as_scoped(&server) .await; let result = manager.get_secret_from_google_secret_manager("key").await; - match status { - 404 => assert_eq!(result.unwrap(), None), - 200 => assert!(matches!( - result, - Err(Error::MissingPayload | Error::Base64(_)) - )), - status => assert!(matches!(result, Err(Error::Status(actual)) if actual == status)), + match expected { + ExpectedReadFailure::Missing => assert_eq!(result.unwrap(), None), + ExpectedReadFailure::Status(expected) => { + assert!(matches!(result, Err(Error::Status(actual)) if actual == expected)); + } + ExpectedReadFailure::MissingPayload => { + assert!(matches!(result, Err(Error::MissingPayload))); + } + ExpectedReadFailure::Base64 => assert!(matches!(result, Err(Error::Base64(_)))), + ExpectedReadFailure::Utf8 => assert!(matches!(result, Err(Error::Utf8))), } drop(failing); Mock::given(path( @@ -109,7 +142,7 @@ async fn failed_or_missing_reads_are_not_cached( } } -#[rstest::rstest] +#[rstest] #[case::always_read(true, Duration::from_secs(60))] #[case::expired_cache(false, Duration::from_millis(1))] #[tokio::test] @@ -141,7 +174,7 @@ async fn always_read_and_expired_cache_fetch_again( } } -#[test] +#[rstest] fn google_manager_requires_host_license_and_project_configuration() { assert!(matches!( GoogleSecretManager::new(Arc::new(|_: &str| None), false), @@ -155,13 +188,29 @@ fn google_manager_requires_host_license_and_project_configuration() { )); } -#[rstest::rstest] -#[case("true")] -#[case("null")] -#[case("\"text\"")] -#[case("{\"key\":1}")] +#[rstest] +#[case::provider_specific("GOOGLE_SECRET_MANAGER_REFRESH_INTERVAL")] +#[case::shared("SECRET_MANAGER_REFRESH_INTERVAL")] +fn google_manager_rejects_invalid_refresh_intervals(#[case] variable: &'static str) { + let environment = Arc::new(move |name: &str| match name { + "GOOGLE_SECRET_MANAGER_PROJECT_ID" => Some("project".to_owned()), + name if name == variable => Some("not-a-number".to_owned()), + _ => None, + }); + + assert!(matches!( + GoogleSecretManager::new(environment, true), + Err(Error::RefreshInterval) + )); +} + +#[rstest] +#[case::boolean("true")] +#[case::null("null")] +#[case::string("\"text\"")] +#[case::object("{\"key\":1}")] #[tokio::test] -async fn cache_preserves_raw_values(#[case] raw: &str) { +async fn cache_preserves_raw_values(default_ttl: Duration, #[case] raw: &str) { let server = MockServer::start().await; Mock::given(path( "/v1/projects/project/secrets/key/versions/latest:access", @@ -173,7 +222,7 @@ async fn cache_preserves_raw_values(#[case] raw: &str) { .expect(1) .mount(&server) .await; - let manager = manager(&server, false, Duration::from_secs(60)); + let manager = manager(&server, false, default_ttl); for _ in 0..2 { assert_eq!( manager diff --git a/litellm-rust/crates/secrets-hashicorp/src/error.rs b/litellm-rust/crates/secrets-hashicorp/src/error.rs index e26033af085..664f9ab18dd 100644 --- a/litellm-rust/crates/secrets-hashicorp/src/error.rs +++ b/litellm-rust/crates/secrets-hashicorp/src/error.rs @@ -4,6 +4,8 @@ pub enum Error { EnterpriseRequired, #[error("invalid secret name")] InvalidSecretName(#[from] litellm_secrets_types::Error), + #[error("HashiCorp Vault received an incompatible operation context")] + InvalidOperationContext, #[error("HashiCorp Vault client failed")] Client( #[from] @@ -29,6 +31,8 @@ pub enum Error { MalformedPayload, #[error("HashiCorp Vault secret value is not a string")] NonStringValue, + #[error("HashiCorp Vault operation timed out")] + Timeout, #[error("invalid HashiCorp Vault refresh interval")] RefreshInterval, } diff --git a/litellm-rust/crates/secrets-hashicorp/src/secret_manager.rs b/litellm-rust/crates/secrets-hashicorp/src/secret_manager.rs index 3ad9a549438..a4d1efb4ff6 100644 --- a/litellm-rust/crates/secrets-hashicorp/src/secret_manager.rs +++ b/litellm-rust/crates/secrets-hashicorp/src/secret_manager.rs @@ -1,13 +1,15 @@ use std::{ collections::HashMap, fmt, + future::Future, sync::Arc, time::{Duration, Instant}, }; use litellm_core_utils::settings::Lookup; use litellm_secrets_types::{ - BaseSecretManager, SecretValue, async_rotate_secret, validate_secret_name, + BaseSecretManager, HashicorpOperationContext, SecretOperationContext, SecretValue, + SecretWriteContext, async_rotate_secret, validate_secret_name, }; use moka::future::Cache; use rustify::errors::ClientError as RustifyClientError; @@ -31,17 +33,23 @@ struct CachedClient { expires_at: Option, } -#[derive(Clone, Debug, PartialEq, Eq)] +#[derive(Clone, Debug, Hash, PartialEq, Eq)] pub struct SecretLocation { pub namespace: Option, pub mount: String, pub path: String, } +#[derive(Clone, Hash, PartialEq, Eq)] +struct CacheKey { + location: SecretLocation, + data_key: String, +} + #[derive(Clone)] pub struct HashicorpVault { config: HashicorpVaultConfig, - cache: Cache, + cache: Cache, auth_client: Arc>>, } @@ -71,7 +79,7 @@ impl HashicorpVault { if !enterprise_enabled { return Err(Error::EnterpriseRequired); } - let cache: Cache = Cache::builder() + let cache: Cache = Cache::builder() .max_capacity(CACHE_CAPACITY) .time_to_live(config.refresh_interval) .build(); @@ -83,9 +91,21 @@ impl HashicorpVault { } pub fn secret_location(&self, secret_name: &str) -> Result { + self.secret_location_with_context(secret_name, &SecretOperationContext::default()) + } + + pub fn secret_location_with_context( + &self, + secret_name: &str, + context: &SecretOperationContext, + ) -> Result { validate_secret_name(secret_name).map_err(Error::InvalidSecretName)?; + let operation: Option<&HashicorpOperationContext> = hashicorp_context(context)?; let path: String = [ - self.config.path_prefix.clone(), + operation + .and_then(|operation| operation.path_prefix.as_deref()) + .and_then(path_component) + .or_else(|| self.config.path_prefix.clone()), Some(secret_name.to_owned()), ] .into_iter() @@ -94,7 +114,10 @@ impl HashicorpVault { .join("/"); Ok(SecretLocation { namespace: self.config.secret_namespace().map(str::to_owned), - mount: self.config.mount.clone(), + mount: operation + .and_then(|operation| operation.mount.as_deref()) + .and_then(path_component) + .unwrap_or_else(|| self.config.mount.clone()), path, }) } @@ -104,19 +127,37 @@ impl HashicorpVault { } pub async fn async_read_secret(&self, secret_name: &str) -> Result, Error> { - let location: SecretLocation = self.secret_location(secret_name)?; - let cache_key: String = cache_key(&location); + self.async_read_secret_with_context(secret_name, &SecretOperationContext::default()) + .await + } + + pub async fn async_read_secret_with_context( + &self, + secret_name: &str, + context: &SecretOperationContext, + ) -> Result, Error> { + let location: SecretLocation = self.secret_location_with_context(secret_name, context)?; + let data_key: String = data_key(context)?; + let cache_key = CacheKey { + location: location.clone(), + data_key: data_key.clone(), + }; if let Some(value) = self.cache.get(&cache_key).await { return Ok(Some(value)); } - let client: Arc = self.vault_client().await?; - let data: HashMap = + let data: Option> = with_timeout(context, async { + let client: Arc = self.vault_client().await?; match kv2::read(client.as_ref(), &location.mount, &location.path).await { - Ok(data) => data, - Err(error) if api_status(&error) == Some(404) => return Ok(None), - Err(error) => return Err(map_api_error(error, ErrorContext::Read)), - }; - let Some(value) = data.get("key") else { + Ok(data) => Ok(Some(data)), + Err(error) if api_status(&error) == Some(404) => Ok(None), + Err(error) => Err(map_api_error(error, ErrorContext::Read)), + } + }) + .await?; + let Some(data) = data else { + return Ok(None); + }; + let Some(value) = data.get(&data_key) else { return Ok(None); }; let value: &str = value.as_str().ok_or(Error::NonStringValue)?; @@ -131,11 +172,29 @@ impl HashicorpVault { value: SecretValue, description: Option<&str>, ) -> Result { - let location: SecretLocation = self.secret_location(secret_name)?; - let cache_key: String = cache_key(&location); - let data: HashMap = match description { + self.async_write_secret_with_context( + secret_name, + &value, + &SecretWriteContext { + description: description.map(str::to_owned), + ..SecretWriteContext::default() + }, + ) + .await + } + + pub async fn async_write_secret_with_context( + &self, + secret_name: &str, + value: &SecretValue, + context: &SecretWriteContext, + ) -> Result { + let location: SecretLocation = + self.secret_location_with_context(secret_name, &context.operation)?; + let data_key: String = data_key(&context.operation)?; + let data: HashMap = match context.description.as_deref() { Some(description) => [ - ("key".to_owned(), Value::String(value.expose().to_owned())), + (data_key, Value::String(value.expose().to_owned())), ( "description".to_owned(), Value::String(description.to_owned()), @@ -143,27 +202,41 @@ impl HashicorpVault { ] .into_iter() .collect(), - None => [("key".to_owned(), Value::String(value.expose().to_owned()))] + None => [(data_key, Value::String(value.expose().to_owned()))] .into_iter() .collect(), }; - let client: Arc = self.vault_client().await?; - let metadata = kv2::set(client.as_ref(), &location.mount, &location.path, &data) - .await - .map_err(|error| map_api_error(error, ErrorContext::Secret))?; - self.cache.invalidate(&cache_key).await; + let metadata = with_timeout(&context.operation, async { + let client: Arc = self.vault_client().await?; + kv2::set(client.as_ref(), &location.mount, &location.path, &data) + .await + .map_err(|error| map_api_error(error, ErrorContext::Secret)) + }) + .await?; + self.cache.invalidate_all(); serde_json::to_value(metadata) .map_err(|source| Error::Client(ClientError::JsonParseError { source })) } pub async fn async_delete_secret(&self, secret_name: &str) -> Result<(), Error> { - let location: SecretLocation = self.secret_location(secret_name)?; - let cache_key: String = cache_key(&location); - let client: Arc = self.vault_client().await?; - kv2::delete_latest(client.as_ref(), &location.mount, &location.path) + self.async_delete_secret_with_context(secret_name, &SecretOperationContext::default()) .await - .map_err(|error| map_api_error(error, ErrorContext::Secret))?; - self.cache.invalidate(&cache_key).await; + } + + pub async fn async_delete_secret_with_context( + &self, + secret_name: &str, + context: &SecretOperationContext, + ) -> Result<(), Error> { + let location: SecretLocation = self.secret_location_with_context(secret_name, context)?; + with_timeout(context, async { + let client: Arc = self.vault_client().await?; + kv2::delete_latest(client.as_ref(), &location.mount, &location.path) + .await + .map_err(|error| map_api_error(error, ErrorContext::Secret)) + }) + .await?; + self.cache.invalidate_all(); Ok(()) } @@ -173,7 +246,23 @@ impl HashicorpVault { new_name: &str, value: &SecretValue, ) -> Result { - async_rotate_secret(self, current_name, new_name, value).await + self.async_rotate_secret_with_context( + current_name, + new_name, + value, + &SecretOperationContext::default(), + ) + .await + } + + pub async fn async_rotate_secret_with_context( + &self, + current_name: &str, + new_name: &str, + value: &SecretValue, + context: &SecretOperationContext, + ) -> Result { + async_rotate_secret(self, current_name, new_name, value, context).await } async fn vault_client(&self) -> Result, Error> { @@ -254,25 +343,30 @@ impl BaseSecretManager for HashicorpVault { type WriteResponse = Value; type DeleteResponse = (); - async fn async_read_secret(&self, name: &str) -> Result, Error> { - HashicorpVault::async_read_secret(self, name).await + async fn async_read_secret( + &self, + name: &str, + context: &SecretOperationContext, + ) -> Result, Error> { + HashicorpVault::async_read_secret_with_context(self, name, context).await } async fn async_write_secret( &self, name: &str, value: &SecretValue, - description: Option<&str>, + context: &SecretWriteContext, ) -> Result { - HashicorpVault::async_write_secret(self, name, value.clone(), description).await + HashicorpVault::async_write_secret_with_context(self, name, value, context).await } async fn async_delete_secret( &self, name: &str, - _recovery_window_in_days: i64, + _recovery_window_in_days: Option, + context: &SecretOperationContext, ) -> Result<(), Error> { - HashicorpVault::async_delete_secret(self, name).await + HashicorpVault::async_delete_secret_with_context(self, name, context).await } } @@ -283,11 +377,42 @@ enum ErrorContext { Secret, } -fn cache_key(location: &SecretLocation) -> String { - format!( - "{:?}/{}/{}", - location.namespace, location.mount, location.path - ) +fn hashicorp_context( + context: &SecretOperationContext, +) -> Result, Error> { + match context { + SecretOperationContext::Hashicorp(context) => Ok(Some(context)), + SecretOperationContext::Default => Ok(None), + SecretOperationContext::Aws(_) | SecretOperationContext::Cyberark(_) => { + Err(Error::InvalidOperationContext) + } + } +} + +fn path_component(value: &str) -> Option { + let value: &str = value.trim().trim_matches('/'); + (!value.is_empty()).then(|| value.to_owned()) +} + +fn data_key(context: &SecretOperationContext) -> Result { + Ok(hashicorp_context(context)? + .and_then(|context| context.data_key.as_deref()) + .map(str::trim) + .filter(|data_key| !data_key.is_empty()) + .map(str::to_owned) + .unwrap_or_else(|| "key".to_owned())) +} + +async fn with_timeout( + context: &SecretOperationContext, + operation: impl Future>, +) -> Result { + match context.timeout() { + Some(timeout) => tokio::time::timeout(timeout, operation) + .await + .map_err(|_| Error::Timeout)?, + None => operation.await, + } } fn identity_for(tls: Option<&TlsCertAuth>) -> Result, Error> { diff --git a/litellm-rust/crates/secrets-hashicorp/tests/secret_manager.rs b/litellm-rust/crates/secrets-hashicorp/tests/secret_manager.rs index c52db46e41e..6b322de4d8d 100644 --- a/litellm-rust/crates/secrets-hashicorp/tests/secret_manager.rs +++ b/litellm-rust/crates/secrets-hashicorp/tests/secret_manager.rs @@ -2,7 +2,11 @@ use std::{collections::HashMap, sync::Arc, time::Duration}; use litellm_core_utils::settings::Lookup; use litellm_secrets_hashicorp::{Error, HashicorpVault, HashicorpVaultConfig}; -use litellm_secrets_types::SecretValue; +use litellm_secrets_types::{ + AwsOperationContext, BaseSecretManager, CyberarkOperationContext, HashicorpOperationContext, + SecretOperationContext, SecretValue, SecretWriteContext, +}; +use rstest::{fixture, rstest}; use serde::Deserialize; use serde_json::json; use wiremock::{ @@ -69,8 +73,14 @@ fn read_response(data: serde_json::Value) -> serde_json::Value { }) } +#[fixture] +fn token_values() -> Vec<(&'static str, &'static str)> { + vec![("HCP_VAULT_TOKEN", "token")] +} + +#[rstest] #[tokio::test] -async fn token_reads_use_vault_headers_and_cache_values() { +async fn token_reads_use_vault_headers_and_cache_values(token_values: Vec<(&str, &str)>) { let server: MockServer = MockServer::start().await; Mock::given(method("GET")) .and(path("/v1/secret/data/name")) @@ -81,7 +91,7 @@ async fn token_reads_use_vault_headers_and_cache_values() { .expect(1) .mount(&server) .await; - let manager: HashicorpVault = manager(&server, &[("HCP_VAULT_TOKEN", "token")]); + let manager: HashicorpVault = manager(&server, &token_values); assert_eq!( manager @@ -109,6 +119,7 @@ async fn token_reads_use_vault_headers_and_cache_values() { ); } +#[rstest] #[tokio::test] async fn namespace_mount_and_prefix_are_sanitized_in_the_url() { let server: MockServer = MockServer::start().await; @@ -138,7 +149,7 @@ async fn namespace_mount_and_prefix_are_sanitized_in_the_url() { assert!(manager.async_read_secret("name").await.unwrap().is_some()); } -#[test] +#[rstest] fn trailing_address_slashes_are_removed() { let environment: Arc = Arc::new(|name: &str| match name { "HCP_VAULT_ADDR" => Some("http://vault.test:8200///".to_owned()), @@ -159,9 +170,9 @@ fn trailing_address_slashes_are_removed() { ); } -#[rstest::rstest] -#[case("-1")] -#[case("not-a-number")] +#[rstest] +#[case::negative("-1")] +#[case::not_a_number("not-a-number")] fn invalid_refresh_intervals_are_rejected(#[case] value: &str) { let environment: Arc = Arc::new(move |name: &str| match name { "HCP_VAULT_REFRESH_INTERVAL" => Some(value.to_owned()), @@ -174,6 +185,7 @@ fn invalid_refresh_intervals_are_rejected(#[case] value: &str) { )); } +#[rstest] #[tokio::test] async fn approle_login_uses_namespace_and_reuses_the_token() { let server: MockServer = MockServer::start().await; @@ -216,6 +228,7 @@ async fn approle_login_uses_namespace_and_reuses_the_token() { assert!(manager.async_read_secret("name-2").await.unwrap().is_none()); } +#[rstest] #[tokio::test] async fn approle_tokens_expire_after_the_vault_lease() { let server: MockServer = MockServer::start().await; @@ -246,6 +259,7 @@ async fn approle_tokens_expire_after_the_vault_lease() { assert!(manager.async_read_secret("second").await.unwrap().is_some()); } +#[rstest] #[tokio::test] async fn tls_login_posts_the_role_and_uses_the_client_identity() { let server: MockServer = MockServer::start().await; @@ -343,21 +357,29 @@ async fn tls_login_posts_the_role_and_uses_the_client_identity() { assert!(login_bodies.contains(&json!({}))); } -#[rstest::rstest] -#[case::missing(404, json!({"errors": ["missing"]}), 0)] -#[case::malformed(200, json!({"data": "invalid"}), 1)] -#[case::missing_key(200, json!({}), 0)] -#[case::non_string(200, json!({"key": 1}), 2)] +#[derive(Clone, Copy)] +enum ExpectedRead { + Missing, + Malformed, + NonString, +} + +#[rstest] +#[case::missing(404, json!({"errors": ["missing"]}), ExpectedRead::Missing)] +#[case::malformed(200, json!({"data": "invalid"}), ExpectedRead::Malformed)] +#[case::missing_key(200, json!({}), ExpectedRead::Missing)] +#[case::non_string(200, json!({"key": 1}), ExpectedRead::NonString)] #[tokio::test] async fn read_responses_distinguish_absence_and_malformed_payloads( + token_values: Vec<(&str, &str)>, #[case] status: u16, #[case] body: serde_json::Value, - #[case] expected: u8, + #[case] expected: ExpectedRead, ) { let server: MockServer = MockServer::start().await; Mock::given(method("GET")) .respond_with(ResponseTemplate::new(status).set_body_json( - if status == 200 && expected != 1 { + if status == 200 && !matches!(expected, ExpectedRead::Malformed) { read_response(body) } else { body @@ -366,20 +388,19 @@ async fn read_responses_distinguish_absence_and_malformed_payloads( .expect(1) .mount(&server) .await; - let result: Result, Error> = - manager(&server, &[("HCP_VAULT_TOKEN", "token")]) - .async_read_secret("name") - .await; + let result: Result, Error> = manager(&server, &token_values) + .async_read_secret("name") + .await; match expected { - 0 => assert!(result.unwrap().is_none()), - 1 => assert!(matches!(result, Err(Error::MalformedPayload))), - 2 => assert!(matches!(result, Err(Error::NonStringValue))), - _ => unreachable!(), + ExpectedRead::Missing => assert!(result.unwrap().is_none()), + ExpectedRead::Malformed => assert!(matches!(result, Err(Error::MalformedPayload))), + ExpectedRead::NonString => assert!(matches!(result, Err(Error::NonStringValue))), } } +#[rstest] #[tokio::test] -async fn write_and_delete_invalidate_the_read_cache() { +async fn write_and_delete_invalidate_the_read_cache(token_values: Vec<(&str, &str)>) { let server: MockServer = MockServer::start().await; Mock::given(method("GET")) .and(path("/v1/secret/data/name")) @@ -418,7 +439,7 @@ async fn write_and_delete_invalidate_the_read_cache() { .expect(1) .mount(&server) .await; - let manager: HashicorpVault = manager(&server, &[("HCP_VAULT_TOKEN", "token")]); + let manager: HashicorpVault = manager(&server, &token_values); assert!(manager.async_read_secret("name").await.unwrap().is_some()); assert!( @@ -431,6 +452,313 @@ async fn write_and_delete_invalidate_the_read_cache() { manager.async_delete_secret("name").await.unwrap(); } +#[rstest] +#[tokio::test] +async fn base_manager_context_overrides_vault_location_and_data_key( + token_values: Vec<(&str, &str)>, +) { + let server: MockServer = MockServer::start().await; + Mock::given(method("GET")) + .and(path("/v1/alternate/data/managed/name")) + .respond_with( + ResponseTemplate::new(200).set_body_json(read_response(json!({"api_token": "value"}))), + ) + .expect(2) + .mount(&server) + .await; + Mock::given(method("POST")) + .and(path("/v1/alternate/data/managed/name")) + .and(body_json(json!({ + "data": {"api_token": "updated", "description": "Managed key"} + }))) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({ + "data": { + "created_time": "", + "deletion_time": "", + "custom_metadata": null, + "destroyed": false, + "version": 2 + }, + "lease_id": "", + "lease_duration": 0, + "renewable": false, + "request_id": "", + "warnings": null, + "wrap_info": null + }))) + .expect(1) + .mount(&server) + .await; + Mock::given(method("DELETE")) + .and(path("/v1/alternate/data/managed/name")) + .respond_with(ResponseTemplate::new(204)) + .expect(1) + .mount(&server) + .await; + let manager: HashicorpVault = manager(&server, &token_values); + let operation = SecretOperationContext::Hashicorp(HashicorpOperationContext { + mount: Some(" /alternate/ ".to_owned()), + path_prefix: Some(" /managed/ ".to_owned()), + data_key: Some("api_token".to_owned()), + ..HashicorpOperationContext::default() + }); + let write_context = SecretWriteContext { + description: Some("Managed key".to_owned()), + operation: operation.clone(), + ..SecretWriteContext::default() + }; + + assert_eq!( + BaseSecretManager::async_read_secret(&manager, "name", &operation) + .await + .unwrap() + .unwrap() + .expose(), + "value" + ); + BaseSecretManager::async_write_secret( + &manager, + "name", + &SecretValue::new("updated"), + &write_context, + ) + .await + .unwrap(); + assert!( + BaseSecretManager::async_read_secret(&manager, "name", &operation) + .await + .unwrap() + .is_some() + ); + BaseSecretManager::async_delete_secret(&manager, "name", None, &operation) + .await + .unwrap(); +} + +#[rstest] +#[tokio::test] +async fn reads_cache_each_data_key_for_the_same_vault_path(token_values: Vec<(&str, &str)>) { + let server: MockServer = MockServer::start().await; + Mock::given(method("GET")) + .and(path("/v1/secret/data/name")) + .respond_with( + ResponseTemplate::new(200).set_body_json(read_response(json!({ + "key": "primary", + "alternate": "secondary" + }))), + ) + .expect(2) + .mount(&server) + .await; + let manager: HashicorpVault = manager(&server, &token_values); + let alternate = SecretOperationContext::Hashicorp(HashicorpOperationContext { + data_key: Some("alternate".to_owned()), + ..HashicorpOperationContext::default() + }); + + assert_eq!( + manager + .async_read_secret("name") + .await + .unwrap() + .unwrap() + .expose(), + "primary" + ); + assert_eq!( + BaseSecretManager::async_read_secret(&manager, "name", &alternate) + .await + .unwrap() + .unwrap() + .expose(), + "secondary" + ); + assert_eq!( + manager + .async_read_secret("name") + .await + .unwrap() + .unwrap() + .expose(), + "primary" + ); +} + +#[rstest] +#[tokio::test] +async fn base_manager_context_timeout_limits_vault_io(token_values: Vec<(&str, &str)>) { + let server: MockServer = MockServer::start().await; + Mock::given(method("GET")) + .and(path("/v1/secret/data/name")) + .respond_with( + ResponseTemplate::new(200) + .set_delay(Duration::from_millis(100)) + .set_body_json(read_response(json!({"key": "value"}))), + ) + .expect(1) + .mount(&server) + .await; + let manager: HashicorpVault = manager(&server, &token_values); + let context = SecretOperationContext::Hashicorp(HashicorpOperationContext { + timeout: Some(Duration::from_millis(10)), + ..HashicorpOperationContext::default() + }); + + assert!(matches!( + BaseSecretManager::async_read_secret(&manager, "name", &context).await, + Err(Error::Timeout) + )); +} + +#[rstest] +#[case::aws(SecretOperationContext::Aws(AwsOperationContext::default()))] +#[case::cyberark(SecretOperationContext::Cyberark(CyberarkOperationContext::default()))] +#[tokio::test] +async fn foreign_contexts_cannot_access_vault_secrets( + token_values: Vec<(&str, &str)>, + #[case] context: SecretOperationContext, + #[values(false, true)] cached: bool, +) { + let server = MockServer::start().await; + Mock::given(method("GET")) + .and(path("/v1/secret/data/name")) + .respond_with( + ResponseTemplate::new(200).set_body_json(read_response(json!({"key": "value"}))), + ) + .expect(u64::from(cached)) + .mount(&server) + .await; + let manager = manager(&server, &token_values); + if cached { + assert!(manager.async_read_secret("name").await.unwrap().is_some()); + } + assert!(matches!( + BaseSecretManager::async_read_secret(&manager, "name", &context).await, + Err(Error::InvalidOperationContext) + )); + assert!(matches!( + BaseSecretManager::async_write_secret( + &manager, + "name", + &SecretValue::new("replacement"), + &SecretWriteContext { + operation: context.clone(), + ..SecretWriteContext::default() + }, + ) + .await, + Err(Error::InvalidOperationContext) + )); + assert!(matches!( + BaseSecretManager::async_delete_secret(&manager, "name", None, &context).await, + Err(Error::InvalidOperationContext) + )); + assert!(matches!( + manager + .async_rotate_secret_with_context( + "name", + "new", + &SecretValue::new("replacement"), + &context + ) + .await, + Err(Error::InvalidOperationContext) + )); + assert_eq!( + server.received_requests().await.unwrap().len(), + usize::from(cached) + ); +} + +#[rstest] +#[tokio::test] +async fn rotation_applies_timeout_to_each_request(token_values: Vec<(&str, &str)>) { + let server = MockServer::start().await; + let timeout = Duration::from_secs(1); + let delay = timeout / 2; + Mock::given(method("GET")) + .and(path("/v1/alternate/data/managed/current")) + .respond_with( + ResponseTemplate::new(200) + .set_delay(delay) + .set_body_json(read_response(json!({"api_token": "original"}))), + ) + .mount(&server) + .await; + Mock::given(method("POST")) + .and(path("/v1/alternate/data/managed/new")) + .and(body_json(json!({ + "data": {"api_token": "replacement", "description": "Rotated from current"} + }))) + .respond_with( + ResponseTemplate::new(200) + .set_delay(delay) + .set_body_json(json!({ + "data": { + "created_time": "", + "deletion_time": "", + "custom_metadata": null, + "destroyed": false, + "version": 1 + }, + "lease_id": "", + "lease_duration": 0, + "renewable": false, + "request_id": "", + "warnings": null, + "wrap_info": null + })), + ) + .mount(&server) + .await; + Mock::given(method("GET")) + .and(path("/v1/alternate/data/managed/new")) + .respond_with( + ResponseTemplate::new(200) + .set_delay(delay) + .set_body_json(read_response(json!({"api_token": "replacement"}))), + ) + .mount(&server) + .await; + Mock::given(method("DELETE")) + .and(path("/v1/alternate/data/managed/current")) + .respond_with(ResponseTemplate::new(204).set_delay(delay)) + .mount(&server) + .await; + let manager = manager(&server, &token_values); + let context = SecretOperationContext::Hashicorp(HashicorpOperationContext { + timeout: Some(timeout), + mount: Some("alternate".to_owned()), + path_prefix: Some("managed".to_owned()), + data_key: Some("api_token".to_owned()), + }); + + manager + .async_rotate_secret_with_context( + "current", + "new", + &SecretValue::new("replacement"), + &context, + ) + .await + .unwrap(); + let requests = server.received_requests().await.unwrap(); + let operations: Vec<_> = requests + .iter() + .map(|request| (request.method.as_str(), request.url.path())) + .collect(); + assert_eq!( + operations, + [ + ("GET", "/v1/alternate/data/managed/current"), + ("POST", "/v1/alternate/data/managed/new"), + ("GET", "/v1/alternate/data/managed/new"), + ("DELETE", "/v1/alternate/data/managed/current"), + ] + ); +} + +#[rstest] #[tokio::test] async fn no_auth_and_invalid_names_fail_without_requests() { let server: MockServer = MockServer::start().await; @@ -447,6 +775,7 @@ async fn no_auth_and_invalid_names_fail_without_requests() { assert!(server.received_requests().await.unwrap().is_empty()); } +#[rstest] #[tokio::test] async fn debug_output_redacts_authentication_values() { let server: MockServer = MockServer::start().await; @@ -468,14 +797,18 @@ struct ParityCase { secret_name: String, } -#[test] -fn configuration_matches_python_parity_fixture() { - let cases: Vec = serde_json::from_str(include_str!(concat!( +#[fixture] +fn parity_cases() -> Vec { + serde_json::from_str(include_str!(concat!( env!("CARGO_MANIFEST_DIR"), "/../../../tests/test_litellm/secret_managers/hashicorp_vault_parity.json" ))) - .unwrap(); - for case in cases { + .unwrap() +} + +#[rstest] +fn configuration_matches_python_parity_fixture(parity_cases: Vec) { + for case in parity_cases { let values: HashMap = case.env.clone(); let environment: Arc = Arc::new(move |name: &str| values.get(name).cloned()); @@ -521,6 +854,7 @@ fn configuration_matches_python_parity_fixture() { } } +#[rstest] #[tokio::test] #[ignore] async fn live_vault_round_trip() { diff --git a/litellm-rust/crates/secrets-types/src/base_secret_manager.rs b/litellm-rust/crates/secrets-types/src/base_secret_manager.rs index d71bce64221..e8aed9f280e 100644 --- a/litellm-rust/crates/secrets-types/src/base_secret_manager.rs +++ b/litellm-rust/crates/secrets-types/src/base_secret_manager.rs @@ -1,4 +1,4 @@ -use crate::{Error, SecretValue}; +use crate::{Error, SecretOperationContext, SecretValue, SecretWriteContext}; pub fn validate_secret_name(name: &str) -> Result<(), Error> { if name.split('/').any(|segment| segment == "..") @@ -20,17 +20,22 @@ pub trait BaseSecretManager { type WriteResponse; type DeleteResponse; - async fn async_read_secret(&self, name: &str) -> Result, Self::Error>; + async fn async_read_secret( + &self, + name: &str, + context: &SecretOperationContext, + ) -> Result, Self::Error>; async fn async_write_secret( &self, name: &str, value: &SecretValue, - description: Option<&str>, + context: &SecretWriteContext, ) -> Result; async fn async_delete_secret( &self, name: &str, - recovery_window_in_days: i64, + recovery_window_in_days: Option, + context: &SecretOperationContext, ) -> Result; } @@ -39,20 +44,31 @@ pub async fn async_rotate_secret( current_name: &str, new_name: &str, value: &SecretValue, + context: &SecretOperationContext, ) -> Result { - if manager.async_read_secret(current_name).await?.is_none() { + if manager + .async_read_secret(current_name, context) + .await? + .is_none() + { return Err(Error::CurrentSecretMissing.into()); } let response = manager .async_write_secret( new_name, value, - Some(&format!("Rotated from {current_name}")), + &SecretWriteContext::rotated_from(current_name, context.clone()), ) .await?; - if manager.async_read_secret(new_name).await?.is_none() { + if manager + .async_read_secret(new_name, context) + .await? + .is_none() + { return Err(Error::NewSecretMissing.into()); } - manager.async_delete_secret(current_name, 7).await?; + manager + .async_delete_secret(current_name, Some(7), context) + .await?; Ok(response) } diff --git a/litellm-rust/crates/secrets-types/src/context.rs b/litellm-rust/crates/secrets-types/src/context.rs new file mode 100644 index 00000000000..126815ab32b --- /dev/null +++ b/litellm-rust/crates/secrets-types/src/context.rs @@ -0,0 +1,65 @@ +use std::{collections::BTreeMap, time::Duration}; + +use crate::SecretValue; + +#[derive(Clone, Debug, Default, Eq, PartialEq)] +pub enum SecretOperationContext { + #[default] + Default, + Aws(AwsOperationContext), + Hashicorp(HashicorpOperationContext), + Cyberark(CyberarkOperationContext), +} + +impl SecretOperationContext { + pub fn timeout(&self) -> Option { + match self { + Self::Default => None, + Self::Aws(context) => context.timeout, + Self::Hashicorp(context) => context.timeout, + Self::Cyberark(context) => context.timeout, + } + } +} + +#[derive(Clone, Debug, Default, Eq, PartialEq)] +pub struct AwsOperationContext { + pub timeout: Option, + pub region_name: Option, + pub role_name: Option, + pub session_name: Option, + pub external_id: Option, + pub profile_name: Option, + pub web_identity_token: Option, + pub sts_endpoint: Option, +} + +#[derive(Clone, Debug, Default, Eq, PartialEq)] +pub struct HashicorpOperationContext { + pub timeout: Option, + pub mount: Option, + pub path_prefix: Option, + pub data_key: Option, +} + +#[derive(Clone, Debug, Default, Eq, PartialEq)] +pub struct CyberarkOperationContext { + pub timeout: Option, +} + +#[derive(Clone, Debug, Default, Eq, PartialEq)] +pub struct SecretWriteContext { + pub description: Option, + pub tags: BTreeMap, + pub operation: SecretOperationContext, +} + +impl SecretWriteContext { + pub fn rotated_from(current_name: &str, operation: SecretOperationContext) -> Self { + Self { + description: Some(format!("Rotated from {current_name}")), + tags: BTreeMap::new(), + operation, + } + } +} diff --git a/litellm-rust/crates/secrets-types/src/lib.rs b/litellm-rust/crates/secrets-types/src/lib.rs index 0823ed13c06..d330fc7bbea 100644 --- a/litellm-rust/crates/secrets-types/src/lib.rs +++ b/litellm-rust/crates/secrets-types/src/lib.rs @@ -2,11 +2,16 @@ mod base_secret_manager; mod config; +mod context; mod error; mod value; pub use base_secret_manager::{BaseSecretManager, async_rotate_secret, validate_secret_name}; pub use config::{AccessMode, KeyManagementSettings, KeyManagementSystem}; +pub use context::{ + AwsOperationContext, CyberarkOperationContext, HashicorpOperationContext, + SecretOperationContext, SecretWriteContext, +}; pub use error::Error; pub use litellm_auth_types::SecretValue; pub use value::Secret; diff --git a/litellm-rust/crates/secrets-types/tests/config.rs b/litellm-rust/crates/secrets-types/tests/config.rs index 4a5f17bc68a..45fd9d0ee03 100644 --- a/litellm-rust/crates/secrets-types/tests/config.rs +++ b/litellm-rust/crates/secrets-types/tests/config.rs @@ -1,15 +1,20 @@ use litellm_secrets_types::{ AccessMode, KeyManagementSettings, KeyManagementSystem, Secret, SecretValue, }; +use rstest::{fixture, rstest}; use serde_json::json; -#[test] -fn config_preserves_defaults_nulls_and_serialized_names() { - let empty: KeyManagementSettings = serde_json::from_value(json!({})).unwrap(); - assert_eq!(empty, KeyManagementSettings::default()); - assert_eq!(empty.access_mode, AccessMode::ReadOnly); - assert_eq!(empty.store_virtual_keys, Some(false)); - assert_eq!(empty.prefix_for_stored_virtual_keys, "litellm/"); +#[fixture] +fn default_settings() -> KeyManagementSettings { + serde_json::from_value(json!({})).unwrap() +} + +#[rstest] +fn config_preserves_defaults_nulls_and_serialized_names(default_settings: KeyManagementSettings) { + assert_eq!(default_settings, KeyManagementSettings::default()); + assert_eq!(default_settings.access_mode, AccessMode::ReadOnly); + assert_eq!(default_settings.store_virtual_keys, Some(false)); + assert_eq!(default_settings.prefix_for_stored_virtual_keys, "litellm/"); let configured: KeyManagementSettings = serde_json::from_value(json!({ "hosted_keys": [], "store_virtual_keys": null, "access_mode": "write_only", "aws_web_identity_token": "private-token", "aws_external_id": "private-id", @@ -29,7 +34,7 @@ fn config_preserves_defaults_nulls_and_serialized_names() { ); } -#[rstest::rstest] +#[rstest] #[case::aws_kms("aws_kms", KeyManagementSystem::AwsKms)] #[case::aws_secret_manager("aws_secret_manager", KeyManagementSystem::AwsSecretManager)] #[case::google_kms("google_kms", KeyManagementSystem::GoogleKms)] @@ -50,11 +55,32 @@ fn key_management_system_serialization_round_trips( assert_eq!(serde_json::to_value(system).unwrap(), name); } -#[test] -fn secret_debug_never_exposes_values() { - assert!( - !format!("{:?}", Secret::String(SecretValue::new("sensitive-value"))) - .contains("sensitive-value") - ); - assert!(!format!("{:?}", Secret::Bool(true)).contains("true")); +#[rstest] +#[case::read_only(AccessMode::ReadOnly, true)] +#[case::write_only(AccessMode::WriteOnly, false)] +#[case::read_and_write(AccessMode::ReadAndWrite, true)] +fn access_mode_reports_readability(#[case] mode: AccessMode, #[case] expected: bool) { + assert_eq!(mode.readable(), expected); +} + +#[rstest] +#[case::string(json!("value"), Secret::String(SecretValue::new("value")))] +#[case::boolean(json!(true), Secret::Bool(true))] +#[case::number(json!(7), Secret::Json(json!(7)))] +#[case::null(json!(null), Secret::Json(json!(null)))] +#[case::array(json!(["value"]), Secret::Json(json!(["value"])))] +#[case::object(json!({"key": "value"}), Secret::Json(json!({"key": "value"})))] +fn secret_conversion_preserves_json_types( + #[case] value: serde_json::Value, + #[case] expected: Secret, +) { + assert_eq!(Secret::from_json(value), expected); +} + +#[rstest] +#[case::string(Secret::String(SecretValue::new("sensitive-value")), "sensitive-value")] +#[case::boolean(Secret::Bool(true), "true")] +#[case::json(Secret::Json(json!({"private": "value"})), "private")] +fn secret_debug_never_exposes_values(#[case] secret: Secret, #[case] sensitive: &str) { + assert!(!format!("{secret:?}").contains(sensitive)); } diff --git a/litellm-rust/crates/secrets-types/tests/context.rs b/litellm-rust/crates/secrets-types/tests/context.rs new file mode 100644 index 00000000000..7c9e91b3a9e --- /dev/null +++ b/litellm-rust/crates/secrets-types/tests/context.rs @@ -0,0 +1,110 @@ +use std::{collections::BTreeMap, time::Duration}; + +use litellm_secrets_types::{ + AwsOperationContext, CyberarkOperationContext, HashicorpOperationContext, + SecretOperationContext, SecretValue, SecretWriteContext, +}; +use rstest::{fixture, rstest}; + +#[fixture] +fn timeout() -> Duration { + Duration::from_secs(30) +} + +#[fixture] +fn aws_context(timeout: Duration) -> SecretOperationContext { + SecretOperationContext::Aws(AwsOperationContext { + timeout: Some(timeout), + region_name: Some("us-west-2".into()), + role_name: Some("role".into()), + external_id: Some(SecretValue::new("external-id")), + web_identity_token: Some(SecretValue::new("web-identity-token")), + ..AwsOperationContext::default() + }) +} + +#[rstest] +fn operation_context_preserves_backend_specific_values_and_redacts_secrets( + aws_context: SecretOperationContext, + timeout: Duration, +) { + assert_eq!(aws_context.timeout(), Some(timeout)); + assert_eq!( + aws_context, + SecretOperationContext::Aws(AwsOperationContext { + timeout: Some(timeout), + region_name: Some("us-west-2".into()), + role_name: Some("role".into()), + external_id: Some(SecretValue::new("external-id")), + web_identity_token: Some(SecretValue::new("web-identity-token")), + ..AwsOperationContext::default() + }) + ); + let debug = format!("{aws_context:?}"); + assert!(!debug.contains("external-id")); + assert!(!debug.contains("web-identity-token")); +} + +#[rstest] +#[case::default(SecretOperationContext::Default, None)] +#[case::aws( + SecretOperationContext::Aws(AwsOperationContext { + timeout: Some(Duration::from_secs(1)), + ..AwsOperationContext::default() + }), + Some(Duration::from_secs(1)) +)] +#[case::hashicorp( + SecretOperationContext::Hashicorp(HashicorpOperationContext { + timeout: Some(Duration::from_secs(2)), + ..HashicorpOperationContext::default() + }), + Some(Duration::from_secs(2)) +)] +#[case::cyberark( + SecretOperationContext::Cyberark(CyberarkOperationContext { + timeout: Some(Duration::from_secs(3)), + }), + Some(Duration::from_secs(3)) +)] +fn operation_context_reports_each_backend_timeout( + #[case] context: SecretOperationContext, + #[case] expected: Option, +) { + assert_eq!(context.timeout(), expected); +} + +#[rstest] +fn write_context_keeps_tags_separate_from_the_operation_context() { + let context = SecretWriteContext { + description: Some("Managed virtual key".into()), + tags: BTreeMap::from([("team".into(), "team-id".into())]), + operation: SecretOperationContext::Hashicorp(HashicorpOperationContext { + mount: Some("secret".into()), + path_prefix: Some("teams".into()), + data_key: Some("api_key".into()), + ..HashicorpOperationContext::default() + }), + }; + + assert_eq!(context.tags.get("team"), Some(&"team-id".into())); + assert_eq!(context.operation.timeout(), None); + assert_eq!( + context.operation, + SecretOperationContext::Hashicorp(HashicorpOperationContext { + mount: Some("secret".into()), + path_prefix: Some("teams".into()), + data_key: Some("api_key".into()), + ..HashicorpOperationContext::default() + }) + ); +} + +#[rstest] +fn rotation_write_context_preserves_the_operation_context(aws_context: SecretOperationContext) { + let context = SecretWriteContext::rotated_from("current", aws_context.clone()); + + assert_eq!(context.description.as_deref(), Some("Rotated from current")); + assert!(context.tags.is_empty()); + assert_eq!(context.operation, aws_context); +} diff --git a/litellm-rust/crates/secrets-types/tests/rotation.rs b/litellm-rust/crates/secrets-types/tests/rotation.rs index 48a5304ece8..1bc5e0a9eeb 100644 --- a/litellm-rust/crates/secrets-types/tests/rotation.rs +++ b/litellm-rust/crates/secrets-types/tests/rotation.rs @@ -1,12 +1,15 @@ use std::sync::atomic::{AtomicUsize, Ordering}; use litellm_secrets_types::{ - BaseSecretManager, Error, SecretValue, async_rotate_secret, validate_secret_name, + BaseSecretManager, Error, HashicorpOperationContext, SecretOperationContext, SecretValue, + SecretWriteContext, async_rotate_secret, validate_secret_name, }; +use rstest::{fixture, rstest}; struct Manager { step: AtomicUsize, absent_at: Option, + operation: SecretOperationContext, } impl BaseSecretManager for Manager { @@ -14,7 +17,12 @@ impl BaseSecretManager for Manager { type WriteResponse = &'static str; type DeleteResponse = (); - async fn async_read_secret(&self, name: &str) -> Result, Error> { + async fn async_read_secret( + &self, + name: &str, + context: &SecretOperationContext, + ) -> Result, Error> { + assert_eq!(context, &self.operation); let step = self.step.fetch_add(1, Ordering::SeqCst); assert_eq!(name, if step == 0 { "old" } else { "new" }); Ok((self.absent_at != Some(step)).then(|| SecretValue::new("value"))) @@ -24,35 +32,54 @@ impl BaseSecretManager for Manager { &self, name: &str, value: &SecretValue, - description: Option<&str>, + context: &SecretWriteContext, ) -> Result { assert_eq!(self.step.fetch_add(1, Ordering::SeqCst), 1); assert_eq!(name, "new"); assert_eq!(value.expose(), "replacement"); - assert_eq!(description, Some("Rotated from old")); + assert_eq!(context.description.as_deref(), Some("Rotated from old")); + assert!(context.tags.is_empty()); + assert_eq!(context.operation, self.operation); Ok("provider-response") } async fn async_delete_secret( &self, name: &str, - recovery_window_in_days: i64, + recovery_window_in_days: Option, + context: &SecretOperationContext, ) -> Result<(), Error> { assert_eq!(self.step.fetch_add(1, Ordering::SeqCst), 3); assert_eq!(name, "old"); - assert_eq!(recovery_window_in_days, 7); + assert_eq!(recovery_window_in_days, Some(7)); + assert_eq!(context, &self.operation); Ok(()) } } +#[fixture] +fn replacement() -> SecretValue { + SecretValue::new("replacement") +} + +#[rstest] +#[case::default(SecretOperationContext::Default)] +#[case::backend_specific(SecretOperationContext::Hashicorp(HashicorpOperationContext { + mount: Some("alternate".to_owned()), + ..HashicorpOperationContext::default() +}))] #[tokio::test] -async fn rotation_verifies_before_deleting_and_returns_provider_response() { +async fn rotation_verifies_before_deleting_and_returns_provider_response( + replacement: SecretValue, + #[case] operation: SecretOperationContext, +) { let manager = Manager { step: AtomicUsize::new(0), absent_at: None, + operation: operation.clone(), }; assert_eq!( - async_rotate_secret(&manager, "old", "new", &SecretValue::new("replacement")) + async_rotate_secret(&manager, "old", "new", &replacement, &operation,) .await .unwrap(), "provider-response" @@ -60,11 +87,12 @@ async fn rotation_verifies_before_deleting_and_returns_provider_response() { assert_eq!(manager.step.load(Ordering::SeqCst), 4); } -#[rstest::rstest] +#[rstest] #[case::current_secret_missing(0, Error::CurrentSecretMissing, 1)] #[case::new_secret_missing(2, Error::NewSecretMissing, 3)] #[tokio::test] async fn missing_old_or_new_value_stops_rotation_before_deletion( + replacement: SecretValue, #[case] absent_at: usize, #[case] expected: Error, #[case] calls: usize, @@ -72,32 +100,57 @@ async fn missing_old_or_new_value_stops_rotation_before_deletion( let manager = Manager { step: AtomicUsize::new(0), absent_at: Some(absent_at), + operation: SecretOperationContext::default(), }; assert_eq!( - async_rotate_secret(&manager, "old", "new", &SecretValue::new("replacement")) - .await - .unwrap_err(), + async_rotate_secret( + &manager, + "old", + "new", + &replacement, + &SecretOperationContext::default(), + ) + .await + .unwrap_err(), expected ); assert_eq!(manager.step.load(Ordering::SeqCst), calls); } -#[rstest::rstest] +#[rstest] #[case::parent("..")] -#[case::parent_prefix("../x")] -#[case::parent_segment("x/../y")] -#[case::parent_suffix("x/..")] -#[case::line_feed("line\n")] -#[case::next_line("\u{85}")] -#[case::line_separator("\u{2028}")] -#[case::paragraph_separator("\u{2029}")] +#[case::parent_prefix("../../../other-app/creds")] +#[case::nested_parent_prefix("litellm/../../secret")] +#[case::parent_segment("foo/../bar")] +#[case::parent_suffix("foo/..")] +#[case::single_parent_prefix("../foo")] +#[case::line_feed("foo\nbar")] +#[case::carriage_return("foo\rbar")] +#[case::tab("foo\tbar")] +#[case::null("foo\0bar")] +#[case::delete("foo\u{7f}bar")] +#[case::next_line("foo\u{85}bar")] +#[case::line_separator("foo\u{2028}bar")] +#[case::paragraph_separator("foo\u{2029}bar")] fn names_reject_path_traversal_and_control_characters(#[case] name: &str) { assert_eq!(validate_secret_name(name), Err(Error::UnsafeSecretName)); } -#[rstest::rstest] +#[rstest] +#[case::plain_alias("plain-alias")] +#[case::key_with_digits("my-key-123")] +#[case::service_path("prod/my-service-key")] +#[case::email_path("team/user@example.com")] +#[case::colon("foo: bar")] +#[case::spaced_fragment("foo # bar")] +#[case::query("foo?evil=1")] +#[case::fragment("foo#bar")] +#[case::long_alias(&"a".repeat(500))] #[case::embedded_double_dot("release-1.0..2")] -#[case::path_separator("folder/key")] +#[case::middle_double_dot("my..key")] +#[case::leading_double_dot("..foo")] +#[case::trailing_double_dot("foo..")] +#[case::version_double_dot("v2.0..1-beta")] #[case::empty("")] #[case::three_dots("...")] fn names_allow_safe_values(#[case] name: &str) { diff --git a/litellm-rust/crates/secrets/src/handler.rs b/litellm-rust/crates/secrets/src/handler.rs index 8762b8e7405..0ae84821883 100644 --- a/litellm-rust/crates/secrets/src/handler.rs +++ b/litellm-rust/crates/secrets/src/handler.rs @@ -2,7 +2,10 @@ use std::{future::Future, pin::Pin, sync::Arc}; use litellm_core_utils::settings::Lookup; -use crate::{Error, KeyManagementSettings, KeyManagementSystem, Secret, SecretValue}; +use crate::{Error, KeyManagementSettings, KeyManagementSystem, Secret}; + +#[cfg(any(feature = "aws", feature = "google"))] +use crate::SecretValue; pub trait ExternalSecretManager: Send + Sync { fn system(&self) -> KeyManagementSystem; @@ -17,7 +20,6 @@ pub trait ExternalSecretManager: Send + Sync { #[derive(Clone)] pub enum SecretManager { - Local, External(Arc), #[cfg(feature = "aws")] AwsKms(crate::aws::AwsKms), @@ -38,7 +40,6 @@ pub enum SecretManager { impl SecretManager { pub fn system(&self) -> KeyManagementSystem { match self { - Self::Local => KeyManagementSystem::Local, Self::External(manager) => manager.system(), #[cfg(feature = "aws")] Self::AwsKms(_) => KeyManagementSystem::AwsKms, @@ -65,10 +66,6 @@ pub async fn get_secret_from_manager( environment: &(dyn Lookup + Send + Sync), ) -> Result, Error> { match client { - SecretManager::Local => Ok(environment - .get(secret_name) - .map(SecretValue::new) - .map(Secret::String)), SecretManager::External(manager) => { manager .read_secret(secret_name, _settings, environment) @@ -117,10 +114,9 @@ pub async fn get_secret_from_manager( .map(|value| value.map(Secret::String)) .map_err(Error::from), #[cfg(feature = "azure")] - SecretManager::AzureKeyVault(client) => client - .get_secret_from_azure_key_vault(secret_name) - .await - .map_err(Error::from), + SecretManager::AzureKeyVault(client) => { + client.get_secret(secret_name).await.map_err(Error::from) + } #[cfg(feature = "cyberark")] SecretManager::Cyberark(client) => client .async_read_secret(secret_name) diff --git a/litellm-rust/crates/secrets/tests/resolution.rs b/litellm-rust/crates/secrets/tests/resolution.rs index 3a826092d72..98b34a2c72e 100644 --- a/litellm-rust/crates/secrets/tests/resolution.rs +++ b/litellm-rust/crates/secrets/tests/resolution.rs @@ -1,19 +1,14 @@ use std::sync::Arc; use litellm_secrets::{ - Error, KeyManagementSettings, OidcResolver, Secret, SecretManager, SecretManagerState, - SecretResolver, SecretValue, secret_manager_would_be_consulted, + Error, OidcResolver, Secret, SecretManagerState, SecretResolver, SecretValue, + secret_manager_would_be_consulted, }; -fn resolver(value: Option<&str>, configured: bool) -> SecretResolver { - let state = if configured { - SecretManagerState::new(SecretManager::Local, KeyManagementSettings::default()) - } else { - SecretManagerState::default() - }; +fn resolver(value: Option<&str>) -> SecretResolver { let value = value.map(str::to_owned); SecretResolver::new( - Arc::new(state), + Arc::new(SecretManagerState::default()), Arc::new(move |_: &str| value.clone()), OidcResolver::default(), ) @@ -30,9 +25,8 @@ fn resolver(value: Option<&str>, configured: bool) -> SecretResolver { async fn conversion_is_explicit_and_independent_of_manager_configuration( #[case] input: &str, #[case] boolean: Option, - #[values(false, true)] configured: bool, ) { - let resolver = resolver(Some(input), configured); + let resolver = resolver(Some(input)); assert_eq!( resolver.get_secret("key", None).await.unwrap(), Some(Secret::String(SecretValue::new(input))) @@ -62,8 +56,8 @@ async fn conversion_is_explicit_and_independent_of_manager_configuration( #[rstest::rstest] #[tokio::test] -async fn defaults_apply_only_to_absence(#[values(false, true)] configured: bool) { - let missing = resolver(None, configured); +async fn defaults_apply_only_to_absence() { + let missing = resolver(None); assert_eq!(missing.get_secret("key", None).await.unwrap(), None); assert_eq!( missing.get_secret_bool("key", Some(false)).await.unwrap(), @@ -92,7 +86,7 @@ async fn defaults_apply_only_to_absence(#[values(false, true)] configured: bool) ); } assert_eq!( - resolver(Some(""), configured) + resolver(Some("")) .get_secret_str("key", Some(SecretValue::new("default"))) .await .unwrap() @@ -103,12 +97,8 @@ async fn defaults_apply_only_to_absence(#[values(false, true)] configured: bool) } #[tokio::test] -async fn prefix_is_removed_once_and_local_manager_is_not_consulted() { - let state = SecretManagerState::new(SecretManager::Local, KeyManagementSettings::default()); - assert_eq!( - state.system(), - Some(litellm_secrets::KeyManagementSystem::Local) - ); +async fn prefix_is_removed_once_and_resolved_from_environment() { + let state = SecretManagerState::default(); assert!(!secret_manager_would_be_consulted( &state, "os.environ/os.environ/KEY" @@ -131,7 +121,7 @@ async fn prefix_is_removed_once_and_local_manager_is_not_consulted() { #[tokio::test] async fn resolver_future_can_run_on_a_tokio_worker() { - let resolver = resolver(Some("worker-value"), false); + let resolver = resolver(Some("worker-value")); let result = tokio::spawn(async move { resolver.get_secret_str("KEY", None).await }) .await .unwrap() @@ -142,7 +132,9 @@ async fn resolver_future_can_run_on_a_tokio_worker() { #[cfg(feature = "aws")] mod aws { use super::*; - use litellm_secrets::{AccessMode, FailurePolicy, aws::AwsSecretsManagerV2}; + use litellm_secrets::{ + AccessMode, FailurePolicy, KeyManagementSettings, SecretManager, aws::AwsSecretsManagerV2, + }; use wiremock::{Mock, MockServer, ResponseTemplate, matchers::method}; fn state(server: &MockServer, settings: KeyManagementSettings) -> SecretManagerState { @@ -319,7 +311,9 @@ mod aws { #[case::failure(503)] #[tokio::test] async fn google_resolver_distinguishes_absence_from_failure(#[case] status: u16) { - use litellm_secrets::{FailurePolicy, google::GoogleSecretManager}; + use litellm_secrets::{ + FailurePolicy, KeyManagementSettings, SecretManager, google::GoogleSecretManager, + }; use wiremock::{Mock, MockServer, ResponseTemplate, matchers::method}; let server = MockServer::start().await; Mock::given(method("GET"))