Skip to main content

azure_core/http/policies/auth/
bearer_token_policy.rs

1// Copyright (c) Microsoft Corporation. All rights reserved.
2// Licensed under the MIT License.
3
4use 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/// Authentication policy for a bearer token.
21#[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    /// Creates a new `BearerTokenAuthorizationPolicy`.
30    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    /// Sets a callback for `send` to invoke once on each request it receives, before sending the request.
44    ///
45    /// See [`OnRequest`] for more details. When not set, the policy authorizes each request using the credential
46    /// and scopes specified to `new`.
47    pub fn with_on_request(mut self, on_request: Arc<dyn OnRequest>) -> Self {
48        self.on_request = on_request;
49        self
50    }
51
52    /// Sets a callback to invoke upon receiving a 401 Unauthorized response with an authentication challenge.
53    ///
54    /// See [`OnChallenge`] for more details. When not set, `send` returns 401 responses without attempting to
55    /// handle their challenges.
56    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                        // this span covers the request which received the 401 response
94                        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/// Callback [`BearerTokenAuthorizationPolicy`] invokes when it receives a 401 Unauthorized response with an authentication challenge (WWW-Authenticate header).
111#[async_trait]
112pub trait OnChallenge: std::fmt::Debug + Send + Sync {
113    /// Called when [`BearerTokenAuthorizationPolicy`] receives a 401 Unauthorized response with a challenge.
114    ///
115    /// Implementations are responsible for parsing authentication parameters from the challenge, authorizing the request via the provided [`Authorizer`],
116    /// and indicating whether the policy should retry the request.
117    ///
118    /// # Arguments
119    /// * `context` - The request context
120    /// * `request` - The HTTP request that received the challenge
121    /// * `authorizer` - Helper used to acquire an access token and set the request's authorization header
122    /// * `headers` - The 401 response's headers
123    ///
124    /// # Returns
125    /// * `Ok` when the callback handled the challenge and [`BearerTokenAuthorizationPolicy`] should retry the request.
126    /// * `Err` when an error occurred while handling the challenge.
127    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/// Callback [`BearerTokenAuthorizationPolicy`] invokes on every request it receives, before sending the request.
137#[async_trait]
138pub trait OnRequest: std::fmt::Debug + Send + Sync {
139    /// Invoked once on every [`BearerTokenAuthorizationPolicy::send`] invocation, before the policy sends the request.
140    ///
141    /// `send` doesn't call this method before retrying a request after an authentication challenge (see [`OnChallenge`]
142    /// for more about challenge handling). Implementations are responsible for authorizing each request via the provided
143    /// [`Authorizer`]. The policy sends the request when this method returns Ok.
144    ///
145    /// # Arguments
146    /// * `context` - The request context
147    /// * `request` - The HTTP request being sent
148    /// * `authorizer` - Helper used to acquire an access token and set the request's authorization header
149    async fn on_request(
150        &self,
151        context: &mut Context,
152        request: &mut Request,
153        authorizer: &dyn Authorizer,
154    ) -> Result<()>;
155}
156
157/// Helper trait used by [`OnChallenge`] and [`OnRequest`] to authorize requests. This trait is sealed and cannot
158/// be implemented outside of this module.
159#[async_trait]
160pub trait Authorizer: crate::private::Sealed + std::fmt::Debug + Send + Sync {
161    /// Acquire an access token for the provided scopes and options, and set the request's authorization header.
162    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                // cache is empty. Upgrade the lock and acquire a token, provided another thread hasn't already done so
204                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                // token is expired or within its refresh window. Upgrade the lock and
212                // acquire a new token, provided another thread hasn't already done so
213                let expires_on = token.expires_on;
214                drop(access_token);
215                let mut access_token = self.access_token.write().await;
216                // access_token shouldn't be None here, but check anyway to guarantee unwrap won't panic
217                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                            // propagate this error because we can't proceed without a new token
228                            return Err(e);
229                        }
230                        Err(_) => {
231                            // ignore this error because the cached token is still valid
232                        }
233                    }
234                }
235            }
236            Some(_) => {
237                // do nothing; cached token is valid and not within its refresh window
238                drop(access_token); // release the read lock so we don't try to acquire it while still holding it
239            }
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    // ensure the number of get_token() calls matches the number of tokens
334    // in a test case i.e., that the policy called get_token() as expected
335    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        // this mock's get_token() will return an error because it has no tokens
361        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                // e.g. if this is the first request, we expect 1 get_token call and tokens[0] in the header
408                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                    // First request gets 401
632                    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                    // Retry with new token succeeds
643                    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}