use std::sync::Arc;
use agent_framework_core::error::{Error, Result};
use async_trait::async_trait;
use azure_core::credentials::TokenCredential as AzureTokenCredential;
use crate::TokenCredential;
pub const AZURE_OPENAI_SCOPE: &str = "https://cognitiveservices.azure.com/.default";
pub const FOUNDRY_SCOPE: &str = "https://ai.azure.com/.default";
#[derive(Debug, Clone)]
pub struct SdkTokenCredential {
inner: Arc<dyn AzureTokenCredential>,
default_scope: String,
}
impl SdkTokenCredential {
pub fn new(inner: Arc<dyn AzureTokenCredential>, default_scope: impl Into<String>) -> Self {
Self {
inner,
default_scope: default_scope.into(),
}
}
pub fn azure_openai(inner: Arc<dyn AzureTokenCredential>) -> Self {
Self::new(inner, AZURE_OPENAI_SCOPE)
}
pub fn foundry(inner: Arc<dyn AzureTokenCredential>) -> Self {
Self::new(inner, FOUNDRY_SCOPE)
}
pub fn default_scope(&self) -> &str {
&self.default_scope
}
pub fn inner(&self) -> &Arc<dyn AzureTokenCredential> {
&self.inner
}
}
#[async_trait]
impl TokenCredential for SdkTokenCredential {
async fn get_token(&self) -> Result<String> {
self.get_token_for_scope(&self.default_scope).await
}
async fn get_token_for_scope(&self, scope: &str) -> Result<String> {
let token = self.inner.get_token(&[scope], None).await.map_err(|e| {
Error::service(format!(
"Azure SDK credential failed for scope '{scope}': {e}"
))
})?;
Ok(token.token.secret().to_string())
}
}
#[cfg(test)]
mod tests {
use super::*;
use azure_core::credentials::AccessToken;
use azure_core::time::OffsetDateTime;
use std::sync::Mutex;
#[derive(Debug, Default)]
struct RecordingCredential {
scopes: Mutex<Vec<String>>,
fail: bool,
}
#[async_trait]
impl AzureTokenCredential for RecordingCredential {
async fn get_token(
&self,
scopes: &[&str],
_options: Option<azure_core::credentials::TokenRequestOptions<'_>>,
) -> azure_core::Result<AccessToken> {
self.scopes.lock().unwrap().push(scopes.join(","));
if self.fail {
return Err(azure_core::Error::with_message(
azure_core::error::ErrorKind::Credential,
"no credential available",
));
}
Ok(AccessToken::new(
"token-value",
OffsetDateTime::now_utc() + azure_core::time::Duration::hours(1),
))
}
}
#[tokio::test]
async fn get_token_uses_the_default_scope() {
let recorder = Arc::new(RecordingCredential::default());
let cred = SdkTokenCredential::azure_openai(recorder.clone());
assert_eq!(cred.get_token().await.unwrap(), "token-value");
assert_eq!(
recorder.scopes.lock().unwrap().as_slice(),
[AZURE_OPENAI_SCOPE]
);
}
#[tokio::test]
async fn get_token_for_scope_overrides_the_default() {
let recorder = Arc::new(RecordingCredential::default());
let cred = SdkTokenCredential::azure_openai(recorder.clone());
cred.get_token_for_scope(FOUNDRY_SCOPE).await.unwrap();
assert_eq!(recorder.scopes.lock().unwrap().as_slice(), [FOUNDRY_SCOPE]);
}
#[tokio::test]
async fn credential_failure_maps_to_a_service_error_naming_the_scope() {
let cred = SdkTokenCredential::foundry(Arc::new(RecordingCredential {
fail: true,
..Default::default()
}));
let err = cred.get_token().await.unwrap_err();
let rendered = err.to_string();
assert!(
rendered.contains(FOUNDRY_SCOPE),
"error should name the scope, got: {rendered}"
);
}
#[test]
fn real_sdk_credentials_fit_the_trait() {
let secret = azure_identity::ClientSecretCredential::new(
"cc7d0b33-84c6-4bd1-bbfc-1b5b1cd8ca3a",
"client-id".into(),
"client-secret".into(),
None,
)
.expect("client-secret credential constructs");
let boxed: Box<dyn TokenCredential> = Box::new(SdkTokenCredential::azure_openai(secret));
let _ = boxed;
}
}