1#![doc = include_str!("../README.md")]
5#![cfg_attr(docsrs, feature(doc_cfg))]
6#![warn(missing_docs)]
7
8mod app_service_managed_identity_credential;
9mod azure_cli_credential;
10mod azure_developer_cli_credential;
11mod azure_pipelines_credential;
12mod cache;
13mod client_assertion_credential;
14#[cfg(feature = "client_certificate")]
15mod client_certificate_credential;
16mod client_secret_credential;
17mod developer_tools_credential;
18mod env;
19mod imds_managed_identity_credential;
20mod managed_identity_credential;
21mod process;
22mod virtual_machine_managed_identity_credential;
23mod workload_identity_credential;
24
25pub use azure_cli_credential::*;
26pub use azure_developer_cli_credential::*;
27pub use azure_pipelines_credential::*;
28pub use client_assertion_credential::*;
29#[cfg(feature = "client_certificate")]
30pub use client_certificate_credential::*;
31pub use client_secret_credential::*;
32pub use developer_tools_credential::*;
33pub use managed_identity_credential::*;
34pub use process::{new_executor, Executor};
35pub use workload_identity_credential::*;
36
37pub(crate) use app_service_managed_identity_credential::*;
38pub(crate) use cache::TokenCache;
39pub(crate) use imds_managed_identity_credential::*;
40pub(crate) use virtual_machine_managed_identity_credential::*;
41
42use crate::env::Env;
43use azure_core::{
44 cloud::CloudConfiguration,
45 credentials::AccessToken,
46 error::ErrorKind,
47 http::{RawResponse, Url},
48 time::{Duration, OffsetDateTime},
49 Error, Result,
50};
51use serde::Deserialize;
52use std::borrow::Cow;
53
54#[derive(Debug, Default, Deserialize)]
55#[serde(default)]
56struct EntraIdErrorResponse<'a> {
57 error_codes: Vec<i32>,
58 error_description: &'a str,
59}
60
61#[derive(Debug, Default, Deserialize)]
62#[serde(default)]
63struct EntraIdTokenResponse {
64 token_type: String,
65 expires_in: i64,
68 ext_expires_in: i64,
69 access_token: String,
70}
71
72fn deserialize<'a, T>(res: &'a RawResponse) -> Result<T>
73where
74 T: serde::Deserialize<'a>,
75{
76 serde_json::from_slice(res.body().as_ref()).map_err(Into::into)
77}
78
79fn handle_entra_response(response: RawResponse) -> Result<AccessToken> {
80 let status = response.status();
81 if status.is_success() {
82 let token_response: EntraIdTokenResponse = deserialize(&response)?;
83 return Ok(AccessToken::new(
84 token_response.access_token,
85 OffsetDateTime::now_utc() + Duration::seconds(token_response.expires_in),
86 ));
87 }
88
89 let error_response: EntraIdErrorResponse<'_> = deserialize(&response)?;
90 let error_code = if error_response.error_codes.is_empty() {
91 None
92 } else {
93 Some(
94 error_response
95 .error_codes
96 .iter()
97 .map(i32::to_string)
98 .collect::<Vec<_>>()
99 .join(","),
100 )
101 };
102 let error_description = error_response.error_description.to_owned();
103
104 Err(Error::new(
105 ErrorKind::HttpResponse {
106 status,
107 error_code,
108 raw_response: Some(Box::new(response)),
109 },
110 error_description,
111 ))
112}
113
114fn validate_not_empty<C>(value: &str, message: C) -> Result<()>
115where
116 C: Into<Cow<'static, str>>,
117{
118 if value.is_empty() {
119 return Err(Error::with_message(ErrorKind::Credential, message));
120 }
121
122 Ok(())
123}
124
125const AZURE_AUTHORITY_HOST_ENV_KEY: &str = "AZURE_AUTHORITY_HOST";
126const AZURE_PUBLIC_CLOUD: &str = "https://login.microsoftonline.com";
127
128fn get_authority_host(env: Option<Env>, cloud: Option<&CloudConfiguration>) -> Result<Url> {
129 let authority_host = match cloud {
130 None => env
131 .unwrap_or_default()
132 .var(AZURE_AUTHORITY_HOST_ENV_KEY)
133 .unwrap_or_else(|_| AZURE_PUBLIC_CLOUD.to_string()),
134 Some(CloudConfiguration::Custom(config)) => config.authority_host.clone(),
135 Some(CloudConfiguration::AzureGovernment) => "https://login.microsoftonline.us".to_string(),
136 Some(CloudConfiguration::AzureChina) => "https://login.chinacloudapi.cn".to_string(),
137 Some(CloudConfiguration::AzurePublic) => AZURE_PUBLIC_CLOUD.to_string(),
138 _ => {
140 return Err(Error::with_message(
141 ErrorKind::Other,
142 format!("unexpected cloud configuration: {:?}", cloud),
143 ))
144 }
145 };
146
147 let url = Url::parse(&authority_host)?;
148 if url.scheme() != "https" {
149 return Err(Error::with_message(
150 ErrorKind::Other,
151 format!("authority host doesn't use HTTPS scheme: {authority_host}"),
152 ));
153 }
154 Ok(url)
155}
156
157const TSG_LINK_ERROR_TEXT: &str =
158 "\nTo troubleshoot, visit https://aka.ms/azsdk/rust/identity/troubleshoot";
159
160fn authentication_error(credential_name: &str, err: Error) -> Error {
163 let link_fragment = match credential_name {
164 stringify!(AzureCliCredential) => "#azure-cli",
165 stringify!(AzureDeveloperCliCredential) => "#azd",
166 stringify!(AzurePipelinesCredential) => "#apc",
167 stringify!(ClientCertificateCredential) => "#client-cert",
168 stringify!(ClientSecretCredential) => "#client-secret",
169 stringify!(ManagedIdentityCredential) => "#managed-id",
170 stringify!(WorkloadIdentityCredential) => "#workload",
171 _ => "",
172 };
173 const WHITESPACE: &[char; 3] = &['\t', '\x0c', ' '];
174
175 let err_str = err.to_string();
176 let err_str = err_str.trim_matches(WHITESPACE);
177 let separator = if err_str.starts_with('\n') { "" } else { " " };
178 let mut message = format!("{credential_name} authentication failed.{separator}{err_str}");
179 if !link_fragment.is_empty() {
180 message.push_str(TSG_LINK_ERROR_TEXT);
181 message.push_str(link_fragment);
182 }
183 Error::with_error(ErrorKind::Credential, err, message)
184}
185
186#[test]
187fn test_validate_not_empty() {
188 assert!(validate_not_empty("", "it's empty").is_err());
189 assert!(validate_not_empty(" ", "it's not empty").is_ok());
190 assert!(validate_not_empty("not empty", "it's not empty").is_ok());
191}
192
193fn validate_scope(scope: &str) -> Result<()> {
194 if scope.is_empty()
195 || !scope.chars().all(|c| {
196 c.is_alphanumeric() || c == '.' || c == '-' || c == '_' || c == ':' || c == '/'
197 })
198 {
199 return Err(Error::with_message(
200 ErrorKind::Credential,
201 format!("invalid scope {scope}"),
202 ));
203 }
204
205 Ok(())
206}
207
208#[test]
209fn test_validate_scope() {
210 assert!(validate_scope("").is_err());
211 assert!(validate_scope("invalid_scope@id").is_err());
212 assert!(validate_scope("A-1b_2c:3d/4.z").is_ok());
213 assert!(validate_scope("http://vault.azure.net").is_ok());
214}
215
216fn validate_subscription(subscription: &str) -> Result<()> {
217 if subscription.is_empty()
218 || !subscription
219 .chars()
220 .all(|c| c.is_alphanumeric() || c == '.' || c == '-' || c == '_' || c == ' ')
221 {
222 return Err(Error::with_message(
223 ErrorKind::Credential,
224 format!("invalid subscription {subscription}. If this is the name of a subscription, use its ID instead"),
225 ));
226 }
227
228 Ok(())
229}
230
231#[test]
232fn test_validate_subscription() {
233 assert!(validate_subscription("").is_err());
234 assert!(validate_subscription("invalid_subscription@id").is_err());
235 assert!(validate_subscription("A-1b_2c 3.z").is_ok());
236 assert!(validate_subscription("7b795fb9-09d3-42f4-a494-38864f99ba3c").is_ok());
237}
238
239fn validate_tenant_id(tenant_id: &str) -> Result<()> {
240 if tenant_id.is_empty()
241 || !tenant_id
242 .chars()
243 .all(|c| c.is_alphanumeric() || c == '.' || c == '-')
244 {
245 return Err(Error::with_message(
246 ErrorKind::Credential,
247 format!("invalid tenant ID {tenant_id}. You can locate your tenantID by following the instructions listed here: https://learn.microsoft.com/partner-center/find-ids-and-domain-names"),
248 ));
249 }
250
251 Ok(())
252}
253
254#[test]
255fn test_validate_tenant_id() {
256 assert!(validate_tenant_id("").is_err());
257 assert!(validate_tenant_id("invalid_tenant@id").is_err());
258 assert!(validate_tenant_id("A-1.z").is_ok());
259 assert!(validate_tenant_id("7b795fb9-09d3-42f4-a494-38864f99ba3c").is_ok());
260}
261
262#[cfg(test)]
263mod tests {
264 use super::*;
265 use crate::{env::Env, process::Executor};
266 use async_trait::async_trait;
267 use azure_core::{
268 cloud::{CloudConfiguration, CustomConfiguration},
269 error::ErrorKind,
270 http::{headers::Headers, AsyncRawResponse, RawResponse, Request, StatusCode},
271 Bytes, Error, Result,
272 };
273 use std::{
274 ffi::OsStr,
275 process::Output,
276 sync::{
277 atomic::{AtomicUsize, Ordering},
278 Arc, Mutex,
279 },
280 };
281
282 pub const FAKE_CLIENT_ID: &str = "fake-client";
283 pub const FAKE_PUBLIC_CLOUD_AUTHORITY: &str = "https://login.microsoftonline.com/fake-tenant";
284 pub const FAKE_TENANT_ID: &str = "fake-tenant";
285 pub const FAKE_TOKEN: &str = "***";
286 pub const LIVE_TEST_RESOURCE: &str = "https://management.azure.com";
287 pub const LIVE_TEST_SCOPES: &[&str] = &["https://management.azure.com/.default"];
288
289 pub type RunCallback = Arc<dyn Fn(&OsStr, &[&OsStr]) + Send + Sync>;
290
291 #[derive(Default)]
292 pub struct MockExecutor {
293 call_count: AtomicUsize,
294 error: Option<std::io::Error>,
295 on_run: Option<RunCallback>,
296 output: Mutex<Option<Output>>,
297 }
298
299 impl std::fmt::Debug for MockExecutor {
300 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
301 f.debug_struct("MockExecutor").finish()
302 }
303 }
304
305 impl MockExecutor {
306 pub fn with_error(err: std::io::Error) -> Arc<Self> {
307 Arc::new(Self {
308 error: Some(err),
309 ..Default::default()
310 })
311 }
312
313 pub fn with_output(
314 exit_code: i32,
315 stdout: &str,
316 stderr: &str,
317 on_run: Option<RunCallback>,
318 ) -> Arc<Self> {
319 let output = Output {
320 status: {
321 #[cfg(windows)]
322 {
323 std::os::windows::process::ExitStatusExt::from_raw(
324 exit_code.try_into().unwrap(),
325 )
326 }
327 #[cfg(unix)]
328 {
329 std::os::unix::process::ExitStatusExt::from_raw(exit_code)
330 }
331 },
332 stdout: stdout.as_bytes().to_vec(),
333 stderr: stderr.as_bytes().to_vec(),
334 };
335 Arc::new(Self {
336 on_run,
337 output: Mutex::new(Some(output)),
338 call_count: AtomicUsize::new(0),
339 ..Default::default()
340 })
341 }
342
343 pub fn call_count(&self) -> usize {
344 self.call_count.load(Ordering::SeqCst)
345 }
346 }
347
348 #[async_trait]
349 impl Executor for MockExecutor {
350 async fn run(&self, program: &OsStr, args: &[&OsStr]) -> std::io::Result<Output> {
351 self.call_count.fetch_add(1, Ordering::SeqCst);
352
353 if let Some(on_run) = &self.on_run {
354 on_run(program, args);
355 }
356 if let Some(err) = &self.error {
357 return Err(std::io::Error::new(err.kind(), err.to_string()));
358 }
359 let output = self.output.lock().unwrap();
360 match output.as_ref() {
361 Some(output) => Ok(output.clone()),
362 None => panic!("MockExecutor output not configured"),
363 }
364 }
365 }
366
367 pub fn token_response() -> AsyncRawResponse {
368 AsyncRawResponse::from_bytes(
369 StatusCode::Ok,
370 Headers::default(),
371 Bytes::from(format!(
372 r#"{{"access_token":"{FAKE_TOKEN}","expires_in":3600,"token_type":"Bearer"}}"#,
373 )),
374 )
375 }
376
377 pub type RequestCallback = Arc<dyn Fn(&Request) -> Result<()> + Send + Sync>;
378
379 pub struct MockSts {
380 responses: Mutex<Vec<AsyncRawResponse>>,
381 on_request: Option<RequestCallback>,
382 }
383
384 impl std::fmt::Debug for MockSts {
385 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
386 f.debug_struct("MockSts").finish()
387 }
388 }
389
390 impl MockSts {
391 pub fn new(responses: Vec<AsyncRawResponse>, on_request: Option<RequestCallback>) -> Self {
392 Self {
393 responses: Mutex::new(responses),
394 on_request,
395 }
396 }
397 }
398
399 #[async_trait::async_trait]
400 impl azure_core::http::HttpClient for MockSts {
401 async fn execute_request(&self, request: &Request) -> Result<AsyncRawResponse> {
402 self.on_request.as_ref().map_or(Ok(()), |f| f(request))?;
403 let mut responses = self.responses.lock().unwrap();
404 if responses.is_empty() {
405 Err(Error::with_message(
406 ErrorKind::Other,
407 "No more mock responses",
408 ))
409 } else {
410 Ok(responses.remove(0)) }
412 }
413 }
414
415 pub fn cloud_configuration_cases() -> Vec<(CloudConfiguration, String)> {
416 let custom_host = "https://login.contoso.local/".to_string();
417
418 let mut custom_no_trailing_slash = CustomConfiguration::default();
419 custom_no_trailing_slash.authority_host = custom_host.trim_end_matches('/').to_string();
420
421 let mut custom_trailing_slash = CustomConfiguration::default();
422 custom_trailing_slash.authority_host = custom_host;
423
424 vec![
425 (
426 CloudConfiguration::AzurePublic,
427 FAKE_PUBLIC_CLOUD_AUTHORITY.to_string(),
428 ),
429 (
430 CloudConfiguration::AzureGovernment,
431 format!("https://login.microsoftonline.us/{FAKE_TENANT_ID}"),
432 ),
433 (
434 CloudConfiguration::AzureChina,
435 format!("https://login.chinacloudapi.cn/{FAKE_TENANT_ID}"),
436 ),
437 (
438 CloudConfiguration::Custom(custom_trailing_slash),
439 format!("https://login.contoso.local/{FAKE_TENANT_ID}"),
440 ),
441 (
442 CloudConfiguration::Custom(custom_no_trailing_slash),
443 format!("https://login.contoso.local/{FAKE_TENANT_ID}"),
444 ),
445 ]
446 }
447
448 #[test]
449 fn cloud_configuration_overrides_env() {
450 let mut config = CustomConfiguration::default();
451 config.authority_host = "https://custom".to_string();
452 let cloud = CloudConfiguration::Custom(config);
453
454 let env = Env::from(&[(crate::AZURE_AUTHORITY_HOST_ENV_KEY, "https://env")][..]);
455
456 let authority = get_authority_host(Some(env), Some(&cloud)).unwrap();
457 assert_eq!(authority.as_str(), "https://custom/"); }
459
460 #[test]
461 fn insecure_authority_host() {
462 let authority_host = "http://insecure";
463 let env = Env::from(&[(crate::AZURE_AUTHORITY_HOST_ENV_KEY, authority_host)][..]);
464 let err = get_authority_host(Some(env), None).unwrap_err();
465 assert!(err.to_string().contains("HTTPS"));
466
467 let mut config = CustomConfiguration::default();
468 config.authority_host = authority_host.to_string();
469 let cloud = CloudConfiguration::Custom(config);
470 let err = get_authority_host(None, Some(&cloud)).unwrap_err();
471 assert!(err.to_string().contains("HTTPS"));
472 }
473
474 #[test]
475 fn entra_error() {
476 let response = RawResponse::from_bytes(
477 StatusCode::BadRequest,
478 Headers::default(),
479 Bytes::from_static(br#"{"error_codes":[123,456],"error_description":"bad news"}"#),
480 );
481
482 let err = handle_entra_response(response).unwrap_err();
483 match err.kind() {
484 ErrorKind::HttpResponse {
485 status,
486 error_code,
487 raw_response,
488 } => {
489 assert_eq!(*status, StatusCode::BadRequest);
490 assert_eq!(error_code.as_deref(), Some("123,456"));
491 assert!(raw_response.is_some());
492 }
493 other => panic!("unexpected error kind: {other:?}"),
494 }
495
496 let inner = err.into_inner().expect("expected inner error");
497 assert_eq!(inner.to_string(), "bad news");
498 }
499}