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 <yujong@berri.ai>
This commit is contained in:
devin-ai-integration[bot] 2026-09-22 09:02:15 -07:00 committed by GitHub
parent 4082523596
commit d6ffc554ec
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
28 changed files with 1839 additions and 486 deletions

View file

@ -3065,6 +3065,7 @@ dependencies = [
"serde_json",
"serde_with",
"sha2 0.10.9",
"thiserror 2.0.19",
"tokio",
"tokio-tungstenite",
"url",

View file

@ -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

View file

@ -131,7 +131,7 @@ impl RouteHost for OcrRouteHost {
fn classify(&self, py: Python<'_>, error: Error) -> PyResult<PyErr> {
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);
}

View file

@ -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<PyBaseException>);
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<PyErr> {
let Error::ExternalManager(source) = error else {
return None;
};
source
.downcast_ref::<PythonSecretError>()
.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<PyAny>,
system: Option<KeyManagementSystem>,
/// The `key_manager` name Python's handler dispatches on.
key_manager: &'static str,
settings: Option<Py<PyAny>>,
}
@ -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<KeyManagementSystem>,
#[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::<String>()
.unwrap(),
"custom"
);
} else {
assert!(optional_params.is_none());
}
});
}

View file

@ -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<PyBaseException>);
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<PyErr> {
let Error::ExternalManager(source) = error else {
return None;
};
source
.downcast_ref::<PythonSecretError>()
.map(|error| PyErr::from_value(error.0.clone_ref(py).into_bound(py).into_any()))
}

View file

@ -1,3 +1,6 @@
pub(crate) mod callback;
pub(crate) mod config;
mod error;
pub(crate) mod resolved;
pub(crate) use error::python_error;

View file

@ -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")]

View file

