Skip to main content

reqsign_core/
signer.rs

1// Licensed to the Apache Software Foundation (ASF) under one
2// or more contributor license agreements.  See the NOTICE file
3// distributed with this work for additional information
4// regarding copyright ownership.  The ASF licenses this file
5// to you under the Apache License, Version 2.0 (the
6// "License"); you may not use this file except in compliance
7// with the License.  You may obtain a copy of the License at
8//
9//   http://www.apache.org/licenses/LICENSE-2.0
10//
11// Unless required by applicable law or agreed to in writing,
12// software distributed under the License is distributed on an
13// "AS IS" BASIS, WITHOUT WARRANTIES OR CONDITIONS OF ANY
14// KIND, either express or implied.  See the License for the
15// specific language governing permissions and limitations
16// under the License.
17
18use crate::Context;
19use crate::Error;
20use crate::ProvideCredential;
21use crate::ProvideCredentialDyn;
22use crate::Result;
23use crate::SignRequest;
24use crate::SignRequestDyn;
25use crate::SigningCredential;
26use mea::mutex::Mutex;
27use std::any::type_name;
28use std::fmt::{Debug, Formatter};
29use std::sync::Arc;
30use std::time::Duration;
31
32/// Loads credentials and atomically signs request heads.
33///
34/// The service-specific [`SignRequest`] runs against a private candidate. Only the
35/// candidate URI and headers are committed after successful signing.
36#[derive(Clone)]
37pub struct Signer<K: SigningCredential> {
38    ctx: Context,
39    loader: Arc<dyn ProvideCredentialDyn<Credential = K>>,
40    builder: Arc<dyn SignRequestDyn<Credential = K>>,
41    credential: Arc<Mutex<Option<K>>>,
42}
43
44impl<K: SigningCredential> Debug for Signer<K> {
45    fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
46        f.debug_struct("Signer")
47            .field("credential_type", &type_name::<K>())
48            .finish_non_exhaustive()
49    }
50}
51
52impl<K: SigningCredential> Signer<K> {
53    /// Create a new signer.
54    pub fn new(
55        ctx: Context,
56        loader: impl ProvideCredential<Credential = K>,
57        builder: impl SignRequest<Credential = K>,
58    ) -> Self {
59        Self {
60            ctx,
61
62            loader: Arc::new(loader),
63            builder: Arc::new(builder),
64            credential: Arc::new(Mutex::new(None)),
65        }
66    }
67
68    /// Replace the context while keeping the credential provider, request signer,
69    /// and shared credential cache.
70    pub fn with_context(mut self, ctx: Context) -> Self {
71        self.ctx = ctx;
72        self
73    }
74
75    /// Replace the credential provider while keeping context and request signer,
76    /// and create an isolated empty credential cache.
77    pub fn with_credential_provider(
78        mut self,
79        provider: impl ProvideCredential<Credential = K>,
80    ) -> Self {
81        self.loader = Arc::new(provider);
82        self.credential = Arc::new(Mutex::new(None));
83        self
84    }
85
86    /// Replace the request signer while keeping context and credential provider.
87    pub fn with_request_signer(mut self, signer: impl SignRequest<Credential = K>) -> Self {
88        self.builder = Arc::new(signer);
89        self
90    }
91
92    /// Sign a wire-ready request head.
93    ///
94    /// The request URI must satisfy the input contract of the configured
95    /// [`SignRequest`]. Built-in signers require an authority and expect path and query
96    /// data to be percent-encoded exactly once before this call. Signing does not
97    /// perform general-purpose URI encoding for the caller.
98    ///
99    /// If credential loading or request signing returns an error, `req` is unchanged.
100    /// On success, only `req.uri` and `req.headers` may change; the method, version, and
101    /// extensions retain their input values.
102    ///
103    /// `expires_in` is a service-specific validity input and does not universally
104    /// select query authentication. The configured service signer and credential type
105    /// determine how it is interpreted.
106    ///
107    /// Cached credentials must be fresh according to [`SigningCredential::is_valid`]
108    /// and usable through [`SignRequest::required_valid_until`]. A refreshed credential
109    /// only needs to satisfy the exact operation deadline. Credential refresh is
110    /// serialized per shared cache. A failed refresh is not cached, so the next waiting
111    /// or later caller can retry. Provider errors are returned without internal retry or
112    /// fallback to the previous cached credential. Request signing runs after refresh
113    /// coordination has completed and remains concurrent.
114    pub async fn sign(
115        &self,
116        req: &mut http::request::Parts,
117        expires_in: Option<Duration>,
118    ) -> Result<()> {
119        let credential = self.credential(expires_in).await?;
120
121        let mut candidate = req.clone();
122        self.builder
123            .sign_request_dyn(&self.ctx, &mut candidate, Some(&credential), expires_in)
124            .await?;
125
126        req.uri = candidate.uri;
127        req.headers = candidate.headers;
128        Ok(())
129    }
130
131    async fn credential(&self, expires_in: Option<Duration>) -> Result<K> {
132        let mut cached = self.credential.lock().await;
133        if let Some(credential) = cached.as_ref() {
134            if credential.is_valid()
135                && credential.is_valid_at(
136                    self.builder
137                        .required_valid_until_dyn(credential, expires_in),
138                )
139            {
140                return Ok(credential.clone());
141            }
142        }
143
144        let credential = self
145            .loader
146            .provide_credential_dyn(&self.ctx)
147            .await?
148            .ok_or_else(|| {
149                Error::credential_invalid("failed to load signing credential")
150                    .with_context(format!("credential_type: {}", type_name::<K>()))
151            })?;
152
153        *cached = Some(credential.clone());
154        drop(cached);
155
156        self.validate_refreshed_credential(credential, expires_in)
157    }
158
159    fn validate_refreshed_credential(
160        &self,
161        credential: K,
162        expires_in: Option<Duration>,
163    ) -> Result<K> {
164        let required_until = self
165            .builder
166            .required_valid_until_dyn(&credential, expires_in);
167        if !credential.is_valid_at(required_until) {
168            return Err(Error::credential_invalid(
169                "refreshed signing credential expires before the requested operation deadline",
170            )
171            .with_context(format!("credential_type: {}", type_name::<K>()))
172            .with_context(format!("required_valid_until: {required_until}")));
173        }
174
175        Ok(credential)
176    }
177}
178
179#[cfg(test)]
180mod tests {
181    use super::*;
182    use crate::time::Timestamp;
183    use crate::{ErrorKind, ProvideCredential, SignRequest};
184    use futures::channel::oneshot;
185    use futures::future::{join_all, pending};
186    use futures::poll;
187    use http::{HeaderValue, Method, Request, Version};
188    use std::collections::VecDeque;
189    use std::sync::Mutex as StdMutex;
190    use std::sync::atomic::{AtomicUsize, Ordering};
191
192    #[derive(Clone, Debug)]
193    struct TestCredential;
194
195    impl SigningCredential for TestCredential {
196        fn is_valid(&self) -> bool {
197            true
198        }
199    }
200
201    #[derive(Debug)]
202    struct StaticProvider;
203
204    impl ProvideCredential for StaticProvider {
205        type Credential = TestCredential;
206
207        async fn provide_credential(&self, _ctx: &Context) -> Result<Option<Self::Credential>> {
208            Ok(Some(TestCredential))
209        }
210    }
211
212    #[derive(Clone, Debug, PartialEq, Eq)]
213    struct Extension(&'static str);
214
215    #[derive(Debug)]
216    struct MutatingSigner {
217        fail: bool,
218    }
219
220    impl SignRequest for MutatingSigner {
221        type Credential = TestCredential;
222
223        async fn sign_request(
224            &self,
225            _ctx: &Context,
226            req: &mut http::request::Parts,
227            _credential: Option<&Self::Credential>,
228            _expires_in: Option<Duration>,
229        ) -> Result<()> {
230            req.method = Method::POST;
231            req.uri = "https://signed.example.com/result?auth=1"
232                .parse()
233                .expect("URI must parse");
234            req.version = Version::HTTP_2;
235            req.headers.clear();
236            req.headers
237                .insert("authorization", HeaderValue::from_static("signed"));
238            req.extensions.insert(Extension("candidate"));
239
240            if self.fail {
241                Err(Error::unexpected("injected signing failure"))
242            } else {
243                Ok(())
244            }
245        }
246    }
247
248    #[derive(Clone, Debug)]
249    struct ExpiringCredential {
250        generation: u8,
251        fresh: bool,
252        expires_at: Timestamp,
253        required_until: Timestamp,
254    }
255
256    impl SigningCredential for ExpiringCredential {
257        fn is_valid(&self) -> bool {
258            self.fresh
259        }
260
261        fn is_valid_at(&self, timestamp: Timestamp) -> bool {
262            self.expires_at > timestamp
263        }
264    }
265
266    type ControlledResponse = Result<Option<ExpiringCredential>>;
267    type ControlledProviderParts = (
268        ControlledProvider,
269        Arc<AtomicUsize>,
270        Vec<oneshot::Sender<ControlledResponse>>,
271    );
272
273    #[derive(Debug)]
274    struct SequenceProvider {
275        responses: StdMutex<VecDeque<Result<Option<ExpiringCredential>>>>,
276        calls: Arc<AtomicUsize>,
277    }
278
279    impl SequenceProvider {
280        fn new(
281            responses: impl IntoIterator<Item = Result<Option<ExpiringCredential>>>,
282        ) -> (Self, Arc<AtomicUsize>) {
283            let calls = Arc::new(AtomicUsize::new(0));
284            (
285                Self {
286                    responses: StdMutex::new(responses.into_iter().collect()),
287                    calls: calls.clone(),
288                },
289                calls,
290            )
291        }
292    }
293
294    impl ProvideCredential for SequenceProvider {
295        type Credential = ExpiringCredential;
296
297        async fn provide_credential(&self, _ctx: &Context) -> Result<Option<Self::Credential>> {
298            self.calls.fetch_add(1, Ordering::SeqCst);
299            self.responses
300                .lock()
301                .expect("lock poisoned")
302                .pop_front()
303                .unwrap_or(Ok(None))
304        }
305    }
306
307    struct ControlledProvider {
308        responses: StdMutex<VecDeque<oneshot::Receiver<ControlledResponse>>>,
309        calls: Arc<AtomicUsize>,
310    }
311
312    impl Debug for ControlledProvider {
313        fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
314            f.debug_struct("ControlledProvider").finish_non_exhaustive()
315        }
316    }
317
318    impl ControlledProvider {
319        fn new(invocations: usize) -> ControlledProviderParts {
320            let calls = Arc::new(AtomicUsize::new(0));
321            let (senders, receivers) = (0..invocations)
322                .map(|_| oneshot::channel())
323                .unzip::<_, _, Vec<_>, VecDeque<_>>();
324            (
325                Self {
326                    responses: StdMutex::new(receivers),
327                    calls: calls.clone(),
328                },
329                calls,
330                senders,
331            )
332        }
333    }
334
335    impl ProvideCredential for ControlledProvider {
336        type Credential = ExpiringCredential;
337
338        async fn provide_credential(&self, _ctx: &Context) -> Result<Option<Self::Credential>> {
339            self.calls.fetch_add(1, Ordering::SeqCst);
340            let response = self
341                .responses
342                .lock()
343                .expect("lock poisoned")
344                .pop_front()
345                .expect("controlled response must exist");
346            response
347                .await
348                .map_err(|_| Error::unexpected("controlled response sender was dropped"))?
349        }
350    }
351
352    struct ControlledRequestSigner {
353        started: Arc<AtomicUsize>,
354        releases: StdMutex<VecDeque<oneshot::Receiver<()>>>,
355    }
356
357    impl Debug for ControlledRequestSigner {
358        fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
359            f.debug_struct("ControlledRequestSigner")
360                .finish_non_exhaustive()
361        }
362    }
363
364    impl ControlledRequestSigner {
365        fn new(count: usize) -> (Self, Arc<AtomicUsize>, Vec<oneshot::Sender<()>>) {
366            let started = Arc::new(AtomicUsize::new(0));
367            let (senders, receivers) = (0..count)
368                .map(|_| oneshot::channel())
369                .unzip::<_, _, Vec<_>, VecDeque<_>>();
370            (
371                Self {
372                    started: started.clone(),
373                    releases: StdMutex::new(receivers),
374                },
375                started,
376                senders,
377            )
378        }
379    }
380
381    impl SignRequest for ControlledRequestSigner {
382        type Credential = ExpiringCredential;
383
384        fn required_valid_until(
385            &self,
386            credential: &Self::Credential,
387            _expires_in: Option<Duration>,
388        ) -> Timestamp {
389            credential.required_until
390        }
391
392        async fn sign_request(
393            &self,
394            _ctx: &Context,
395            req: &mut http::request::Parts,
396            credential: Option<&Self::Credential>,
397            _expires_in: Option<Duration>,
398        ) -> Result<()> {
399            self.started.fetch_add(1, Ordering::SeqCst);
400            let release = self
401                .releases
402                .lock()
403                .expect("lock poisoned")
404                .pop_front()
405                .expect("signing release must exist");
406            release
407                .await
408                .map_err(|_| Error::unexpected("signing release sender was dropped"))?;
409            req.headers.insert(
410                "x-credential-generation",
411                credential
412                    .expect("credential must be present")
413                    .generation
414                    .to_string()
415                    .parse()?,
416            );
417            Ok(())
418        }
419    }
420
421    struct CancellationProvider {
422        calls: Arc<AtomicUsize>,
423        credential: ExpiringCredential,
424    }
425
426    impl Debug for CancellationProvider {
427        fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
428            f.debug_struct("CancellationProvider")
429                .finish_non_exhaustive()
430        }
431    }
432
433    impl ProvideCredential for CancellationProvider {
434        type Credential = ExpiringCredential;
435
436        async fn provide_credential(&self, _ctx: &Context) -> Result<Option<Self::Credential>> {
437            if self.calls.fetch_add(1, Ordering::SeqCst) == 0 {
438                pending::<()>().await;
439            }
440            Ok(Some(self.credential.clone()))
441        }
442    }
443
444    const CREDENTIAL_SECRET: &str = "credential-secret-must-not-leak";
445
446    #[derive(Clone, Debug)]
447    struct SecretCredential {
448        secret: &'static str,
449    }
450
451    impl SigningCredential for SecretCredential {
452        fn is_valid(&self) -> bool {
453            !self.secret.is_empty()
454        }
455
456        fn is_valid_at(&self, _timestamp: Timestamp) -> bool {
457            false
458        }
459    }
460
461    #[derive(Debug)]
462    struct SecretProvider;
463
464    impl ProvideCredential for SecretProvider {
465        type Credential = SecretCredential;
466
467        async fn provide_credential(&self, _ctx: &Context) -> Result<Option<Self::Credential>> {
468            Ok(Some(SecretCredential {
469                secret: CREDENTIAL_SECRET,
470            }))
471        }
472    }
473
474    #[derive(Debug)]
475    struct SecretRequestSigner;
476
477    impl SignRequest for SecretRequestSigner {
478        type Credential = SecretCredential;
479
480        async fn sign_request(
481            &self,
482            _ctx: &Context,
483            _req: &mut http::request::Parts,
484            _credential: Option<&Self::Credential>,
485            _expires_in: Option<Duration>,
486        ) -> Result<()> {
487            Ok(())
488        }
489    }
490
491    #[derive(Debug)]
492    struct OperationSigner;
493
494    impl SignRequest for OperationSigner {
495        type Credential = ExpiringCredential;
496
497        fn required_valid_until(
498            &self,
499            credential: &Self::Credential,
500            _expires_in: Option<Duration>,
501        ) -> Timestamp {
502            credential.required_until
503        }
504
505        async fn sign_request(
506            &self,
507            _ctx: &Context,
508            req: &mut http::request::Parts,
509            credential: Option<&Self::Credential>,
510            expires_in: Option<Duration>,
511        ) -> Result<()> {
512            let credential = credential.expect("credential must be present");
513            if !credential.is_valid_at(self.required_valid_until(credential, expires_in)) {
514                return Err(Error::credential_invalid(
515                    "credential is not valid for operation",
516                ));
517            }
518            req.headers.insert(
519                "x-credential-generation",
520                credential.generation.to_string().parse()?,
521            );
522            Ok(())
523        }
524    }
525
526    fn request_parts() -> http::request::Parts {
527        let mut parts = Request::get("https://example.com/original?x=%2F")
528            .version(Version::HTTP_11)
529            .header("x-original", "value")
530            .body(())
531            .expect("request must build")
532            .into_parts()
533            .0;
534        parts.extensions.insert(Extension("caller"));
535        parts
536    }
537
538    fn set_cached_credential<K: SigningCredential>(signer: &Signer<K>, credential: K) {
539        *signer
540            .credential
541            .try_lock()
542            .expect("credential cache must be unlocked") = Some(credential);
543    }
544
545    #[test]
546    fn failure_leaves_entire_request_head_unchanged() {
547        let signer = Signer::new(
548            Context::new(),
549            StaticProvider,
550            MutatingSigner { fail: true },
551        );
552        let mut parts = request_parts();
553        let original = parts.clone();
554
555        let result = futures::executor::block_on(signer.sign(&mut parts, None));
556
557        assert!(result.is_err());
558        assert_eq!(parts.method, original.method);
559        assert_eq!(parts.uri, original.uri);
560        assert_eq!(parts.version, original.version);
561        assert_eq!(parts.headers, original.headers);
562        assert_eq!(
563            parts.extensions.get::<Extension>(),
564            original.extensions.get::<Extension>()
565        );
566    }
567
568    #[test]
569    fn success_commits_only_uri_and_headers() {
570        let signer = Signer::new(
571            Context::new(),
572            StaticProvider,
573            MutatingSigner { fail: false },
574        );
575        let mut parts = request_parts();
576        let original = parts.clone();
577
578        futures::executor::block_on(signer.sign(&mut parts, None)).expect("signing must succeed");
579
580        assert_eq!(parts.method, original.method);
581        assert_eq!(parts.version, original.version);
582        assert_eq!(
583            parts.extensions.get::<Extension>(),
584            original.extensions.get::<Extension>()
585        );
586        assert_eq!(
587            parts.uri,
588            "https://signed.example.com/result?auth=1"
589                .parse::<http::Uri>()
590                .expect("URI must parse")
591        );
592        assert_eq!(
593            parts.headers.get("authorization"),
594            Some(&HeaderValue::from_static("signed"))
595        );
596        assert!(!parts.headers.contains_key("x-original"));
597    }
598
599    #[test]
600    fn concurrent_cold_start_invokes_provider_once() {
601        futures::executor::block_on(async {
602            let base = Timestamp::from_second(500).expect("timestamp must be valid");
603            let credential = ExpiringCredential {
604                generation: 1,
605                fresh: true,
606                expires_at: base + Duration::from_secs(30),
607                required_until: base + Duration::from_secs(10),
608            };
609            let (provider, calls, mut responses) = ControlledProvider::new(1);
610            let signer = Signer::new(Context::new(), provider, OperationSigner);
611            let signers = (0..8).map(|_| signer.clone()).collect::<Vec<_>>();
612            let mut requests = (0..8).map(|_| request_parts()).collect::<Vec<_>>();
613            let mut batch = Box::pin(join_all(
614                requests
615                    .iter_mut()
616                    .zip(signers.iter())
617                    .map(|(request, signer)| signer.sign(request, None)),
618            ));
619
620            assert!(poll!(&mut batch).is_pending());
621            assert_eq!(calls.load(Ordering::SeqCst), 1);
622
623            responses
624                .remove(0)
625                .send(Ok(Some(credential)))
626                .expect("controlled response must be received");
627            for result in batch.await {
628                result.expect("concurrent cold-start signing must succeed");
629            }
630            assert_eq!(calls.load(Ordering::SeqCst), 1);
631        });
632    }
633
634    #[test]
635    fn concurrent_stale_refresh_invokes_provider_once() {
636        futures::executor::block_on(async {
637            let base = Timestamp::from_second(600).expect("timestamp must be valid");
638            let cached = ExpiringCredential {
639                generation: 1,
640                fresh: false,
641                expires_at: base + Duration::from_secs(30),
642                required_until: base + Duration::from_secs(10),
643            };
644            let refreshed = ExpiringCredential {
645                generation: 2,
646                fresh: true,
647                expires_at: base + Duration::from_secs(30),
648                required_until: base + Duration::from_secs(10),
649            };
650            let (provider, calls, mut responses) = ControlledProvider::new(1);
651            let signer = Signer::new(Context::new(), provider, OperationSigner);
652            set_cached_credential(&signer, cached);
653            let mut requests = (0..8).map(|_| request_parts()).collect::<Vec<_>>();
654            let mut batch = Box::pin(join_all(
655                requests
656                    .iter_mut()
657                    .map(|request| signer.sign(request, None)),
658            ));
659
660            assert!(poll!(&mut batch).is_pending());
661            assert_eq!(calls.load(Ordering::SeqCst), 1);
662
663            responses
664                .remove(0)
665                .send(Ok(Some(refreshed)))
666                .expect("controlled response must be received");
667            for result in batch.await {
668                result.expect("concurrent stale refresh must succeed");
669            }
670            assert_eq!(calls.load(Ordering::SeqCst), 1);
671            for request in requests {
672                assert_eq!(
673                    request.headers.get("x-credential-generation"),
674                    Some(&HeaderValue::from_static("2"))
675                );
676            }
677        });
678    }
679
680    #[test]
681    fn concurrent_refresh_failure_allows_waiter_retry() {
682        futures::executor::block_on(async {
683            let base = Timestamp::from_second(700).expect("timestamp must be valid");
684            let refreshed = ExpiringCredential {
685                generation: 2,
686                fresh: true,
687                expires_at: base + Duration::from_secs(30),
688                required_until: base + Duration::from_secs(10),
689            };
690            let (provider, calls, mut responses) = ControlledProvider::new(2);
691            let signer = Signer::new(Context::new(), provider, OperationSigner);
692            let mut requests = (0..8).map(|_| request_parts()).collect::<Vec<_>>();
693            let mut batch = Box::pin(join_all(
694                requests
695                    .iter_mut()
696                    .map(|request| signer.sign(request, None)),
697            ));
698
699            assert!(poll!(&mut batch).is_pending());
700            assert_eq!(calls.load(Ordering::SeqCst), 1);
701
702            responses
703                .remove(0)
704                .send(Err(Error::rate_limited("injected refresh failure")
705                    .with_context("refresh_generation: 1")))
706                .expect("controlled failure must be received");
707
708            assert!(poll!(&mut batch).is_pending());
709            assert_eq!(calls.load(Ordering::SeqCst), 2);
710            responses
711                .remove(0)
712                .send(Ok(Some(refreshed)))
713                .expect("controlled recovery must be received");
714
715            let results = batch.await;
716            let errors = results
717                .iter()
718                .filter_map(|result| result.as_ref().err())
719                .collect::<Vec<_>>();
720            assert_eq!(errors.len(), 1);
721            assert_eq!(errors[0].kind(), ErrorKind::RateLimited);
722            assert_eq!(errors[0].to_string(), "injected refresh failure");
723            assert_eq!(errors[0].context(), &["refresh_generation: 1"]);
724            assert!(errors[0].is_retryable());
725            assert_eq!(results.iter().filter(|result| result.is_ok()).count(), 7);
726            assert_eq!(calls.load(Ordering::SeqCst), 2);
727
728            let mut later_request = request_parts();
729            signer
730                .sign(&mut later_request, None)
731                .await
732                .expect("later caller must reuse the recovered credential");
733            assert_eq!(calls.load(Ordering::SeqCst), 2);
734        });
735    }
736
737    #[test]
738    fn refreshed_credential_is_checked_for_exact_operation_deadline() {
739        futures::executor::block_on(async {
740            let base = Timestamp::from_second(800).expect("timestamp must be valid");
741            let credential = ExpiringCredential {
742                generation: 1,
743                fresh: true,
744                expires_at: base + Duration::from_secs(10),
745                required_until: base + Duration::from_secs(10),
746            };
747            let (provider, calls, mut responses) = ControlledProvider::new(1);
748            let signer = Signer::new(Context::new(), provider, OperationSigner);
749            let mut request = request_parts();
750            let mut signing = Box::pin(signer.sign(&mut request, None));
751
752            assert!(poll!(&mut signing).is_pending());
753            responses
754                .remove(0)
755                .send(Ok(Some(credential)))
756                .expect("controlled response must be received");
757            let error = signing
758                .await
759                .expect_err("exact deadline must reject the credential");
760            assert_eq!(error.kind(), ErrorKind::CredentialInvalid);
761            assert!(
762                error
763                    .to_string()
764                    .contains("expires before the requested operation deadline")
765            );
766            assert_eq!(calls.load(Ordering::SeqCst), 1);
767        });
768    }
769
770    #[test]
771    fn distinct_credential_caches_do_not_block_each_other() {
772        futures::executor::block_on(async {
773            let base = Timestamp::from_second(900).expect("timestamp must be valid");
774            let credential_a = ExpiringCredential {
775                generation: 1,
776                fresh: true,
777                expires_at: base + Duration::from_secs(30),
778                required_until: base + Duration::from_secs(10),
779            };
780            let credential_b = ExpiringCredential {
781                generation: 2,
782                fresh: true,
783                expires_at: base + Duration::from_secs(30),
784                required_until: base + Duration::from_secs(10),
785            };
786            let (provider_a, calls_a, mut responses_a) = ControlledProvider::new(1);
787            let (provider_b, calls_b, mut responses_b) = ControlledProvider::new(1);
788            let signer_a = Signer::new(Context::new(), provider_a, OperationSigner);
789            let signer_b = Signer::new(Context::new(), provider_b, OperationSigner);
790            let mut request_a = request_parts();
791            let mut request_b = request_parts();
792            let mut future_a = Box::pin(signer_a.sign(&mut request_a, None));
793            let mut future_b = Box::pin(signer_b.sign(&mut request_b, None));
794
795            assert!(poll!(&mut future_a).is_pending());
796            assert!(poll!(&mut future_b).is_pending());
797            assert_eq!(calls_a.load(Ordering::SeqCst), 1);
798            assert_eq!(calls_b.load(Ordering::SeqCst), 1);
799
800            responses_b
801                .remove(0)
802                .send(Ok(Some(credential_b)))
803                .expect("second cache response must be received");
804            future_b
805                .await
806                .expect("second cache must complete while first is blocked");
807            assert!(poll!(&mut future_a).is_pending());
808
809            responses_a
810                .remove(0)
811                .send(Ok(Some(credential_a)))
812                .expect("first cache response must be received");
813            future_a.await.expect("first cache must complete");
814        });
815    }
816
817    #[test]
818    fn request_signing_remains_concurrent_after_refresh() {
819        futures::executor::block_on(async {
820            let base = Timestamp::from_second(950).expect("timestamp must be valid");
821            let credential = ExpiringCredential {
822                generation: 1,
823                fresh: true,
824                expires_at: base + Duration::from_secs(30),
825                required_until: base + Duration::from_secs(10),
826            };
827            let (provider, calls, mut responses) = ControlledProvider::new(1);
828            let (request_signer, started, releases) = ControlledRequestSigner::new(8);
829            let signer = Signer::new(Context::new(), provider, request_signer);
830            let mut requests = (0..8).map(|_| request_parts()).collect::<Vec<_>>();
831            let mut batch = Box::pin(join_all(
832                requests
833                    .iter_mut()
834                    .map(|request| signer.sign(request, None)),
835            ));
836
837            assert!(poll!(&mut batch).is_pending());
838            assert_eq!(calls.load(Ordering::SeqCst), 1);
839            assert_eq!(started.load(Ordering::SeqCst), 0);
840            responses
841                .remove(0)
842                .send(Ok(Some(credential)))
843                .expect("controlled response must be received");
844
845            assert!(poll!(&mut batch).is_pending());
846            assert_eq!(started.load(Ordering::SeqCst), 8);
847            for release in releases {
848                release.send(()).expect("signing release must be received");
849            }
850            for result in batch.await {
851                result.expect("concurrent request signing must succeed");
852            }
853        });
854    }
855
856    #[test]
857    fn cache_sharing_and_reset_contract_is_preserved() {
858        let signer = Signer::new(
859            Context::new(),
860            StaticProvider,
861            MutatingSigner { fail: false },
862        );
863        let clone = signer.clone();
864        let with_context = signer.clone().with_context(Context::new());
865        let with_request_signer = signer
866            .clone()
867            .with_request_signer(MutatingSigner { fail: false });
868        let with_provider = signer.clone().with_credential_provider(StaticProvider);
869
870        assert!(Arc::ptr_eq(&signer.credential, &clone.credential));
871        assert!(Arc::ptr_eq(&signer.credential, &with_context.credential));
872        assert!(Arc::ptr_eq(
873            &signer.credential,
874            &with_request_signer.credential
875        ));
876        assert!(!Arc::ptr_eq(&signer.credential, &with_provider.credential));
877    }
878
879    #[test]
880    fn cancelled_refresh_releases_lock_and_waiter_retries() {
881        futures::executor::block_on(async {
882            let base = Timestamp::from_second(975).expect("timestamp must be valid");
883            let calls = Arc::new(AtomicUsize::new(0));
884            let provider = CancellationProvider {
885                calls: calls.clone(),
886                credential: ExpiringCredential {
887                    generation: 2,
888                    fresh: true,
889                    expires_at: base + Duration::from_secs(30),
890                    required_until: base + Duration::from_secs(10),
891                },
892            };
893            let signer = Signer::new(Context::new(), provider, OperationSigner);
894            let mut leader_request = request_parts();
895            let mut waiter_request = request_parts();
896            let mut leader = Box::pin(signer.sign(&mut leader_request, None));
897            let mut waiter = Box::pin(signer.sign(&mut waiter_request, None));
898
899            assert!(poll!(&mut leader).is_pending());
900            assert!(poll!(&mut waiter).is_pending());
901            assert_eq!(calls.load(Ordering::SeqCst), 1);
902            drop(leader);
903
904            waiter
905                .await
906                .expect("waiter must retry after leader cancellation");
907            assert_eq!(calls.load(Ordering::SeqCst), 2);
908            assert_eq!(
909                waiter_request.headers.get("x-credential-generation"),
910                Some(&HeaderValue::from_static("2"))
911            );
912        });
913    }
914
915    #[test]
916    fn credential_values_are_redacted_from_debug_and_validation_errors() {
917        let signer = Signer::new(Context::new(), SecretProvider, SecretRequestSigner);
918        let mut request = request_parts();
919        let error = futures::executor::block_on(signer.sign(&mut request, None))
920            .expect_err("unusable credential must fail validation");
921
922        assert!(!format!("{signer:?}").contains(CREDENTIAL_SECRET));
923        assert!(!format!("{error:?}").contains(CREDENTIAL_SECRET));
924        assert!(!error.to_string().contains(CREDENTIAL_SECRET));
925    }
926
927    #[cfg(not(target_arch = "wasm32"))]
928    #[test]
929    fn sign_future_remains_send_on_native_targets() {
930        fn assert_send<T: Send>(_future: T) {}
931
932        let signer = Signer::new(
933            Context::new(),
934            StaticProvider,
935            MutatingSigner { fail: false },
936        );
937        let mut request = request_parts();
938        assert_send(signer.sign(&mut request, None));
939    }
940
941    #[test]
942    fn refreshes_cached_credential_for_operation_requirement() {
943        let base = Timestamp::from_second(1_000).expect("timestamp must be valid");
944        let cached = ExpiringCredential {
945            generation: 1,
946            fresh: true,
947            expires_at: base + Duration::from_secs(20),
948            required_until: base + Duration::from_secs(30),
949        };
950        let refreshed = ExpiringCredential {
951            generation: 2,
952            fresh: true,
953            expires_at: base + Duration::from_secs(20),
954            required_until: base + Duration::from_secs(10),
955        };
956        let (provider, calls) = SequenceProvider::new([Ok(Some(refreshed))]);
957        let signer = Signer::new(Context::new(), provider, OperationSigner);
958        set_cached_credential(&signer, cached);
959
960        let mut parts = request_parts();
961        futures::executor::block_on(signer.sign(&mut parts, None))
962            .expect("refreshed credential must satisfy the recomputed requirement");
963
964        assert_eq!(calls.load(Ordering::SeqCst), 1);
965        assert_eq!(
966            parts.headers.get("x-credential-generation"),
967            Some(&HeaderValue::from_static("2"))
968        );
969    }
970
971    #[test]
972    fn uses_refreshed_credential_that_is_usable_but_not_fresh() {
973        let base = Timestamp::from_second(2_000).expect("timestamp must be valid");
974        let credential = ExpiringCredential {
975            generation: 1,
976            fresh: false,
977            expires_at: base + Duration::from_secs(30),
978            required_until: base + Duration::from_secs(10),
979        };
980        let (provider, calls) =
981            SequenceProvider::new([Ok(Some(credential.clone())), Ok(Some(credential))]);
982        let signer = Signer::new(Context::new(), provider, OperationSigner);
983
984        for _ in 0..2 {
985            let mut parts = request_parts();
986            futures::executor::block_on(signer.sign(&mut parts, None))
987                .expect("usable refreshed credential must be accepted");
988        }
989
990        assert_eq!(calls.load(Ordering::SeqCst), 2);
991    }
992
993    #[test]
994    fn refresh_error_does_not_fall_back_and_caller_can_retry() {
995        let base = Timestamp::from_second(3_000).expect("timestamp must be valid");
996        let cached = ExpiringCredential {
997            generation: 1,
998            fresh: false,
999            expires_at: base + Duration::from_secs(30),
1000            required_until: base + Duration::from_secs(10),
1001        };
1002        let refreshed = ExpiringCredential {
1003            generation: 2,
1004            fresh: true,
1005            expires_at: base + Duration::from_secs(30),
1006            required_until: base + Duration::from_secs(10),
1007        };
1008        let (provider, calls) = SequenceProvider::new([
1009            Err(Error::unexpected("injected refresh failure")),
1010            Ok(Some(refreshed)),
1011        ]);
1012        let signer = Signer::new(Context::new(), provider, OperationSigner);
1013        set_cached_credential(&signer, cached);
1014
1015        let mut parts = request_parts();
1016        let original = parts.clone();
1017        let err = futures::executor::block_on(signer.sign(&mut parts, None))
1018            .expect_err("refresh error must be returned");
1019        assert_eq!(err.kind(), ErrorKind::Unexpected);
1020        assert_eq!(parts.uri, original.uri);
1021        assert_eq!(parts.headers, original.headers);
1022        assert_eq!(calls.load(Ordering::SeqCst), 1);
1023
1024        futures::executor::block_on(signer.sign(&mut parts, None))
1025            .expect("caller retry must attempt refresh again");
1026        assert_eq!(calls.load(Ordering::SeqCst), 2);
1027        assert_eq!(
1028            parts.headers.get("x-credential-generation"),
1029            Some(&HeaderValue::from_static("2"))
1030        );
1031    }
1032
1033    #[test]
1034    fn missing_refresh_does_not_fall_back_and_caller_can_retry() {
1035        let base = Timestamp::from_second(4_000).expect("timestamp must be valid");
1036        let cached = ExpiringCredential {
1037            generation: 1,
1038            fresh: false,
1039            expires_at: base + Duration::from_secs(30),
1040            required_until: base + Duration::from_secs(10),
1041        };
1042        let refreshed = ExpiringCredential {
1043            generation: 2,
1044            fresh: true,
1045            expires_at: base + Duration::from_secs(30),
1046            required_until: base + Duration::from_secs(10),
1047        };
1048        let (provider, calls) = SequenceProvider::new([Ok(None), Ok(Some(refreshed))]);
1049        let signer = Signer::new(Context::new(), provider, OperationSigner);
1050        set_cached_credential(&signer, cached);
1051        let mut parts = request_parts();
1052        let original = parts.clone();
1053
1054        let err = futures::executor::block_on(signer.sign(&mut parts, None))
1055            .expect_err("missing credential must fail");
1056
1057        assert_eq!(err.kind(), ErrorKind::CredentialInvalid);
1058        assert_eq!(calls.load(Ordering::SeqCst), 1);
1059        assert_eq!(parts.uri, original.uri);
1060        assert_eq!(parts.headers, original.headers);
1061
1062        futures::executor::block_on(signer.sign(&mut parts, None))
1063            .expect("caller retry must attempt refresh again");
1064        assert_eq!(calls.load(Ordering::SeqCst), 2);
1065        assert_eq!(
1066            parts.headers.get("x-credential-generation"),
1067            Some(&HeaderValue::from_static("2"))
1068        );
1069    }
1070
1071    #[test]
1072    fn debug_is_opaque() {
1073        let signer = Signer::new(
1074            Context::new(),
1075            StaticProvider,
1076            MutatingSigner { fail: false },
1077        );
1078        set_cached_credential(&signer, TestCredential);
1079
1080        let debug = format!("{signer:?}");
1081        assert!(debug.starts_with("Signer"));
1082        assert!(!debug.contains("StaticProvider"));
1083        assert!(!debug.contains("MutatingSigner"));
1084        assert!(!debug.contains("credential:"));
1085    }
1086}