1use crate::{
5 credentials::{AccessToken, TokenCredential, TokenRequestOptions},
6 error::ErrorKind,
7 http::{
8 headers::{Headers, AUTHORIZATION, WWW_AUTHENTICATE},
9 policies::{Policy, PolicyResult, ERROR_TYPE_ATTRIBUTE},
10 ClientMethodOptions, Context, Request, StatusCode,
11 },
12 time::{Duration, OffsetDateTime},
13 tracing::Span,
14 Error, Result,
15};
16use async_lock::RwLock;
17use async_trait::async_trait;
18use std::sync::Arc;
19
20#[derive(Debug, Clone)]
22pub struct BearerTokenAuthorizationPolicy {
23 authorizer: Arc<BearerTokenAuthorizer>,
24 on_request: Arc<dyn OnRequest>,
25 on_challenge: Option<Arc<dyn OnChallenge>>,
26}
27
28impl BearerTokenAuthorizationPolicy {
29 pub fn new<A, B>(credential: Arc<dyn TokenCredential>, scopes: A) -> Self
31 where
32 A: IntoIterator<Item = B>,
33 B: Into<String>,
34 {
35 let scopes: Vec<String> = scopes.into_iter().map(|s| s.into()).collect();
36 Self {
37 authorizer: Arc::new(BearerTokenAuthorizer::new(credential)),
38 on_request: Arc::new(DefaultOnRequest { scopes }),
39 on_challenge: None,
40 }
41 }
42
43 pub fn with_on_request(mut self, on_request: Arc<dyn OnRequest>) -> Self {
48 self.on_request = on_request;
49 self
50 }
51
52 pub fn with_on_challenge(mut self, on_challenge: Arc<dyn OnChallenge>) -> Self {
57 self.on_challenge = Some(on_challenge);
58 self
59 }
60}
61
62#[async_trait]
63impl Policy for BearerTokenAuthorizationPolicy {
64 async fn send(
65 &self,
66 ctx: &Context,
67 request: &mut Request,
68 next: &[Arc<dyn Policy>],
69 ) -> PolicyResult {
70 if request.url().scheme() != "https" {
71 return Err(Error::with_message(
72 ErrorKind::Other,
73 "authorized requests are not permitted for non-TLS protected endpoints",
74 ));
75 }
76
77 let mut ctx = ctx.to_borrowed();
78 self.on_request
79 .on_request(&mut ctx, request, self.authorizer.as_ref())
80 .await?;
81
82 let mut response = next[0].send(&ctx, request, &next[1..]).await?;
83
84 if response.status() == StatusCode::Unauthorized {
85 self.authorizer.invalidate_cache().await;
86 if let Some(ref callback) = self.on_challenge {
87 if response.headers().get_str(&WWW_AUTHENTICATE).is_ok() {
88 callback
89 .on_challenge(&ctx, request, self.authorizer.as_ref(), response.headers())
90 .await?;
91 request.body_mut().reset().await?;
92 if let Some(span) = ctx.value::<Arc<dyn Span>>() {
93 if span.is_recording() {
95 span.set_attribute(
96 ERROR_TYPE_ATTRIBUTE,
97 response.status().to_string().into(),
98 );
99 }
100 }
101 response = next[0].send(&ctx, request, &next[1..]).await?
102 }
103 }
104 }
105
106 Ok(response)
107 }
108}
109
110#[async_trait]
112pub trait OnChallenge: std::fmt::Debug + Send + Sync {
113 async fn on_challenge(
128 &self,
129 context: &Context,
130 request: &mut Request,
131 authorizer: &dyn Authorizer,
132 headers: &Headers,
133 ) -> Result<()>;
134}
135
136#[async_trait]
138pub trait OnRequest: std::fmt::Debug + Send + Sync {
139 async fn on_request(
150 &self,
151 context: &mut Context,
152 request: &mut Request,
153 authorizer: &dyn Authorizer,
154 ) -> Result<()>;
155}
156
157#[async_trait]
160pub trait Authorizer: crate::private::Sealed + std::fmt::Debug + Send + Sync {
161 async fn authorize(
163 &self,
164 request: &mut Request,
165 scopes: &[&str],
166 options: TokenRequestOptions<'_>,
167 ) -> Result<()>;
168}
169
170#[derive(Debug)]
171struct BearerTokenAuthorizer {
172 access_token: Arc<RwLock<Option<AccessToken>>>,
173 credential: Arc<dyn TokenCredential>,
174}
175
176impl BearerTokenAuthorizer {
177 fn new(credential: Arc<dyn TokenCredential>) -> Self {
178 Self {
179 access_token: Arc::new(RwLock::new(None)),
180 credential,
181 }
182 }
183
184 async fn invalidate_cache(&self) {
185 let mut access_token = self.access_token.write().await;
186 *access_token = None;
187 }
188}
189
190impl crate::private::Sealed for BearerTokenAuthorizer {}
191
192#[async_trait]
193impl Authorizer for BearerTokenAuthorizer {
194 async fn authorize(
195 &self,
196 request: &mut Request,
197 scopes: &[&str],
198 options: TokenRequestOptions<'_>,
199 ) -> Result<()> {
200 let access_token = self.access_token.read().await;
201 match access_token.as_ref() {
202 None => {
203 drop(access_token);
205 let mut access_token = self.access_token.write().await;
206 if access_token.is_none() {
207 *access_token = Some(self.credential.get_token(scopes, Some(options)).await?);
208 }
209 }
210 Some(token) if should_refresh(&token.expires_on) => {
211 let expires_on = token.expires_on;
214 drop(access_token);
215 let mut access_token = self.access_token.write().await;
216 if access_token.is_none() || access_token.as_ref().unwrap().expires_on == expires_on
218 {
219 match self.credential.get_token(scopes, Some(options)).await {
220 Ok(new_token) => {
221 *access_token = Some(new_token);
222 }
223 Err(e)
224 if access_token.is_none()
225 || expires_on <= OffsetDateTime::now_utc() =>
226 {
227 return Err(e);
229 }
230 Err(_) => {
231 }
233 }
234 }
235 }
236 Some(_) => {
237 drop(access_token); }
240 }
241
242 let access_token = self.access_token.read().await;
243 let token = access_token
244 .as_ref()
245 .ok_or_else(|| {
246 Error::with_message(
247 ErrorKind::Credential,
248 "The request failed due to an error while fetching the access token.",
249 )
250 })?
251 .token
252 .secret();
253 request.insert_header(AUTHORIZATION, format!("Bearer {token}"));
254
255 Ok(())
256 }
257}
258
259fn should_refresh(expires_on: &OffsetDateTime) -> bool {
260 *expires_on <= OffsetDateTime::now_utc() + Duration::minutes(5)
261}
262
263#[derive(Debug, Default)]
264struct DefaultOnRequest {
265 scopes: Vec<String>,
266}
267
268#[async_trait]
269impl OnRequest for DefaultOnRequest {
270 async fn on_request(
271 &self,
272 context: &mut Context,
273 request: &mut Request,
274 authorizer: &dyn Authorizer,
275 ) -> Result<()> {
276 let options = TokenRequestOptions {
277 method_options: ClientMethodOptions {
278 context: context.clone(),
279 },
280 };
281 let scopes: Vec<&str> = self.scopes.iter().map(String::as_str).collect();
282 authorizer.authorize(request, &scopes, options).await
283 }
284}
285
286#[cfg(test)]
287mod tests {
288 use super::*;
289 use crate::{
290 credentials::{Secret, TokenCredential, TokenRequestOptions},
291 http::{
292 headers::{HeaderName, HeaderValue, Headers, AUTHORIZATION},
293 policies::{Policy, TransportPolicy},
294 AsyncRawResponse, ClientMethodOptions, Method, Request, StatusCode, Transport,
295 },
296 time::{Duration, OffsetDateTime},
297 tracing::{SpanKind, TracerProvider},
298 Bytes, Result,
299 };
300 use async_trait::async_trait;
301 use azure_core_test::{
302 http::MockHttpClient,
303 tracing::{
304 check_instrumentation_result, ExpectedSpanInformation, ExpectedTracerInformation,
305 MockTracingProvider,
306 },
307 };
308 use futures::FutureExt;
309 use std::sync::{
310 atomic::{AtomicUsize, Ordering},
311 Arc,
312 };
313
314 #[derive(Debug, Clone)]
315 struct MockCredential {
316 calls: Arc<AtomicUsize>,
317 tokens: Arc<[AccessToken]>,
318 }
319
320 impl MockCredential {
321 fn new(tokens: &[AccessToken]) -> Self {
322 Self {
323 calls: Arc::new(AtomicUsize::new(0)),
324 tokens: tokens.into(),
325 }
326 }
327
328 fn get_token_calls(&self) -> usize {
329 self.calls.load(Ordering::SeqCst)
330 }
331 }
332
333 impl Drop for MockCredential {
336 fn drop(&mut self) {
337 if !self.tokens.is_empty() {
338 assert_eq!(self.tokens.len(), self.calls.load(Ordering::SeqCst));
339 }
340 }
341 }
342
343 #[async_trait]
344 impl TokenCredential for MockCredential {
345 async fn get_token(
346 &self,
347 _: &[&str],
348 _: Option<TokenRequestOptions<'_>>,
349 ) -> Result<AccessToken> {
350 let i = self.calls.fetch_add(1, Ordering::SeqCst);
351 self.tokens
352 .get(i)
353 .ok_or_else(|| Error::with_message(ErrorKind::Credential, "no more mock tokens"))
354 .cloned()
355 }
356 }
357
358 #[tokio::test]
359 async fn authn_error() {
360 let credential = MockCredential::new(&[]);
362 let policy = BearerTokenAuthorizationPolicy::new(Arc::new(credential), ["scope"]);
363 let client = MockHttpClient::new(|_| panic!("expected an error from get_token"));
364 let transport = Arc::new(TransportPolicy::new(Transport::new(Arc::new(client))));
365 let mut req = Request::new("https://localhost".parse().unwrap(), Method::Get);
366
367 let err = policy
368 .send(
369 &Context::default(),
370 &mut req,
371 std::slice::from_ref(&(transport.clone() as Arc<dyn Policy>)),
372 )
373 .await
374 .expect_err("request should fail");
375
376 assert_eq!(ErrorKind::Credential, *err.kind());
377 }
378
379 #[tokio::test]
380 async fn http_scheme_rejected() {
381 let credential = MockCredential::new(&[]);
382 let policy = BearerTokenAuthorizationPolicy::new(Arc::new(credential), ["scope"]);
383 let client = MockHttpClient::new(|_| panic!("transport should not be reached"));
384 let transport = Arc::new(TransportPolicy::new(Transport::new(Arc::new(client))));
385 let mut req = Request::new("http://localhost".parse().unwrap(), Method::Get);
386
387 let err = policy
388 .send(
389 &Context::default(),
390 &mut req,
391 std::slice::from_ref(&(transport.clone() as Arc<dyn Policy>)),
392 )
393 .await
394 .expect_err("request should fail for non-https url");
395
396 assert_eq!(err.kind(), &ErrorKind::Other);
397 assert!(err.to_string().contains("TLS"));
398 }
399
400 async fn run_test(tokens: &[AccessToken]) {
401 let credential = Arc::new(MockCredential::new(tokens));
402 let policy = BearerTokenAuthorizationPolicy::new(credential.clone(), ["scope"]);
403 let client = Arc::new(MockHttpClient::new(move |actual| {
404 let credential = credential.clone();
405 async move {
406 let authz = actual.headers().get_str(&AUTHORIZATION)?;
407 let i = credential.get_token_calls().saturating_sub(1);
409 let expected = &credential.tokens[i];
410
411 assert_eq!(format!("Bearer {}", expected.token.secret()), authz);
412
413 Ok(AsyncRawResponse::from_bytes(
414 StatusCode::Ok,
415 Headers::new(),
416 Bytes::new(),
417 ))
418 }
419 .boxed()
420 }));
421 let transport = Arc::new(TransportPolicy::new(Transport::new(client)));
422
423 let mut handles = vec![];
424 for _ in 0..4 {
425 let policy = policy.clone();
426 let transport = transport.clone();
427 let handle = tokio::spawn(async move {
428 let ctx = Context::default();
429 let mut req = Request::new("https://localhost".parse().unwrap(), Method::Get);
430 policy
431 .send(
432 &ctx,
433 &mut req,
434 std::slice::from_ref(&(transport.clone() as Arc<dyn Policy>)),
435 )
436 .await
437 .expect("successful request");
438 });
439 handles.push(handle);
440 }
441
442 for handle in handles {
443 tokio::time::timeout(Duration::seconds(2).try_into().unwrap(), handle)
444 .await
445 .expect("task timed out after 2 seconds")
446 .expect("completed task");
447 }
448 }
449
450 #[tokio::test]
451 async fn caches_token() {
452 run_test(&[AccessToken {
453 token: Secret::new("fake".to_string()),
454 expires_on: OffsetDateTime::now_utc() + Duration::seconds(3600),
455 }])
456 .await;
457 }
458
459 #[tokio::test]
460 async fn refreshes_token() {
461 run_test(&[
462 AccessToken {
463 token: Secret::new("1".to_string()),
464 expires_on: OffsetDateTime::now_utc() - Duration::seconds(1),
465 },
466 AccessToken {
467 token: Secret::new("2".to_string()),
468 expires_on: OffsetDateTime::now_utc() + Duration::seconds(3600),
469 },
470 ])
471 .await;
472 }
473
474 #[derive(Debug)]
475 struct TestOnChallenge {
476 calls: Arc<AtomicUsize>,
477 error: Option<Error>,
478 }
479
480 #[async_trait]
481 impl OnChallenge for TestOnChallenge {
482 async fn on_challenge(
483 &self,
484 context: &Context,
485 request: &mut Request,
486 authorizer: &dyn Authorizer,
487 _headers: &Headers,
488 ) -> Result<()> {
489 self.calls.fetch_add(1, Ordering::SeqCst);
490 if let Some(ref e) = self.error {
491 return Err(Error::with_message(e.kind().clone(), e.to_string()));
492 }
493 let options = TokenRequestOptions {
494 method_options: ClientMethodOptions {
495 context: context.clone(),
496 },
497 };
498 authorizer.authorize(request, &["scope"], options).await?;
499 Ok(())
500 }
501 }
502
503 #[tokio::test]
504 async fn on_challenge_error() {
505 let calls = Arc::new(AtomicUsize::new(0));
506 let on_challenge = Arc::new(TestOnChallenge {
507 calls: calls.clone(),
508 error: Some(Error::with_message(
509 ErrorKind::Other,
510 "something went wrong",
511 )),
512 });
513
514 let credential = Arc::new(MockCredential::new(&[AccessToken {
515 token: Secret::new("fake".to_string()),
516 expires_on: OffsetDateTime::now_utc() + Duration::seconds(3600),
517 }]));
518
519 let policy = BearerTokenAuthorizationPolicy::new(credential, ["scope"])
520 .with_on_challenge(on_challenge);
521
522 let client = MockHttpClient::new(|_| {
523 async {
524 let mut headers = Headers::new();
525 headers.insert(WWW_AUTHENTICATE, "Bearer challenge");
526 Ok(AsyncRawResponse::from_bytes(
527 StatusCode::Unauthorized,
528 headers,
529 Bytes::new(),
530 ))
531 }
532 .boxed()
533 });
534 let transport = Arc::new(TransportPolicy::new(Transport::new(Arc::new(client))));
535
536 let mut req = Request::new("https://localhost".parse().unwrap(), Method::Get);
537 let err = policy
538 .send(
539 &Context::default(),
540 &mut req,
541 std::slice::from_ref(&(transport as Arc<dyn Policy>)),
542 )
543 .await
544 .expect_err("request should fail");
545
546 assert_eq!(ErrorKind::Other, *err.kind());
547 assert_eq!("something went wrong", err.to_string());
548 assert_eq!(1, calls.load(Ordering::SeqCst));
549 }
550
551 #[tokio::test]
552 async fn on_challenge_not_called_without_header() {
553 let calls = Arc::new(AtomicUsize::new(0));
554 let on_challenge = Arc::new(TestOnChallenge {
555 calls: calls.clone(),
556 error: None,
557 });
558
559 let credential = Arc::new(MockCredential::new(&[AccessToken {
560 token: Secret::new("fake".to_string()),
561 expires_on: OffsetDateTime::now_utc() + Duration::seconds(3600),
562 }]));
563
564 let policy = BearerTokenAuthorizationPolicy::new(credential, ["scope"])
565 .with_on_challenge(on_challenge);
566
567 let client = MockHttpClient::new(|_| {
568 async {
569 Ok(AsyncRawResponse::from_bytes(
570 StatusCode::Unauthorized,
571 Headers::new(),
572 Bytes::new(),
573 ))
574 }
575 .boxed()
576 });
577 let transport = Arc::new(TransportPolicy::new(Transport::new(Arc::new(client))));
578
579 let mut req = Request::new("https://localhost".parse().unwrap(), Method::Get);
580 let response = policy
581 .send(
582 &Context::default(),
583 &mut req,
584 std::slice::from_ref(&(transport as Arc<dyn Policy>)),
585 )
586 .await
587 .expect("successful request");
588
589 assert_eq!(StatusCode::Unauthorized, response.status());
590 assert_eq!(0, calls.load(Ordering::SeqCst));
591 }
592
593 #[tokio::test]
594 async fn on_challenge_with_retry() {
595 let on_challenge_calls = Arc::new(AtomicUsize::new(0));
596 let on_challenge = Arc::new(TestOnChallenge {
597 calls: on_challenge_calls.clone(),
598 error: None,
599 });
600
601 let on_request_calls = Arc::new(AtomicUsize::new(0));
602 let on_request = Arc::new(TestOnRequest {
603 calls: on_request_calls.clone(),
604 error: None,
605 });
606
607 let credential = Arc::new(MockCredential::new(&[
608 AccessToken {
609 token: Secret::new("first".to_string()),
610 expires_on: OffsetDateTime::now_utc() + Duration::seconds(3600),
611 },
612 AccessToken {
613 token: Secret::new("second".to_string()),
614 expires_on: OffsetDateTime::now_utc() + Duration::seconds(3600),
615 },
616 ]));
617
618 let policy = BearerTokenAuthorizationPolicy::new(credential.clone(), ["scope"])
619 .with_on_request(on_request)
620 .with_on_challenge(on_challenge);
621
622 let request_count = Arc::new(AtomicUsize::new(0));
623 let request_count_clone = request_count.clone();
624
625 let client = MockHttpClient::new(move |actual| {
626 let count = request_count_clone.fetch_add(1, Ordering::SeqCst);
627 async move {
628 let authz = actual.headers().get_str(&AUTHORIZATION)?;
629
630 if count == 0 {
631 assert_eq!("Bearer first", authz);
633 Ok(AsyncRawResponse::from_bytes(
634 StatusCode::Unauthorized,
635 Headers::from(std::collections::HashMap::from([(
636 WWW_AUTHENTICATE,
637 HeaderValue::from("Bearer challenge".to_string()),
638 )])),
639 Bytes::new(),
640 ))
641 } else {
642 assert_eq!("Bearer second", authz);
644 Ok(AsyncRawResponse::from_bytes(
645 StatusCode::Ok,
646 Headers::new(),
647 Bytes::new(),
648 ))
649 }
650 }
651 .boxed()
652 });
653 let transport = Arc::new(TransportPolicy::new(Transport::new(Arc::new(client))));
654
655 let mut req = Request::new("https://localhost".parse().unwrap(), Method::Get);
656 let response = policy
657 .send(
658 &Context::default(),
659 &mut req,
660 std::slice::from_ref(&(transport as Arc<dyn Policy>)),
661 )
662 .await
663 .expect("successful request");
664
665 assert_eq!(StatusCode::Ok, response.status());
666 assert_eq!(1, on_request_calls.load(Ordering::SeqCst));
667 assert_eq!(1, on_challenge_calls.load(Ordering::SeqCst));
668 assert_eq!(2, request_count.load(Ordering::SeqCst));
669 assert_eq!(2, credential.get_token_calls());
670 }
671
672 #[derive(Debug)]
673 struct TestOnRequest {
674 calls: Arc<AtomicUsize>,
675 error: Option<Error>,
676 }
677
678 #[async_trait]
679 impl OnRequest for TestOnRequest {
680 async fn on_request(
681 &self,
682 _: &mut Context,
683 request: &mut Request,
684 authorizer: &dyn Authorizer,
685 ) -> Result<()> {
686 let calls = self.calls.fetch_add(1, Ordering::SeqCst) + 1;
687 request.insert_header("on-request-calls", calls.to_string());
688
689 if let Some(ref e) = self.error {
690 Err(Error::with_message(e.kind().clone(), e.to_string()))
691 } else {
692 authorizer
693 .authorize(
694 request,
695 &["scope"],
696 TokenRequestOptions {
697 method_options: ClientMethodOptions {
698 context: Context::default(),
699 },
700 },
701 )
702 .await?;
703 Ok(())
704 }
705 }
706 }
707
708 #[tokio::test]
709 async fn on_request() {
710 let called = Arc::new(AtomicUsize::new(0));
711 let on_request = Arc::new(TestOnRequest {
712 calls: called.clone(),
713 error: None,
714 });
715
716 let credential = Arc::new(MockCredential::new(&[AccessToken::new(
717 "token",
718 OffsetDateTime::now_utc() + Duration::seconds(3600),
719 )]));
720
721 let policy =
722 BearerTokenAuthorizationPolicy::new(credential, ["scope"]).with_on_request(on_request);
723
724 let client = MockHttpClient::new(|actual| {
725 async {
726 assert_eq!(
727 "1",
728 actual
729 .headers()
730 .get_str(&HeaderName::from_static("on-request-calls"))?,
731 "on_request should have set the test header to 1",
732 );
733 Ok(AsyncRawResponse::from_bytes(
734 StatusCode::Ok,
735 Headers::new(),
736 Bytes::new(),
737 ))
738 }
739 .boxed()
740 });
741 let transport: Arc<dyn Policy> =
742 Arc::new(TransportPolicy::new(Transport::new(Arc::new(client))));
743
744 let ctx = Context::default();
745 let mut req = Request::new("https://localhost".parse().unwrap(), Method::Get);
746 req.insert_header("on-request-calls", stringify!(TestOnRequest));
747 policy
748 .send(&ctx, &mut req, std::slice::from_ref(&transport))
749 .await
750 .expect("successful request");
751
752 assert_eq!(1, called.load(Ordering::SeqCst));
753 }
754 #[tokio::test]
755 async fn on_request_error() {
756 let calls = Arc::new(AtomicUsize::new(0));
757 let on_request = Arc::new(TestOnRequest {
758 calls: calls.clone(),
759 error: Some(Error::with_message(
760 ErrorKind::Other,
761 "something went wrong",
762 )),
763 });
764
765 let credential = Arc::new(MockCredential::new(&[]));
766
767 let policy =
768 BearerTokenAuthorizationPolicy::new(credential, ["scope"]).with_on_request(on_request);
769
770 let client =
771 MockHttpClient::new(|_| panic!("request should not be sent when on_request errors"));
772 let transport: Arc<dyn Policy> =
773 Arc::new(TransportPolicy::new(Transport::new(Arc::new(client))));
774
775 let ctx = Context::default();
776 let mut req = Request::new("https://localhost".parse().unwrap(), Method::Get);
777
778 let err = policy
779 .send(&ctx, &mut req, std::slice::from_ref(&transport))
780 .await
781 .expect_err("request should fail");
782
783 assert_eq!(ErrorKind::Other, *err.kind());
784 assert_eq!("something went wrong", err.to_string());
785 assert_eq!(1, calls.load(Ordering::SeqCst));
786 }
787
788 #[tokio::test]
789 async fn resets_stream_for_retry_after_challenge() {
790 use crate::{http::Body, stream::BytesStream};
791 use futures::StreamExt;
792
793 let on_challenge_calls = Arc::new(AtomicUsize::new(0));
794 let on_challenge = Arc::new(TestOnChallenge {
795 calls: on_challenge_calls.clone(),
796 error: None,
797 });
798
799 let credential = Arc::new(MockCredential::new(&[
800 AccessToken {
801 token: Secret::new("first".to_string()),
802 expires_on: OffsetDateTime::now_utc() + Duration::seconds(3600),
803 },
804 AccessToken {
805 token: Secret::new("second".to_string()),
806 expires_on: OffsetDateTime::now_utc() + Duration::seconds(3600),
807 },
808 ]));
809
810 let policy = BearerTokenAuthorizationPolicy::new(credential, ["scope"])
811 .with_on_challenge(on_challenge);
812
813 let request_count = Arc::new(AtomicUsize::new(0));
814 let request_count_clone = request_count.clone();
815
816 let client = MockHttpClient::new(move |actual| {
817 let count = request_count_clone.fetch_add(1, Ordering::SeqCst);
818 async move {
819 match actual.body() {
820 Body::SeekableStream(stream) => {
821 let mut stream = stream.clone();
822 let mut collected = Vec::new();
823 while let Some(chunk) = stream.next().await {
824 let chunk = chunk?;
825 collected.extend_from_slice(&chunk);
826 }
827 assert_eq!(b"test data", collected.as_slice());
828 }
829 _ => unreachable!("body is a SeekableStream"),
830 }
831
832 if count == 0 {
833 Ok(AsyncRawResponse::from_bytes(
834 StatusCode::Unauthorized,
835 Headers::from(std::collections::HashMap::from([(
836 WWW_AUTHENTICATE,
837 HeaderValue::from("Bearer challenge".to_string()),
838 )])),
839 Bytes::new(),
840 ))
841 } else {
842 Ok(AsyncRawResponse::from_bytes(
843 StatusCode::Ok,
844 Headers::new(),
845 Bytes::new(),
846 ))
847 }
848 }
849 .boxed()
850 });
851 let transport = Arc::new(TransportPolicy::new(Transport::new(Arc::new(client))));
852
853 let mut req = Request::new("https://localhost".parse().unwrap(), Method::Get);
854 let stream = BytesStream::new(b"test data".as_slice());
855 req.set_body(Body::SeekableStream(Box::new(stream)));
856
857 let res = policy
858 .send(
859 &Context::default(),
860 &mut req,
861 std::slice::from_ref(&(transport as Arc<dyn Policy>)),
862 )
863 .await
864 .expect("policy should reset the body stream and succeed on retry");
865
866 assert_eq!(StatusCode::Ok, res.status());
867 assert_eq!(1, on_challenge_calls.load(Ordering::SeqCst));
868 assert_eq!(2, request_count.load(Ordering::SeqCst));
869 }
870
871 #[tokio::test]
872 async fn challenged_request_tracing() {
873 let on_challenge = Arc::new(TestOnChallenge {
874 calls: Arc::new(AtomicUsize::new(0)),
875 error: None,
876 });
877 let credential = Arc::new(MockCredential::new(
878 &(0..2)
879 .map(|_| AccessToken {
880 token: Secret::new("token".to_string()),
881 expires_on: OffsetDateTime::now_utc() + Duration::seconds(3600),
882 })
883 .collect::<Vec<_>>(),
884 ));
885 let policy = BearerTokenAuthorizationPolicy::new(credential, ["scope"])
886 .with_on_challenge(on_challenge);
887 let client = MockHttpClient::new(|_| {
888 async {
889 Ok(AsyncRawResponse::from_bytes(
890 StatusCode::Unauthorized,
891 Headers::from(std::collections::HashMap::from([(
892 WWW_AUTHENTICATE,
893 HeaderValue::from("Bearer challenge".to_string()),
894 )])),
895 Bytes::new(),
896 ))
897 }
898 .boxed()
899 });
900 let transport = Arc::new(TransportPolicy::new(Transport::new(Arc::new(client))));
901 let provider = Arc::new(MockTracingProvider::new());
902 let tracer = provider.get_tracer(None, "test_crate", None);
903 {
904 let span = tracer.start_span("test_span".into(), SpanKind::Internal, vec![]);
905 let mut ctx = Context::default();
906 ctx.insert(span);
907 policy
908 .send(
909 &ctx,
910 &mut Request::new("https://localhost".parse().unwrap(), Method::Get),
911 std::slice::from_ref(&(transport as Arc<dyn Policy>)),
912 )
913 .await
914 .expect("successful request");
915 }
916 check_instrumentation_result(
917 provider,
918 vec![ExpectedTracerInformation {
919 name: "test_crate",
920 version: None,
921 namespace: None,
922 spans: vec![ExpectedSpanInformation {
923 span_name: "test_span",
924 kind: SpanKind::Internal,
925 attributes: vec![(ERROR_TYPE_ATTRIBUTE, "401".into())],
926 ..Default::default()
927 }],
928 }],
929 );
930 }
931}