feat(rust): add typed secret managers and shared auth adapters

This commit is contained in:
Yujong Lee 2026-09-20 16:09:20 -07:00
parent 3fd2dd635b
commit 82bc67b122
41 changed files with 4459 additions and 34 deletions

View file

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

File diff suppressed because it is too large Load diff

View file

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

View file

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

View file

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

View file

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

View 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"));
}
}

View 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"

View 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()
}
}

View 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),
}

View 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))))
}

View 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};

View 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>())
}

View 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()
);
}
}

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

View 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"

View 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)),
)
}

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

View 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)))
}

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

View 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)),
}
}

View 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()
);
}
}

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

View 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

View 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)
}

View 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)
}

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

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

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

View 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"));
}

View 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(()));
}

View 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

View 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),
}

View 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)
}

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

View 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())))
}

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

View 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))
}

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

View 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,
);
}
}

View 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"
);
}