1use 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#[derive(Debug, Clone)]
15#[non_exhaustive]
16pub enum UserAssignedId {
17 ClientId(String),
19 ObjectId(String),
21 ResourceId(String),
23}
24
25pub 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#[derive(Clone, Default)]
38pub struct ManagedIdentityCredentialOptions {
39 pub user_assigned_id: Option<UserAssignedId>,
42
43 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 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 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 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}