mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-25 01:02:15 +00:00
feat(rust): add typed secret managers and shared auth adapters
This commit is contained in:
parent
3fd2dd635b
commit
82bc67b122
41 changed files with 4459 additions and 34 deletions
7
.github/workflows/test-rust.yml
vendored
7
.github/workflows/test-rust.yml
vendored
|
|
@ -127,6 +127,13 @@ jobs:
|
|||
cargo check -p litellm-python-bridge --locked --no-default-features --features "abi3${features:+,$features}"
|
||||
done
|
||||
|
||||
- name: Test secret manager feature combinations
|
||||
run: |
|
||||
cargo test -p litellm-auth-gcp --locked --no-default-features
|
||||
for features in '' aws google aws,google; do
|
||||
cargo test -p litellm-secrets --locked --no-default-features --features "$features"
|
||||
done
|
||||
|
||||
rust-wheel:
|
||||
runs-on: ubuntu-latest
|
||||
timeout-minutes: 30
|
||||
|
|
|
|||
1037
litellm-rust/Cargo.lock
generated
1037
litellm-rust/Cargo.lock
generated
File diff suppressed because it is too large
Load diff
|
|
@ -18,6 +18,10 @@ litellm-auth-types = { path = "crates/auth-types" }
|
|||
litellm-auth-aws = { path = "crates/auth-aws" }
|
||||
litellm-auth-azure = { path = "crates/auth-azure" }
|
||||
litellm-auth-gcp = { path = "crates/auth-gcp" }
|
||||
litellm-secrets = { path = "crates/secrets" }
|
||||
litellm-secrets-types = { path = "crates/secrets-types" }
|
||||
litellm-secrets-aws = { path = "crates/secrets-aws" }
|
||||
litellm-secrets-google = { path = "crates/secrets-google" }
|
||||
litellm-http = { path = "crates/http" }
|
||||
litellm-llms = { path = "crates/llms" }
|
||||
litellm-types = { path = "crates/types" }
|
||||
|
|
@ -32,6 +36,8 @@ litellm-host-python = { path = "crates/host-python" }
|
|||
|
||||
bytes = "1"
|
||||
http = "1"
|
||||
google-cloud-auth = { version = "1.16.0", default-features = false }
|
||||
jsonwebtoken = { version = "11.1.0", default-features = false }
|
||||
hyper-util = { version = "0.1.20", default-features = false, features = ["client-proxy"] }
|
||||
proptest = "1.7.0"
|
||||
pyo3 = "0.29.2"
|
||||
|
|
|
|||
|
|
@ -3,6 +3,8 @@ pub const AWS_SECRET_ACCESS_KEY: &str = "AWS_SECRET_ACCESS_KEY";
|
|||
pub const AWS_SESSION_TOKEN: &str = "AWS_SESSION_TOKEN";
|
||||
pub const AWS_REGION_NAME: &str = "AWS_REGION_NAME";
|
||||
pub const AWS_REGION: &str = "AWS_REGION";
|
||||
pub const AWS_DEFAULT_REGION: &str = "AWS_DEFAULT_REGION";
|
||||
pub const AWS_BEDROCK_RUNTIME_ENDPOINT: &str = "AWS_BEDROCK_RUNTIME_ENDPOINT";
|
||||
pub const AWS_SESSION_NAME: &str = "AWS_SESSION_NAME";
|
||||
pub const AWS_PROFILE_NAME: &str = "AWS_PROFILE_NAME";
|
||||
pub const AWS_ROLE_NAME: &str = "AWS_ROLE_NAME";
|
||||
|
|
|
|||
|
|
@ -5,6 +5,9 @@ edition.workspace = true
|
|||
license.workspace = true
|
||||
repository.workspace = true
|
||||
|
||||
[features]
|
||||
google-sdk = ["dep:google-cloud-auth", "dep:http"]
|
||||
|
||||
[dependencies]
|
||||
litellm-auth-types.workspace = true
|
||||
|
||||
|
|
@ -14,3 +17,5 @@ sha2.workspace = true
|
|||
tokio.workspace = true
|
||||
|
||||
gcp_auth = "0.12.7"
|
||||
google-cloud-auth = { workspace = true, optional = true }
|
||||
http = { workspace = true, optional = true }
|
||||
|
|
|
|||
|
|
@ -8,6 +8,11 @@ use moka::future::Cache;
|
|||
use serde_json::{Map, Value};
|
||||
use sha2::{Digest, Sha256};
|
||||
|
||||
#[cfg(feature = "google-sdk")]
|
||||
mod sdk;
|
||||
#[cfg(feature = "google-sdk")]
|
||||
pub use sdk::GoogleCredentials;
|
||||
|
||||
const CLOUD_PLATFORM_SCOPE: &str = "https://www.googleapis.com/auth/cloud-platform";
|
||||
const GOOGLE_OAUTH_TOKEN_ENDPOINT: &str = "https://oauth2.googleapis.com/token";
|
||||
const GOOGLE_APPLICATION_CREDENTIALS_ENV: &str = "GOOGLE_APPLICATION_CREDENTIALS";
|
||||
|
|
@ -26,19 +31,31 @@ pub struct VertexConfig {
|
|||
}
|
||||
|
||||
impl VertexConfig {
|
||||
pub fn new(
|
||||
credentials: Option<Sourced<SecretValue>>,
|
||||
project_id: Option<String>,
|
||||
location: Option<String>,
|
||||
) -> Self {
|
||||
Self {
|
||||
credentials: credentials.filter(|value| !value.value().expose().trim().is_empty()),
|
||||
project_id: project_id.filter(|value| !value.trim().is_empty()),
|
||||
location: location.filter(|value| !value.trim().is_empty()),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn from_sourced_optional_params(
|
||||
params: &Map<String, Value>,
|
||||
sources: &BTreeMap<String, InputSource>,
|
||||
) -> Result<Self, Error> {
|
||||
Ok(Self {
|
||||
credentials: optional_credentials(
|
||||
Ok(Self::new(
|
||||
optional_credentials(
|
||||
params,
|
||||
sources,
|
||||
&["vertex_credentials", "vertex_ai_credentials"],
|
||||
)?,
|
||||
project_id: optional_string(params, &["vertex_project", "vertex_ai_project"])?,
|
||||
location: optional_string(params, &["vertex_location", "vertex_ai_location"])?,
|
||||
})
|
||||
optional_string(params, &["vertex_project", "vertex_ai_project"])?,
|
||||
optional_string(params, &["vertex_location", "vertex_ai_location"])?,
|
||||
))
|
||||
}
|
||||
|
||||
pub fn or_configured(self, project_id: Option<&str>, location: Option<&str>) -> Self {
|
||||
|
|
@ -469,6 +486,39 @@ mod tests {
|
|||
assert_eq!(config.location(), Some("alias-location"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn typed_config_preserves_source_and_empty_value_fallback() {
|
||||
let configured = VertexConfig::new(
|
||||
Some(Sourced::new(
|
||||
SecretValue::new("inline-json"),
|
||||
InputSource::Request,
|
||||
)),
|
||||
Some("project".into()),
|
||||
Some("location".into()),
|
||||
);
|
||||
assert!(matches!(
|
||||
credential_source(&configured, &|_| Some("environment-json".into())),
|
||||
CredentialSource::Inline(value) if value.expose() == "inline-json"
|
||||
));
|
||||
let empty = VertexConfig::new(
|
||||
Some(Sourced::new(SecretValue::new(" "), InputSource::Request)),
|
||||
Some(" ".into()),
|
||||
Some(" ".into()),
|
||||
);
|
||||
assert!(matches!(
|
||||
credential_source(&empty, &|_| None),
|
||||
CredentialSource::Adc
|
||||
));
|
||||
assert_eq!(
|
||||
get_vertex_ai_project(&empty, &|_| Some("env-project".into())).as_deref(),
|
||||
Some("env-project")
|
||||
);
|
||||
assert_eq!(
|
||||
get_vertex_ai_location(&empty, &|_| Some("env-location".into())).as_deref(),
|
||||
Some("env-location")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn project_and_location_prefer_input_then_environment() {
|
||||
let configured =
|
||||
|
|
|
|||
106
litellm-rust/crates/auth-gcp/src/sdk.rs
Normal file
106
litellm-rust/crates/auth-gcp/src/sdk.rs
Normal file
|
|
@ -0,0 +1,106 @@
|
|||
use std::sync::Arc;
|
||||
|
||||
use google_cloud_auth::credentials::{CacheableResource, CredentialsProvider, EntityTag};
|
||||
use google_cloud_auth::errors::CredentialsError;
|
||||
use http::{Extensions, HeaderMap, HeaderName, HeaderValue};
|
||||
use litellm_auth_types::Error;
|
||||
|
||||
use crate::{VertexAuth, VertexConfig};
|
||||
|
||||
type EnvironmentLookup = dyn Fn(&str) -> Option<String> + Send + Sync;
|
||||
|
||||
pub struct GoogleCredentials {
|
||||
auth: VertexAuth,
|
||||
config: VertexConfig,
|
||||
environment: Arc<EnvironmentLookup>,
|
||||
}
|
||||
|
||||
impl GoogleCredentials {
|
||||
pub fn new(config: VertexConfig, environment: Arc<EnvironmentLookup>) -> Self {
|
||||
Self {
|
||||
auth: VertexAuth::default(),
|
||||
config,
|
||||
environment,
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn request_headers(&self) -> Result<HeaderMap, Error> {
|
||||
let response = self
|
||||
.auth
|
||||
.validate_environment(Vec::new(), None, &self.config, &|name| {
|
||||
(self.environment)(name)
|
||||
})
|
||||
.await?;
|
||||
response
|
||||
.headers
|
||||
.into_iter()
|
||||
.map(|(key, value)| {
|
||||
let name =
|
||||
HeaderName::from_bytes(key.as_bytes()).map_err(|_| Error::InvalidHeader)?;
|
||||
let value = HeaderValue::from_str(&value).map_err(|_| Error::InvalidHeader)?;
|
||||
Ok((name, value))
|
||||
})
|
||||
.collect()
|
||||
}
|
||||
}
|
||||
|
||||
impl CredentialsProvider for GoogleCredentials {
|
||||
async fn headers(
|
||||
&self,
|
||||
_: Extensions,
|
||||
) -> Result<CacheableResource<HeaderMap>, CredentialsError> {
|
||||
self.request_headers()
|
||||
.await
|
||||
.map(|data| CacheableResource::New {
|
||||
entity_tag: EntityTag::new(),
|
||||
data,
|
||||
})
|
||||
.map_err(|_| CredentialsError::from_msg(false, "Google authentication failed"))
|
||||
}
|
||||
|
||||
async fn universe_domain(&self) -> Option<String> {
|
||||
None
|
||||
}
|
||||
}
|
||||
|
||||
impl std::fmt::Debug for GoogleCredentials {
|
||||
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
f.debug_struct("GoogleCredentials").finish_non_exhaustive()
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[tokio::test]
|
||||
async fn sdk_and_http_credentials_share_token_resolution_and_redaction() {
|
||||
let credentials = GoogleCredentials::new(
|
||||
VertexConfig::new(None, Some("project".into()), None),
|
||||
Arc::new(|name| (name == "VERTEX_AI_API_KEY").then(|| "private-token".into())),
|
||||
);
|
||||
let direct = credentials.request_headers().await.unwrap();
|
||||
let CacheableResource::New { data, .. } =
|
||||
credentials.headers(Extensions::new()).await.unwrap()
|
||||
else {
|
||||
panic!("first request did not return headers");
|
||||
};
|
||||
assert_eq!(direct, data);
|
||||
assert_eq!(data[http::header::AUTHORIZATION], "Bearer private-token");
|
||||
assert!(!format!("{credentials:?}").contains("private-token"));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn invalid_token_headers_return_a_redacted_sdk_error() {
|
||||
let credentials = GoogleCredentials::new(
|
||||
VertexConfig::new(None, Some("project".into()), None),
|
||||
Arc::new(|name| (name == "VERTEX_AI_API_KEY").then(|| "private\nvalue".into())),
|
||||
);
|
||||
assert_eq!(
|
||||
credentials.request_headers().await.unwrap_err(),
|
||||
Error::InvalidHeader
|
||||
);
|
||||
let error = credentials.headers(Extensions::new()).await.unwrap_err();
|
||||
assert!(!format!("{error:?}").contains("private"));
|
||||
}
|
||||
}
|
||||
24
litellm-rust/crates/secrets-aws/Cargo.toml
Normal file
24
litellm-rust/crates/secrets-aws/Cargo.toml
Normal file
|
|
@ -0,0 +1,24 @@
|
|||
[package]
|
||||
name = "litellm-secrets-aws"
|
||||
version = "0.1.0"
|
||||
edition.workspace = true
|
||||
license.workspace = true
|
||||
repository.workspace = true
|
||||
|
||||
[dependencies]
|
||||
litellm-auth-aws.workspace = true
|
||||
litellm-secrets-types.workspace = true
|
||||
litellm-core-utils.workspace = true
|
||||
serde_json.workspace = true
|
||||
thiserror.workspace = true
|
||||
tracing = "0.1"
|
||||
veil.workspace = true
|
||||
aws-sdk-kms = "1.120.0"
|
||||
aws-sdk-secretsmanager = "1.117.0"
|
||||
aws-credential-types = "1.3.0"
|
||||
|
||||
[dev-dependencies]
|
||||
base64.workspace = true
|
||||
rstest.workspace = true
|
||||
tokio.workspace = true
|
||||
wiremock = "0.6.5"
|
||||
79
litellm-rust/crates/secrets-aws/src/auth.rs
Normal file
79
litellm-rust/crates/secrets-aws/src/auth.rs
Normal file
|
|
@ -0,0 +1,79 @@
|
|||
use std::sync::Arc;
|
||||
|
||||
use aws_credential_types::provider::{ProvideCredentials, error::CredentialsError, future};
|
||||
use litellm_auth_aws::{
|
||||
AwsAuthConfig,
|
||||
constants::{AWS_DEFAULT_REGION, AWS_REGION, AWS_REGION_NAME},
|
||||
resolve_credentials,
|
||||
};
|
||||
use litellm_core_utils::settings::Lookup;
|
||||
use litellm_secrets_types::KeyManagementSettings;
|
||||
|
||||
use crate::Error;
|
||||
|
||||
#[derive(Clone)]
|
||||
pub(crate) struct Credentials {
|
||||
config: AwsAuthConfig,
|
||||
environment: Arc<dyn Lookup + Send + Sync>,
|
||||
}
|
||||
|
||||
impl Credentials {
|
||||
pub(crate) fn new(
|
||||
settings: &KeyManagementSettings,
|
||||
environment: Arc<dyn Lookup + Send + Sync>,
|
||||
) -> Self {
|
||||
Self {
|
||||
config: AwsAuthConfig {
|
||||
region_name: region(settings, environment.as_ref()).ok(),
|
||||
role_name: settings.aws_role_name.clone(),
|
||||
session_name: settings.aws_session_name.clone(),
|
||||
external_id: settings
|
||||
.aws_external_id
|
||||
.as_ref()
|
||||
.map(|v| v.expose().to_owned()),
|
||||
profile_name: settings.aws_profile_name.clone(),
|
||||
web_identity_token: settings
|
||||
.aws_web_identity_token
|
||||
.as_ref()
|
||||
.map(|v| v.expose().to_owned()),
|
||||
sts_endpoint: settings.aws_sts_endpoint.clone(),
|
||||
..Default::default()
|
||||
},
|
||||
environment,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl ProvideCredentials for Credentials {
|
||||
fn provide_credentials<'a>(&'a self) -> future::ProvideCredentials<'a>
|
||||
where
|
||||
Self: 'a,
|
||||
{
|
||||
future::ProvideCredentials::new(async {
|
||||
resolve_credentials(self.config.clone(), &|name| self.environment.get(name))
|
||||
.await
|
||||
.map_err(|_| {
|
||||
CredentialsError::provider_error("secret manager authentication failed")
|
||||
})
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn region(
|
||||
settings: &KeyManagementSettings,
|
||||
environment: &dyn Lookup,
|
||||
) -> Result<String, Error> {
|
||||
settings
|
||||
.aws_region_name
|
||||
.clone()
|
||||
.or_else(|| environment.get(AWS_REGION_NAME))
|
||||
.or_else(|| environment.get(AWS_REGION))
|
||||
.or_else(|| environment.get(AWS_DEFAULT_REGION))
|
||||
.ok_or(Error::MissingRegion)
|
||||
}
|
||||
|
||||
impl std::fmt::Debug for Credentials {
|
||||
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
|
||||
f.debug_struct("Credentials").finish_non_exhaustive()
|
||||
}
|
||||
}
|
||||
29
litellm-rust/crates/secrets-aws/src/error.rs
Normal file
29
litellm-rust/crates/secrets-aws/src/error.rs
Normal file
|
|
@ -0,0 +1,29 @@
|
|||
use aws_sdk_secretsmanager::error::SdkError;
|
||||
|
||||
#[derive(thiserror::Error, veil::Redact)]
|
||||
pub enum Error {
|
||||
#[error("AWS authentication failed")]
|
||||
Auth(#[from] #[redact] litellm_auth_aws::Error),
|
||||
#[error("AWS region is not configured")]
|
||||
MissingRegion,
|
||||
#[error("KMS response has no plaintext")]
|
||||
MissingPlaintext,
|
||||
#[error("AWS request timed out")]
|
||||
Timeout,
|
||||
#[error("AWS KMS decrypt failed")]
|
||||
Decrypt(#[from] #[redact] Box<SdkError<aws_sdk_kms::operation::decrypt::DecryptError>>),
|
||||
#[error("AWS Secrets Manager request preparation failed")]
|
||||
Read(#[from] #[redact] Box<SdkError<aws_sdk_secretsmanager::operation::get_secret_value::GetSecretValueError>>),
|
||||
#[error("AWS Secrets Manager create failed")]
|
||||
Create(#[from] #[redact] Box<SdkError<aws_sdk_secretsmanager::operation::create_secret::CreateSecretError>>),
|
||||
#[error("AWS Secrets Manager update failed")]
|
||||
Put(#[from] #[redact] Box<SdkError<aws_sdk_secretsmanager::operation::put_secret_value::PutSecretValueError>>),
|
||||
#[error("AWS Secrets Manager delete failed")]
|
||||
Delete(#[from] #[redact] Box<SdkError<aws_sdk_secretsmanager::operation::delete_secret::DeleteSecretError>>),
|
||||
#[error("AWS Secrets Manager replication failed")]
|
||||
Replicate(#[from] #[redact] Box<SdkError<aws_sdk_secretsmanager::operation::replicate_secret_to_regions::ReplicateSecretToRegionsError>>),
|
||||
#[error("primary secret is not a JSON object")]
|
||||
PrimarySecret,
|
||||
#[error(transparent)]
|
||||
Operation(#[from] litellm_secrets_types::Error),
|
||||
}
|
||||
63
litellm-rust/crates/secrets-aws/src/kms.rs
Normal file
63
litellm-rust/crates/secrets-aws/src/kms.rs
Normal file
|
|
@ -0,0 +1,63 @@
|
|||
use litellm_auth_aws::constants::AWS_REGION_NAME;
|
||||
use std::sync::Arc;
|
||||
|
||||
use aws_sdk_kms::{
|
||||
Client,
|
||||
config::{BehaviorVersion, Region},
|
||||
primitives::Blob,
|
||||
};
|
||||
use litellm_core_utils::settings::Lookup;
|
||||
use litellm_secrets_types::KeyManagementSettings;
|
||||
|
||||
use crate::{Error, auth};
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct AwsKms {
|
||||
client: Client,
|
||||
}
|
||||
|
||||
impl AwsKms {
|
||||
pub fn new(client: Client) -> Self {
|
||||
Self { client }
|
||||
}
|
||||
|
||||
pub async fn decrypt(&self, ciphertext: Vec<u8>) -> Result<Vec<u8>, Error> {
|
||||
let response = self
|
||||
.client
|
||||
.decrypt()
|
||||
.ciphertext_blob(Blob::new(ciphertext))
|
||||
.send()
|
||||
.await
|
||||
.map_err(|error| Error::Decrypt(Box::new(error)))?;
|
||||
Ok(response
|
||||
.plaintext
|
||||
.ok_or(Error::MissingPlaintext)?
|
||||
.into_inner())
|
||||
}
|
||||
}
|
||||
|
||||
pub fn validate_environment(environment: &dyn Lookup) -> Result<(), Error> {
|
||||
environment
|
||||
.get(AWS_REGION_NAME)
|
||||
.map(|_| ())
|
||||
.ok_or(Error::MissingRegion)
|
||||
}
|
||||
|
||||
pub fn load_aws_kms(
|
||||
use_aws_kms: Option<bool>,
|
||||
settings: &KeyManagementSettings,
|
||||
environment: Arc<dyn Lookup + Send + Sync>,
|
||||
) -> Result<Option<AwsKms>, Error> {
|
||||
if use_aws_kms != Some(true) {
|
||||
return Ok(None);
|
||||
}
|
||||
if settings.aws_region_name.is_none() {
|
||||
validate_environment(environment.as_ref())?;
|
||||
}
|
||||
let config = aws_sdk_kms::Config::builder()
|
||||
.behavior_version(BehaviorVersion::latest())
|
||||
.region(Region::new(auth::region(settings, environment.as_ref())?))
|
||||
.credentials_provider(auth::Credentials::new(settings, environment))
|
||||
.build();
|
||||
Ok(Some(AwsKms::new(Client::from_conf(config))))
|
||||
}
|
||||
10
litellm-rust/crates/secrets-aws/src/lib.rs
Normal file
10
litellm-rust/crates/secrets-aws/src/lib.rs
Normal file
|
|
@ -0,0 +1,10 @@
|
|||
#![forbid(unsafe_code)]
|
||||
|
||||
mod auth;
|
||||
mod error;
|
||||
pub mod kms;
|
||||
pub mod secret_manager;
|
||||
|
||||
pub use error::Error;
|
||||
pub use kms::{AwsKms, load_aws_kms};
|
||||
pub use secret_manager::{AwsSecretWriteSettings, AwsSecretsManagerV2, RotationResponse};
|
||||
297
litellm-rust/crates/secrets-aws/src/secret_manager.rs
Normal file
297
litellm-rust/crates/secrets-aws/src/secret_manager.rs
Normal file
|
|
@ -0,0 +1,297 @@
|
|||
use litellm_auth_aws::constants::AWS_BEDROCK_RUNTIME_ENDPOINT;
|
||||
use std::{collections::BTreeMap, sync::Arc};
|
||||
|
||||
use aws_sdk_secretsmanager::{
|
||||
Client,
|
||||
config::{BehaviorVersion, Region},
|
||||
operation::{
|
||||
create_secret::CreateSecretOutput, delete_secret::DeleteSecretOutput,
|
||||
put_secret_value::PutSecretValueOutput,
|
||||
replicate_secret_to_regions::ReplicateSecretToRegionsOutput,
|
||||
},
|
||||
types::{ReplicaRegionType, Tag},
|
||||
};
|
||||
use litellm_auth_aws::constants::{
|
||||
AWS_ACCESS_KEY_ID, AWS_REGION, AWS_REGION_NAME, AWS_SECRET_ACCESS_KEY,
|
||||
};
|
||||
use litellm_core_utils::settings::Lookup;
|
||||
use litellm_secrets_types::{
|
||||
BaseSecretManager, KeyManagementSettings, Secret, SecretValue, async_rotate_secret,
|
||||
};
|
||||
use serde_json::Value;
|
||||
|
||||
use crate::{Error, auth};
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct AwsSecretsManagerV2 {
|
||||
client: Client,
|
||||
write_settings: AwsSecretWriteSettings,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Default)]
|
||||
pub struct AwsSecretWriteSettings {
|
||||
pub kms_key_id: Option<String>,
|
||||
pub tags: Option<BTreeMap<String, String>>,
|
||||
pub replica_regions: Option<Vec<String>>,
|
||||
}
|
||||
|
||||
impl From<&KeyManagementSettings> for AwsSecretWriteSettings {
|
||||
fn from(settings: &KeyManagementSettings) -> Self {
|
||||
Self {
|
||||
kms_key_id: settings.kms_key_id.clone(),
|
||||
tags: settings.tags.clone(),
|
||||
replica_regions: settings.replica_regions.clone(),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
pub enum RotationResponse {
|
||||
Created(CreateSecretOutput),
|
||||
Updated(PutSecretValueOutput),
|
||||
}
|
||||
|
||||
impl AwsSecretsManagerV2 {
|
||||
pub fn new(client: Client, write_settings: AwsSecretWriteSettings) -> Self {
|
||||
Self {
|
||||
client,
|
||||
write_settings,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn load_aws_secret_manager(
|
||||
use_aws_secret_manager: Option<bool>,
|
||||
settings: KeyManagementSettings,
|
||||
environment: Arc<dyn Lookup + Send + Sync>,
|
||||
) -> Result<Option<Self>, Error> {
|
||||
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(),
|
||||
};
|
||||
Ok(Some(Self::new(
|
||||
Client::from_conf(config),
|
||||
(&settings).into(),
|
||||
)))
|
||||
}
|
||||
|
||||
pub async fn read_secret_for_resolver(
|
||||
&self,
|
||||
name: &str,
|
||||
primary_name: Option<&str>,
|
||||
environment: &(dyn Lookup + Sync),
|
||||
) -> Result<Option<Secret>, Error> {
|
||||
if bootstrap_key(name) {
|
||||
return Ok(environment
|
||||
.get(name)
|
||||
.map(SecretValue::new)
|
||||
.map(Secret::String));
|
||||
}
|
||||
match primary_name.filter(|name| !name.is_empty()) {
|
||||
None => self
|
||||
.async_read_secret(name)
|
||||
.await
|
||||
.map(|value| value.map(Secret::String)),
|
||||
Some(primary) => {
|
||||
let value = if bootstrap_key(primary) {
|
||||
environment.get(primary).map(SecretValue::new)
|
||||
} else {
|
||||
self.async_read_secret(primary).await?
|
||||
};
|
||||
let object: Value = serde_json::from_str(
|
||||
value
|
||||
.as_ref()
|
||||
.map(SecretValue::expose)
|
||||
.filter(|v| !v.is_empty())
|
||||
.unwrap_or("{}"),
|
||||
)
|
||||
.map_err(|_| Error::PrimarySecret)?;
|
||||
let object = object.as_object().ok_or(Error::PrimarySecret)?;
|
||||
Ok(object.get(name).cloned().and_then(Secret::from_json))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn async_read_secret(&self, name: &str) -> Result<Option<SecretValue>, Error> {
|
||||
match self.client.get_secret_value().secret_id(name).send().await {
|
||||
Ok(response) => Ok(response.secret_string.map(SecretValue::new)),
|
||||
Err(error)
|
||||
if matches!(
|
||||
&error,
|
||||
aws_sdk_secretsmanager::error::SdkError::TimeoutError(_)
|
||||
) || matches!(&error, aws_sdk_secretsmanager::error::SdkError::DispatchFailure(failure) if failure.is_timeout()) =>
|
||||
{
|
||||
Err(Error::Timeout)
|
||||
}
|
||||
Err(error) if request_preparation_failed(&error) => Err(Error::Read(Box::new(error))),
|
||||
Err(_) => {
|
||||
tracing::error!("AWS secret read failed");
|
||||
Ok(None)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn async_write_secret(
|
||||
&self,
|
||||
name: &str,
|
||||
value: &SecretValue,
|
||||
description: Option<&str>,
|
||||
) -> Result<CreateSecretOutput, Error> {
|
||||
let response = self
|
||||
.client
|
||||
.create_secret()
|
||||
.name(name)
|
||||
.secret_string(value.expose())
|
||||
.set_description(description.filter(|v| !v.is_empty()).map(str::to_owned))
|
||||
.set_kms_key_id(
|
||||
self.write_settings
|
||||
.kms_key_id
|
||||
.clone()
|
||||
.filter(|v| !v.is_empty()),
|
||||
)
|
||||
.set_tags(self.write_settings.tags.as_ref().map(|tags| {
|
||||
tags.iter()
|
||||
.map(|(key, value)| Tag::builder().key(key).value(value).build())
|
||||
.collect()
|
||||
}))
|
||||
.send()
|
||||
.await
|
||||
.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()
|
||||
{
|
||||
tracing::warn!("secret created but replication failed");
|
||||
}
|
||||
Ok(response)
|
||||
}
|
||||
|
||||
pub async fn async_replicate_secret(
|
||||
&self,
|
||||
name: &str,
|
||||
regions: &[String],
|
||||
) -> Result<Option<ReplicateSecretToRegionsOutput>, Error> {
|
||||
if regions.is_empty() {
|
||||
return Ok(None);
|
||||
}
|
||||
self.client
|
||||
.replicate_secret_to_regions()
|
||||
.secret_id(name)
|
||||
.set_add_replica_regions(Some(
|
||||
regions
|
||||
.iter()
|
||||
.map(|region| ReplicaRegionType::builder().region(region).build())
|
||||
.collect(),
|
||||
))
|
||||
.send()
|
||||
.await
|
||||
.map(Some)
|
||||
.map_err(|error| Error::Replicate(Box::new(error)))
|
||||
}
|
||||
|
||||
pub async fn async_put_secret_value(
|
||||
&self,
|
||||
name: &str,
|
||||
value: &SecretValue,
|
||||
) -> Result<PutSecretValueOutput, Error> {
|
||||
self.client
|
||||
.put_secret_value()
|
||||
.secret_id(name)
|
||||
.secret_string(value.expose())
|
||||
.send()
|
||||
.await
|
||||
.map_err(|error| Error::Put(Box::new(error)))
|
||||
}
|
||||
|
||||
pub async fn async_delete_secret(
|
||||
&self,
|
||||
name: &str,
|
||||
recovery_window_in_days: i64,
|
||||
) -> Result<DeleteSecretOutput, Error> {
|
||||
self.client
|
||||
.delete_secret()
|
||||
.secret_id(name)
|
||||
.recovery_window_in_days(recovery_window_in_days)
|
||||
.send()
|
||||
.await
|
||||
.map_err(|error| Error::Delete(Box::new(error)))
|
||||
}
|
||||
|
||||
pub async fn async_rotate_secret(
|
||||
&self,
|
||||
current_name: &str,
|
||||
new_name: &str,
|
||||
value: &SecretValue,
|
||||
) -> Result<RotationResponse, Error> {
|
||||
if current_name == new_name {
|
||||
return self
|
||||
.async_put_secret_value(current_name, value)
|
||||
.await
|
||||
.map(RotationResponse::Updated);
|
||||
}
|
||||
async_rotate_secret(self, current_name, new_name, value)
|
||||
.await
|
||||
.map(RotationResponse::Created)
|
||||
}
|
||||
}
|
||||
|
||||
impl BaseSecretManager for AwsSecretsManagerV2 {
|
||||
type Error = Error;
|
||||
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_write_secret(
|
||||
&self,
|
||||
name: &str,
|
||||
value: &SecretValue,
|
||||
description: Option<&str>,
|
||||
) -> Result<CreateSecretOutput, Error> {
|
||||
self.async_write_secret(name, value, description).await
|
||||
}
|
||||
|
||||
async fn async_delete_secret(
|
||||
&self,
|
||||
name: &str,
|
||||
recovery_window_in_days: i64,
|
||||
) -> Result<DeleteSecretOutput, Error> {
|
||||
self.async_delete_secret(name, recovery_window_in_days)
|
||||
.await
|
||||
}
|
||||
}
|
||||
|
||||
fn bootstrap_key(name: &str) -> bool {
|
||||
matches!(
|
||||
name,
|
||||
AWS_ACCESS_KEY_ID
|
||||
| AWS_SECRET_ACCESS_KEY
|
||||
| AWS_REGION_NAME
|
||||
| AWS_REGION
|
||||
| AWS_BEDROCK_RUNTIME_ENDPOINT
|
||||
)
|
||||
}
|
||||
|
||||
fn request_preparation_failed(
|
||||
error: &aws_sdk_secretsmanager::error::SdkError<
|
||||
aws_sdk_secretsmanager::operation::get_secret_value::GetSecretValueError,
|
||||
>,
|
||||
) -> bool {
|
||||
matches!(
|
||||
error,
|
||||
aws_sdk_secretsmanager::error::SdkError::ConstructionFailure(_)
|
||||
) || std::iter::successors(Some(error as &(dyn std::error::Error + 'static)), |error| {
|
||||
error.source()
|
||||
})
|
||||
.any(|source| source.is::<aws_credential_types::provider::error::CredentialsError>())
|
||||
}
|
||||
59
litellm-rust/crates/secrets-aws/tests/kms.rs
Normal file
59
litellm-rust/crates/secrets-aws/tests/kms.rs
Normal file
|
|
@ -0,0 +1,59 @@
|
|||
use aws_sdk_kms::{
|
||||
Client,
|
||||
config::{BehaviorVersion, Credentials, Region, retry::RetryConfig},
|
||||
};
|
||||
use base64::{Engine, engine::general_purpose::STANDARD};
|
||||
use litellm_secrets_aws::AwsKms;
|
||||
use wiremock::{
|
||||
Mock, MockServer, ResponseTemplate,
|
||||
matchers::{body_json, header},
|
||||
};
|
||||
|
||||
#[tokio::test]
|
||||
async fn kms_decrypt_calls_the_sdk_without_applying_lookup_policy() {
|
||||
let server = MockServer::start().await;
|
||||
let plaintext = " private-value\n";
|
||||
Mock::given(header("x-amz-target", "TrentService.Decrypt"))
|
||||
.and(body_json(
|
||||
serde_json::json!({"CiphertextBlob": STANDARD.encode("encrypted")}),
|
||||
))
|
||||
.respond_with(
|
||||
ResponseTemplate::new(200)
|
||||
.set_body_json(serde_json::json!({"Plaintext": STANDARD.encode(plaintext)})),
|
||||
)
|
||||
.expect(1)
|
||||
.mount(&server)
|
||||
.await;
|
||||
let client = Client::from_conf(
|
||||
aws_sdk_kms::Config::builder()
|
||||
.behavior_version(BehaviorVersion::latest())
|
||||
.region(Region::new("us-east-1"))
|
||||
.credentials_provider(Credentials::new("test", "test", None, None, "test"))
|
||||
.endpoint_url(server.uri())
|
||||
.retry_config(RetryConfig::disabled())
|
||||
.build(),
|
||||
);
|
||||
let manager = AwsKms::new(client);
|
||||
assert_eq!(
|
||||
manager.decrypt(b"encrypted".to_vec()).await.unwrap(),
|
||||
plaintext.as_bytes()
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn disabled_kms_loader_does_not_require_environment_configuration() {
|
||||
use litellm_secrets_aws::load_aws_kms;
|
||||
use litellm_secrets_types::KeyManagementSettings;
|
||||
use std::sync::Arc;
|
||||
for enabled in [None, Some(false)] {
|
||||
assert!(
|
||||
load_aws_kms(
|
||||
enabled,
|
||||
&KeyManagementSettings::default(),
|
||||
Arc::new(|_: &str| None)
|
||||
)
|
||||
.unwrap()
|
||||
.is_none()
|
||||
);
|
||||
}
|
||||
}
|
||||
292
litellm-rust/crates/secrets-aws/tests/secret_manager.rs
Normal file
292
litellm-rust/crates/secrets-aws/tests/secret_manager.rs
Normal file
|
|
@ -0,0 +1,292 @@
|
|||
use std::sync::atomic::{AtomicUsize, Ordering};
|
||||
|
||||
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 serde_json::json;
|
||||
use wiremock::{
|
||||
Mock, MockServer, ResponseTemplate,
|
||||
matchers::{body_partial_json, header},
|
||||
};
|
||||
|
||||
fn manager(server: &MockServer, settings: KeyManagementSettings) -> AwsSecretsManagerV2 {
|
||||
let client = Client::from_conf(
|
||||
aws_sdk_secretsmanager::Config::builder()
|
||||
.behavior_version(BehaviorVersion::latest())
|
||||
.region(Region::new("us-east-1"))
|
||||
.credentials_provider(Credentials::new("test", "test", None, None, "test"))
|
||||
.endpoint_url(server.uri())
|
||||
.retry_config(RetryConfig::disabled())
|
||||
.build(),
|
||||
);
|
||||
AwsSecretsManagerV2::new(client, (&settings).into())
|
||||
}
|
||||
|
||||
#[rstest::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(
|
||||
#[case] name: &str,
|
||||
#[case] expected: Option<&str>,
|
||||
) {
|
||||
let server = MockServer::start().await;
|
||||
Mock::given(header("x-amz-target", "secretsmanager.GetSecretValue"))
|
||||
.and(body_partial_json(json!({"SecretId":"primary"})))
|
||||
.respond_with(
|
||||
ResponseTemplate::new(200).set_body_json(
|
||||
json!({"SecretString":json!({"KEY":"value", "BOOL":true}).to_string()}),
|
||||
),
|
||||
)
|
||||
.expect(1)
|
||||
.mount(&server)
|
||||
.await;
|
||||
let manager = manager(&server, KeyManagementSettings::default());
|
||||
assert_eq!(
|
||||
manager
|
||||
.read_secret_for_resolver(name, Some("primary"), &|_: &str| None)
|
||||
.await
|
||||
.unwrap()
|
||||
.and_then(|v| v.as_str().map(str::to_owned))
|
||||
.as_deref(),
|
||||
expected
|
||||
);
|
||||
}
|
||||
|
||||
#[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) {
|
||||
let server = MockServer::start().await;
|
||||
let manager = manager(&server, KeyManagementSettings::default());
|
||||
assert_eq!(
|
||||
manager
|
||||
.read_secret_for_resolver(name, Some("primary"), &|_: &str| Some("bootstrap".into()))
|
||||
.await
|
||||
.unwrap()
|
||||
.unwrap()
|
||||
.as_str()
|
||||
.unwrap(),
|
||||
"bootstrap"
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn failed_read_returns_none_but_invalid_primary_json_is_an_error() {
|
||||
let server = MockServer::start().await;
|
||||
Mock::given(body_partial_json(json!({"SecretId":"missing"})))
|
||||
.respond_with(
|
||||
ResponseTemplate::new(400).set_body_json(json!({"__type":"ResourceNotFoundException"})),
|
||||
)
|
||||
.mount(&server)
|
||||
.await;
|
||||
Mock::given(body_partial_json(json!({"SecretId":"invalid"})))
|
||||
.respond_with(ResponseTemplate::new(200).set_body_json(json!({"SecretString":"not-json"})))
|
||||
.mount(&server)
|
||||
.await;
|
||||
let manager = manager(&server, KeyManagementSettings::default());
|
||||
assert!(
|
||||
manager
|
||||
.async_read_secret("missing")
|
||||
.await
|
||||
.unwrap()
|
||||
.is_none()
|
||||
);
|
||||
assert!(matches!(
|
||||
manager
|
||||
.read_secret_for_resolver("KEY", Some("invalid"), &|_: &str| None)
|
||||
.await,
|
||||
Err(Error::PrimarySecret)
|
||||
));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn same_name_rotation_uses_put_and_returns_its_response() {
|
||||
let server = MockServer::start().await;
|
||||
Mock::given(header("x-amz-target", "secretsmanager.PutSecretValue"))
|
||||
.and(body_partial_json(
|
||||
json!({"SecretId":"key", "SecretString":"replacement"}),
|
||||
))
|
||||
.respond_with(
|
||||
ResponseTemplate::new(200).set_body_json(json!({"Name":"key", "VersionId":"version"})),
|
||||
)
|
||||
.expect(1)
|
||||
.mount(&server)
|
||||
.await;
|
||||
let response = manager(&server, KeyManagementSettings::default())
|
||||
.async_rotate_secret("key", "key", &SecretValue::new("replacement"))
|
||||
.await
|
||||
.unwrap();
|
||||
match response {
|
||||
RotationResponse::Updated(output) => assert_eq!(output.version_id(), Some("version")),
|
||||
_ => panic!("rotation created a second secret"),
|
||||
}
|
||||
assert_eq!(server.received_requests().await.unwrap().len(), 1);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn renamed_rotation_reads_creates_verifies_then_deletes() {
|
||||
let server = MockServer::start().await;
|
||||
let step = AtomicUsize::new(0);
|
||||
Mock::given(wiremock::matchers::method("POST"))
|
||||
.respond_with(move |request: &wiremock::Request| {
|
||||
let body: serde_json::Value = request.body_json().unwrap();
|
||||
let action = request
|
||||
.headers
|
||||
.get("x-amz-target")
|
||||
.unwrap()
|
||||
.to_str()
|
||||
.unwrap();
|
||||
match step.fetch_add(1, Ordering::SeqCst) {
|
||||
0 => {
|
||||
assert_eq!(action, "secretsmanager.GetSecretValue");
|
||||
assert_eq!(body["SecretId"], "old");
|
||||
ResponseTemplate::new(200).set_body_json(json!({"SecretString":"old-value"}))
|
||||
}
|
||||
1 => {
|
||||
assert_eq!(action, "secretsmanager.CreateSecret");
|
||||
assert_eq!(body["Name"], "new");
|
||||
assert_eq!(body["Description"], "Rotated from old");
|
||||
assert_eq!(body["SecretString"], "replacement");
|
||||
ResponseTemplate::new(200).set_body_json(json!({"Name":"new"}))
|
||||
}
|
||||
2 => {
|
||||
assert_eq!(action, "secretsmanager.GetSecretValue");
|
||||
assert_eq!(body["SecretId"], "new");
|
||||
ResponseTemplate::new(200).set_body_json(json!({"SecretString":"replacement"}))
|
||||
}
|
||||
3 => {
|
||||
assert_eq!(action, "secretsmanager.DeleteSecret");
|
||||
assert_eq!(body["SecretId"], "old");
|
||||
assert_eq!(body["RecoveryWindowInDays"], 7);
|
||||
ResponseTemplate::new(200).set_body_json(json!({"Name":"old"}))
|
||||
}
|
||||
_ => panic!("unexpected request"),
|
||||
}
|
||||
})
|
||||
.expect(4)
|
||||
.mount(&server)
|
||||
.await;
|
||||
assert!(matches!(
|
||||
manager(&server, KeyManagementSettings::default())
|
||||
.async_rotate_secret("old", "new", &SecretValue::new("replacement"))
|
||||
.await
|
||||
.unwrap(),
|
||||
RotationResponse::Created(_)
|
||||
));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn creation_passes_tags_and_kms_and_survives_replication_failure() {
|
||||
let server = MockServer::start().await;
|
||||
Mock::given(header("x-amz-target", "secretsmanager.CreateSecret"))
|
||||
.and(body_partial_json(json!({"Name":"key", "SecretString":"value", "KmsKeyId":"kms-key", "Tags":[{"Key":"stage", "Value":"test"}]})))
|
||||
.respond_with(ResponseTemplate::new(200).set_body_json(json!({"Name":"key"}))).expect(1).mount(&server).await;
|
||||
Mock::given(header(
|
||||
"x-amz-target",
|
||||
"secretsmanager.ReplicateSecretToRegions",
|
||||
))
|
||||
.and(body_partial_json(
|
||||
json!({"SecretId":"key", "AddReplicaRegions":[{"Region":"replica-region"}]}),
|
||||
))
|
||||
.respond_with(
|
||||
ResponseTemplate::new(400).set_body_json(json!({"__type":"InvalidRequestException"})),
|
||||
)
|
||||
.expect(1)
|
||||
.mount(&server)
|
||||
.await;
|
||||
let settings = KeyManagementSettings {
|
||||
kms_key_id: Some("kms-key".into()),
|
||||
tags: Some(std::collections::BTreeMap::from([(
|
||||
"stage".into(),
|
||||
"test".into(),
|
||||
)])),
|
||||
replica_regions: Some(vec!["replica-region".into()]),
|
||||
..Default::default()
|
||||
};
|
||||
let manager = manager(&server, settings);
|
||||
assert_eq!(
|
||||
manager
|
||||
.async_write_secret("key", &SecretValue::new("value"), None)
|
||||
.await
|
||||
.unwrap()
|
||||
.name(),
|
||||
Some("key")
|
||||
);
|
||||
assert!(
|
||||
manager
|
||||
.async_replicate_secret("key", &[])
|
||||
.await
|
||||
.unwrap()
|
||||
.is_none()
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn credential_failures_are_not_swallowed_as_missing_secrets() {
|
||||
use aws_credential_types::provider::{ProvideCredentials, error::CredentialsError, future};
|
||||
#[derive(Debug)]
|
||||
struct FailedCredentials;
|
||||
impl ProvideCredentials for FailedCredentials {
|
||||
fn provide_credentials<'a>(&'a self) -> future::ProvideCredentials<'a>
|
||||
where
|
||||
Self: 'a,
|
||||
{
|
||||
future::ProvideCredentials::ready(Err(CredentialsError::provider_error(
|
||||
"private-auth-detail",
|
||||
)))
|
||||
}
|
||||
}
|
||||
let server = MockServer::start().await;
|
||||
let config = aws_sdk_secretsmanager::Config::builder()
|
||||
.behavior_version(BehaviorVersion::latest())
|
||||
.region(Region::new("us-east-1"))
|
||||
.credentials_provider(FailedCredentials)
|
||||
.endpoint_url(server.uri())
|
||||
.retry_config(RetryConfig::disabled())
|
||||
.build();
|
||||
let manager = AwsSecretsManagerV2::new(Client::from_conf(config), Default::default());
|
||||
let error = manager.async_read_secret("key").await.unwrap_err();
|
||||
assert!(!format!("{error:?}").contains("private-auth-detail"));
|
||||
assert!(matches!(error, Error::Read(_)));
|
||||
assert!(server.received_requests().await.unwrap().is_empty());
|
||||
}
|
||||
|
||||
#[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(
|
||||
ResponseTemplate::new(200)
|
||||
.set_delay(Duration::from_secs(1))
|
||||
.set_body_json(json!({"SecretString":"late"})),
|
||||
)
|
||||
.mount(&server)
|
||||
.await;
|
||||
let config = aws_sdk_secretsmanager::Config::builder()
|
||||
.behavior_version(BehaviorVersion::latest())
|
||||
.region(Region::new("us-east-1"))
|
||||
.credentials_provider(Credentials::new("test", "test", None, None, "test"))
|
||||
.endpoint_url(server.uri())
|
||||
.retry_config(RetryConfig::disabled())
|
||||
.timeout_config(
|
||||
aws_sdk_secretsmanager::config::timeout::TimeoutConfig::builder()
|
||||
.operation_timeout(Duration::from_millis(30))
|
||||
.build(),
|
||||
)
|
||||
.build();
|
||||
let manager = AwsSecretsManagerV2::new(Client::from_conf(config), Default::default());
|
||||
assert!(matches!(
|
||||
manager.async_read_secret("key").await,
|
||||
Err(Error::Timeout)
|
||||
));
|
||||
}
|
||||
28
litellm-rust/crates/secrets-google/Cargo.toml
Normal file
28
litellm-rust/crates/secrets-google/Cargo.toml
Normal file
|
|
@ -0,0 +1,28 @@
|
|||
[package]
|
||||
name = "litellm-secrets-google"
|
||||
version = "0.1.0"
|
||||
edition.workspace = true
|
||||
license.workspace = true
|
||||
repository.workspace = true
|
||||
|
||||
[dependencies]
|
||||
litellm-auth-gcp = { workspace = true, features = ["google-sdk"] }
|
||||
litellm-secrets-types.workspace = true
|
||||
litellm-auth-types.workspace = true
|
||||
litellm-core-utils.workspace = true
|
||||
base64.workspace = true
|
||||
serde_json.workspace = true
|
||||
thiserror.workspace = true
|
||||
moka.workspace = true
|
||||
veil.workspace = true
|
||||
google-cloud-kms-v1 = "1.14.0"
|
||||
google-cloud-gax = { version = "1.14.0", default-features = false }
|
||||
percent-encoding = "2.3"
|
||||
serde.workspace = true
|
||||
reqwest.workspace = true
|
||||
|
||||
[dev-dependencies]
|
||||
google-cloud-auth.workspace = true
|
||||
rstest.workspace = true
|
||||
tokio.workspace = true
|
||||
wiremock = "0.6.5"
|
||||
21
litellm-rust/crates/secrets-google/src/auth.rs
Normal file
21
litellm-rust/crates/secrets-google/src/auth.rs
Normal file
|
|
@ -0,0 +1,21 @@
|
|||
use std::sync::Arc;
|
||||
|
||||
use litellm_auth_gcp::{GoogleCredentials, VertexConfig};
|
||||
use litellm_auth_types::{InputSource, Sourced};
|
||||
use litellm_core_utils::settings::Lookup;
|
||||
use litellm_secrets_types::SecretValue;
|
||||
|
||||
pub(crate) fn credentials(
|
||||
project: Option<String>,
|
||||
credentials: Option<SecretValue>,
|
||||
environment: Arc<dyn Lookup + Send + Sync>,
|
||||
) -> GoogleCredentials {
|
||||
GoogleCredentials::new(
|
||||
VertexConfig::new(
|
||||
credentials.map(|value| Sourced::new(value, InputSource::Environment)),
|
||||
project,
|
||||
None,
|
||||
),
|
||||
Arc::new(move |name| environment.get(name)),
|
||||
)
|
||||
}
|
||||
43
litellm-rust/crates/secrets-google/src/error.rs
Normal file
43
litellm-rust/crates/secrets-google/src/error.rs
Normal file
|
|
@ -0,0 +1,43 @@
|
|||
#[derive(thiserror::Error, veil::Redact)]
|
||||
pub enum Error {
|
||||
#[error("Google KMS client configuration failed")]
|
||||
Client(
|
||||
#[from]
|
||||
#[redact]
|
||||
google_cloud_gax::client_builder::Error,
|
||||
),
|
||||
#[error("Google authentication failed")]
|
||||
Auth(
|
||||
#[from]
|
||||
#[redact]
|
||||
litellm_auth_types::Error,
|
||||
),
|
||||
#[error("Google KMS request failed")]
|
||||
Kms(
|
||||
#[from]
|
||||
#[redact]
|
||||
google_cloud_gax::error::Error,
|
||||
),
|
||||
#[error("Google Secret Manager HTTP request failed")]
|
||||
Http(
|
||||
#[from]
|
||||
#[redact]
|
||||
reqwest::Error,
|
||||
),
|
||||
#[error("Google Secret Manager returned HTTP {0}")]
|
||||
Status(u16),
|
||||
#[error("Google Secret Manager returned no payload")]
|
||||
MissingPayload,
|
||||
#[error("required environment variable is missing: {0}")]
|
||||
MissingEnvironment(&'static str),
|
||||
#[error("invalid refresh interval")]
|
||||
RefreshInterval,
|
||||
#[error("payload is not valid base64")]
|
||||
Base64(#[from] base64::DecodeError),
|
||||
#[error("decrypted value is not UTF-8")]
|
||||
Utf8,
|
||||
#[error("invalid Google Secret Manager endpoint")]
|
||||
Endpoint,
|
||||
#[error("Google Secret Manager requires an enterprise license")]
|
||||
EnterpriseRequired,
|
||||
}
|
||||
67
litellm-rust/crates/secrets-google/src/kms.rs
Normal file
67
litellm-rust/crates/secrets-google/src/kms.rs
Normal file
|
|
@ -0,0 +1,67 @@
|
|||
use std::sync::Arc;
|
||||
|
||||
use google_cloud_kms_v1::client::KeyManagementService;
|
||||
use litellm_core_utils::settings::Lookup;
|
||||
use litellm_secrets_types::SecretValue;
|
||||
|
||||
use crate::{Error, auth};
|
||||
|
||||
const GOOGLE_APPLICATION_CREDENTIALS: &str = "GOOGLE_APPLICATION_CREDENTIALS";
|
||||
const GOOGLE_KMS_RESOURCE_NAME: &str = "GOOGLE_KMS_RESOURCE_NAME";
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct GoogleKms {
|
||||
client: KeyManagementService,
|
||||
resource_name: String,
|
||||
}
|
||||
|
||||
impl GoogleKms {
|
||||
pub fn new(client: KeyManagementService, resource_name: String) -> Self {
|
||||
Self {
|
||||
client,
|
||||
resource_name,
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn decrypt(&self, ciphertext: Vec<u8>) -> Result<Vec<u8>, Error> {
|
||||
let response = self
|
||||
.client
|
||||
.decrypt()
|
||||
.set_name(&self.resource_name)
|
||||
.set_ciphertext(ciphertext)
|
||||
.send()
|
||||
.await?;
|
||||
Ok(response.plaintext.to_vec())
|
||||
}
|
||||
}
|
||||
|
||||
pub fn validate_environment(environment: &dyn Lookup) -> Result<(), Error> {
|
||||
for key in [GOOGLE_APPLICATION_CREDENTIALS, GOOGLE_KMS_RESOURCE_NAME] {
|
||||
if environment.get(key).is_none() {
|
||||
return Err(Error::MissingEnvironment(key));
|
||||
}
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
pub async fn load_google_kms(
|
||||
use_google_kms: Option<bool>,
|
||||
environment: Arc<dyn Lookup + Send + Sync>,
|
||||
) -> Result<Option<GoogleKms>, Error> {
|
||||
if use_google_kms != Some(true) {
|
||||
return Ok(None);
|
||||
}
|
||||
validate_environment(environment.as_ref())?;
|
||||
let credentials = environment
|
||||
.get(GOOGLE_APPLICATION_CREDENTIALS)
|
||||
.ok_or(Error::MissingEnvironment(GOOGLE_APPLICATION_CREDENTIALS))?;
|
||||
let resource_name = environment
|
||||
.get(GOOGLE_KMS_RESOURCE_NAME)
|
||||
.ok_or(Error::MissingEnvironment(GOOGLE_KMS_RESOURCE_NAME))?;
|
||||
let credentials = auth::credentials(None, Some(SecretValue::new(credentials)), environment);
|
||||
let client = KeyManagementService::builder()
|
||||
.with_credentials(credentials)
|
||||
.build()
|
||||
.await?;
|
||||
Ok(Some(GoogleKms::new(client, resource_name)))
|
||||
}
|
||||
10
litellm-rust/crates/secrets-google/src/lib.rs
Normal file
10
litellm-rust/crates/secrets-google/src/lib.rs
Normal file
|
|
@ -0,0 +1,10 @@
|
|||
#![forbid(unsafe_code)]
|
||||
|
||||
mod auth;
|
||||
mod error;
|
||||
pub mod kms;
|
||||
pub mod secret_manager;
|
||||
|
||||
pub use error::Error;
|
||||
pub use kms::{GoogleKms, load_google_kms};
|
||||
pub use secret_manager::GoogleSecretManager;
|
||||
169
litellm-rust/crates/secrets-google/src/secret_manager.rs
Normal file
169
litellm-rust/crates/secrets-google/src/secret_manager.rs
Normal file
|
|
@ -0,0 +1,169 @@
|
|||
use std::{sync::Arc, time::Duration};
|
||||
|
||||
use base64::{Engine, engine::general_purpose::STANDARD};
|
||||
use litellm_core_utils::settings::Lookup;
|
||||
use litellm_secrets_types::{Secret, SecretValue};
|
||||
use moka::future::Cache;
|
||||
use serde::Deserialize;
|
||||
|
||||
use litellm_auth_gcp::GoogleCredentials;
|
||||
|
||||
use crate::{Error, auth};
|
||||
|
||||
const GOOGLE_SECRET_MANAGER_PROJECT_ID: &str = "GOOGLE_SECRET_MANAGER_PROJECT_ID";
|
||||
const GOOGLE_SECRET_MANAGER_REFRESH_INTERVAL: &str = "GOOGLE_SECRET_MANAGER_REFRESH_INTERVAL";
|
||||
const SECRET_MANAGER_REFRESH_INTERVAL: &str = "SECRET_MANAGER_REFRESH_INTERVAL";
|
||||
const GOOGLE_SECRET_MANAGER_ALWAYS_READ_SECRET_MANAGER: &str =
|
||||
"GOOGLE_SECRET_MANAGER_ALWAYS_READ_SECRET_MANAGER";
|
||||
const GCS_PATH_SERVICE_ACCOUNT: &str = "GCS_PATH_SERVICE_ACCOUNT";
|
||||
const DEFAULT_REFRESH_INTERVAL: Duration = Duration::from_secs(86400);
|
||||
const DEFAULT_CACHE_TTL: Duration = Duration::from_secs(600);
|
||||
const CACHE_CAPACITY: u64 = 200;
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct GoogleSecretManager {
|
||||
client: reqwest::Client,
|
||||
credentials: Arc<GoogleCredentials>,
|
||||
endpoint: reqwest::Url,
|
||||
project: String,
|
||||
cache: Cache<String, Option<SecretValue>>,
|
||||
always_read: bool,
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
struct Response {
|
||||
payload: Option<Payload>,
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
struct Payload {
|
||||
data: Option<String>,
|
||||
}
|
||||
|
||||
impl GoogleSecretManager {
|
||||
pub fn with_client(
|
||||
client: reqwest::Client,
|
||||
endpoint: reqwest::Url,
|
||||
project: String,
|
||||
environment: Arc<dyn Lookup + Send + Sync>,
|
||||
refresh_interval: Option<Duration>,
|
||||
always_read: bool,
|
||||
) -> Result<Self, Error> {
|
||||
let credentials = auth::credentials(
|
||||
Some(project.clone()),
|
||||
environment
|
||||
.get(GCS_PATH_SERVICE_ACCOUNT)
|
||||
.map(SecretValue::new),
|
||||
environment,
|
||||
);
|
||||
let ttl = refresh_interval
|
||||
.filter(|ttl| !ttl.is_zero())
|
||||
.unwrap_or(DEFAULT_CACHE_TTL);
|
||||
let cache = Cache::builder()
|
||||
.max_capacity(CACHE_CAPACITY)
|
||||
.time_to_live(ttl)
|
||||
.build();
|
||||
Ok(Self {
|
||||
client,
|
||||
credentials: Arc::new(credentials),
|
||||
endpoint,
|
||||
project,
|
||||
cache,
|
||||
always_read,
|
||||
})
|
||||
}
|
||||
|
||||
pub fn new(
|
||||
environment: Arc<dyn Lookup + Send + Sync>,
|
||||
enterprise_enabled: bool,
|
||||
) -> Result<Self, Error> {
|
||||
if !enterprise_enabled {
|
||||
return Err(Error::EnterpriseRequired);
|
||||
}
|
||||
let project = environment
|
||||
.get(GOOGLE_SECRET_MANAGER_PROJECT_ID)
|
||||
.ok_or(Error::MissingEnvironment(GOOGLE_SECRET_MANAGER_PROJECT_ID))?;
|
||||
let ttl = environment
|
||||
.get(GOOGLE_SECRET_MANAGER_REFRESH_INTERVAL)
|
||||
.filter(|v| !v.is_empty())
|
||||
.map(|v| v.parse::<i64>().map_err(|_| Error::RefreshInterval))
|
||||
.transpose()?
|
||||
.unwrap_or(
|
||||
environment
|
||||
.get(SECRET_MANAGER_REFRESH_INTERVAL)
|
||||
.map(|v| v.parse::<i64>().map_err(|_| Error::RefreshInterval))
|
||||
.transpose()?
|
||||
.unwrap_or(DEFAULT_REFRESH_INTERVAL.as_secs() as i64),
|
||||
);
|
||||
let always_read = environment
|
||||
.get(GOOGLE_SECRET_MANAGER_ALWAYS_READ_SECRET_MANAGER)
|
||||
.is_some_and(|v| v.eq_ignore_ascii_case("true"));
|
||||
Self::with_client(
|
||||
reqwest::Client::new(),
|
||||
reqwest::Url::parse("https://secretmanager.googleapis.com").expect("static URL"),
|
||||
project,
|
||||
environment,
|
||||
Some(if ttl < 0 {
|
||||
Duration::from_nanos(1)
|
||||
} else {
|
||||
Duration::from_secs(ttl as u64)
|
||||
}),
|
||||
always_read,
|
||||
)
|
||||
}
|
||||
|
||||
pub async fn get_secret_from_google_secret_manager(
|
||||
&self,
|
||||
name: &str,
|
||||
) -> Result<Option<Secret>, Error> {
|
||||
if !self.always_read
|
||||
&& let Some(cached) = self.cache.get(name).await
|
||||
{
|
||||
return Ok(cached.and_then(cached_secret));
|
||||
}
|
||||
let url = self
|
||||
.endpoint
|
||||
.join(&format!(
|
||||
"/v1/projects/{}/secrets/{}/versions/latest:access",
|
||||
percent_encoding::utf8_percent_encode(
|
||||
&self.project,
|
||||
percent_encoding::NON_ALPHANUMERIC
|
||||
),
|
||||
percent_encoding::utf8_percent_encode(name, percent_encoding::NON_ALPHANUMERIC)
|
||||
))
|
||||
.map_err(|_| Error::Endpoint)?;
|
||||
let response = self
|
||||
.client
|
||||
.get(url)
|
||||
.headers(self.credentials.request_headers().await?)
|
||||
.send()
|
||||
.await?;
|
||||
if response.status() != reqwest::StatusCode::OK {
|
||||
self.cache.insert(name.to_owned(), None).await;
|
||||
return Err(Error::Status(response.status().as_u16()));
|
||||
}
|
||||
let response: Response = response.json().await?;
|
||||
let Some(data) = response.payload.and_then(|payload| payload.data) else {
|
||||
self.cache.insert(name.to_owned(), None).await;
|
||||
return Err(Error::MissingPayload);
|
||||
};
|
||||
let filtered: String = data
|
||||
.chars()
|
||||
.filter(|c| c.is_ascii_alphanumeric() || matches!(c, '+' | '/' | '='))
|
||||
.collect();
|
||||
let bytes = STANDARD.decode(filtered)?;
|
||||
let plaintext = String::from_utf8(bytes).map_err(|_| Error::Utf8)?;
|
||||
let value = SecretValue::new(plaintext);
|
||||
self.cache
|
||||
.insert(name.to_owned(), Some(value.clone()))
|
||||
.await;
|
||||
Ok(Some(Secret::String(value)))
|
||||
}
|
||||
}
|
||||
|
||||
fn cached_secret(value: SecretValue) -> Option<Secret> {
|
||||
match serde_json::from_str(value.expose()) {
|
||||
Ok(json) => Secret::from_json(json),
|
||||
Err(_) => Some(Secret::String(value)),
|
||||
}
|
||||
}
|
||||
49
litellm-rust/crates/secrets-google/tests/kms.rs
Normal file
49
litellm-rust/crates/secrets-google/tests/kms.rs
Normal file
|
|
@ -0,0 +1,49 @@
|
|||
use base64::{Engine, engine::general_purpose::STANDARD};
|
||||
use google_cloud_kms_v1::client::KeyManagementService;
|
||||
use litellm_secrets_google::GoogleKms;
|
||||
use wiremock::{
|
||||
Mock, MockServer, ResponseTemplate,
|
||||
matchers::{body_json, path},
|
||||
};
|
||||
|
||||
#[tokio::test]
|
||||
async fn google_kms_decrypts_using_the_configured_resource() {
|
||||
let server = MockServer::start().await;
|
||||
let resource = "projects/project/locations/global/keyRings/ring/cryptoKeys/key";
|
||||
Mock::given(path(format!("/v1/{resource}:decrypt")))
|
||||
.and(body_json(
|
||||
serde_json::json!({"ciphertext":STANDARD.encode("encrypted")}),
|
||||
))
|
||||
.respond_with(
|
||||
ResponseTemplate::new(200)
|
||||
.set_body_json(serde_json::json!({"plaintext":STANDARD.encode(" value\n")})),
|
||||
)
|
||||
.expect(1)
|
||||
.mount(&server)
|
||||
.await;
|
||||
let client = KeyManagementService::builder()
|
||||
.with_endpoint(server.uri())
|
||||
.with_credentials(google_cloud_auth::credentials::anonymous::Builder::new().build())
|
||||
.with_retry_policy(google_cloud_gax::retry_policy::NeverRetry)
|
||||
.build()
|
||||
.await
|
||||
.unwrap();
|
||||
let manager = GoogleKms::new(client, resource.into());
|
||||
assert_eq!(
|
||||
manager.decrypt(b"encrypted".to_vec()).await.unwrap(),
|
||||
b" value\n"
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn disabled_google_kms_loader_does_not_require_environment_configuration() {
|
||||
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()
|
||||
);
|
||||
}
|
||||
}
|
||||
176
litellm-rust/crates/secrets-google/tests/secret_manager.rs
Normal file
176
litellm-rust/crates/secrets-google/tests/secret_manager.rs
Normal file
|
|
@ -0,0 +1,176 @@
|
|||
use std::{sync::Arc, time::Duration};
|
||||
|
||||
use base64::{Engine, engine::general_purpose::STANDARD};
|
||||
use litellm_secrets_google::{Error, GoogleSecretManager};
|
||||
use litellm_secrets_types::Secret;
|
||||
use wiremock::{
|
||||
Mock, MockServer, ResponseTemplate,
|
||||
matchers::{header, path},
|
||||
};
|
||||
|
||||
fn manager(server: &MockServer, always_read: bool, ttl: Duration) -> GoogleSecretManager {
|
||||
GoogleSecretManager::with_client(
|
||||
reqwest::Client::new(),
|
||||
server.uri().parse().unwrap(),
|
||||
"project".into(),
|
||||
Arc::new(|name: &str| (name == "VERTEX_AI_API_KEY").then(|| "token".into())),
|
||||
Some(ttl),
|
||||
always_read,
|
||||
)
|
||||
.unwrap()
|
||||
}
|
||||
|
||||
#[rstest::rstest]
|
||||
#[case::nonempty("private-value")]
|
||||
#[case::empty("")]
|
||||
#[tokio::test]
|
||||
async fn successful_reads_use_auth_latest_version_and_cache_including_empty_values(
|
||||
#[case] value: &str,
|
||||
) {
|
||||
let server = MockServer::start().await;
|
||||
Mock::given(path(
|
||||
"/v1/projects/project/secrets/key/versions/latest:access",
|
||||
))
|
||||
.and(header("authorization", "Bearer token"))
|
||||
.respond_with(
|
||||
ResponseTemplate::new(200)
|
||||
.set_body_json(serde_json::json!({"payload":{"data":STANDARD.encode(value)}})),
|
||||
)
|
||||
.expect(1)
|
||||
.mount(&server)
|
||||
.await;
|
||||
let manager = manager(&server, false, Duration::from_secs(60));
|
||||
for _ in 0..2 {
|
||||
assert_eq!(
|
||||
manager
|
||||
.get_secret_from_google_secret_manager("key")
|
||||
.await
|
||||
.unwrap()
|
||||
.unwrap()
|
||||
.as_str()
|
||||
.unwrap(),
|
||||
value
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[rstest::rstest]
|
||||
#[case::not_found(ResponseTemplate::new(404))]
|
||||
#[case::missing_payload(
|
||||
ResponseTemplate::new(200).set_body_json(serde_json::json!({"payload":{}}))
|
||||
)]
|
||||
#[tokio::test]
|
||||
async fn negative_cache_returns_none_after_initial_error(#[case] response: ResponseTemplate) {
|
||||
let server = MockServer::start().await;
|
||||
Mock::given(path(
|
||||
"/v1/projects/project/secrets/key/versions/latest:access",
|
||||
))
|
||||
.respond_with(response)
|
||||
.expect(1)
|
||||
.mount(&server)
|
||||
.await;
|
||||
let manager = manager(&server, false, Duration::from_secs(60));
|
||||
assert!(matches!(
|
||||
manager.get_secret_from_google_secret_manager("key").await,
|
||||
Err(Error::Status(404) | Error::MissingPayload)
|
||||
));
|
||||
assert!(
|
||||
manager
|
||||
.get_secret_from_google_secret_manager("key")
|
||||
.await
|
||||
.unwrap()
|
||||
.is_none()
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest::rstest]
|
||||
#[case::always_read(true, Duration::from_secs(60))]
|
||||
#[case::expired_cache(false, Duration::from_millis(1))]
|
||||
#[tokio::test]
|
||||
async fn always_read_and_expired_cache_fetch_again(
|
||||
#[case] always_read: bool,
|
||||
#[case] ttl: Duration,
|
||||
) {
|
||||
let server = MockServer::start().await;
|
||||
Mock::given(path(
|
||||
"/v1/projects/project/secrets/key/versions/latest:access",
|
||||
))
|
||||
.respond_with(
|
||||
ResponseTemplate::new(200)
|
||||
.set_body_json(serde_json::json!({"payload":{"data":STANDARD.encode("value")}})),
|
||||
)
|
||||
.expect(2)
|
||||
.mount(&server)
|
||||
.await;
|
||||
let manager = manager(&server, always_read, ttl);
|
||||
for _ in 0..2 {
|
||||
tokio::time::sleep(Duration::from_millis(5)).await;
|
||||
assert!(
|
||||
manager
|
||||
.get_secret_from_google_secret_manager("key")
|
||||
.await
|
||||
.unwrap()
|
||||
.is_some()
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn google_manager_requires_host_license_and_project_configuration() {
|
||||
assert!(matches!(
|
||||
GoogleSecretManager::new(Arc::new(|_: &str| None), false),
|
||||
Err(Error::EnterpriseRequired)
|
||||
));
|
||||
assert!(matches!(
|
||||
GoogleSecretManager::new(Arc::new(|_: &str| None), true),
|
||||
Err(Error::MissingEnvironment(
|
||||
"GOOGLE_SECRET_MANAGER_PROJECT_ID"
|
||||
))
|
||||
));
|
||||
}
|
||||
|
||||
#[rstest::rstest]
|
||||
#[case::boolean("true", Some(Secret::Bool(true)))]
|
||||
#[case::null("null", None)]
|
||||
#[case::string(
|
||||
"\"text\"",
|
||||
Some(Secret::String(litellm_secrets_types::SecretValue::new("text")))
|
||||
)]
|
||||
#[case::object(
|
||||
"{\"key\":1}",
|
||||
Secret::from_json(serde_json::json!({"key":1}))
|
||||
)]
|
||||
#[tokio::test]
|
||||
async fn cached_values_preserve_python_json_conversion(
|
||||
#[case] raw: &str,
|
||||
#[case] expected: Option<Secret>,
|
||||
) {
|
||||
let server = MockServer::start().await;
|
||||
Mock::given(path(
|
||||
"/v1/projects/project/secrets/key/versions/latest:access",
|
||||
))
|
||||
.respond_with(
|
||||
ResponseTemplate::new(200)
|
||||
.set_body_json(serde_json::json!({"payload":{"data":STANDARD.encode(raw)}})),
|
||||
)
|
||||
.expect(1)
|
||||
.mount(&server)
|
||||
.await;
|
||||
let manager = manager(&server, false, Duration::from_secs(60));
|
||||
assert_eq!(
|
||||
manager
|
||||
.get_secret_from_google_secret_manager("key")
|
||||
.await
|
||||
.unwrap()
|
||||
.unwrap()
|
||||
.as_str(),
|
||||
Some(raw)
|
||||
);
|
||||
assert_eq!(
|
||||
manager
|
||||
.get_secret_from_google_secret_manager("key")
|
||||
.await
|
||||
.unwrap(),
|
||||
expected
|
||||
);
|
||||
}
|
||||
17
litellm-rust/crates/secrets-types/Cargo.toml
Normal file
17
litellm-rust/crates/secrets-types/Cargo.toml
Normal file
|
|
@ -0,0 +1,17 @@
|
|||
[package]
|
||||
name = "litellm-secrets-types"
|
||||
version = "0.1.0"
|
||||
edition.workspace = true
|
||||
license.workspace = true
|
||||
repository.workspace = true
|
||||
|
||||
[dependencies]
|
||||
litellm-auth-types.workspace = true
|
||||
serde.workspace = true
|
||||
serde_json.workspace = true
|
||||
thiserror.workspace = true
|
||||
veil.workspace = true
|
||||
|
||||
[dev-dependencies]
|
||||
rstest.workspace = true
|
||||
tokio.workspace = true
|
||||
58
litellm-rust/crates/secrets-types/src/base_secret_manager.rs
Normal file
58
litellm-rust/crates/secrets-types/src/base_secret_manager.rs
Normal file
|
|
@ -0,0 +1,58 @@
|
|||
use crate::{Error, SecretValue};
|
||||
|
||||
pub fn validate_secret_name(name: &str) -> Result<(), Error> {
|
||||
if name.split('/').any(|segment| segment == "..")
|
||||
|| name
|
||||
.chars()
|
||||
.any(|c| c.is_control() || matches!(c, '\u{2028}' | '\u{2029}'))
|
||||
{
|
||||
return Err(Error::UnsafeSecretName);
|
||||
}
|
||||
Ok(())
|
||||
}
|
||||
|
||||
#[expect(
|
||||
async_fn_in_trait,
|
||||
reason = "closed backend dispatch does not require Send bounds on generic rotation"
|
||||
)]
|
||||
pub trait BaseSecretManager {
|
||||
type Error: From<Error>;
|
||||
type WriteResponse;
|
||||
type DeleteResponse;
|
||||
|
||||
async fn async_read_secret(&self, name: &str) -> Result<Option<SecretValue>, Self::Error>;
|
||||
async fn async_write_secret(
|
||||
&self,
|
||||
name: &str,
|
||||
value: &SecretValue,
|
||||
description: Option<&str>,
|
||||
) -> Result<Self::WriteResponse, Self::Error>;
|
||||
async fn async_delete_secret(
|
||||
&self,
|
||||
name: &str,
|
||||
recovery_window_in_days: i64,
|
||||
) -> Result<Self::DeleteResponse, Self::Error>;
|
||||
}
|
||||
|
||||
pub async fn async_rotate_secret<M: BaseSecretManager>(
|
||||
manager: &M,
|
||||
current_name: &str,
|
||||
new_name: &str,
|
||||
value: &SecretValue,
|
||||
) -> Result<M::WriteResponse, M::Error> {
|
||||
if manager.async_read_secret(current_name).await?.is_none() {
|
||||
return Err(Error::CurrentSecretMissing.into());
|
||||
}
|
||||
let response = manager
|
||||
.async_write_secret(
|
||||
new_name,
|
||||
value,
|
||||
Some(&format!("Rotated from {current_name}")),
|
||||
)
|
||||
.await?;
|
||||
if manager.async_read_secret(new_name).await?.is_none() {
|
||||
return Err(Error::NewSecretMissing.into());
|
||||
}
|
||||
manager.async_delete_secret(current_name, 7).await?;
|
||||
Ok(response)
|
||||
}
|
||||
92
litellm-rust/crates/secrets-types/src/config.rs
Normal file
92
litellm-rust/crates/secrets-types/src/config.rs
Normal file
|
|
@ -0,0 +1,92 @@
|
|||
use std::collections::BTreeMap;
|
||||
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
use crate::SecretValue;
|
||||
|
||||
#[derive(Clone, Copy, Debug, Deserialize, Eq, PartialEq, Serialize)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub enum KeyManagementSystem {
|
||||
GoogleKms,
|
||||
AzureKeyVault,
|
||||
AwsSecretManager,
|
||||
GoogleSecretManager,
|
||||
HashicorpVault,
|
||||
Cyberark,
|
||||
Local,
|
||||
AwsKms,
|
||||
Custom,
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, Default, Deserialize, Eq, PartialEq, Serialize)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub enum AccessMode {
|
||||
#[default]
|
||||
ReadOnly,
|
||||
WriteOnly,
|
||||
ReadAndWrite,
|
||||
}
|
||||
|
||||
impl AccessMode {
|
||||
pub fn readable(self) -> bool {
|
||||
matches!(self, Self::ReadOnly | Self::ReadAndWrite)
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Deserialize, Serialize, PartialEq, Eq)]
|
||||
#[serde(default)]
|
||||
pub struct KeyManagementSettings {
|
||||
pub hosted_keys: Option<Vec<String>>,
|
||||
pub store_virtual_keys: Option<bool>,
|
||||
pub prefix_for_stored_virtual_keys: String,
|
||||
pub access_mode: AccessMode,
|
||||
pub primary_secret_name: Option<String>,
|
||||
pub description: Option<String>,
|
||||
pub tags: Option<BTreeMap<String, String>>,
|
||||
pub kms_key_id: Option<String>,
|
||||
pub custom_secret_manager: Option<String>,
|
||||
pub aws_region_name: Option<String>,
|
||||
pub aws_role_name: Option<String>,
|
||||
pub aws_session_name: Option<String>,
|
||||
#[serde(serialize_with = "serialize_secret")]
|
||||
pub aws_external_id: Option<SecretValue>,
|
||||
pub aws_profile_name: Option<String>,
|
||||
#[serde(serialize_with = "serialize_secret")]
|
||||
pub aws_web_identity_token: Option<SecretValue>,
|
||||
pub aws_sts_endpoint: Option<String>,
|
||||
pub replica_regions: Option<Vec<String>>,
|
||||
}
|
||||
|
||||
impl Default for KeyManagementSettings {
|
||||
fn default() -> Self {
|
||||
Self {
|
||||
hosted_keys: None,
|
||||
store_virtual_keys: Some(false),
|
||||
prefix_for_stored_virtual_keys: "litellm/".into(),
|
||||
access_mode: AccessMode::ReadOnly,
|
||||
primary_secret_name: None,
|
||||
description: None,
|
||||
tags: None,
|
||||
kms_key_id: None,
|
||||
custom_secret_manager: None,
|
||||
aws_region_name: None,
|
||||
aws_role_name: None,
|
||||
aws_session_name: None,
|
||||
aws_external_id: None,
|
||||
aws_profile_name: None,
|
||||
aws_web_identity_token: None,
|
||||
aws_sts_endpoint: None,
|
||||
replica_regions: None,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn serialize_secret<S: serde::Serializer>(
|
||||
value: &Option<SecretValue>,
|
||||
serializer: S,
|
||||
) -> Result<S::Ok, S::Error> {
|
||||
value
|
||||
.as_ref()
|
||||
.map(SecretValue::expose)
|
||||
.serialize(serializer)
|
||||
}
|
||||
9
litellm-rust/crates/secrets-types/src/error.rs
Normal file
9
litellm-rust/crates/secrets-types/src/error.rs
Normal file
|
|
@ -0,0 +1,9 @@
|
|||
#[derive(Debug, thiserror::Error, PartialEq, Eq)]
|
||||
pub enum Error {
|
||||
#[error("secret name contains an unsafe path segment or control character")]
|
||||
UnsafeSecretName,
|
||||
#[error("current secret was not found")]
|
||||
CurrentSecretMissing,
|
||||
#[error("new secret could not be verified")]
|
||||
NewSecretMissing,
|
||||
}
|
||||
12
litellm-rust/crates/secrets-types/src/lib.rs
Normal file
12
litellm-rust/crates/secrets-types/src/lib.rs
Normal file
|
|
@ -0,0 +1,12 @@
|
|||
#![forbid(unsafe_code)]
|
||||
|
||||
mod base_secret_manager;
|
||||
mod config;
|
||||
mod error;
|
||||
mod value;
|
||||
|
||||
pub use base_secret_manager::{BaseSecretManager, async_rotate_secret, validate_secret_name};
|
||||
pub use config::{AccessMode, KeyManagementSettings, KeyManagementSystem};
|
||||
pub use error::Error;
|
||||
pub use litellm_auth_types::SecretValue;
|
||||
pub use value::Secret;
|
||||
32
litellm-rust/crates/secrets-types/src/value.rs
Normal file
32
litellm-rust/crates/secrets-types/src/value.rs
Normal file
|
|
@ -0,0 +1,32 @@
|
|||
use crate::SecretValue;
|
||||
|
||||
#[derive(Clone, PartialEq, Eq, veil::Redact)]
|
||||
pub enum Secret {
|
||||
String(SecretValue),
|
||||
Bool(#[redact] bool),
|
||||
Json(#[redact] serde_json::Value),
|
||||
}
|
||||
|
||||
impl From<SecretValue> for Secret {
|
||||
fn from(value: SecretValue) -> Self {
|
||||
Self::String(value)
|
||||
}
|
||||
}
|
||||
|
||||
impl Secret {
|
||||
pub fn from_json(value: serde_json::Value) -> Option<Self> {
|
||||
match value {
|
||||
serde_json::Value::Null => None,
|
||||
serde_json::Value::String(value) => Some(Self::String(SecretValue::new(value))),
|
||||
serde_json::Value::Bool(value) => Some(Self::Bool(value)),
|
||||
value => Some(Self::Json(value)),
|
||||
}
|
||||
}
|
||||
|
||||
pub fn as_str(&self) -> Option<&str> {
|
||||
match self {
|
||||
Self::String(value) => Some(value.expose()),
|
||||
Self::Bool(_) | Self::Json(_) => None,
|
||||
}
|
||||
}
|
||||
}
|
||||
60
litellm-rust/crates/secrets-types/tests/config.rs
Normal file
60
litellm-rust/crates/secrets-types/tests/config.rs
Normal file
|
|
@ -0,0 +1,60 @@
|
|||
use litellm_secrets_types::{
|
||||
AccessMode, KeyManagementSettings, KeyManagementSystem, Secret, SecretValue,
|
||||
};
|
||||
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/");
|
||||
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",
|
||||
"tags": {"stage": "test"}, "replica_regions": ["test-region"]
|
||||
}))
|
||||
.unwrap();
|
||||
assert!(!configured.access_mode.readable());
|
||||
assert_eq!(configured.store_virtual_keys, None);
|
||||
assert_eq!(configured.hosted_keys.as_deref(), Some([].as_slice()));
|
||||
assert!(!format!("{configured:?}").contains("private-"));
|
||||
let serialized = serde_json::to_value(&configured).unwrap();
|
||||
assert_eq!(serialized["access_mode"], "write_only");
|
||||
assert_eq!(serialized["aws_web_identity_token"], "private-token");
|
||||
assert_eq!(
|
||||
serde_json::from_value::<KeyManagementSettings>(serialized).unwrap(),
|
||||
configured
|
||||
);
|
||||
}
|
||||
|
||||
#[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)]
|
||||
#[case::google_secret_manager("google_secret_manager", KeyManagementSystem::GoogleSecretManager)]
|
||||
#[case::azure_key_vault("azure_key_vault", KeyManagementSystem::AzureKeyVault)]
|
||||
#[case::hashicorp_vault("hashicorp_vault", KeyManagementSystem::HashicorpVault)]
|
||||
#[case::cyberark("cyberark", KeyManagementSystem::Cyberark)]
|
||||
#[case::custom("custom", KeyManagementSystem::Custom)]
|
||||
#[case::local("local", KeyManagementSystem::Local)]
|
||||
fn key_management_system_serialization_round_trips(
|
||||
#[case] name: &str,
|
||||
#[case] system: KeyManagementSystem,
|
||||
) {
|
||||
assert_eq!(
|
||||
serde_json::from_value::<KeyManagementSystem>(json!(name)).unwrap(),
|
||||
system
|
||||
);
|
||||
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"));
|
||||
}
|
||||
105
litellm-rust/crates/secrets-types/tests/rotation.rs
Normal file
105
litellm-rust/crates/secrets-types/tests/rotation.rs
Normal file
|
|
@ -0,0 +1,105 @@
|
|||
use std::sync::atomic::{AtomicUsize, Ordering};
|
||||
|
||||
use litellm_secrets_types::{
|
||||
BaseSecretManager, Error, SecretValue, async_rotate_secret, validate_secret_name,
|
||||
};
|
||||
|
||||
struct Manager {
|
||||
step: AtomicUsize,
|
||||
absent_at: Option<usize>,
|
||||
}
|
||||
|
||||
impl BaseSecretManager for Manager {
|
||||
type Error = Error;
|
||||
type WriteResponse = &'static str;
|
||||
type DeleteResponse = ();
|
||||
|
||||
async fn async_read_secret(&self, name: &str) -> Result<Option<SecretValue>, Error> {
|
||||
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")))
|
||||
}
|
||||
|
||||
async fn async_write_secret(
|
||||
&self,
|
||||
name: &str,
|
||||
value: &SecretValue,
|
||||
description: Option<&str>,
|
||||
) -> 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"));
|
||||
Ok("provider-response")
|
||||
}
|
||||
|
||||
async fn async_delete_secret(
|
||||
&self,
|
||||
name: &str,
|
||||
recovery_window_in_days: i64,
|
||||
) -> Result<(), Error> {
|
||||
assert_eq!(self.step.fetch_add(1, Ordering::SeqCst), 3);
|
||||
assert_eq!(name, "old");
|
||||
assert_eq!(recovery_window_in_days, 7);
|
||||
Ok(())
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn rotation_verifies_before_deleting_and_returns_provider_response() {
|
||||
let manager = Manager {
|
||||
step: AtomicUsize::new(0),
|
||||
absent_at: None,
|
||||
};
|
||||
assert_eq!(
|
||||
async_rotate_secret(&manager, "old", "new", &SecretValue::new("replacement"))
|
||||
.await
|
||||
.unwrap(),
|
||||
"provider-response"
|
||||
);
|
||||
assert_eq!(manager.step.load(Ordering::SeqCst), 4);
|
||||
}
|
||||
|
||||
#[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(
|
||||
#[case] absent_at: usize,
|
||||
#[case] expected: Error,
|
||||
#[case] calls: usize,
|
||||
) {
|
||||
let manager = Manager {
|
||||
step: AtomicUsize::new(0),
|
||||
absent_at: Some(absent_at),
|
||||
};
|
||||
assert_eq!(
|
||||
async_rotate_secret(&manager, "old", "new", &SecretValue::new("replacement"))
|
||||
.await
|
||||
.unwrap_err(),
|
||||
expected
|
||||
);
|
||||
assert_eq!(manager.step.load(Ordering::SeqCst), calls);
|
||||
}
|
||||
|
||||
#[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}")]
|
||||
fn names_reject_path_traversal_and_control_characters(#[case] name: &str) {
|
||||
assert_eq!(validate_secret_name(name), Err(Error::UnsafeSecretName));
|
||||
}
|
||||
|
||||
#[rstest::rstest]
|
||||
#[case::embedded_double_dot("release-1.0..2")]
|
||||
#[case::path_separator("folder/key")]
|
||||
#[case::empty("")]
|
||||
#[case::three_dots("...")]
|
||||
fn names_allow_safe_values(#[case] name: &str) {
|
||||
assert_eq!(validate_secret_name(name), Ok(()));
|
||||
}
|
||||
37
litellm-rust/crates/secrets/Cargo.toml
Normal file
37
litellm-rust/crates/secrets/Cargo.toml
Normal file
|
|
@ -0,0 +1,37 @@
|
|||
[package]
|
||||
name = "litellm-secrets"
|
||||
version = "0.1.0"
|
||||
edition.workspace = true
|
||||
license.workspace = true
|
||||
repository.workspace = true
|
||||
|
||||
[features]
|
||||
default = []
|
||||
aws = ["dep:litellm-secrets-aws"]
|
||||
google = ["dep:litellm-secrets-google"]
|
||||
|
||||
[dependencies]
|
||||
litellm-secrets-types.workspace = true
|
||||
litellm-secrets-aws = { workspace = true, optional = true }
|
||||
litellm-secrets-google = { workspace = true, optional = true }
|
||||
litellm-core-utils.workspace = true
|
||||
base64.workspace = true
|
||||
serde.workspace = true
|
||||
strum.workspace = true
|
||||
jsonwebtoken.workspace = true
|
||||
serde_json.workspace = true
|
||||
thiserror.workspace = true
|
||||
tracing = "0.1"
|
||||
reqwest.workspace = true
|
||||
moka.workspace = true
|
||||
tokio = { workspace = true, features = ["fs"] }
|
||||
|
||||
rustpython-parser = { version = "0.4.0", default-features = false, features = ["num-bigint"] }
|
||||
|
||||
[dev-dependencies]
|
||||
rstest.workspace = true
|
||||
wiremock = "0.6.5"
|
||||
tempfile = "3"
|
||||
aws-sdk-kms = "1.120.0"
|
||||
google-cloud-kms-v1 = "1.14.0"
|
||||
google-cloud-auth.workspace = true
|
||||
39
litellm-rust/crates/secrets/src/error.rs
Normal file
39
litellm-rust/crates/secrets/src/error.rs
Normal file
|
|
@ -0,0 +1,39 @@
|
|||
use crate::KeyManagementSystem;
|
||||
|
||||
#[derive(Debug, thiserror::Error)]
|
||||
pub enum Error {
|
||||
#[error("encrypted environment value is missing")]
|
||||
MissingCiphertext,
|
||||
#[error("ciphertext is not valid base64 for the configured manager")]
|
||||
InvalidCiphertext,
|
||||
#[error("decrypted value is not UTF-8")]
|
||||
Utf8,
|
||||
#[error("secret manager backend is not compiled: {0:?}")]
|
||||
UnsupportedBackend(KeyManagementSystem),
|
||||
#[error("configured secret manager does not match its backend")]
|
||||
BackendMismatch,
|
||||
#[error("unsupported OIDC provider or missing build feature")]
|
||||
UnsupportedOidc,
|
||||
#[error("OIDC reference requires a provider and audience")]
|
||||
InvalidOidc,
|
||||
#[error("OIDC environment variable is missing")]
|
||||
MissingEnvironment,
|
||||
#[error("OIDC request failed")]
|
||||
OidcHttp,
|
||||
#[error("OIDC provider returned HTTP {0}")]
|
||||
OidcStatus(u16),
|
||||
#[error("OIDC response is invalid")]
|
||||
OidcResponse,
|
||||
#[error("OIDC file path must be absolute and within the credential allowlist")]
|
||||
UnsafeOidcPath,
|
||||
#[error("OIDC file could not be read")]
|
||||
OidcFile,
|
||||
#[error("secret manager returned no secret")]
|
||||
MissingSecret,
|
||||
#[cfg(feature = "aws")]
|
||||
#[error(transparent)]
|
||||
Aws(#[from] litellm_secrets_aws::Error),
|
||||
#[cfg(feature = "google")]
|
||||
#[error(transparent)]
|
||||
Google(#[from] litellm_secrets_google::Error),
|
||||
}
|
||||
118
litellm-rust/crates/secrets/src/handler.rs
Normal file
118
litellm-rust/crates/secrets/src/handler.rs
Normal file
|
|
@ -0,0 +1,118 @@
|
|||
use litellm_core_utils::settings::Lookup;
|
||||
|
||||
use crate::{Error, KeyManagementSettings, KeyManagementSystem, Secret, SecretValue};
|
||||
|
||||
#[derive(Clone)]
|
||||
pub enum SecretManager {
|
||||
Local,
|
||||
#[cfg(feature = "aws")]
|
||||
AwsKms(crate::aws::AwsKms),
|
||||
#[cfg(feature = "aws")]
|
||||
AwsSecretsManagerV2(crate::aws::AwsSecretsManagerV2),
|
||||
#[cfg(feature = "google")]
|
||||
GoogleKms(crate::google::GoogleKms),
|
||||
#[cfg(feature = "google")]
|
||||
GoogleSecretManager(crate::google::GoogleSecretManager),
|
||||
}
|
||||
|
||||
impl SecretManager {
|
||||
pub fn system(&self) -> KeyManagementSystem {
|
||||
match self {
|
||||
Self::Local => KeyManagementSystem::Local,
|
||||
#[cfg(feature = "aws")]
|
||||
Self::AwsKms(_) => KeyManagementSystem::AwsKms,
|
||||
#[cfg(feature = "aws")]
|
||||
Self::AwsSecretsManagerV2(_) => KeyManagementSystem::AwsSecretManager,
|
||||
#[cfg(feature = "google")]
|
||||
Self::GoogleKms(_) => KeyManagementSystem::GoogleKms,
|
||||
#[cfg(feature = "google")]
|
||||
Self::GoogleSecretManager(_) => KeyManagementSystem::GoogleSecretManager,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn get_secret_from_manager(
|
||||
client: &SecretManager,
|
||||
secret_name: &str,
|
||||
_settings: &KeyManagementSettings,
|
||||
environment: &(dyn Lookup + Send + Sync),
|
||||
) -> Result<Option<Secret>, Error> {
|
||||
match client {
|
||||
SecretManager::Local => Ok(environment
|
||||
.get(secret_name)
|
||||
.map(SecretValue::new)
|
||||
.map(Secret::String)),
|
||||
#[cfg(feature = "aws")]
|
||||
SecretManager::AwsKms(client) => {
|
||||
let ciphertext = environment
|
||||
.get(secret_name)
|
||||
.ok_or(Error::MissingCiphertext)?;
|
||||
let plaintext = client
|
||||
.decrypt(decode_ciphertext(&ciphertext, Base64Mode::Permissive)?)
|
||||
.await?;
|
||||
let value = String::from_utf8(plaintext).map_err(|_| Error::Utf8)?;
|
||||
Ok(Some(Secret::String(SecretValue::new(value.trim()))))
|
||||
}
|
||||
#[cfg(feature = "google")]
|
||||
SecretManager::GoogleKms(client) => {
|
||||
let ciphertext = environment
|
||||
.get(secret_name)
|
||||
.ok_or(Error::MissingCiphertext)?;
|
||||
let plaintext = client
|
||||
.decrypt(decode_ciphertext(&ciphertext, Base64Mode::Canonical)?)
|
||||
.await?;
|
||||
let value = String::from_utf8(plaintext).map_err(|_| Error::Utf8)?;
|
||||
Ok(Some(Secret::String(SecretValue::new(value))))
|
||||
}
|
||||
#[cfg(feature = "aws")]
|
||||
SecretManager::AwsSecretsManagerV2(client) => client
|
||||
.read_secret_for_resolver(
|
||||
secret_name,
|
||||
_settings.primary_secret_name.as_deref(),
|
||||
environment,
|
||||
)
|
||||
.await
|
||||
.map_err(Error::from),
|
||||
#[cfg(feature = "google")]
|
||||
SecretManager::GoogleSecretManager(client) => client
|
||||
.get_secret_from_google_secret_manager(secret_name)
|
||||
.await?
|
||||
.map(Some)
|
||||
.ok_or(Error::MissingSecret),
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(any(feature = "aws", feature = "google"))]
|
||||
#[derive(Clone, Copy)]
|
||||
enum Base64Mode {
|
||||
#[cfg(feature = "google")]
|
||||
Canonical,
|
||||
#[cfg(feature = "aws")]
|
||||
Permissive,
|
||||
}
|
||||
|
||||
#[cfg(any(feature = "aws", feature = "google"))]
|
||||
fn decode_ciphertext(value: &str, mode: Base64Mode) -> Result<Vec<u8>, Error> {
|
||||
use base64::{Engine, engine::general_purpose::STANDARD};
|
||||
let canonical = match mode {
|
||||
#[cfg(feature = "google")]
|
||||
Base64Mode::Canonical => true,
|
||||
#[cfg(feature = "aws")]
|
||||
Base64Mode::Permissive => false,
|
||||
};
|
||||
let encoded = if canonical {
|
||||
value.to_owned()
|
||||
} else {
|
||||
value
|
||||
.chars()
|
||||
.filter(|c| c.is_ascii_alphanumeric() || matches!(c, '+' | '/' | '='))
|
||||
.collect()
|
||||
};
|
||||
let ciphertext = STANDARD
|
||||
.decode(&encoded)
|
||||
.map_err(|_| Error::InvalidCiphertext)?;
|
||||
if canonical && STANDARD.encode(&ciphertext) != encoded {
|
||||
return Err(Error::InvalidCiphertext);
|
||||
}
|
||||
Ok(ciphertext)
|
||||
}
|
||||
21
litellm-rust/crates/secrets/src/lib.rs
Normal file
21
litellm-rust/crates/secrets/src/lib.rs
Normal file
|
|
@ -0,0 +1,21 @@
|
|||
#![forbid(unsafe_code)]
|
||||
|
||||
mod error;
|
||||
mod handler;
|
||||
mod oidc;
|
||||
mod resolver;
|
||||
mod state;
|
||||
|
||||
pub use error::Error;
|
||||
pub use handler::{SecretManager, get_secret_from_manager};
|
||||
pub use litellm_secrets_types::{
|
||||
AccessMode, KeyManagementSettings, KeyManagementSystem, Secret, SecretValue,
|
||||
};
|
||||
pub use oidc::{OidcProvider, OidcReference, OidcResolver};
|
||||
pub use resolver::SecretResolver;
|
||||
pub use state::{SecretManagerState, secret_manager_would_be_consulted};
|
||||
|
||||
#[cfg(feature = "aws")]
|
||||
pub use litellm_secrets_aws as aws;
|
||||
#[cfg(feature = "google")]
|
||||
pub use litellm_secrets_google as google;
|
||||
264
litellm-rust/crates/secrets/src/oidc.rs
Normal file
264
litellm-rust/crates/secrets/src/oidc.rs
Normal file
|
|
@ -0,0 +1,264 @@
|
|||
use std::{
|
||||
path::Path,
|
||||
time::{Duration, SystemTime, UNIX_EPOCH},
|
||||
};
|
||||
|
||||
use jsonwebtoken::dangerous::insecure_decode_claims;
|
||||
use litellm_core_utils::settings::Lookup;
|
||||
use moka::future::Cache;
|
||||
use serde::Deserialize;
|
||||
|
||||
use crate::{Error, SecretValue};
|
||||
|
||||
const GOOGLE_TOKEN_MAX_TTL: Duration = Duration::from_secs(3540);
|
||||
const GITHUB_TOKEN_TTL: Duration = Duration::from_secs(295);
|
||||
const TOKEN_EXPIRY_MARGIN_SECONDS: f64 = 60.0;
|
||||
const CIRCLE_OIDC_TOKEN: &str = "CIRCLE_OIDC_TOKEN";
|
||||
const CIRCLE_OIDC_TOKEN_V2: &str = "CIRCLE_OIDC_TOKEN_V2";
|
||||
const AZURE_FEDERATED_TOKEN_FILE: &str = "AZURE_FEDERATED_TOKEN_FILE";
|
||||
const ACTIONS_ID_TOKEN_REQUEST_URL: &str = "ACTIONS_ID_TOKEN_REQUEST_URL";
|
||||
const ACTIONS_ID_TOKEN_REQUEST_TOKEN: &str = "ACTIONS_ID_TOKEN_REQUEST_TOKEN";
|
||||
const OIDC_ALLOWED_CREDENTIAL_DIRS: &str = "LITELLM_OIDC_ALLOWED_CREDENTIAL_DIRS";
|
||||
const DEFAULT_CREDENTIAL_DIRS: &str = "/var/run/secrets,/run/secrets";
|
||||
|
||||
#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq, strum::EnumString, strum::AsRefStr)]
|
||||
#[strum(serialize_all = "snake_case")]
|
||||
pub enum OidcProvider {
|
||||
Google,
|
||||
#[strum(serialize = "circleci")]
|
||||
CircleCi,
|
||||
#[strum(serialize = "circleci_v2")]
|
||||
CircleCiV2,
|
||||
Github,
|
||||
Azure,
|
||||
File,
|
||||
Env,
|
||||
EnvPath,
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
|
||||
pub struct OidcReference<'a> {
|
||||
pub provider: OidcProvider,
|
||||
pub audience: &'a str,
|
||||
}
|
||||
|
||||
impl<'a> TryFrom<&'a str> for OidcReference<'a> {
|
||||
type Error = Error;
|
||||
|
||||
fn try_from(reference: &'a str) -> Result<Self, Error> {
|
||||
let (provider, audience) = reference
|
||||
.strip_prefix("oidc/")
|
||||
.and_then(|body| body.split_once('/'))
|
||||
.ok_or(Error::InvalidOidc)?;
|
||||
Ok(Self {
|
||||
provider: provider.parse().map_err(|_| Error::UnsupportedOidc)?,
|
||||
audience,
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
struct OidcTokenClaims {
|
||||
exp: Option<NumericDate>,
|
||||
}
|
||||
|
||||
#[derive(Deserialize)]
|
||||
#[serde(untagged)]
|
||||
enum NumericDate {
|
||||
Number(f64),
|
||||
String(String),
|
||||
}
|
||||
|
||||
impl NumericDate {
|
||||
fn seconds(self) -> Option<f64> {
|
||||
match self {
|
||||
Self::Number(value) => Some(value),
|
||||
Self::String(value) => value.parse().ok(),
|
||||
}
|
||||
.filter(|value| value.is_finite())
|
||||
}
|
||||
}
|
||||
|
||||
pub struct OidcResolver {
|
||||
client: reqwest::Client,
|
||||
google_identity_endpoint: reqwest::Url,
|
||||
cache: Cache<String, (SecretValue, SystemTime)>,
|
||||
clock: fn() -> SystemTime,
|
||||
}
|
||||
|
||||
impl Default for OidcResolver {
|
||||
fn default() -> Self {
|
||||
Self::new(
|
||||
reqwest::Client::builder().timeout(Duration::from_secs(600)).connect_timeout(Duration::from_secs(5)).build().expect("HTTP client configuration"),
|
||||
reqwest::Url::parse("http://metadata.google.internal/computeMetadata/v1/instance/service-accounts/default/identity").expect("static URL"),
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
impl OidcResolver {
|
||||
pub fn new(client: reqwest::Client, google_identity_endpoint: reqwest::Url) -> Self {
|
||||
Self {
|
||||
client,
|
||||
google_identity_endpoint,
|
||||
cache: Cache::builder()
|
||||
.max_capacity(200)
|
||||
.time_to_live(GOOGLE_TOKEN_MAX_TTL)
|
||||
.build(),
|
||||
clock: SystemTime::now,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn with_clock(self, clock: fn() -> SystemTime) -> Self {
|
||||
Self { clock, ..self }
|
||||
}
|
||||
|
||||
pub async fn resolve(
|
||||
&self,
|
||||
reference: &str,
|
||||
environment: &(dyn Lookup + Send + Sync),
|
||||
) -> Result<Option<SecretValue>, Error> {
|
||||
let OidcReference { provider, audience } = reference.try_into()?;
|
||||
match provider {
|
||||
OidcProvider::CircleCi => required_env(environment, CIRCLE_OIDC_TOKEN)
|
||||
.map(SecretValue::new)
|
||||
.map(Some),
|
||||
OidcProvider::CircleCiV2 => required_env(environment, CIRCLE_OIDC_TOKEN_V2)
|
||||
.map(SecretValue::new)
|
||||
.map(Some),
|
||||
OidcProvider::Env => required_env(environment, audience)
|
||||
.map(SecretValue::new)
|
||||
.map(Some),
|
||||
OidcProvider::EnvPath => read_file(&required_env(environment, audience)?)
|
||||
.await
|
||||
.map(Some),
|
||||
OidcProvider::File => read_allowed_file(audience, environment).await.map(Some),
|
||||
OidcProvider::Azure => {
|
||||
if let Some(path) = environment.get(AZURE_FEDERATED_TOKEN_FILE) {
|
||||
return read_file(&path).await.map(Some);
|
||||
}
|
||||
Err(Error::UnsupportedOidc)
|
||||
}
|
||||
OidcProvider::Github => {
|
||||
let url = required_env(environment, ACTIONS_ID_TOKEN_REQUEST_URL)?;
|
||||
let authorization = required_env(environment, ACTIONS_ID_TOKEN_REQUEST_TOKEN)?;
|
||||
if let Some(value) = self.cached(reference).await {
|
||||
return Ok(Some(value));
|
||||
}
|
||||
let response = self
|
||||
.client
|
||||
.get(url)
|
||||
.query(&[("audience", audience)])
|
||||
.bearer_auth(authorization)
|
||||
.header("Accept", "application/json; api-version=2.0")
|
||||
.send()
|
||||
.await
|
||||
.map_err(|_| Error::OidcHttp)?;
|
||||
if response.status() != reqwest::StatusCode::OK {
|
||||
return Err(Error::OidcStatus(response.status().as_u16()));
|
||||
}
|
||||
#[derive(Deserialize)]
|
||||
struct Token {
|
||||
value: Option<SecretValue>,
|
||||
}
|
||||
let token: Token = response.json().await.map_err(|_| Error::OidcResponse)?;
|
||||
if let Some(value) = &token.value {
|
||||
self.cache
|
||||
.insert(
|
||||
reference.to_owned(),
|
||||
(value.clone(), (self.clock)() + GITHUB_TOKEN_TTL),
|
||||
)
|
||||
.await;
|
||||
}
|
||||
Ok(token.value)
|
||||
}
|
||||
OidcProvider::Google => {
|
||||
if !cfg!(feature = "google") {
|
||||
return Err(Error::UnsupportedOidc);
|
||||
}
|
||||
if let Some(value) = self.cached(reference).await {
|
||||
return Ok(Some(value));
|
||||
}
|
||||
let response = self
|
||||
.client
|
||||
.get(self.google_identity_endpoint.clone())
|
||||
.query(&[("audience", audience)])
|
||||
.header("Metadata-Flavor", "Google")
|
||||
.send()
|
||||
.await
|
||||
.map_err(|_| Error::OidcHttp)?;
|
||||
if response.status() != reqwest::StatusCode::OK {
|
||||
return Err(Error::OidcStatus(response.status().as_u16()));
|
||||
}
|
||||
let token = response.text().await.map_err(|_| Error::OidcResponse)?;
|
||||
let now = (self.clock)();
|
||||
let ttl = oidc_token_cache_ttl(&token, now, GOOGLE_TOKEN_MAX_TTL);
|
||||
let value = SecretValue::new(token);
|
||||
if let Some(ttl) = ttl.filter(|ttl| !ttl.is_zero()) {
|
||||
self.cache
|
||||
.insert(reference.to_owned(), (value.clone(), now + ttl))
|
||||
.await;
|
||||
}
|
||||
Ok(Some(value))
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
async fn cached(&self, reference: &str) -> Option<SecretValue> {
|
||||
self.cache
|
||||
.get(reference)
|
||||
.await
|
||||
.and_then(|(value, expires)| ((self.clock)() < expires).then_some(value))
|
||||
}
|
||||
}
|
||||
|
||||
fn required_env(environment: &dyn Lookup, name: &str) -> Result<String, Error> {
|
||||
environment.get(name).ok_or(Error::MissingEnvironment)
|
||||
}
|
||||
|
||||
async fn read_file(path: &str) -> Result<SecretValue, Error> {
|
||||
tokio::fs::read_to_string(path)
|
||||
.await
|
||||
.map(|value| SecretValue::new(value.replace("\r\n", "\n").replace('\r', "\n")))
|
||||
.map_err(|_| Error::OidcFile)
|
||||
}
|
||||
|
||||
async fn read_allowed_file(
|
||||
path: &str,
|
||||
environment: &(dyn Lookup + Sync),
|
||||
) -> Result<SecretValue, Error> {
|
||||
if !Path::new(path).is_absolute() {
|
||||
return Err(Error::UnsafeOidcPath);
|
||||
}
|
||||
let resolved = tokio::fs::canonicalize(path)
|
||||
.await
|
||||
.map_err(|_| Error::OidcFile)?;
|
||||
let allowed = environment
|
||||
.get(OIDC_ALLOWED_CREDENTIAL_DIRS)
|
||||
.filter(|v| !v.is_empty())
|
||||
.unwrap_or_else(|| DEFAULT_CREDENTIAL_DIRS.into());
|
||||
for directory in allowed.split(',').map(str::trim).filter(|d| !d.is_empty()) {
|
||||
if let Ok(directory) = tokio::fs::canonicalize(directory).await
|
||||
&& resolved.starts_with(directory)
|
||||
{
|
||||
return tokio::fs::read_to_string(&resolved)
|
||||
.await
|
||||
.map(|value| SecretValue::new(value.replace("\r\n", "\n").replace('\r', "\n")))
|
||||
.map_err(|_| Error::OidcFile);
|
||||
}
|
||||
}
|
||||
Err(Error::UnsafeOidcPath)
|
||||
}
|
||||
|
||||
fn oidc_token_cache_ttl(token: &str, now: SystemTime, max_ttl: Duration) -> Option<Duration> {
|
||||
let fallback = Some(max_ttl);
|
||||
let Ok(claims) = insecure_decode_claims::<OidcTokenClaims>(token) else {
|
||||
return fallback;
|
||||
};
|
||||
let Some(exp) = claims.exp.and_then(NumericDate::seconds) else {
|
||||
return fallback;
|
||||
};
|
||||
let seconds = exp.trunc()
|
||||
- now.duration_since(UNIX_EPOCH).ok()?.as_secs() as f64
|
||||
- TOKEN_EXPIRY_MARGIN_SECONDS;
|
||||
(seconds > 0.0).then(|| Duration::from_secs_f64(seconds.min(max_ttl.as_secs_f64())))
|
||||
}
|
||||
140
litellm-rust/crates/secrets/src/resolver.rs
Normal file
140
litellm-rust/crates/secrets/src/resolver.rs
Normal file
|
|
@ -0,0 +1,140 @@
|
|||
use std::sync::Arc;
|
||||
|
||||
use litellm_core_utils::settings::{Lookup, ProcessEnvironment};
|
||||
|
||||
use crate::{Error, OidcResolver, Secret, SecretManagerState, SecretValue};
|
||||
|
||||
use crate::state::{LookupTarget, normalize_secret_name};
|
||||
|
||||
pub struct SecretResolver {
|
||||
state: Arc<SecretManagerState>,
|
||||
environment: Arc<dyn Lookup + Send + Sync>,
|
||||
oidc: OidcResolver,
|
||||
}
|
||||
|
||||
impl Default for SecretResolver {
|
||||
fn default() -> Self {
|
||||
Self::new(
|
||||
Arc::new(SecretManagerState::default()),
|
||||
Arc::new(ProcessEnvironment),
|
||||
OidcResolver::default(),
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
impl SecretResolver {
|
||||
pub fn new(
|
||||
state: Arc<SecretManagerState>,
|
||||
environment: Arc<dyn Lookup + Send + Sync>,
|
||||
oidc: OidcResolver,
|
||||
) -> Self {
|
||||
Self {
|
||||
state,
|
||||
environment,
|
||||
oidc,
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn get_secret(
|
||||
&self,
|
||||
name: &str,
|
||||
_default_value: Option<Secret>,
|
||||
) -> Result<Option<Secret>, Error> {
|
||||
let name = normalize_secret_name(name);
|
||||
if name.starts_with("oidc/") {
|
||||
return self
|
||||
.oidc
|
||||
.resolve(name, self.environment.as_ref())
|
||||
.await
|
||||
.map(|value| value.map(Secret::String));
|
||||
}
|
||||
if !self.state.readable() {
|
||||
return Ok(self
|
||||
.environment
|
||||
.get(name)
|
||||
.map(|value| match str_to_bool(&value) {
|
||||
Some(value) => Secret::Bool(value),
|
||||
None => Secret::String(SecretValue::new(value)),
|
||||
}));
|
||||
}
|
||||
let result = match self.state.lookup_target(name) {
|
||||
LookupTarget::Environment => Ok(self.environment_secret(name)),
|
||||
LookupTarget::Manager { backend, settings } => {
|
||||
crate::get_secret_from_manager(backend, name, settings, self.environment.as_ref())
|
||||
.await
|
||||
}
|
||||
};
|
||||
let value = match result {
|
||||
Ok(value) => value,
|
||||
Err(_) => {
|
||||
tracing::error!("secret manager lookup failed; falling back to environment");
|
||||
self.environment_secret(name)
|
||||
}
|
||||
};
|
||||
Ok(value.and_then(managed_secret))
|
||||
}
|
||||
|
||||
fn environment_secret(&self, name: &str) -> Option<Secret> {
|
||||
self.environment
|
||||
.get(name)
|
||||
.map(SecretValue::new)
|
||||
.map(Secret::String)
|
||||
}
|
||||
|
||||
pub async fn get_secret_str(
|
||||
&self,
|
||||
name: &str,
|
||||
default_value: Option<Secret>,
|
||||
) -> Result<Option<SecretValue>, Error> {
|
||||
Ok(match self.get_secret(name, default_value).await? {
|
||||
Some(Secret::String(value)) => Some(value),
|
||||
Some(Secret::Bool(_) | Secret::Json(_)) | None => None,
|
||||
})
|
||||
}
|
||||
|
||||
pub async fn get_secret_bool(
|
||||
&self,
|
||||
name: &str,
|
||||
default_value: Option<bool>,
|
||||
) -> Result<Option<bool>, Error> {
|
||||
Ok(
|
||||
match self
|
||||
.get_secret(name, default_value.map(Secret::Bool))
|
||||
.await?
|
||||
{
|
||||
Some(Secret::Bool(value)) => Some(value),
|
||||
Some(Secret::String(value)) => str_to_bool(value.expose()),
|
||||
Some(Secret::Json(_)) | None => None,
|
||||
},
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
fn str_to_bool(value: &str) -> Option<bool> {
|
||||
match value.trim().to_ascii_lowercase().as_str() {
|
||||
"true" => Some(true),
|
||||
"false" => Some(false),
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
fn literal_bool(value: &str) -> Option<bool> {
|
||||
use rustpython_parser::{Parse, ast};
|
||||
match ast::Expr::parse(value.trim_start_matches([' ', '\t']), "<secret>").ok()? {
|
||||
ast::Expr::Constant(node) => match node.value {
|
||||
ast::Constant::Bool(value) => Some(value),
|
||||
_ => None,
|
||||
},
|
||||
_ => None,
|
||||
}
|
||||
}
|
||||
|
||||
fn managed_secret(value: Secret) -> Option<Secret> {
|
||||
match value {
|
||||
Secret::String(value) => Some(match literal_bool(value.expose()) {
|
||||
Some(boolean) => Secret::Bool(boolean),
|
||||
None => Secret::String(value),
|
||||
}),
|
||||
Secret::Bool(_) | Secret::Json(_) => None,
|
||||
}
|
||||
}
|
||||
105
litellm-rust/crates/secrets/src/state.rs
Normal file
105
litellm-rust/crates/secrets/src/state.rs
Normal file
|
|
@ -0,0 +1,105 @@
|
|||
use crate::{Error, KeyManagementSettings, KeyManagementSystem, SecretManager};
|
||||
|
||||
pub(crate) enum LookupTarget<'a> {
|
||||
Environment,
|
||||
Manager {
|
||||
backend: &'a SecretManager,
|
||||
settings: &'a KeyManagementSettings,
|
||||
},
|
||||
}
|
||||
|
||||
pub(crate) fn normalize_secret_name(name: &str) -> &str {
|
||||
name.strip_prefix("os.environ/").unwrap_or(name)
|
||||
}
|
||||
|
||||
#[derive(Clone, Default)]
|
||||
pub struct SecretManagerState {
|
||||
system: Option<KeyManagementSystem>,
|
||||
settings: Option<KeyManagementSettings>,
|
||||
backend: Option<SecretManager>,
|
||||
}
|
||||
|
||||
impl SecretManagerState {
|
||||
pub fn new(
|
||||
system: Option<KeyManagementSystem>,
|
||||
settings: Option<KeyManagementSettings>,
|
||||
backend: Option<SecretManager>,
|
||||
) -> Result<Self, Error> {
|
||||
if let Some(system) = system {
|
||||
let available = match system {
|
||||
KeyManagementSystem::Local => true,
|
||||
KeyManagementSystem::AwsKms | KeyManagementSystem::AwsSecretManager => {
|
||||
cfg!(feature = "aws")
|
||||
}
|
||||
KeyManagementSystem::GoogleKms | KeyManagementSystem::GoogleSecretManager => {
|
||||
cfg!(feature = "google")
|
||||
}
|
||||
KeyManagementSystem::AzureKeyVault
|
||||
| KeyManagementSystem::HashicorpVault
|
||||
| KeyManagementSystem::Cyberark
|
||||
| KeyManagementSystem::Custom => false,
|
||||
};
|
||||
if !available {
|
||||
return Err(Error::UnsupportedBackend(system));
|
||||
}
|
||||
if let Some(backend) = &backend
|
||||
&& system != backend.system()
|
||||
{
|
||||
return Err(Error::BackendMismatch);
|
||||
}
|
||||
}
|
||||
Ok(Self {
|
||||
system,
|
||||
settings,
|
||||
backend,
|
||||
})
|
||||
}
|
||||
|
||||
pub fn system(&self) -> Option<KeyManagementSystem> {
|
||||
self.system
|
||||
}
|
||||
pub fn settings(&self) -> Option<&KeyManagementSettings> {
|
||||
self.settings.as_ref()
|
||||
}
|
||||
pub fn backend(&self) -> Option<&SecretManager> {
|
||||
self.backend.as_ref()
|
||||
}
|
||||
|
||||
pub(crate) fn readable(&self) -> bool {
|
||||
self.backend.is_some()
|
||||
&& self
|
||||
.settings
|
||||
.as_ref()
|
||||
.is_some_and(|settings| settings.access_mode.readable())
|
||||
}
|
||||
|
||||
pub(crate) fn lookup_target(&self, name: &str) -> LookupTarget<'_> {
|
||||
match (&self.backend, &self.settings) {
|
||||
(Some(backend), Some(settings))
|
||||
if settings.access_mode.readable()
|
||||
&& hosts_secret(settings, name)
|
||||
&& self
|
||||
.system
|
||||
.is_some_and(|system| system != KeyManagementSystem::Local) =>
|
||||
{
|
||||
LookupTarget::Manager { backend, settings }
|
||||
}
|
||||
_ => LookupTarget::Environment,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
pub fn secret_manager_would_be_consulted(state: &SecretManagerState, name: &str) -> bool {
|
||||
state.readable()
|
||||
&& state
|
||||
.settings
|
||||
.as_ref()
|
||||
.is_some_and(|settings| hosts_secret(settings, normalize_secret_name(name)))
|
||||
}
|
||||
|
||||
fn hosts_secret(settings: &KeyManagementSettings, name: &str) -> bool {
|
||||
settings
|
||||
.hosted_keys
|
||||
.as_ref()
|
||||
.is_none_or(|keys| keys.iter().any(|key| key == name))
|
||||
}
|
||||
107
litellm-rust/crates/secrets/tests/handler.rs
Normal file
107
litellm-rust/crates/secrets/tests/handler.rs
Normal file
|
|
@ -0,0 +1,107 @@
|
|||
#[cfg(feature = "aws")]
|
||||
#[tokio::test]
|
||||
async fn aws_handler_reads_ciphertext_decodes_trims_and_redacts() {
|
||||
use aws_sdk_kms::{
|
||||
Client,
|
||||
config::{BehaviorVersion, Credentials, Region},
|
||||
};
|
||||
use base64::{Engine, engine::general_purpose::STANDARD};
|
||||
use litellm_secrets::{
|
||||
Error, KeyManagementSettings, SecretManager, aws::AwsKms, get_secret_from_manager,
|
||||
};
|
||||
use wiremock::{Mock, MockServer, ResponseTemplate, matchers::body_json};
|
||||
|
||||
let server = MockServer::start().await;
|
||||
Mock::given(body_json(
|
||||
serde_json::json!({"CiphertextBlob": STANDARD.encode("encrypted")}),
|
||||
))
|
||||
.respond_with(
|
||||
ResponseTemplate::new(200)
|
||||
.set_body_json(serde_json::json!({"Plaintext":STANDARD.encode(" value\n")})),
|
||||
)
|
||||
.expect(1)
|
||||
.mount(&server)
|
||||
.await;
|
||||
let client = Client::from_conf(
|
||||
aws_sdk_kms::Config::builder()
|
||||
.behavior_version(BehaviorVersion::latest())
|
||||
.region(Region::new("us-east-1"))
|
||||
.credentials_provider(Credentials::new("test", "test", None, None, "test"))
|
||||
.endpoint_url(server.uri())
|
||||
.build(),
|
||||
);
|
||||
let manager = SecretManager::AwsKms(AwsKms::new(client));
|
||||
let settings = KeyManagementSettings::default();
|
||||
let value = get_secret_from_manager(&manager, "KEY", &settings, &|name: &str| {
|
||||
assert_eq!(name, "KEY");
|
||||
Some(format!(" {}\n", STANDARD.encode("encrypted")))
|
||||
})
|
||||
.await
|
||||
.unwrap()
|
||||
.unwrap();
|
||||
assert_eq!(value.as_str(), Some("value"));
|
||||
assert!(!format!("{value:?}").contains("value"));
|
||||
assert!(matches!(
|
||||
get_secret_from_manager(&manager, "KEY", &settings, &|_: &str| None).await,
|
||||
Err(Error::MissingCiphertext)
|
||||
));
|
||||
assert!(matches!(
|
||||
get_secret_from_manager(&manager, "KEY", &settings, &|_: &str| Some("abc".into())).await,
|
||||
Err(Error::InvalidCiphertext)
|
||||
));
|
||||
}
|
||||
|
||||
#[cfg(feature = "google")]
|
||||
#[tokio::test]
|
||||
async fn google_handler_requires_canonical_base64_and_preserves_plaintext_whitespace() {
|
||||
use base64::{Engine, engine::general_purpose::STANDARD};
|
||||
use google_cloud_kms_v1::client::KeyManagementService;
|
||||
use litellm_secrets::{
|
||||
Error, KeyManagementSettings, SecretManager, get_secret_from_manager, google::GoogleKms,
|
||||
};
|
||||
use wiremock::{
|
||||
Mock, MockServer, ResponseTemplate,
|
||||
matchers::{body_json, path},
|
||||
};
|
||||
|
||||
let server = MockServer::start().await;
|
||||
let resource = "projects/project/locations/global/keyRings/ring/cryptoKeys/key";
|
||||
Mock::given(path(format!("/v1/{resource}:decrypt")))
|
||||
.and(body_json(
|
||||
serde_json::json!({"ciphertext":STANDARD.encode("encrypted")}),
|
||||
))
|
||||
.respond_with(
|
||||
ResponseTemplate::new(200)
|
||||
.set_body_json(serde_json::json!({"plaintext":STANDARD.encode(" value\n")})),
|
||||
)
|
||||
.expect(1)
|
||||
.mount(&server)
|
||||
.await;
|
||||
let client = KeyManagementService::builder()
|
||||
.with_endpoint(server.uri())
|
||||
.with_credentials(google_cloud_auth::credentials::anonymous::Builder::new().build())
|
||||
.build()
|
||||
.await
|
||||
.unwrap();
|
||||
let manager = SecretManager::GoogleKms(GoogleKms::new(client, resource.into()));
|
||||
let settings = KeyManagementSettings::default();
|
||||
let value = get_secret_from_manager(&manager, "KEY", &settings, &|_: &str| {
|
||||
Some(STANDARD.encode("encrypted"))
|
||||
})
|
||||
.await
|
||||
.unwrap()
|
||||
.unwrap();
|
||||
assert_eq!(value.as_str(), Some(" value\n"));
|
||||
assert!(matches!(
|
||||
get_secret_from_manager(&manager, "KEY", &settings, &|_: &str| Some(format!(
|
||||
" {}",
|
||||
STANDARD.encode("encrypted")
|
||||
)))
|
||||
.await,
|
||||
Err(Error::InvalidCiphertext)
|
||||
));
|
||||
assert!(matches!(
|
||||
get_secret_from_manager(&manager, "KEY", &settings, &|_: &str| None).await,
|
||||
Err(Error::MissingCiphertext)
|
||||
));
|
||||
}
|
||||
295
litellm-rust/crates/secrets/tests/oidc.rs
Normal file
295
litellm-rust/crates/secrets/tests/oidc.rs
Normal file
|
|
@ -0,0 +1,295 @@
|
|||
use std::{collections::BTreeMap, sync::Arc};
|
||||
|
||||
use litellm_core_utils::settings::Lookup;
|
||||
use litellm_secrets::{Error, OidcResolver, Secret, SecretManagerState, SecretResolver};
|
||||
use wiremock::{
|
||||
Mock, MockServer, ResponseTemplate,
|
||||
matchers::{header, method, path, query_param},
|
||||
};
|
||||
|
||||
fn environment(pairs: &[(&str, &str)]) -> Arc<dyn Lookup + Send + Sync> {
|
||||
let values: BTreeMap<String, String> = pairs
|
||||
.iter()
|
||||
.map(|(k, v)| (k.to_string(), v.to_string()))
|
||||
.collect();
|
||||
Arc::new(move |name: &str| values.get(name).cloned())
|
||||
}
|
||||
|
||||
#[rstest::rstest]
|
||||
#[case::environment("oidc/env/TOKEN", "true")]
|
||||
#[case::circleci("oidc/circleci/audience", "circle")]
|
||||
#[case::circleci_v2("oidc/circleci_v2/audience", "circle-v2")]
|
||||
#[tokio::test]
|
||||
async fn environment_sources_resolve_expected_value(
|
||||
#[case] reference: &str,
|
||||
#[case] expected: &str,
|
||||
) {
|
||||
let env = environment(&[
|
||||
("TOKEN", "true"),
|
||||
("CIRCLE_OIDC_TOKEN", "circle"),
|
||||
("CIRCLE_OIDC_TOKEN_V2", "circle-v2"),
|
||||
]);
|
||||
assert_eq!(
|
||||
OidcResolver::default()
|
||||
.resolve(reference, env.as_ref())
|
||||
.await
|
||||
.unwrap()
|
||||
.unwrap()
|
||||
.expose(),
|
||||
expected
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn environment_sources_bypass_boolean_conversion_and_defaults() {
|
||||
let env = environment(&[("TOKEN", "true")]);
|
||||
let oidc = OidcResolver::default();
|
||||
let resolver = SecretResolver::new(Arc::new(SecretManagerState::default()), env, oidc);
|
||||
assert_eq!(
|
||||
resolver
|
||||
.get_secret_str("os.environ/oidc/env/TOKEN", None)
|
||||
.await
|
||||
.unwrap()
|
||||
.unwrap()
|
||||
.expose(),
|
||||
"true"
|
||||
);
|
||||
assert_eq!(
|
||||
resolver
|
||||
.get_secret_bool("oidc/env/TOKEN", None)
|
||||
.await
|
||||
.unwrap(),
|
||||
Some(true)
|
||||
);
|
||||
assert!(matches!(
|
||||
resolver
|
||||
.get_secret("oidc/env/MISSING", Some(Secret::Bool(true)))
|
||||
.await,
|
||||
Err(Error::MissingEnvironment)
|
||||
));
|
||||
assert!(matches!(
|
||||
resolver.get_secret("oidc/invalid", None).await,
|
||||
Err(Error::InvalidOidc)
|
||||
));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn github_requests_are_authenticated_cached_and_revalidate_environment() {
|
||||
let server = MockServer::start().await;
|
||||
Mock::given(method("GET"))
|
||||
.and(path("/token"))
|
||||
.and(query_param("audience", "https://service/oidc/path"))
|
||||
.and(header("authorization", "Bearer request-token"))
|
||||
.and(header("accept", "application/json; api-version=2.0"))
|
||||
.respond_with(
|
||||
ResponseTemplate::new(200).set_body_json(serde_json::json!({"value":"identity-token"})),
|
||||
)
|
||||
.expect(1)
|
||||
.mount(&server)
|
||||
.await;
|
||||
let env = environment(&[
|
||||
(
|
||||
"ACTIONS_ID_TOKEN_REQUEST_URL",
|
||||
&format!("{}/token", server.uri()),
|
||||
),
|
||||
("ACTIONS_ID_TOKEN_REQUEST_TOKEN", "request-token"),
|
||||
]);
|
||||
let oidc = OidcResolver::default();
|
||||
for _ in 0..2 {
|
||||
assert_eq!(
|
||||
oidc.resolve("oidc/github/https://service/oidc/path", env.as_ref())
|
||||
.await
|
||||
.unwrap()
|
||||
.unwrap()
|
||||
.expose(),
|
||||
"identity-token"
|
||||
);
|
||||
}
|
||||
assert!(matches!(
|
||||
oidc.resolve(
|
||||
"oidc/github/https://service/oidc/path",
|
||||
environment(&[]).as_ref()
|
||||
)
|
||||
.await,
|
||||
Err(Error::MissingEnvironment)
|
||||
));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn file_allowlist_resolves_symlinks_while_environment_paths_remain_explicit() {
|
||||
let allowed = tempfile::tempdir().unwrap();
|
||||
let outside = tempfile::tempdir().unwrap();
|
||||
let token = allowed.path().join("token");
|
||||
let private = outside.path().join("private");
|
||||
std::fs::write(&token, "token\r\n").unwrap();
|
||||
std::fs::write(&private, "outside").unwrap();
|
||||
let env = environment(&[
|
||||
(
|
||||
"LITELLM_OIDC_ALLOWED_CREDENTIAL_DIRS",
|
||||
allowed.path().to_str().unwrap(),
|
||||
),
|
||||
("PATH_TOKEN", private.to_str().unwrap()),
|
||||
("AZURE_FEDERATED_TOKEN_FILE", token.to_str().unwrap()),
|
||||
]);
|
||||
let oidc = OidcResolver::default();
|
||||
assert_eq!(
|
||||
oidc.resolve(&format!("oidc/file/{}", token.display()), env.as_ref())
|
||||
.await
|
||||
.unwrap()
|
||||
.unwrap()
|
||||
.expose(),
|
||||
"token\n"
|
||||
);
|
||||
assert!(matches!(
|
||||
oidc.resolve("oidc/file/relative", env.as_ref()).await,
|
||||
Err(Error::UnsafeOidcPath)
|
||||
));
|
||||
assert!(matches!(
|
||||
oidc.resolve(&format!("oidc/file/{}", private.display()), env.as_ref())
|
||||
.await,
|
||||
Err(Error::UnsafeOidcPath)
|
||||
));
|
||||
assert_eq!(
|
||||
oidc.resolve("oidc/env_path/PATH_TOKEN", env.as_ref())
|
||||
.await
|
||||
.unwrap()
|
||||
.unwrap()
|
||||
.expose(),
|
||||
"outside"
|
||||
);
|
||||
assert_eq!(
|
||||
oidc.resolve("oidc/azure/scope", env.as_ref())
|
||||
.await
|
||||
.unwrap()
|
||||
.unwrap()
|
||||
.expose(),
|
||||
"token\n"
|
||||
);
|
||||
#[cfg(unix)]
|
||||
{
|
||||
let link = allowed.path().join("link");
|
||||
std::os::unix::fs::symlink(&private, &link).unwrap();
|
||||
assert!(matches!(
|
||||
oidc.resolve(&format!("oidc/file/{}", link.display()), env.as_ref())
|
||||
.await,
|
||||
Err(Error::UnsafeOidcPath)
|
||||
));
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(feature = "google")]
|
||||
#[rstest::rstest]
|
||||
#[case::at_refresh_boundary(serde_json::json!(1060), 2)]
|
||||
#[case::beyond_refresh_boundary(serde_json::json!(1061), 1)]
|
||||
#[case::already_expired(serde_json::json!(999), 2)]
|
||||
#[case::string_expiry(serde_json::json!("999"), 2)]
|
||||
#[case::fractional_expiry(serde_json::json!(1060.9), 2)]
|
||||
#[case::negative_expiry(serde_json::json!(-1), 2)]
|
||||
#[case::null_expiry(serde_json::Value::Null, 1)]
|
||||
#[case::unreadable_expiry(serde_json::json!("invalid"), 1)]
|
||||
#[case::nonfinite_expiry(serde_json::json!("NaN"), 1)]
|
||||
#[tokio::test]
|
||||
async fn google_expiry_caps_cache_and_preserves_audience(
|
||||
#[case] expiry: serde_json::Value,
|
||||
#[case] calls: u64,
|
||||
) {
|
||||
use base64::{Engine, engine::general_purpose::URL_SAFE_NO_PAD};
|
||||
use std::time::{Duration, SystemTime, UNIX_EPOCH};
|
||||
fn now() -> SystemTime {
|
||||
UNIX_EPOCH + Duration::from_secs(1000)
|
||||
}
|
||||
let server = MockServer::start().await;
|
||||
let token = format!(
|
||||
"{}.{}.signature",
|
||||
URL_SAFE_NO_PAD.encode(serde_json::json!({"alg":"RS256","typ":"JWT"}).to_string()),
|
||||
URL_SAFE_NO_PAD.encode(serde_json::json!({"exp":expiry}).to_string())
|
||||
);
|
||||
Mock::given(method("GET"))
|
||||
.and(header("metadata-flavor", "Google"))
|
||||
.and(query_param("audience", "https://service/oidc/path"))
|
||||
.respond_with(ResponseTemplate::new(200).set_body_string(&token))
|
||||
.expect(calls)
|
||||
.mount(&server)
|
||||
.await;
|
||||
let oidc =
|
||||
OidcResolver::new(reqwest::Client::new(), server.uri().parse().unwrap()).with_clock(now);
|
||||
for _ in 0..2 {
|
||||
assert_eq!(
|
||||
oidc.resolve(
|
||||
"oidc/google/https://service/oidc/path",
|
||||
environment(&[]).as_ref()
|
||||
)
|
||||
.await
|
||||
.unwrap()
|
||||
.unwrap()
|
||||
.expose(),
|
||||
token
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(not(feature = "google"))]
|
||||
#[tokio::test]
|
||||
async fn google_oidc_requires_its_build_feature() {
|
||||
assert!(matches!(
|
||||
OidcResolver::default()
|
||||
.resolve("oidc/google/audience", environment(&[]).as_ref())
|
||||
.await,
|
||||
Err(Error::UnsupportedOidc)
|
||||
));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn azure_oidc_without_a_token_file_requires_an_unimplemented_backend() {
|
||||
assert!(matches!(
|
||||
OidcResolver::default()
|
||||
.resolve("oidc/azure/scope", environment(&[]).as_ref())
|
||||
.await,
|
||||
Err(Error::UnsupportedOidc)
|
||||
));
|
||||
}
|
||||
|
||||
#[rstest::rstest]
|
||||
#[case::missing_prefix("env/TOKEN", false)]
|
||||
#[case::missing_audience_separator("oidc/env", false)]
|
||||
#[case::unknown_provider("oidc/unknown/TOKEN", true)]
|
||||
#[tokio::test]
|
||||
async fn invalid_references_fail_before_environment_lookup(
|
||||
#[case] reference: &str,
|
||||
#[case] unsupported: bool,
|
||||
) {
|
||||
let error = OidcResolver::default()
|
||||
.resolve(reference, &|_: &str| {
|
||||
panic!("invalid reference reached environment lookup")
|
||||
})
|
||||
.await
|
||||
.unwrap_err();
|
||||
assert!(matches!(error, Error::UnsupportedOidc) == unsupported);
|
||||
assert!(matches!(error, Error::InvalidOidc) != unsupported);
|
||||
}
|
||||
|
||||
#[cfg(feature = "google")]
|
||||
#[rstest::rstest]
|
||||
#[case::opaque("opaque-token")]
|
||||
#[case::missing_expiry("header.e30.signature")]
|
||||
#[tokio::test]
|
||||
async fn unreadable_expiry_keeps_python_cache_fallback(#[case] token: &str) {
|
||||
let server = MockServer::start().await;
|
||||
Mock::given(method("GET"))
|
||||
.respond_with(ResponseTemplate::new(200).set_body_string(token))
|
||||
.expect(1)
|
||||
.mount(&server)
|
||||
.await;
|
||||
let resolver = OidcResolver::new(reqwest::Client::new(), server.uri().parse().unwrap());
|
||||
for _ in 0..2 {
|
||||
assert_eq!(
|
||||
resolver
|
||||
.resolve("oidc/google/audience", environment(&[]).as_ref())
|
||||
.await
|
||||
.unwrap()
|
||||
.unwrap()
|
||||
.expose(),
|
||||
token,
|
||||
);
|
||||
}
|
||||
}
|
||||
343
litellm-rust/crates/secrets/tests/resolution.rs
Normal file
343
litellm-rust/crates/secrets/tests/resolution.rs
Normal file
|
|
@ -0,0 +1,343 @@
|
|||
use std::sync::Arc;
|
||||
|
||||
use litellm_secrets::{
|
||||
AccessMode, KeyManagementSettings, KeyManagementSystem, OidcResolver, Secret, SecretManager,
|
||||
SecretManagerState, SecretResolver, SecretValue, secret_manager_would_be_consulted,
|
||||
};
|
||||
|
||||
fn resolver(value: Option<&str>, readable: bool) -> SecretResolver {
|
||||
let state = if readable {
|
||||
SecretManagerState::new(
|
||||
Some(KeyManagementSystem::Local),
|
||||
Some(KeyManagementSettings::default()),
|
||||
Some(SecretManager::Local),
|
||||
)
|
||||
.unwrap()
|
||||
} else {
|
||||
SecretManagerState::default()
|
||||
};
|
||||
let value = value.map(str::to_owned);
|
||||
SecretResolver::new(
|
||||
Arc::new(state),
|
||||
Arc::new(move |_: &str| value.clone()),
|
||||
OidcResolver::default(),
|
||||
)
|
||||
}
|
||||
|
||||
#[rstest::rstest]
|
||||
#[case::lowercase_true("true", Some(true), None)]
|
||||
#[case::whitespace_lowercase_false(" FALSE ", Some(false), None)]
|
||||
#[case::python_true("True", Some(true), Some(true))]
|
||||
#[case::python_false("False", Some(false), Some(false))]
|
||||
#[case::parenthesized_python_true("(True)", None, Some(true))]
|
||||
#[case::commented_python_false("False # comment", None, Some(false))]
|
||||
#[case::integer("1", None, None)]
|
||||
#[case::yes("yes", None, None)]
|
||||
#[case::plain_string("secret", None, None)]
|
||||
#[tokio::test]
|
||||
async fn boolean_conversion_preserves_local_and_manager_differences(
|
||||
#[case] input: &str,
|
||||
#[case] local: Option<bool>,
|
||||
#[case] manager: Option<bool>,
|
||||
#[values(false, true)] readable: bool,
|
||||
) {
|
||||
let boolean = if readable { manager } else { local };
|
||||
let resolver = resolver(Some(input), readable);
|
||||
let expected = boolean
|
||||
.map(Secret::Bool)
|
||||
.unwrap_or_else(|| Secret::String(SecretValue::new(input)));
|
||||
assert_eq!(
|
||||
resolver.get_secret("key", None).await.unwrap(),
|
||||
Some(expected)
|
||||
);
|
||||
assert_eq!(
|
||||
resolver
|
||||
.get_secret_str("key", None)
|
||||
.await
|
||||
.unwrap()
|
||||
.map(|v| v.expose().to_owned()),
|
||||
boolean.is_none().then(|| input.to_owned())
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn manager_boolean_conversion_trims_whitespace() {
|
||||
assert_eq!(
|
||||
resolver(Some(" true "), true)
|
||||
.get_secret_bool("key", None)
|
||||
.await
|
||||
.unwrap(),
|
||||
Some(true)
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn missing_values_ignore_defaults_and_prefix_is_removed_before_lookup() {
|
||||
let missing = resolver(None, false);
|
||||
assert_eq!(
|
||||
missing
|
||||
.get_secret("missing", Some(Secret::Bool(true)))
|
||||
.await
|
||||
.unwrap(),
|
||||
None
|
||||
);
|
||||
assert_eq!(
|
||||
missing
|
||||
.get_secret_bool("missing", Some(true))
|
||||
.await
|
||||
.unwrap(),
|
||||
None
|
||||
);
|
||||
let resolver = SecretResolver::new(
|
||||
Arc::new(SecretManagerState::default()),
|
||||
Arc::new(|name: &str| (name == "KEY").then(|| "value".into())),
|
||||
OidcResolver::default(),
|
||||
);
|
||||
assert_eq!(
|
||||
resolver
|
||||
.get_secret_str("os.environ/KEY", None)
|
||||
.await
|
||||
.unwrap()
|
||||
.unwrap()
|
||||
.expose(),
|
||||
"value"
|
||||
);
|
||||
}
|
||||
|
||||
#[rstest::rstest]
|
||||
#[case::all_keys(None)]
|
||||
#[case::no_keys(Some(Vec::new()))]
|
||||
#[case::allowlisted_key(Some(vec!["KEY".into()]))]
|
||||
fn manager_gating_requires_client_readable_settings_and_allowlisted_name(
|
||||
#[values(AccessMode::ReadOnly, AccessMode::WriteOnly, AccessMode::ReadAndWrite)]
|
||||
access_mode: AccessMode,
|
||||
#[values(false, true)] client: bool,
|
||||
#[case] keys: Option<Vec<String>>,
|
||||
) {
|
||||
let expected = client
|
||||
&& access_mode.readable()
|
||||
&& keys
|
||||
.as_ref()
|
||||
.is_none_or(|keys| keys.iter().any(|key| key == "KEY"));
|
||||
let state = SecretManagerState::new(
|
||||
Some(KeyManagementSystem::Local),
|
||||
Some(KeyManagementSettings {
|
||||
access_mode,
|
||||
hosted_keys: keys,
|
||||
..Default::default()
|
||||
}),
|
||||
client.then_some(SecretManager::Local),
|
||||
)
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
secret_manager_would_be_consulted(&state, "os.environ/KEY"),
|
||||
expected
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn manager_gating_requires_settings() {
|
||||
let no_settings = SecretManagerState::new(None, None, Some(SecretManager::Local)).unwrap();
|
||||
assert!(!secret_manager_would_be_consulted(&no_settings, "KEY"));
|
||||
}
|
||||
|
||||
#[cfg(feature = "aws")]
|
||||
#[rstest::rstest]
|
||||
#[case::missing_value(None, None)]
|
||||
#[case::lookup_error(Some("primary".to_owned()), Some("environment-value"))]
|
||||
#[tokio::test]
|
||||
async fn aws_missing_values_do_not_fallback_but_lookup_errors_do(
|
||||
#[case] primary: Option<String>,
|
||||
#[case] expected: Option<&str>,
|
||||
) {
|
||||
use litellm_secrets::aws::AwsSecretsManagerV2;
|
||||
use wiremock::{Mock, MockServer, ResponseTemplate, matchers::body_partial_json};
|
||||
let server = MockServer::start().await;
|
||||
Mock::given(body_partial_json(serde_json::json!({"SecretId":"KEY"})))
|
||||
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({})))
|
||||
.expect(u64::from(primary.is_none()))
|
||||
.mount(&server)
|
||||
.await;
|
||||
Mock::given(body_partial_json(serde_json::json!({"SecretId":"primary"})))
|
||||
.respond_with(
|
||||
ResponseTemplate::new(200)
|
||||
.set_body_json(serde_json::json!({"SecretString":"invalid-json"})),
|
||||
)
|
||||
.expect(u64::from(primary.is_some()))
|
||||
.mount(&server)
|
||||
.await;
|
||||
let endpoint = server.uri();
|
||||
let environment: Arc<dyn litellm_core_utils::settings::Lookup + Send + Sync> =
|
||||
Arc::new(move |name: &str| match name {
|
||||
"AWS_REGION_NAME" => Some("us-east-1".into()),
|
||||
"AWS_ACCESS_KEY_ID" | "AWS_SECRET_ACCESS_KEY" => Some("test".into()),
|
||||
"AWS_BEDROCK_RUNTIME_ENDPOINT" => Some(endpoint.clone()),
|
||||
"KEY" => Some("environment-value".into()),
|
||||
_ => None,
|
||||
});
|
||||
let settings = KeyManagementSettings {
|
||||
primary_secret_name: primary,
|
||||
..Default::default()
|
||||
};
|
||||
let manager = AwsSecretsManagerV2::load_aws_secret_manager(
|
||||
Some(true),
|
||||
settings.clone(),
|
||||
environment.clone(),
|
||||
)
|
||||
.unwrap()
|
||||
.unwrap();
|
||||
let state = SecretManagerState::new(
|
||||
Some(KeyManagementSystem::AwsSecretManager),
|
||||
Some(settings),
|
||||
Some(SecretManager::AwsSecretsManagerV2(manager)),
|
||||
)
|
||||
.unwrap();
|
||||
let resolver = SecretResolver::new(
|
||||
Arc::new(state),
|
||||
environment.clone(),
|
||||
OidcResolver::default(),
|
||||
);
|
||||
assert_eq!(
|
||||
resolver
|
||||
.get_secret_str("os.environ/KEY", None)
|
||||
.await
|
||||
.unwrap()
|
||||
.map(|v| v.expose().to_owned())
|
||||
.as_deref(),
|
||||
expected
|
||||
);
|
||||
}
|
||||
|
||||
#[cfg(feature = "google")]
|
||||
#[rstest::rstest]
|
||||
#[case::hosted_filter(Some(Vec::new()), Some(KeyManagementSystem::GoogleSecretManager))]
|
||||
#[case::negative_cache(None, Some(KeyManagementSystem::GoogleSecretManager))]
|
||||
#[case::missing_system(None, None)]
|
||||
#[case::hosted_nested_prefix(Some(vec!["os.environ/KEY".into()]), Some(KeyManagementSystem::GoogleSecretManager))]
|
||||
#[tokio::test]
|
||||
async fn google_negative_cache_still_falls_back_and_hosted_filter_avoids_io(
|
||||
#[case] hosted_keys: Option<Vec<String>>,
|
||||
#[case] system: Option<KeyManagementSystem>,
|
||||
) {
|
||||
use litellm_secrets::google::GoogleSecretManager;
|
||||
use std::time::Duration;
|
||||
use wiremock::{Mock, MockServer, ResponseTemplate, matchers::path};
|
||||
let server = MockServer::start().await;
|
||||
Mock::given(path(
|
||||
"/v1/projects/project/secrets/os%2Eenviron%2FKEY/versions/latest:access",
|
||||
))
|
||||
.respond_with(ResponseTemplate::new(404))
|
||||
.expect(u64::from(
|
||||
hosted_keys.as_ref().is_none_or(|keys| !keys.is_empty()) && system.is_some(),
|
||||
))
|
||||
.mount(&server)
|
||||
.await;
|
||||
let environment: Arc<dyn litellm_core_utils::settings::Lookup + Send + Sync> =
|
||||
Arc::new(|name: &str| match name {
|
||||
"VERTEX_AI_API_KEY" => Some("token".into()),
|
||||
"os.environ/KEY" => Some("environment-value".into()),
|
||||
_ => None,
|
||||
});
|
||||
let manager = GoogleSecretManager::with_client(
|
||||
reqwest::Client::new(),
|
||||
server.uri().parse().unwrap(),
|
||||
"project".into(),
|
||||
environment.clone(),
|
||||
Some(Duration::from_secs(60)),
|
||||
false,
|
||||
)
|
||||
.unwrap();
|
||||
let settings = KeyManagementSettings {
|
||||
hosted_keys,
|
||||
..Default::default()
|
||||
};
|
||||
let state = SecretManagerState::new(
|
||||
system,
|
||||
Some(settings),
|
||||
Some(SecretManager::GoogleSecretManager(manager)),
|
||||
)
|
||||
.unwrap();
|
||||
let resolver = SecretResolver::new(
|
||||
Arc::new(state),
|
||||
environment.clone(),
|
||||
OidcResolver::default(),
|
||||
);
|
||||
for _ in 0..2 {
|
||||
assert_eq!(
|
||||
resolver
|
||||
.get_secret_str("os.environ/os.environ/KEY", None)
|
||||
.await
|
||||
.unwrap()
|
||||
.unwrap()
|
||||
.expose(),
|
||||
"environment-value"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn resolver_future_can_run_on_a_tokio_worker() {
|
||||
let resolver = resolver(Some("worker-value"), false);
|
||||
let result = tokio::spawn(async move { resolver.get_secret_str("KEY", None).await })
|
||||
.await
|
||||
.unwrap()
|
||||
.unwrap();
|
||||
assert_eq!(result.unwrap().expose(), "worker-value");
|
||||
}
|
||||
|
||||
#[rstest::rstest]
|
||||
#[case::nested_true("((True)) # comment", Some(true))]
|
||||
#[case::commented_false("(False # comment\n)", Some(false))]
|
||||
#[case::boolean_expression("True and False", None)]
|
||||
#[case::string_literal("'True'", None)]
|
||||
#[case::tuple("(True,)", None)]
|
||||
#[case::unary_expression("not False", None)]
|
||||
#[case::multiple_expressions("True\nFalse", None)]
|
||||
#[case::incomplete_expression("(True", None)]
|
||||
#[tokio::test]
|
||||
async fn manager_boolean_literals_follow_python_syntax(
|
||||
#[case] input: &str,
|
||||
#[case] expected: Option<bool>,
|
||||
) {
|
||||
let value = resolver(Some(input), true)
|
||||
.get_secret("key", None)
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
value,
|
||||
Some(
|
||||
expected
|
||||
.map(Secret::Bool)
|
||||
.unwrap_or_else(|| Secret::String(SecretValue::new(input)))
|
||||
)
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn environment_prefix_is_removed_only_once_and_gating_uses_the_same_name() {
|
||||
let name = "os.environ/folder/os.environ/KEY";
|
||||
let state = SecretManagerState::new(
|
||||
Some(KeyManagementSystem::Local),
|
||||
Some(KeyManagementSettings {
|
||||
hosted_keys: Some(vec!["folder/os.environ/KEY".into()]),
|
||||
..Default::default()
|
||||
}),
|
||||
Some(SecretManager::Local),
|
||||
)
|
||||
.unwrap();
|
||||
assert!(secret_manager_would_be_consulted(&state, name));
|
||||
let resolver = SecretResolver::new(
|
||||
Arc::new(state),
|
||||
Arc::new(|name: &str| (name == "folder/os.environ/KEY").then(|| "value".into())),
|
||||
OidcResolver::default(),
|
||||
);
|
||||
assert_eq!(
|
||||
resolver
|
||||
.get_secret_str(name, None)
|
||||
.await
|
||||
.unwrap()
|
||||
.unwrap()
|
||||
.expose(),
|
||||
"value"
|
||||
);
|
||||
}
|
||||
Loading…
Add table
Reference in a new issue