Skip to main content

azure_identity/
client_assertion_credential.rs

1// Copyright (c) Microsoft Corporation. All rights reserved.
2// Licensed under the MIT License.
3
4use crate::{get_authority_host, validate_not_empty, validate_tenant_id, TokenCache};
5use azure_core::{
6    credentials::{AccessToken, TokenCredential, TokenRequestOptions},
7    error::{ErrorKind, ResultExt},
8    http::{
9        headers::{self, content_type},
10        ClientMethodOptions, ClientOptions, Method, Pipeline, PipelineSendOptions, Request, Url,
11    },
12};
13use std::{any::type_name, fmt, str, sync::Arc};
14use url::form_urlencoded;
15
16const ASSERTION_TYPE: &str = "urn:ietf:params:oauth:client-assertion-type:jwt-bearer";
17
18/// Authenticates an application with client assertions.
19///
20/// This credential is for advanced scenarios. `ClientCertificateCredential` has a more convenient API for
21/// the most common assertion scenario, authenticating a service principal with a certificate.
22///
23/// See
24/// [Entra ID documentation](https://learn.microsoft.com/entra/identity-platform/certificate-credentials#assertion-format)
25/// for details of the assertion format.
26pub struct ClientAssertionCredential<C> {
27    name: &'static str,
28    client_id: String,
29    endpoint: Url,
30    assertion: C,
31    cache: TokenCache,
32    pipeline: Pipeline,
33}
34
35impl<C> fmt::Debug for ClientAssertionCredential<C> {
36    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
37        f.debug_struct(type_name::<Self>())
38            .field("client_id", &self.client_id)
39            .field("endpoint", &self.endpoint)
40            .finish_non_exhaustive()
41    }
42}
43
44/// Options for constructing a new [`ClientAssertionCredential`].
45#[derive(Default)]
46pub struct ClientAssertionCredentialOptions {
47    /// Options for the credential's HTTP pipeline.
48    pub client_options: ClientOptions,
49}
50
51impl fmt::Debug for ClientAssertionCredentialOptions {
52    fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
53        f.debug_struct(type_name::<Self>()).finish_non_exhaustive()
54    }
55}
56
57#[async_trait::async_trait]
58/// Represents an entity capable of supplying a client assertion.
59pub trait ClientAssertion: Send + Sync + fmt::Debug {
60    /// Supply the client assertion secret.
61    async fn secret(&self, options: Option<ClientMethodOptions<'_>>) -> azure_core::Result<String>;
62}
63
64impl<C: ClientAssertion> ClientAssertionCredential<C> {
65    /// Create a new `ClientAssertionCredential`.
66    ///
67    /// # Arguments
68    /// - `tenant_id`: The tenant (directory) ID of the service principal.
69    /// - `client_id`: The client (application) ID of the service principal.
70    /// - `assertion`: an implementation of [`ClientAssertion`] that provides assertions to the credential.
71    /// - `options`: Options for configuring the credential. If `None`, the credential uses its default options.
72    ///
73    pub fn new(
74        tenant_id: String,
75        client_id: String,
76        assertion: C,
77        options: Option<ClientAssertionCredentialOptions>,
78    ) -> azure_core::Result<Arc<Self>> {
79        Ok(Arc::new(Self::new_exclusive(
80            tenant_id,
81            client_id,
82            assertion,
83            stringify!(ClientAssertionCredential),
84            options,
85        )?))
86    }
87
88    /// Create a new `ClientAssertionCredential` without wrapping it in an
89    /// `Arc`. Intended for use by other credentials in the crate that will
90    /// themselves be protected by an `Arc`.
91    pub(crate) fn new_exclusive(
92        tenant_id: String,
93        client_id: String,
94        assertion: C,
95        name: &'static str,
96        options: Option<ClientAssertionCredentialOptions>,
97    ) -> azure_core::Result<Self> {
98        validate_tenant_id(&tenant_id)?;
99        validate_not_empty(&client_id, "no client ID specified")?;
100        let options = options.unwrap_or_default();
101        let authority_host = get_authority_host(None, options.client_options.cloud.as_deref())?;
102        let endpoint = authority_host
103            .join(&format!("/{tenant_id}/oauth2/v2.0/token"))
104            .with_context_fn(ErrorKind::DataConversion, || {
105                format!("tenant_id {tenant_id} could not be URL encoded")
106            })?;
107        let pipeline = Pipeline::new(
108            option_env!("CARGO_PKG_NAME"),
109            option_env!("CARGO_PKG_VERSION"),
110            options.client_options,
111            Vec::default(),
112            Vec::default(),
113            None,
114        );
115        Ok(Self {
116            name,
117            client_id,
118            assertion,
119            endpoint,
120            cache: TokenCache::new(),
121            pipeline,
122        })
123    }
124
125    async fn get_token_impl(
126        &self,
127        scopes: &[&str],
128        options: Option<TokenRequestOptions<'_>>,
129    ) -> azure_core::Result<AccessToken> {
130        let mut req = Request::new(self.endpoint.clone(), Method::Post);
131        req.insert_header(
132            headers::CONTENT_TYPE,
133            content_type::APPLICATION_X_WWW_FORM_URLENCODED,
134        );
135        let options = options.unwrap_or_default();
136        let assertion = self
137            .assertion
138            .secret(Some(options.method_options.to_owned()))
139            .await?;
140        let encoded: String = form_urlencoded::Serializer::new(String::new())
141            .append_pair("client_assertion", assertion.as_str())
142            .append_pair("client_assertion_type", ASSERTION_TYPE)
143            .append_pair("client_id", self.client_id.as_str())
144            .append_pair("grant_type", "client_credentials")
145            .append_pair("scope", &scopes.join(" "))
146            .finish();
147        req.set_body(encoded);
148
149        let ctx = options.method_options.context.to_borrowed();
150        let res = self
151            .pipeline
152            .send(
153                &ctx,
154                &mut req,
155                Some(PipelineSendOptions {
156                    skip_checks: true,
157                    ..Default::default()
158                }),
159            )
160            .await?;
161
162        crate::handle_entra_response(res)
163    }
164}
165
166#[async_trait::async_trait]
167impl<C: ClientAssertion> TokenCredential for ClientAssertionCredential<C> {
168    async fn get_token(
169        &self,
170        scopes: &[&str],
171        options: Option<TokenRequestOptions<'_>>,
172    ) -> azure_core::Result<AccessToken> {
173        self.cache
174            .get_token(scopes, options, |s, o| self.get_token_impl(s, o))
175            .await
176            .map_err(|err| crate::authentication_error(self.name, err))
177    }
178}
179
180#[cfg(test)]
181pub(crate) mod tests {
182    use super::*;
183    use crate::tests::*;
184    use azure_core::{
185        http::{
186            headers::{self, content_type, Headers},
187            AsyncRawResponse, Body, Method, RawResponse, Request, StatusCode, Transport,
188        },
189        Bytes,
190    };
191    use std::{collections::HashMap, time::SystemTime};
192    use time::UtcOffset;
193    use url::form_urlencoded;
194
195    pub const FAKE_ASSERTION: &str = "fake assertion";
196
197    pub fn is_valid_request(
198        expected_authority: String,
199        expected_assertion: Option<String>,
200    ) -> impl Fn(&Request) -> azure_core::Result<()> {
201        let expected_url = format!("{expected_authority}/oauth2/v2.0/token");
202        move |req: &Request| {
203            assert_eq!(Method::Post, req.method());
204            assert_eq!(expected_url, req.url().to_string());
205            assert_eq!(
206                content_type::APPLICATION_X_WWW_FORM_URLENCODED.as_str(),
207                req.headers().get_str(&headers::CONTENT_TYPE).unwrap()
208            );
209            let body = match req.body() {
210                Body::Bytes(bytes) => str::from_utf8(bytes).unwrap(),
211                _ => panic!("unexpected body type"),
212            };
213            let actual_params: HashMap<String, String> = form_urlencoded::parse(body.as_bytes())
214                .map(|(k, v)| (k.to_string(), v.to_string()))
215                .collect();
216            let assertion = actual_params
217                .get("client_assertion")
218                .expect("request body should contain client_assertion");
219            match &expected_assertion {
220                Some(expected) => assert_eq!(expected, assertion),
221                None => assert!(
222                    !assertion.is_empty(),
223                    "expected client_assertion to be present"
224                ),
225            }
226            let expected_params = [
227                ("client_assertion_type", ASSERTION_TYPE),
228                ("client_id", FAKE_CLIENT_ID),
229                ("grant_type", "client_credentials"),
230                ("scope", &LIVE_TEST_SCOPES.join(" ")),
231            ];
232            for (key, value) in expected_params.iter() {
233                assert_eq!(
234                    *value,
235                    actual_params
236                        .get(*key)
237                        .unwrap_or_else(|| panic!("no {} in request body", key))
238                );
239            }
240            Ok(())
241        }
242    }
243
244    #[derive(Debug)]
245    struct MockAssertion {}
246
247    #[async_trait::async_trait]
248    impl ClientAssertion for MockAssertion {
249        async fn secret(&self, _: Option<ClientMethodOptions<'_>>) -> azure_core::Result<String> {
250            Ok(FAKE_ASSERTION.to_string())
251        }
252    }
253
254    #[tokio::test]
255    async fn get_token_error() {
256        let body = Bytes::from(
257            r#"{"error":"invalid_request","error_description":"error description from the response","error_codes":[50027],"timestamp":"2025-04-18 16:04:37Z","trace_id":"...","correlation_id":"...","error_uri":"https://login.microsoftonline.com/error?code=50027"}"#,
258        );
259        let mut headers = Headers::default();
260        headers.insert("key", "value");
261        let expected_status = StatusCode::BadRequest;
262        let expected_response =
263            RawResponse::from_bytes(expected_status, headers.clone(), body.clone());
264        let mock_response = AsyncRawResponse::from_bytes(expected_status, headers, body);
265
266        let mock = MockSts::new(
267            vec![mock_response],
268            Some(Arc::new(is_valid_request(
269                FAKE_PUBLIC_CLOUD_AUTHORITY.to_string(),
270                Some(FAKE_ASSERTION.to_string()),
271            ))),
272        );
273        let credential = ClientAssertionCredential::new(
274            FAKE_TENANT_ID.to_string(),
275            FAKE_CLIENT_ID.to_string(),
276            MockAssertion {},
277            Some(ClientAssertionCredentialOptions {
278                client_options: ClientOptions {
279                    transport: Some(Transport::new(Arc::new(mock))),
280                    ..Default::default()
281                },
282            }),
283        )
284        .expect("valid credential");
285
286        let err = credential
287            .get_token(LIVE_TEST_SCOPES, None)
288            .await
289            .expect_err("authentication error");
290        assert!(matches!(err.kind(), ErrorKind::Credential));
291        assert_eq!(
292            "ClientAssertionCredential authentication failed. error description from the response",
293            err.to_string(),
294        );
295        match err
296            .downcast_ref::<azure_core::Error>()
297            .expect("returned error should wrap an azure_core::Error")
298            .kind()
299        {
300            ErrorKind::HttpResponse {
301                error_code: Some(error_code),
302                raw_response: Some(response),
303                status,
304            } => {
305                assert_eq!("50027", error_code);
306                assert_eq!(&expected_response, response.as_ref());
307                assert_eq!(expected_status, *status);
308            }
309            err => panic!("unexpected {:?}", err),
310        };
311    }
312
313    #[tokio::test]
314    async fn get_token_success() {
315        let mock = MockSts::new(
316            vec![token_response()],
317            Some(Arc::new(is_valid_request(
318                FAKE_PUBLIC_CLOUD_AUTHORITY.to_string(),
319                Some(FAKE_ASSERTION.to_string()),
320            ))),
321        );
322        let credential = ClientAssertionCredential::new(
323            FAKE_TENANT_ID.to_string(),
324            FAKE_CLIENT_ID.to_string(),
325            MockAssertion {},
326            Some(ClientAssertionCredentialOptions {
327                client_options: ClientOptions {
328                    transport: Some(Transport::new(Arc::new(mock))),
329                    ..Default::default()
330                },
331            }),
332        )
333        .expect("valid credential");
334
335        let token = credential
336            .get_token(LIVE_TEST_SCOPES, None)
337            .await
338            .expect("token");
339        assert_eq!(FAKE_TOKEN, token.token.secret());
340        assert!(token.expires_on > SystemTime::now());
341        assert_eq!(UtcOffset::UTC, token.expires_on.offset());
342
343        // MockSts will return an error if the credential sends another request
344        let cached_token = credential
345            .get_token(LIVE_TEST_SCOPES, None)
346            .await
347            .expect("cached token");
348        assert_eq!(token.token.secret(), cached_token.token.secret());
349        assert_eq!(token.expires_on, cached_token.expires_on);
350    }
351
352    #[tokio::test]
353    async fn cloud_configuration() {
354        for (cloud, expected_authority) in cloud_configuration_cases() {
355            let mock = MockSts::new(
356                vec![token_response()],
357                Some(Arc::new(is_valid_request(
358                    expected_authority,
359                    Some(FAKE_ASSERTION.to_string()),
360                ))),
361            );
362            let credential = ClientAssertionCredential::new(
363                FAKE_TENANT_ID.to_string(),
364                FAKE_CLIENT_ID.to_string(),
365                MockAssertion {},
366                Some(ClientAssertionCredentialOptions {
367                    client_options: ClientOptions {
368                        transport: Some(Transport::new(Arc::new(mock))),
369                        cloud: Some(Arc::new(cloud)),
370                        ..Default::default()
371                    },
372                }),
373            )
374            .expect("valid credential");
375
376            credential
377                .get_token(LIVE_TEST_SCOPES, None)
378                .await
379                .expect("token");
380        }
381    }
382}