@ -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<Box<ContextClientFactory>>,
write_settings: AwsSecretWriteSettings,
}
#[derive(Clone)]
struct ContextClientFactory {
settings: KeyManagementSettings,
environment: Arc<dyn Lookup + Send + Sync>,
endpoint_url: Option<String>,
}
#[derive(Clone, Debug, Default)]
pub struct AwsSecretWriteSettings {
pub kms_key_id: Option<String>,
@ -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<Option<SecretValue>, 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<Option<SecretValue>, 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<CreateSecretOutput, Error> {
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<String, String>>,
) -> Result<CreateSecretOutput, Error> {
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<Option<ReplicateSecretToRegionsOutput>, 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<Option<ReplicateSecretToRegionsOutput>, 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<PutSecretValueOutput, Error> {
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<PutSecretValueOutput, Error> {
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<u32>,
) -> Result<DeleteSecretOutput, Error> {
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<u32>,
) -> Result<DeleteSecretOutput, Error> {
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<RotationResponse, Error> {
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<RotationResponse, Error> {
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<Client, Error> {
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<Client, Error> {
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<Option<SecretValue>, Error> {
self.async_read_secret(name).await
async fn async_read_secret(
&self,
name: &str,
context: &SecretOperationContext,
) -> Result<Option<SecretValue>, 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<CreateSecretOutput, Error> {
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<u32>,
context: &SecretOperationContext,
) -> Result<DeleteSecretOutput, Error> {
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
}
}

View file

@ -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<bool>) {
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)
));
}

View file

@ -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<dyn litellm_core_utils::settings::Lookup + Send + Sync> =
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(_))

View file

@ -76,10 +76,7 @@ impl AzureKeyVault {
.unwrap_or_default()
}
pub async fn get_secret_from_azure_key_vault(
&self,
name: &str,
) -> Result<Option<Secret>, Error> {
pub async fn get_secret(&self, name: &str) -> Result<Option<Secret>, Error> {
let token = self
.auth
.get_azure_ad_token(&self.inputs, &|key| self.environment.get(key))

View file

@ -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<u16>) {
let server = MockServer::start().await;
@ -73,9 +79,7 @@ async fn handles_statuses(#[case] status: u16, #[case] expected_status: Option<u
.mount(&server)
.await;
let result = manager(&server)
.get_secret_from_azure_key_vault("NAME")
.await;
let result = manager(&server).get_secret("NAME").await;
match expected_status {
None => assert_eq!(result.unwrap(), None),
@ -83,6 +87,7 @@ async fn handles_statuses(#[case] status: u16, #[case] expected_status: Option<u
}
}
#[rstest]
#[tokio::test]
async fn missing_value_is_an_error() {
let server = MockServer::start().await;
@ -93,41 +98,43 @@ async fn missing_value_is_an_error() {
.await;
assert!(matches!(
manager(&server)
.get_secret_from_azure_key_vault("NAME")
.await,
manager(&server).get_secret("NAME").await,
Err(Error::MissingValue)
));
}
#[test]
fn new_validates_vault_environment() {
assert!(matches!(
AzureKeyVault::new(Arc::new(|_: &str| None)),
Err(Error::MissingEnvironment("AZURE_KEY_VAULT_URI"))
));
assert!(matches!(
AzureKeyVault::new(Arc::new(|name: &str| {
(name == "AZURE_KEY_VAULT_URI").then(|| "http://vault.example".to_owned())
})),
Err(Error::VaultUri)
));
assert!(matches!(
AzureKeyVault::new(Arc::new(|name: &str| {
(name == "AZURE_KEY_VAULT_URI").then(|| "vault.example".to_owned())
})),
Err(Error::VaultUri)
));
#[rstest]
#[case::missing(None, true)]
#[case::http(Some("http://vault.example"), false)]
#[case::relative(Some("vault.example"), false)]
#[case::malformed(Some("://"), false)]
fn new_validates_vault_environment(
#[case] uri: Option<&'static str>,
#[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<bool>,
}
#[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) {

View file

@ -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()

View file

@ -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<SecretValue, Error> {
async fn authenticate(&self, context: &SecretOperationContext) -> Result<SecretValue, Error> {
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<String, Error> {
async fn authorization_header(
&self,
context: &SecretOperationContext,
) -> Result<String, Error> {
Ok(format!(
"Token token=\"{}\"",
self.authenticate().await?.expose()
self.authenticate(context).await?.expose()
))
}
pub async fn async_read_secret(&self, name: &str) -> Result<Option<SecretValue>, 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<Option<SecretValue>, 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<u32>,
) -> Result<DeleteOutcome, Error> {
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<u32>,
_context: &SecretOperationContext,
) -> Result<DeleteOutcome, Error> {
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<Option<SecretValue>, Error> {
self.async_read_secret(name).await
async fn async_read_secret(
&self,
name: &str,
context: &SecretOperationContext,
) -> Result<Option<SecretValue>, 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<u32>,
context: &SecretOperationContext,
) -> Result<DeleteOutcome, Error> {
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()));

View file

@ -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<dyn litellm_core_utils::settings::Lookup + Send + Sync> =
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"
);
}

View file

@ -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<bool>,
) {
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());
}

View file

@ -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

View file

@ -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,
}

View file

@ -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<Instant>,
}
#[derive(Clone, Debug, PartialEq, Eq)]
#[derive(Clone, Debug, Hash, PartialEq, Eq)]
pub struct SecretLocation {
pub namespace: Option<String>,
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<String, SecretValue>,
cache: Cache<CacheKey, SecretValue>,
auth_client: Arc<Mutex<Option<CachedClient>>>,
}
@ -71,7 +79,7 @@ impl HashicorpVault {
if !enterprise_enabled {
return Err(Error::EnterpriseRequired);
}
let cache: Cache<String, SecretValue> = Cache::builder()
let cache: Cache<CacheKey, SecretValue> = 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<SecretLocation, Error> {
self.secret_location_with_context(secret_name, &SecretOperationContext::default())
}
pub fn secret_location_with_context(
&self,
secret_name: &str,
context: &SecretOperationContext,
) -> Result<SecretLocation, Error> {
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<Option<SecretValue>, 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<Option<SecretValue>, 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<VaultClient> = self.vault_client().await?;
let data: HashMap<String, Value> =
let data: Option<HashMap<String, Value>> = with_timeout(context, async {
let client: Arc<VaultClient> = 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<Value, Error> {
let location: SecretLocation = self.secret_location(secret_name)?;
let cache_key: String = cache_key(&location);
let data: HashMap<String, Value> = 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<Value, Error> {
let location: SecretLocation =
self.secret_location_with_context(secret_name, &context.operation)?;
let data_key: String = data_key(&context.operation)?;
let data: HashMap<String, Value> = 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<VaultClient> = 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<VaultClient> = 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<VaultClient> = 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<VaultClient> = 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<Value, Error> {
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<Value, Error> {
async_rotate_secret(self, current_name, new_name, value, context).await
}
async fn vault_client(&self) -> Result<Arc<VaultClient>, Error> {
@ -254,25 +343,30 @@ impl BaseSecretManager for HashicorpVault {
type WriteResponse = Value;
type DeleteResponse = ();
async fn async_read_secret(&self, name: &str) -> Result<Option<SecretValue>, Error> {
HashicorpVault::async_read_secret(self, name).await
async fn async_read_secret(
&self,
name: &str,
context: &SecretOperationContext,
) -> Result<Option<SecretValue>, 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<Value, Error> {
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<u32>,
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<Option<&HashicorpOperationContext>, 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<String> {
let value: &str = value.trim().trim_matches('/');
(!value.is_empty()).then(|| value.to_owned())
}
fn data_key(context: &SecretOperationContext) -> Result<String, Error> {
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<T>(
context: &SecretOperationContext,
operation: impl Future<Output = Result<T, Error>>,
) -> Result<T, Error> {
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<Option<Identity>, Error> {

View file

@ -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<dyn Lookup + Send + Sync> = 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<dyn Lookup + Send + Sync> = 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<Option<SecretValue>, Error> =
manager(&server, &[("HCP_VAULT_TOKEN", "token")])
.async_read_secret("name")
.await;
let result: Result<Option<SecretValue>, 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<ParityCase> = serde_json::from_str(include_str!(concat!(
#[fixture]
fn parity_cases() -> Vec<ParityCase> {
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<ParityCase>) {
for case in parity_cases {
let values: HashMap<String, String> = case.env.clone();
let environment: Arc<dyn Lookup + Send + Sync> =
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() {

View file

@ -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<Option<SecretValue>, Self::Error>;
async fn async_read_secret(
&self,
name: &str,
context: &SecretOperationContext,
) -> Result<Option<SecretValue>, Self::Error>;
async fn async_write_secret(
&self,
name: &str,
value: &SecretValue,
description: Option<&str>,
context: &SecretWriteContext,
) -> Result<Self::WriteResponse, Self::Error>;
async fn async_delete_secret(
&self,
name: &str,
recovery_window_in_days: i64,
recovery_window_in_days: Option<u32>,
context: &SecretOperationContext,
) -> Result<Self::DeleteResponse, Self::Error>;
}
@ -39,20 +44,31 @@ pub async fn async_rotate_secret<M: BaseSecretManager>(
current_name: &str,
new_name: &str,
value: &SecretValue,
context: &SecretOperationContext,
) -> Result<M::WriteResponse, M::Error> {
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)
}

View file

@ -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<Duration> {
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<Duration>,
pub region_name: Option<String>,
pub role_name: Option<String>,
pub session_name: Option<String>,
pub external_id: Option<SecretValue>,
pub profile_name: Option<String>,
pub web_identity_token: Option<SecretValue>,
pub sts_endpoint: Option<String>,
}
#[derive(Clone, Debug, Default, Eq, PartialEq)]
pub struct HashicorpOperationContext {
pub timeout: Option<Duration>,
pub mount: Option<String>,
pub path_prefix: Option<String>,
pub data_key: Option<String>,
}
#[derive(Clone, Debug, Default, Eq, PartialEq)]
pub struct CyberarkOperationContext {
pub timeout: Option<Duration>,
}
#[derive(Clone, Debug, Default, Eq, PartialEq)]
pub struct SecretWriteContext {
pub description: Option<String>,
pub tags: BTreeMap<String, String>,
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,
}
}
}

View file

@ -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;

View file

@ -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));
}

View file

@ -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<Duration>,
) {
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);
}

View file

@ -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<usize>,
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<Option<SecretValue>, Error> {
async fn async_read_secret(
&self,
name: &str,
context: &SecretOperationContext,
) -> Result<Option<SecretValue>, 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<Self::WriteResponse, Error> {
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<u32>,
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) {

View file

@ -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<dyn ExternalSecretManager>),
#[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<Option<Secret>, 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)

View file

@ -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<bool>,
#[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"))