1use crate::env::Env;
5use async_lock::{RwLock, RwLockUpgradableReadGuard};
6use azure_core::{
7 credentials::{AccessToken, Secret, TokenCredential, TokenRequestOptions},
8 error::{ErrorKind, ResultExt},
9 http::ClientMethodOptions,
10 Error,
11};
12use futures::channel::oneshot;
13use std::{
14 any::type_name,
15 fmt, fs,
16 path::PathBuf,
17 str,
18 sync::Arc,
19 thread,
20 time::{Duration, Instant},
21};
22
23use super::{ClientAssertion, ClientAssertionCredential, ClientAssertionCredentialOptions};
24
25const AZURE_CLIENT_ID: &str = "AZURE_CLIENT_ID";
26const AZURE_FEDERATED_TOKEN_FILE: &str = "AZURE_FEDERATED_TOKEN_FILE";
27const AZURE_TENANT_ID: &str = "AZURE_TENANT_ID";
28
29pub struct WorkloadIdentityCredential(ClientAssertionCredential<Token>);
31
32impl fmt::Debug for WorkloadIdentityCredential {
33 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
34 f.debug_tuple(type_name::<Self>()).finish_non_exhaustive()
35 }
36}
37
38#[derive(Default)]
40pub struct WorkloadIdentityCredentialOptions {
41 pub credential_options: ClientAssertionCredentialOptions,
43
44 pub client_id: Option<String>,
46
47 pub tenant_id: Option<String>,
49
50 pub token_file_path: Option<PathBuf>,
53
54 #[cfg(test)]
55 pub(crate) env: Env,
56}
57
58impl fmt::Debug for WorkloadIdentityCredentialOptions {
59 fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
60 f.debug_struct(type_name::<Self>())
61 .field("tenant_id", &self.tenant_id)
62 .field("client_id", &self.client_id)
63 .finish_non_exhaustive()
64 }
65}
66
67impl WorkloadIdentityCredential {
68 pub fn new(
70 options: Option<WorkloadIdentityCredentialOptions>,
71 ) -> azure_core::Result<Arc<Self>> {
72 let options = options.unwrap_or_default();
73 #[cfg(test)]
74 let env = options.env;
75 #[cfg(not(test))]
76 let env = Env::default();
77 let tenant_id = match options.tenant_id {
78 Some(id) => id,
79 None => env.var(AZURE_TENANT_ID).with_context_fn(ErrorKind::Credential, || {
80 "no tenant ID specified. Check pod configuration or set tenant_id in the options"
81 })?
82 };
83 crate::validate_tenant_id(&tenant_id)?;
84 let path = match options.token_file_path {
85 Some(path) => path,
86 None => env.var(AZURE_FEDERATED_TOKEN_FILE).map(PathBuf::from).with_context_fn(ErrorKind::Credential, || {
87 "no token file specified. Check pod configuration or set token_file_path in the options"
88 })?
89 };
90 let client_id = match options.client_id {
91 Some(id) => id,
92 None => env.var(AZURE_CLIENT_ID).with_context_fn(ErrorKind::Credential, || {
93 "no client id specified. Check pod configuration or set client_id in the options"
94 })?
95 };
96 Ok(Arc::new(Self(
97 ClientAssertionCredential::<Token>::new_exclusive(
98 tenant_id,
99 client_id,
100 Token::new(path)?,
101 stringify!(WorkloadIdentityCredential),
102 Some(options.credential_options),
103 )?,
104 )))
105 }
106}
107
108#[async_trait::async_trait]
109impl TokenCredential for WorkloadIdentityCredential {
110 async fn get_token(
111 &self,
112 scopes: &[&str],
113 options: Option<TokenRequestOptions<'_>>,
114 ) -> azure_core::Result<AccessToken> {
115 if scopes.is_empty() {
116 return Err(Error::with_message(
117 ErrorKind::Credential,
118 "no scopes specified",
119 ));
120 }
121 self.0.get_token(scopes, options).await
122 }
123}
124
125#[derive(Debug)]
126struct Token {
127 path: PathBuf,
128 cache: Arc<RwLock<FileCache>>,
129}
130
131#[derive(Debug)]
132struct FileCache {
133 token: Secret,
134 last_read: Instant,
135}
136
137impl Token {
138 fn new(path: PathBuf) -> azure_core::Result<Self> {
139 let last_read = Instant::now();
140 let token =
141 std::fs::read_to_string(&path).with_context_fn(ErrorKind::Credential, || {
142 format!(
143 "failed to read federated token from file {}",
144 path.display()
145 )
146 })?;
147
148 Ok(Self {
149 path,
150 cache: Arc::new(RwLock::new(FileCache {
151 token: Secret::new(token),
152 last_read,
153 })),
154 })
155 }
156}
157
158#[async_trait::async_trait]
159impl ClientAssertion for Token {
160 async fn secret(&self, _: Option<ClientMethodOptions<'_>>) -> azure_core::Result<String> {
161 const TIMEOUT: Duration = Duration::from_secs(600);
162
163 let now = Instant::now();
164 let cache = self.cache.upgradable_read().await;
165 if now - cache.last_read > TIMEOUT {
166 let path = self.path.clone();
168 let (tx, rx) = oneshot::channel();
169 thread::spawn(move || {
170 let token =
171 fs::read_to_string(&path).with_context_fn(ErrorKind::Credential, || {
172 format!(
173 "failed to read federated token from file {}",
174 path.display()
175 )
176 });
177 tx.send(token)
178 });
179
180 let mut write_cache = RwLockUpgradableReadGuard::upgrade(cache).await;
181 let token = rx.await.map_err(|err| {
182 azure_core::Error::with_error(ErrorKind::Io, err, "canceled reading certificate")
183 })??;
184
185 write_cache.token = Secret::new(token);
186 write_cache.last_read = now;
187
188 return Ok(write_cache.token.secret().into());
189 }
190
191 Ok(cache.token.secret().into())
192 }
193}
194
195#[cfg(test)]
196mod tests {
197 use super::*;
198 use crate::{
199 client_assertion_credential::tests::{is_valid_request, FAKE_ASSERTION},
200 env::Env,
201 tests::*,
202 };
203 use azure_core::{
204 http::{
205 headers::Headers, AsyncRawResponse, ClientOptions, Method, RawResponse, Request,
206 StatusCode, Transport, Url,
207 },
208 Bytes,
209 };
210 use azure_core_test::recorded;
211 use std::{
212 env,
213 fs::File,
214 io::Write,
215 sync::atomic::{AtomicUsize, Ordering},
216 time::SystemTime,
217 };
218
219 static TEMP_FILE_COUNTER: AtomicUsize = AtomicUsize::new(0);
220
221 pub struct TempFile {
222 pub path: PathBuf,
223 }
224
225 impl TempFile {
226 pub fn new(content: &str) -> Self {
227 let n = TEMP_FILE_COUNTER.fetch_add(1, Ordering::SeqCst);
228 let path = env::temp_dir().join(format!("azure_identity_test_{}", n));
229 File::create(&path)
230 .expect("create temp file")
231 .write_all(content.as_bytes())
232 .expect("write temp file");
233
234 Self { path }
235 }
236 }
237
238 impl Drop for TempFile {
239 fn drop(&mut self) {
240 let _ = fs::remove_file(&self.path);
241 }
242 }
243
244 #[tokio::test]
245 async fn env_vars() {
246 let temp_file = TempFile::new(FAKE_ASSERTION);
247 let mock = MockSts::new(
248 vec![AsyncRawResponse::from_bytes(
249 StatusCode::Ok,
250 Headers::default(),
251 Bytes::from(format!(
252 r#"{{"access_token":"{}","expires_in":3600,"ext_expires_in":3600,"token_type":"Bearer"}}"#,
253 FAKE_TOKEN
254 )),
255 )],
256 Some(Arc::new(is_valid_request(
257 FAKE_PUBLIC_CLOUD_AUTHORITY.to_string(),
258 Some(FAKE_ASSERTION.to_string()),
259 ))),
260 );
261 let cred = WorkloadIdentityCredential::new(Some(WorkloadIdentityCredentialOptions {
262 credential_options: ClientAssertionCredentialOptions {
263 client_options: ClientOptions {
264 transport: Some(Transport::new(Arc::new(mock))),
265 ..Default::default()
266 },
267 },
268 env: Env::from(
269 &[
270 (AZURE_CLIENT_ID, FAKE_CLIENT_ID),
271 (AZURE_TENANT_ID, FAKE_TENANT_ID),
272 (AZURE_FEDERATED_TOKEN_FILE, temp_file.path.to_str().unwrap()),
273 ][..],
274 ),
275 ..Default::default()
276 }))
277 .expect("valid credential");
278
279 let token = cred.get_token(LIVE_TEST_SCOPES, None).await.expect("token");
280 assert_eq!(FAKE_TOKEN, token.token.secret());
281 assert!(token.expires_on > SystemTime::now());
282 }
283
284 #[tokio::test]
285 async fn get_token_error() {
286 let temp_file = TempFile::new(FAKE_ASSERTION);
287 let expected_status = StatusCode::Forbidden;
288 let body = r#"{"error":"invalid_request","error_description":"invalid assertion"}"#;
289 let mut headers = Headers::default();
290 headers.insert("key", "value");
291 let expected_response = RawResponse::from_bytes(expected_status, headers.clone(), body);
292 let mock = MockSts::new(
293 vec![AsyncRawResponse::from_bytes(
294 expected_status,
295 headers.clone(),
296 Bytes::from(body),
297 )],
298 Some(Arc::new(is_valid_request(
299 FAKE_PUBLIC_CLOUD_AUTHORITY.to_string(),
300 Some(FAKE_ASSERTION.to_string()),
301 ))),
302 );
303 let cred = WorkloadIdentityCredential::new(Some(WorkloadIdentityCredentialOptions {
304 credential_options: ClientAssertionCredentialOptions {
305 client_options: ClientOptions {
306 transport: Some(Transport::new(Arc::new(mock))),
307 ..Default::default()
308 },
309 },
310 env: Env::from(
311 &[
312 (AZURE_CLIENT_ID, FAKE_CLIENT_ID),
313 (AZURE_TENANT_ID, FAKE_TENANT_ID),
314 (AZURE_FEDERATED_TOKEN_FILE, temp_file.path.to_str().unwrap()),
315 ][..],
316 ),
317 ..Default::default()
318 }))
319 .expect("valid credential");
320
321 let err = cred
322 .get_token(LIVE_TEST_SCOPES, None)
323 .await
324 .expect_err("expected error");
325
326 assert!(matches!(err.kind(), ErrorKind::Credential));
327 assert_eq!(
328 "WorkloadIdentityCredential authentication failed. invalid assertion\nTo troubleshoot, visit https://aka.ms/azsdk/rust/identity/troubleshoot#workload",
329 err.to_string(),
330 );
331 match err
332 .downcast_ref::<azure_core::Error>()
333 .expect("returned error should wrap an azure_core::Error")
334 .kind()
335 {
336 ErrorKind::HttpResponse {
337 error_code: None,
338 raw_response: Some(response),
339 status,
340 ..
341 } => {
342 assert_eq!(&expected_response, response.as_ref());
343 assert_eq!(expected_status, *status);
344 }
345 kind => panic!("unexpected ErrorKind {:?}", kind),
346 };
347 }
348
349 #[test]
350 fn invalid_tenant_id() {
351 let temp_file = TempFile::new(FAKE_ASSERTION);
352 WorkloadIdentityCredential::new(Some(WorkloadIdentityCredentialOptions {
353 client_id: Some(FAKE_CLIENT_ID.to_string()),
354 tenant_id: Some("not a valid tenant".to_string()),
355 token_file_path: Some(temp_file.path.clone()),
356 ..Default::default()
357 }))
358 .expect_err("invalid tenant ID");
359 }
360
361 #[recorded::test(live)]
362 async fn live() -> azure_core::Result<()> {
363 if env::var("CI_HAS_DEPLOYED_RESOURCES").is_err() {
364 println!("Skipped: workload identity live tests require deployed resources");
365 return Ok(());
366 }
367 let ip = env::var("IDENTITY_AKS_IP").expect("IDENTITY_AKS_IP");
368 let storage_name = env::var("IDENTITY_STORAGE_NAME_USER_ASSIGNED")
369 .expect("IDENTITY_STORAGE_NAME_USER_ASSIGNED");
370
371 let url =
372 format!("http://{ip}:8080/api?test=workload-identity&storage-name={storage_name}");
373 let u = Url::parse(&url).expect("valid URL");
374 let client = azure_core::http::new_http_client(None);
375 let req = Request::new(u, Method::Get);
376
377 let res = client.execute_request(&req).await.expect("response");
378 let status = res.status();
379 let body = res
380 .into_body()
381 .collect_string()
382 .await
383 .expect("body content");
384
385 assert_eq!(StatusCode::Ok, status, "Test app responded with '{body}'");
386
387 Ok(())
388 }
389
390 #[test]
391 fn missing_config() {
392 WorkloadIdentityCredential::new(None).expect_err("missing config");
393 }
394
395 #[tokio::test]
396 async fn no_scopes() {
397 let temp_file = TempFile::new(FAKE_ASSERTION);
398 WorkloadIdentityCredential::new(Some(WorkloadIdentityCredentialOptions {
399 client_id: Some(FAKE_CLIENT_ID.to_string()),
400 tenant_id: Some(FAKE_TENANT_ID.to_string()),
401 token_file_path: Some(temp_file.path.clone()),
402 ..Default::default()
403 }))
404 .expect("valid credential")
405 .get_token(&[], None)
406 .await
407 .expect_err("no scopes specified");
408 }
409
410 #[tokio::test]
411 async fn options_override_env() {
412 let right_file = TempFile::new(FAKE_ASSERTION);
413 let wrong_file = TempFile::new("wrong assertion");
414 let mock = MockSts::new(
415 vec![AsyncRawResponse::from_bytes(
416 StatusCode::Ok,
417 Headers::default(),
418 Bytes::from(format!(
419 r#"{{"access_token":"{}","expires_in":3600,"ext_expires_in":3600,"token_type":"Bearer"}}"#,
420 FAKE_TOKEN
421 )),
422 )],
423 Some(Arc::new(is_valid_request(
424 FAKE_PUBLIC_CLOUD_AUTHORITY.to_string(),
425 Some(FAKE_ASSERTION.to_string()),
426 ))),
427 );
428 let cred = WorkloadIdentityCredential::new(Some(WorkloadIdentityCredentialOptions {
429 client_id: Some(FAKE_CLIENT_ID.to_string()),
430 tenant_id: Some(FAKE_TENANT_ID.to_string()),
431 token_file_path: Some(right_file.path.clone()),
432 credential_options: ClientAssertionCredentialOptions {
433 client_options: ClientOptions {
434 transport: Some(Transport::new(Arc::new(mock))),
435 ..Default::default()
436 },
437 },
438 env: Env::from(
439 &[
440 (AZURE_CLIENT_ID, "wrong-client-id"),
441 (AZURE_TENANT_ID, "wrong-tenant-id"),
442 (
443 AZURE_FEDERATED_TOKEN_FILE,
444 wrong_file.path.to_str().unwrap(),
445 ),
446 ][..],
447 ),
448 }))
449 .expect("valid credential");
450
451 let token = cred.get_token(LIVE_TEST_SCOPES, None).await.expect("token");
452 assert_eq!(FAKE_TOKEN, token.token.secret());
453 assert!(token.expires_on > SystemTime::now());
454 }
455}