use super::SecretResolver;
use crate::error::{CliError, CliResult};
use async_trait::async_trait;
use azure_core::credentials::{AccessToken, Secret, TokenCredential, TokenRequestOptions};
use azure_identity::{ClientSecretCredential, DeveloperToolsCredential, ManagedIdentityCredential};
use std::sync::Arc;
use tokio::sync::OnceCell;
struct ChainedCredential {
sources: Vec<Arc<dyn TokenCredential>>,
}
impl ChainedCredential {
fn new(sources: Vec<Arc<dyn TokenCredential>>) -> Arc<Self> {
Arc::new(Self { sources })
}
}
impl std::fmt::Debug for ChainedCredential {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("ChainedCredential")
.field("sources_len", &self.sources.len())
.finish()
}
}
#[async_trait]
impl TokenCredential for ChainedCredential {
async fn get_token(
&self,
scopes: &[&str],
options: Option<TokenRequestOptions<'_>>,
) -> azure_core::Result<AccessToken> {
let mut errors = Vec::new();
for source in &self.sources {
match source.get_token(scopes, options.clone()).await {
Ok(token) => return Ok(token),
Err(e) => errors.push(e.to_string()),
}
}
Err(azure_core::Error::with_message(
azure_core::error::ErrorKind::Credential,
format!(
"All Azure credential sources failed:\n{}",
errors.join("\n")
),
))
}
}
fn build_credential() -> CliResult<Arc<dyn TokenCredential>> {
let mut sources: Vec<Arc<dyn TokenCredential>> = Vec::new();
if let (Ok(tenant_id), Ok(client_id), Ok(client_secret)) = (
std::env::var("AZURE_TENANT_ID"),
std::env::var("AZURE_CLIENT_ID"),
std::env::var("AZURE_CLIENT_SECRET"),
) {
match ClientSecretCredential::new(&tenant_id, client_id, Secret::new(client_secret), None) {
Ok(cred) => sources.push(cred),
Err(e) => {
tracing::debug!("Azure ClientSecretCredential build failed: {e}");
}
}
}
match ManagedIdentityCredential::new(None) {
Ok(cred) => sources.push(cred),
Err(e) => tracing::debug!("Azure ManagedIdentityCredential build failed: {e}"),
}
match DeveloperToolsCredential::new(None) {
Ok(cred) => sources.push(cred),
Err(e) => tracing::debug!("Azure DeveloperToolsCredential build failed: {e}"),
}
if sources.is_empty() {
return Err(CliError::SecretAuthFailed {
scheme: "azure-kv".into(),
hint: "no Azure credentials available — set AZURE_TENANT_ID / AZURE_CLIENT_ID / \
AZURE_CLIENT_SECRET, use a managed identity, or run `az login`"
.into(),
});
}
Ok(ChainedCredential::new(sources))
}
pub struct AzureKvResolver {
credential: OnceCell<Arc<dyn TokenCredential>>,
}
impl AzureKvResolver {
pub fn new() -> Self {
Self {
credential: OnceCell::new(),
}
}
async fn credential(&self) -> CliResult<Arc<dyn TokenCredential>> {
let cred = self
.credential
.get_or_try_init(|| async { build_credential() })
.await?;
Ok(Arc::clone(cred))
}
}
impl Default for AzureKvResolver {
fn default() -> Self {
Self::new()
}
}
#[async_trait]
impl SecretResolver for AzureKvResolver {
fn scheme(&self) -> &'static str {
"azure-kv"
}
async fn resolve(&self, reference: &str) -> CliResult<String> {
let mut parts = reference.splitn(3, '/');
let vault =
parts
.next()
.filter(|s| !s.is_empty())
.ok_or_else(|| CliError::SecretFetchFailed {
scheme: "azure-kv".into(),
reference: reference.into(),
source: "expected '<vault>/<secret>[/<version>]'".into(),
})?;
let secret_name =
parts
.next()
.filter(|s| !s.is_empty())
.ok_or_else(|| CliError::SecretFetchFailed {
scheme: "azure-kv".into(),
reference: reference.into(),
source: "expected '<vault>/<secret>[/<version>]'".into(),
})?;
let version = parts.next().filter(|s| !s.is_empty());
let credential = self.credential().await?;
let vault_url = format!("https://{vault}.vault.azure.net/");
let client =
azure_security_keyvault_secrets::SecretClient::new(&vault_url, credential, None)
.map_err(|e| CliError::SecretFetchFailed {
scheme: "azure-kv".into(),
reference: reference.into(),
source: Box::new(e),
})?;
let options = version.map(|v| {
azure_security_keyvault_secrets::models::SecretClientGetSecretOptions {
secret_version: Some(v.to_string()),
..Default::default()
}
});
let resp = client.get_secret(secret_name, options).await.map_err(|e| {
CliError::SecretFetchFailed {
scheme: "azure-kv".into(),
reference: reference.into(),
source: Box::new(e),
}
})?;
let secret = resp.into_model().map_err(|e| CliError::SecretFetchFailed {
scheme: "azure-kv".into(),
reference: reference.into(),
source: Box::new(e),
})?;
secret.value.ok_or_else(|| CliError::SecretNotFound {
scheme: "azure-kv".into(),
reference: reference.into(),
})
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn credential_chain_is_built_once() {
let resolver = AzureKvResolver::new();
let first = resolver.credential().await.expect("chain builds");
let second = resolver.credential().await.expect("chain builds");
assert!(
Arc::ptr_eq(&first, &second),
"credential chain should be built at most once and reused"
);
}
#[test]
fn scheme_is_azure_kv() {
assert_eq!(AzureKvResolver::new().scheme(), "azure-kv");
}
#[test]
fn default_constructs_a_resolver() {
let _r = AzureKvResolver::default();
}
#[tokio::test]
async fn empty_reference_is_a_parse_error_before_any_network() {
let resolver = AzureKvResolver::new();
match resolver.resolve("").await.unwrap_err() {
CliError::SecretFetchFailed {
scheme, reference, ..
} => {
assert_eq!(scheme, "azure-kv");
assert_eq!(reference, "");
}
other => panic!("expected SecretFetchFailed, got {other:?}"),
}
}
#[tokio::test]
async fn missing_secret_name_segment_is_a_parse_error() {
let resolver = AzureKvResolver::new();
match resolver.resolve("myvault").await.unwrap_err() {
CliError::SecretFetchFailed {
scheme, reference, ..
} => {
assert_eq!(scheme, "azure-kv");
assert_eq!(reference, "myvault");
}
other => panic!("expected SecretFetchFailed, got {other:?}"),
}
}
#[tokio::test]
async fn empty_secret_name_segment_is_a_parse_error() {
let resolver = AzureKvResolver::new();
match resolver.resolve("myvault/").await.unwrap_err() {
CliError::SecretFetchFailed { reference, .. } => {
assert_eq!(reference, "myvault/");
}
other => panic!("expected SecretFetchFailed, got {other:?}"),
}
}
#[tokio::test]
async fn chained_credential_aggregates_errors_when_all_fail() {
let chain = ChainedCredential::new(Vec::new());
let err = chain.get_token(&["scope"], None).await.unwrap_err();
assert!(
err.to_string()
.contains("All Azure credential sources failed"),
"{err}"
);
}
#[test]
fn chained_credential_debug_reports_source_count() {
let chain = ChainedCredential::new(Vec::new());
let dbg = format!("{chain:?}");
assert!(dbg.contains("ChainedCredential"), "{dbg}");
assert!(dbg.contains("sources_len"), "{dbg}");
}
}