Skip to main content

azure_identity/
workload_identity_credential.rs

1// Copyright (c) Microsoft Corporation. All rights reserved.
2// Licensed under the MIT License.
3
4use 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
29/// Authenticates an [Entra Workload Identity on Kubernetes](https://learn.microsoft.com/azure/aks/workload-identity-overview).
30pub 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/// Options for constructing a new [`WorkloadIdentityCredential`].
39#[derive(Default)]
40pub struct WorkloadIdentityCredentialOptions {
41    /// Options for the [`ClientAssertionCredential`] used by the [`WorkloadIdentityCredential`].
42    pub credential_options: ClientAssertionCredentialOptions,
43
44    /// Client ID of the Entra identity. Defaults to the value of the environment variable `AZURE_CLIENT_ID`.
45    pub client_id: Option<String>,
46
47    /// Tenant ID of the Entra identity. Defaults to the value of the environment variable `AZURE_TENANT_ID`.
48    pub tenant_id: Option<String>,
49
50    /// Path of a file containing a Kubernetes service account token. Defaults to the value of the environment
51    /// variable `AZURE_FEDERATED_TOKEN_FILE`.
52    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    /// Create a new `WorkloadIdentityCredential`.
69    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            // TODO: https://github.com/Azure/azure-sdk-for-rust/issues/2002
167            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}