1use 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
18pub 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#[derive(Default)]
46pub struct ClientAssertionCredentialOptions {
47 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]
58pub trait ClientAssertion: Send + Sync + fmt::Debug {
60 async fn secret(&self, options: Option<ClientMethodOptions<'_>>) -> azure_core::Result<String>;
62}
63
64impl<C: ClientAssertion> ClientAssertionCredential<C> {
65 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 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 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}