Skip to main content

azure_identity/
managed_identity_credential.rs

1// Copyright (c) Microsoft Corporation. All rights reserved.
2// Licensed under the MIT License.
3
4use crate::{
5    authentication_error, env::Env, AppServiceManagedIdentityCredential, ImdsId,
6    VirtualMachineManagedIdentityCredential,
7};
8use azure_core::credentials::{AccessToken, TokenCredential, TokenRequestOptions};
9use azure_core::http::ClientOptions;
10use std::{any::type_name, fmt, sync::Arc};
11use tracing::info;
12
13/// Identifies a specific user-assigned identity for [`ManagedIdentityCredential`] to authenticate.
14#[derive(Debug, Clone)]
15#[non_exhaustive]
16pub enum UserAssignedId {
17    /// The client ID of a user-assigned identity
18    ClientId(String),
19    /// The object or principal ID of a user-assigned identity
20    ObjectId(String),
21    /// The Azure resource ID of a user-assigned identity
22    ResourceId(String),
23}
24
25/// Authenticates a managed identity from Azure App Service or an Azure Virtual Machine.
26pub struct ManagedIdentityCredential {
27    credential: Arc<dyn TokenCredential>,
28}
29
30impl fmt::Debug for ManagedIdentityCredential {
31    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
32        f.debug_struct(type_name::<Self>()).finish_non_exhaustive()
33    }
34}
35
36/// Options for constructing a new [`ManagedIdentityCredential`].
37#[derive(Clone, Default)]
38pub struct ManagedIdentityCredentialOptions {
39    /// Specifies a user-assigned identity the credential should authenticate.
40    /// When `None`, the credential will authenticate a system-assigned identity, if any.
41    pub user_assigned_id: Option<UserAssignedId>,
42
43    /// The [`ClientOptions`] to use for the credential's pipeline.
44    pub client_options: ClientOptions,
45
46    #[cfg(test)]
47    pub(crate) env: Env,
48}
49
50impl fmt::Debug for ManagedIdentityCredentialOptions {
51    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
52        f.debug_struct(type_name::<Self>()).finish_non_exhaustive()
53    }
54}
55
56impl ManagedIdentityCredential {
57    /// Creates a new instance of `ManagedIdentityCredential`.
58    ///
59    /// # Arguments
60    /// * `options`: Options for configuring the credential. If `None`, the credential uses its default options.
61    ///
62    pub fn new(options: Option<ManagedIdentityCredentialOptions>) -> azure_core::Result<Arc<Self>> {
63        let options = options.unwrap_or_default();
64        #[cfg(test)]
65        let env = options.env;
66        #[cfg(not(test))]
67        let env = Env::default();
68        let source = get_source(&env);
69        let id = options
70            .user_assigned_id
71            .clone()
72            .map(Into::into)
73            .unwrap_or(ImdsId::SystemAssigned);
74
75        let credential: Arc<dyn TokenCredential> = match source {
76            ManagedIdentitySource::AppService => {
77                // App Service does accept resource IDs, however this crate's current implementation sends
78                // them in the wrong query parameter: https://github.com/Azure/azure-sdk-for-rust/issues/2407
79                if let ImdsId::MsiResId(_) = id {
80                    return Err(azure_core::Error::with_message_fn(
81                        azure_core::error::ErrorKind::Credential,
82                        || {
83                            "User-assigned resource IDs aren't supported for App Service. Use a client or object ID instead.".to_string()
84                        },
85                    ));
86                }
87                AppServiceManagedIdentityCredential::new(id, options.client_options, env)?
88            }
89            ManagedIdentitySource::Imds => {
90                VirtualMachineManagedIdentityCredential::new(id, options.client_options, env)?
91            }
92            _ => {
93                return Err(azure_core::Error::with_message_fn(
94                    azure_core::error::ErrorKind::Credential,
95                    || format!("{} managed identity isn't supported", source.as_str()),
96                ));
97            }
98        };
99
100        info!(user_assigned_id = ?options.user_assigned_id, "ManagedIdentityCredential will use {} managed identity", source.as_str());
101
102        Ok(Arc::new(Self { credential }))
103    }
104}
105
106#[async_trait::async_trait]
107impl TokenCredential for ManagedIdentityCredential {
108    async fn get_token(
109        &self,
110        scopes: &[&str],
111        options: Option<TokenRequestOptions<'_>>,
112    ) -> azure_core::Result<AccessToken> {
113        if scopes.len() != 1 {
114            return Err(azure_core::Error::with_message(
115                azure_core::error::ErrorKind::Credential,
116                "ManagedIdentityCredential requires exactly one scope".to_string(),
117            ));
118        }
119        self.credential
120            .get_token(scopes, options)
121            .await
122            .map_err(|err| authentication_error(stringify!(ManagedIdentityCredential), err))
123    }
124}
125
126#[derive(Debug, Copy, Clone)]
127enum ManagedIdentitySource {
128    AzureArc,
129    AzureML,
130    AppService,
131    CloudShell,
132    Imds,
133    ServiceFabric,
134}
135
136impl ManagedIdentitySource {
137    pub fn as_str(&self) -> &'static str {
138        match self {
139            ManagedIdentitySource::AzureArc => "Azure Arc",
140            ManagedIdentitySource::AzureML => "Azure ML",
141            ManagedIdentitySource::AppService => "App Service",
142            ManagedIdentitySource::CloudShell => "CloudShell",
143            ManagedIdentitySource::Imds => "IMDS",
144            ManagedIdentitySource::ServiceFabric => "Service Fabric",
145        }
146    }
147}
148
149const IDENTITY_ENDPOINT: &str = "IDENTITY_ENDPOINT";
150const IDENTITY_HEADER: &str = "IDENTITY_HEADER";
151const IDENTITY_SERVER_THUMBPRINT: &str = "IDENTITY_SERVER_THUMBPRINT";
152const IMDS_ENDPOINT: &str = "IMDS_ENDPOINT";
153const MSI_ENDPOINT: &str = "MSI_ENDPOINT";
154const MSI_SECRET: &str = "MSI_SECRET";
155
156fn get_source(env: &Env) -> ManagedIdentitySource {
157    use ManagedIdentitySource::*;
158    if env.var(IDENTITY_ENDPOINT).is_ok() {
159        if env.var(IDENTITY_HEADER).is_ok() {
160            if env.var(IDENTITY_SERVER_THUMBPRINT).is_ok() {
161                return ServiceFabric;
162            }
163            return AppService;
164        } else if env.var(IMDS_ENDPOINT).is_ok() {
165            return AzureArc;
166        }
167    } else if env.var(MSI_ENDPOINT).is_ok() {
168        if env.var(MSI_SECRET).is_ok() {
169            return AzureML;
170        }
171        return CloudShell;
172    }
173    Imds
174}
175
176#[cfg(test)]
177mod tests {
178    use super::*;
179    use crate::{
180        env::Env,
181        tests::{LIVE_TEST_RESOURCE, LIVE_TEST_SCOPES},
182    };
183    use azure_core::http::{
184        AsyncRawResponse, Method, RawResponse, Request, StatusCode, Transport, Url,
185    };
186    use azure_core::time::OffsetDateTime;
187    use azure_core::Bytes;
188    use azure_core::{error::ErrorKind, http::headers::Headers};
189    use azure_core_test::{http::MockHttpClient, recorded};
190    use futures::FutureExt;
191    use std::env;
192    use std::sync::atomic::{AtomicUsize, Ordering};
193    use std::time::{SystemTime, UNIX_EPOCH};
194
195    const EXPIRES_ON: &str = "EXPIRES_ON";
196
197    async fn run_deployed_test(
198        authority: &str,
199        storage_name: &str,
200        id: Option<UserAssignedId>,
201    ) -> azure_core::Result<()> {
202        let id_param = id.map_or("".to_string(), |id| match id {
203            UserAssignedId::ClientId(id) => format!("client-id={id}&"),
204            UserAssignedId::ObjectId(id) => format!("object-id={id}&"),
205            UserAssignedId::ResourceId(id) => format!("resource-id={id}&"),
206        });
207        let url = format!(
208            "http://{authority}/api?test=managed-identity&{id_param}storage-name={storage_name}"
209        );
210        let u = Url::parse(&url).expect("invalid URL");
211        let client = azure_core::http::new_http_client(None);
212        let req = Request::new(u, Method::Get);
213
214        let res = client.execute_request(&req).await.expect("request failed");
215        let status = res.status();
216        let body = res.into_body().collect_string().await?;
217        assert_eq!(StatusCode::Ok, status, "Test app responded with '{body}'");
218
219        Ok(())
220    }
221
222    async fn run_error_response_test(source: ManagedIdentitySource) {
223        let expected_status = StatusCode::ImATeapot;
224        let headers = Headers::default();
225        let content: &str = "is a teapot";
226        let body = Bytes::copy_from_slice(content.as_bytes());
227        let expected_response =
228            RawResponse::from_bytes(expected_status, headers.clone(), body.clone());
229        let mock_headers = headers.clone();
230        let mock_body = body.clone();
231        let mock_client = MockHttpClient::new(move |_| {
232            let headers = mock_headers.clone();
233            let body = mock_body.clone();
234            async move { Ok(AsyncRawResponse::from_bytes(expected_status, headers, body)) }.boxed()
235        });
236        let test_env = match source {
237            ManagedIdentitySource::Imds => Env::from(&[][..]),
238            ManagedIdentitySource::AppService => Env::from(
239                &[
240                    (
241                        IDENTITY_ENDPOINT,
242                        "http://localhost/metadata/identity/oauth2/token",
243                    ),
244                    (IDENTITY_HEADER, "secret"),
245                ][..],
246            ),
247            other => panic!("unsupported managed identity source {:?}", other),
248        };
249        let options = ManagedIdentityCredentialOptions {
250            client_options: ClientOptions {
251                transport: Some(Transport::new(Arc::new(mock_client))),
252                ..Default::default()
253            },
254            env: test_env,
255            ..Default::default()
256        };
257        let credential = ManagedIdentityCredential::new(Some(options)).expect("credential");
258        let err = credential
259            .get_token(LIVE_TEST_SCOPES, None)
260            .await
261            .expect_err("expected error");
262        assert!(matches!(err.kind(), ErrorKind::Credential));
263        assert_eq!(
264            "ManagedIdentityCredential authentication failed. The request failed: is a teapot\nTo troubleshoot, visit https://aka.ms/azsdk/rust/identity/troubleshoot#managed-id",
265            err.to_string(),
266        );
267        match err
268            .downcast_ref::<azure_core::Error>()
269            .expect("returned error should wrap an azure_core::Error")
270            .kind()
271        {
272            ErrorKind::HttpResponse {
273                error_code: None,
274                raw_response: Some(response),
275                status,
276            } => {
277                assert_eq!(response.as_ref(), &expected_response);
278                assert_eq!(expected_status, *status);
279            }
280            err => panic!("unexpected {:?}", err),
281        };
282    }
283
284    async fn run_supported_source_test(
285        env: Env,
286        options: Option<ManagedIdentityCredentialOptions>,
287        expected_source: ManagedIdentitySource,
288        model_request: Request,
289        response_format: String,
290    ) {
291        let actual_source = get_source(&env);
292        assert_eq!(
293            std::mem::discriminant(&actual_source),
294            std::mem::discriminant(&expected_source)
295        );
296        let token_requests = Arc::new(AtomicUsize::new(0));
297        let token_requests_clone = token_requests.clone();
298        let expires_on = SystemTime::now()
299            .duration_since(UNIX_EPOCH)
300            .unwrap()
301            .as_secs()
302            + 3600;
303        let mock_client = MockHttpClient::new(move |actual| {
304            {
305                token_requests_clone.fetch_add(1, Ordering::SeqCst);
306                let expected = model_request.clone();
307                let response_format = response_format.clone();
308                async move {
309                    assert_eq!(expected.method(), actual.method());
310
311                    let mut actual_params: Vec<_> =
312                        actual.url().query_pairs().into_owned().collect();
313                    actual_params.sort();
314                    let mut expected_params: Vec<_> =
315                        expected.url().query_pairs().into_owned().collect();
316                    expected_params.sort();
317                    assert_eq!(expected_params, actual_params);
318
319                    let mut actual_url = actual.url().clone();
320                    actual_url.set_query(None);
321                    let mut expected_url = expected.url().clone();
322                    expected_url.set_query(None);
323                    assert_eq!(actual_url, expected_url);
324
325                    // allow additional headers in the actual request so changing
326                    // the underlying client in the future won't break tests
327                    expected.headers().iter().for_each(|(k, v)| {
328                        assert_eq!(actual.headers().get_str(k).unwrap(), v.as_str())
329                    });
330
331                    Ok(AsyncRawResponse::from_bytes(
332                        StatusCode::Ok,
333                        Headers::default(),
334                        Bytes::from(response_format.replacen(
335                            EXPIRES_ON,
336                            &expires_on.to_string(),
337                            1,
338                        )),
339                    ))
340                }
341            }
342            .boxed()
343        });
344        let mut options = options.unwrap_or_default();
345        options.env = env;
346        options.client_options = ClientOptions {
347            transport: Some(Transport::new(Arc::new(mock_client))),
348            ..Default::default()
349        };
350        let cred = ManagedIdentityCredential::new(Some(options)).expect("credential");
351        for _ in 0..4 {
352            let token = cred.get_token(LIVE_TEST_SCOPES, None).await.expect("token");
353            assert_eq!(token.expires_on.unix_timestamp(), expires_on as i64);
354            assert_eq!(token.token.secret(), "*");
355            assert_eq!(token_requests.load(Ordering::SeqCst), 1);
356        }
357    }
358
359    fn run_unsupported_source_test(env: Env, expected_source: ManagedIdentitySource) {
360        let actual_source = get_source(&env);
361        assert_eq!(
362            std::mem::discriminant(&actual_source),
363            std::mem::discriminant(&expected_source)
364        );
365        let result = ManagedIdentityCredential::new(Some(ManagedIdentityCredentialOptions {
366            env,
367            ..Default::default()
368        }));
369        assert!(
370            matches!(result, Err(ref e) if *e.kind() == azure_core::error::ErrorKind::Credential),
371            "Expected constructor error"
372        );
373    }
374
375    #[recorded::test(live)]
376    async fn aci_user_assigned_live() -> azure_core::Result<()> {
377        if env::var("CI_HAS_DEPLOYED_RESOURCES").is_err() {
378            println!("Skipped: ACI live tests require deployed resources");
379            return Ok(());
380        }
381        let ip = env::var("IDENTITY_ACI_IP_USER_ASSIGNED").expect("IDENTITY_ACI_IP_USER_ASSIGNED");
382        let storage_name = env::var("IDENTITY_STORAGE_NAME_USER_ASSIGNED")
383            .expect("IDENTITY_STORAGE_NAME_USER_ASSIGNED");
384        let client_id = env::var("IDENTITY_USER_ASSIGNED_IDENTITY_CLIENT_ID")
385            .expect("IDENTITY_USER_ASSIGNED_IDENTITY_CLIENT_ID");
386        run_deployed_test(
387            &format!("{}:8080", ip),
388            &storage_name,
389            Some(UserAssignedId::ClientId(client_id)),
390        )
391        .await?;
392
393        Ok(())
394    }
395
396    async fn run_app_service_test(options: Option<ManagedIdentityCredentialOptions>) {
397        let endpoint = "http://localhost/metadata/identity/oauth2/token";
398        let x_id_header = "x-id-header";
399        let mut model = Request::new(endpoint.parse().unwrap(), Method::Get);
400        model.insert_header("x-identity-header", x_id_header);
401        let mut params = Vec::from([
402            ("api-version", "2019-08-01"),
403            ("resource", LIVE_TEST_RESOURCE),
404        ]);
405        if let Some(options) = options.as_ref() {
406            if let Some(ref id) = options.user_assigned_id {
407                match id {
408                    UserAssignedId::ClientId(client_id) => {
409                        params.push(("client_id", client_id));
410                    }
411                    UserAssignedId::ObjectId(object_id) => {
412                        params.push(("object_id", object_id));
413                    }
414                    UserAssignedId::ResourceId(resource_id) => {
415                        params.push(("mi_res_id", resource_id));
416                    }
417                }
418            }
419        }
420        model.url_mut().query_pairs_mut().extend_pairs(params);
421        run_supported_source_test(
422            Env::from(
423                &[
424                    (IDENTITY_ENDPOINT, endpoint),
425                    (IDENTITY_HEADER, x_id_header),
426                ][..],
427            ),
428            options,
429            ManagedIdentitySource::AppService,
430            model,
431            format!(
432                r#"{{"access_token":"*","expires_on":"{}","resource":"{}","token_type":"Bearer"}}"#,
433                EXPIRES_ON, LIVE_TEST_RESOURCE
434            )
435            .to_string(),
436        )
437        .await;
438    }
439
440    #[tokio::test]
441    async fn app_service() {
442        run_app_service_test(None).await;
443    }
444
445    #[tokio::test]
446    async fn app_service_client_id() {
447        run_app_service_test(Some(ManagedIdentityCredentialOptions {
448            user_assigned_id: Some(UserAssignedId::ClientId("expected client ID".to_string())),
449            ..Default::default()
450        }))
451        .await;
452    }
453
454    #[tokio::test]
455    async fn app_service_error_response() {
456        run_error_response_test(ManagedIdentitySource::AppService).await
457    }
458
459    #[tokio::test]
460    async fn app_service_object_id() {
461        run_app_service_test(Some(ManagedIdentityCredentialOptions {
462            user_assigned_id: Some(UserAssignedId::ObjectId("expected object ID".to_string())),
463            ..Default::default()
464        }))
465        .await;
466    }
467
468    #[tokio::test]
469    async fn app_service_resource_id() {
470        let result = ManagedIdentityCredential::new(Some(ManagedIdentityCredentialOptions {
471            env: Env::from(&[(IDENTITY_ENDPOINT, "..."), (IDENTITY_HEADER, "x-id-header")][..]),
472            user_assigned_id: Some(UserAssignedId::ResourceId(
473                "expected resource ID".to_string(),
474            )),
475            ..Default::default()
476        }));
477        assert!(
478            matches!(result, Err(ref e) if *e.kind() == azure_core::error::ErrorKind::Credential),
479            "Expected constructor error"
480        );
481    }
482
483    #[test]
484    fn arc() {
485        run_unsupported_source_test(
486            Env::from(
487                &[
488                    (IDENTITY_ENDPOINT, "http://localhost"),
489                    (IMDS_ENDPOINT, "..."),
490                ][..],
491            ),
492            ManagedIdentitySource::AzureArc,
493        );
494    }
495
496    #[test]
497    fn azure_ml() {
498        run_unsupported_source_test(
499            Env::from(&[(MSI_ENDPOINT, "..."), (MSI_SECRET, "...")][..]),
500            ManagedIdentitySource::AzureML,
501        );
502    }
503
504    #[test]
505    fn cloudshell() {
506        run_unsupported_source_test(
507            Env::from(&[(MSI_ENDPOINT, "http://localhost")][..]),
508            ManagedIdentitySource::CloudShell,
509        );
510    }
511
512    async fn run_imds_live_test(id: Option<UserAssignedId>) -> azure_core::Result<()> {
513        if std::env::var("IDENTITY_IMDS_AVAILABLE").is_err() {
514            println!("Skipped: IMDS isn't available");
515            return Ok(());
516        }
517
518        let credential = ManagedIdentityCredential::new(Some(ManagedIdentityCredentialOptions {
519            user_assigned_id: id,
520            ..Default::default()
521        }))
522        .expect("valid credential");
523
524        let token = credential.get_token(LIVE_TEST_SCOPES, None).await?;
525
526        assert!(!token.token.secret().is_empty());
527        assert_eq!(time::UtcOffset::UTC, token.expires_on.offset());
528        assert!(token.expires_on.unix_timestamp() > OffsetDateTime::now_utc().unix_timestamp());
529
530        Ok(())
531    }
532
533    async fn run_imds_test(options: Option<ManagedIdentityCredentialOptions>) {
534        let mut model = Request::new(
535            "http://169.254.169.254/metadata/identity/oauth2/token"
536                .parse()
537                .unwrap(),
538            Method::Get,
539        );
540        model.insert_header("metadata", "true");
541
542        let mut params = Vec::from([
543            ("api-version", "2019-08-01"),
544            ("resource", LIVE_TEST_RESOURCE),
545        ]);
546        if let Some(options) = options.as_ref() {
547            if let Some(ref id) = options.user_assigned_id {
548                match id {
549                    UserAssignedId::ClientId(client_id) => {
550                        params.push(("client_id", client_id));
551                    }
552                    UserAssignedId::ObjectId(object_id) => {
553                        params.push(("object_id", object_id));
554                    }
555                    UserAssignedId::ResourceId(resource_id) => {
556                        params.push(("msi_res_id", resource_id));
557                    }
558                }
559            }
560        }
561        model.url_mut().query_pairs_mut().extend_pairs(params);
562
563        run_supported_source_test(
564            Env::from(&[][..]),
565            options,
566            ManagedIdentitySource::Imds,
567            model,
568            format!(r#"{{"token_type":"Bearer","expires_in":"85770","expires_on":"{}","ext_expires_in":86399,"access_token":"*","resource":"{}"}}"#, EXPIRES_ON, LIVE_TEST_RESOURCE).to_string(),
569        ).await;
570    }
571
572    #[tokio::test]
573    async fn imds_client_id() {
574        run_imds_test(Some(ManagedIdentityCredentialOptions {
575            user_assigned_id: Some(UserAssignedId::ClientId("expected client ID".to_string())),
576            ..Default::default()
577        }))
578        .await;
579    }
580
581    #[tokio::test]
582    async fn imds_error_response() {
583        run_error_response_test(ManagedIdentitySource::Imds).await
584    }
585
586    #[tokio::test]
587    async fn imds_object_id() {
588        run_imds_test(Some(ManagedIdentityCredentialOptions {
589            user_assigned_id: Some(UserAssignedId::ObjectId("expected object ID".to_string())),
590            ..Default::default()
591        }))
592        .await;
593    }
594
595    #[tokio::test]
596    async fn imds_resource_id() {
597        run_imds_test(Some(ManagedIdentityCredentialOptions {
598            user_assigned_id: Some(UserAssignedId::ResourceId(
599                "expected resource ID".to_string(),
600            )),
601            ..Default::default()
602        }))
603        .await;
604    }
605
606    #[tokio::test]
607    async fn imds_system_assigned() {
608        run_imds_test(None).await;
609    }
610
611    #[recorded::test(live)]
612    async fn imds_system_assigned_live() -> azure_core::Result<()> {
613        run_imds_live_test(None).await
614    }
615
616    #[tokio::test]
617    async fn requires_one_scope() {
618        let credential = ManagedIdentityCredential::new(None).expect("valid credential");
619        for scopes in [&[][..], &["A", "B"][..]].iter() {
620            credential
621                .get_token(scopes, None)
622                .await
623                .expect_err("expected an error, got");
624        }
625    }
626
627    #[test]
628    fn service_fabric() {
629        run_unsupported_source_test(
630            Env::from(
631                &[
632                    (IDENTITY_ENDPOINT, "http://localhost"),
633                    (IDENTITY_HEADER, "..."),
634                    (IDENTITY_SERVER_THUMBPRINT, "..."),
635                ][..],
636            ),
637            ManagedIdentitySource::ServiceFabric,
638        );
639    }
640}