mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-24 00:52:24 +00:00
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:
parent
4082523596
commit
d6ffc554ec
28 changed files with 1839 additions and 486 deletions
1
litellm-rust/Cargo.lock
generated
1
litellm-rust/Cargo.lock
generated
|
|
@ -3065,6 +3065,7 @@ dependencies = [
|
|||
"serde_json",
|
||||
"serde_with",
|
||||
"sha2 0.10.9",
|
||||
"thiserror 2.0.19",
|
||||
"tokio",
|
||||
"tokio-tungstenite",
|
||||
"url",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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());
|
||||
}
|
||||
});
|
||||
}
|
||||
|
||||
|
|
|
|||
27
litellm-rust/crates/python-bridge/src/secrets/error.rs
Normal file
27
litellm-rust/crates/python-bridge/src/secrets/error.rs
Normal 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()))
|
||||
}
|
||||
|
|
@ -1,3 +1,6 @@
|
|||
pub(crate) mod callback;
|
||||
pub(crate) mod config;
|
||||
mod error;
|
||||
pub(crate) mod resolved;
|
||||
|
||||
pub(crate) use error::python_error;
|
||||
|
|
|
|||
|
|
@ -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")]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
));
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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(_))
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
|
|
|
|||
|
|
@ -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) {
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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()));
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
);
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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());
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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> {
|
||||
|
|
|
|||
|
|
@ -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() {
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
}
|
||||
|
|
|
|||
65
litellm-rust/crates/secrets-types/src/context.rs
Normal file
65
litellm-rust/crates/secrets-types/src/context.rs
Normal 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,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -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));
|
||||
}
|
||||
|
|
|
|||
110
litellm-rust/crates/secrets-types/tests/context.rs
Normal file
110
litellm-rust/crates/secrets-types/tests/context.rs
Normal 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);
|
||||
}
|
||||
|
|
@ -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) {
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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"))
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue