1use 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#[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 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 pub fn with_context(mut self, ctx: Context) -> Self {
71 self.ctx = ctx;
72 self
73 }
74
75 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 pub fn with_request_signer(mut self, signer: impl SignRequest<Credential = K>) -> Self {
88 self.builder = Arc::new(signer);
89 self
90 }
91
92 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}