Skip to main content

azure_identity/
lib.rs

1// Copyright (c) Microsoft Corporation. All rights reserved.
2// Licensed under the MIT License.
3
4#![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    // these are i64 to avoid conversion when calling Duration::seconds
66    // (real values are unsigned)
67    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        // need this arm because CloudConfiguration is non-exhaustive
139        _ => {
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
160/// Map an error from a credential's get_token() method to an ErrorKind::Credential error, appending
161/// a link to the troubleshooting guide entry for that credential, if it has one.
162fn 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)) // Use remove(0) to return responses in the correct order
411            }
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/"); // Url::parse adds the trailing slash
458    }
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}