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 std::any::type_name;
27use std::fmt::{Debug, Formatter};
28use std::sync::{Arc, Mutex};
29use std::time::Duration;
30
31#[derive(Clone)]
36pub struct Signer<K: SigningCredential> {
37 ctx: Context,
38 loader: Arc<dyn ProvideCredentialDyn<Credential = K>>,
39 builder: Arc<dyn SignRequestDyn<Credential = K>>,
40 credential: Arc<Mutex<Option<K>>>,
41}
42
43impl<K: SigningCredential> Debug for Signer<K> {
44 fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
45 f.debug_struct("Signer")
46 .field("credential_type", &type_name::<K>())
47 .finish_non_exhaustive()
48 }
49}
50
51impl<K: SigningCredential> Signer<K> {
52 pub fn new(
54 ctx: Context,
55 loader: impl ProvideCredential<Credential = K>,
56 builder: impl SignRequest<Credential = K>,
57 ) -> Self {
58 Self {
59 ctx,
60
61 loader: Arc::new(loader),
62 builder: Arc::new(builder),
63 credential: Arc::new(Mutex::new(None)),
64 }
65 }
66
67 pub fn with_context(mut self, ctx: Context) -> Self {
69 self.ctx = ctx;
70 self
71 }
72
73 pub fn with_credential_provider(
75 mut self,
76 provider: impl ProvideCredential<Credential = K>,
77 ) -> Self {
78 self.loader = Arc::new(provider);
79 self.credential = Arc::new(Mutex::new(None)); self
81 }
82
83 pub fn with_request_signer(mut self, signer: impl SignRequest<Credential = K>) -> Self {
85 self.builder = Arc::new(signer);
86 self
87 }
88
89 pub async fn sign(
109 &self,
110 req: &mut http::request::Parts,
111 expires_in: Option<Duration>,
112 ) -> Result<()> {
113 let credential = self.credential.lock().expect("lock poisoned").clone();
114 let credential = match credential {
115 Some(credential)
116 if credential.is_valid()
117 && credential.is_valid_at(
118 self.builder
119 .required_valid_until_dyn(&credential, expires_in),
120 ) =>
121 {
122 credential
123 }
124 _ => {
125 let credential = self
126 .loader
127 .provide_credential_dyn(&self.ctx)
128 .await?
129 .ok_or_else(|| {
130 Error::credential_invalid("failed to load signing credential")
131 .with_context(format!("credential_type: {}", type_name::<K>()))
132 })?;
133
134 *self.credential.lock().expect("lock poisoned") = Some(credential.clone());
135
136 let required_until = self
137 .builder
138 .required_valid_until_dyn(&credential, expires_in);
139 if !credential.is_valid_at(required_until) {
140 return Err(Error::credential_invalid(
141 "refreshed signing credential expires before the requested operation deadline",
142 )
143 .with_context(format!("credential_type: {}", type_name::<K>()))
144 .with_context(format!("required_valid_until: {required_until}")));
145 }
146
147 credential
148 }
149 };
150
151 let mut candidate = req.clone();
152 self.builder
153 .sign_request_dyn(&self.ctx, &mut candidate, Some(&credential), expires_in)
154 .await?;
155
156 req.uri = candidate.uri;
157 req.headers = candidate.headers;
158 Ok(())
159 }
160}
161
162#[cfg(test)]
163mod tests {
164 use super::*;
165 use crate::time::Timestamp;
166 use crate::{ErrorKind, ProvideCredential, SignRequest};
167 use http::{HeaderValue, Method, Request, Version};
168 use std::collections::VecDeque;
169 use std::sync::atomic::{AtomicUsize, Ordering};
170
171 #[derive(Clone, Debug)]
172 struct TestCredential;
173
174 impl SigningCredential for TestCredential {
175 fn is_valid(&self) -> bool {
176 true
177 }
178 }
179
180 #[derive(Debug)]
181 struct StaticProvider;
182
183 impl ProvideCredential for StaticProvider {
184 type Credential = TestCredential;
185
186 async fn provide_credential(&self, _ctx: &Context) -> Result<Option<Self::Credential>> {
187 Ok(Some(TestCredential))
188 }
189 }
190
191 #[derive(Clone, Debug, PartialEq, Eq)]
192 struct Extension(&'static str);
193
194 #[derive(Debug)]
195 struct MutatingSigner {
196 fail: bool,
197 }
198
199 impl SignRequest for MutatingSigner {
200 type Credential = TestCredential;
201
202 async fn sign_request(
203 &self,
204 _ctx: &Context,
205 req: &mut http::request::Parts,
206 _credential: Option<&Self::Credential>,
207 _expires_in: Option<Duration>,
208 ) -> Result<()> {
209 req.method = Method::POST;
210 req.uri = "https://signed.example.com/result?auth=1"
211 .parse()
212 .expect("URI must parse");
213 req.version = Version::HTTP_2;
214 req.headers.clear();
215 req.headers
216 .insert("authorization", HeaderValue::from_static("signed"));
217 req.extensions.insert(Extension("candidate"));
218
219 if self.fail {
220 Err(Error::unexpected("injected signing failure"))
221 } else {
222 Ok(())
223 }
224 }
225 }
226
227 #[derive(Clone, Debug)]
228 struct ExpiringCredential {
229 generation: u8,
230 fresh: bool,
231 expires_at: Timestamp,
232 required_until: Timestamp,
233 }
234
235 impl SigningCredential for ExpiringCredential {
236 fn is_valid(&self) -> bool {
237 self.fresh
238 }
239
240 fn is_valid_at(&self, timestamp: Timestamp) -> bool {
241 self.expires_at > timestamp
242 }
243 }
244
245 #[derive(Debug)]
246 struct SequenceProvider {
247 responses: Mutex<VecDeque<Result<Option<ExpiringCredential>>>>,
248 calls: Arc<AtomicUsize>,
249 }
250
251 impl SequenceProvider {
252 fn new(
253 responses: impl IntoIterator<Item = Result<Option<ExpiringCredential>>>,
254 ) -> (Self, Arc<AtomicUsize>) {
255 let calls = Arc::new(AtomicUsize::new(0));
256 (
257 Self {
258 responses: Mutex::new(responses.into_iter().collect()),
259 calls: calls.clone(),
260 },
261 calls,
262 )
263 }
264 }
265
266 impl ProvideCredential for SequenceProvider {
267 type Credential = ExpiringCredential;
268
269 async fn provide_credential(&self, _ctx: &Context) -> Result<Option<Self::Credential>> {
270 self.calls.fetch_add(1, Ordering::SeqCst);
271 self.responses
272 .lock()
273 .expect("lock poisoned")
274 .pop_front()
275 .unwrap_or(Ok(None))
276 }
277 }
278
279 #[derive(Debug)]
280 struct OperationSigner;
281
282 impl SignRequest for OperationSigner {
283 type Credential = ExpiringCredential;
284
285 fn required_valid_until(
286 &self,
287 credential: &Self::Credential,
288 _expires_in: Option<Duration>,
289 ) -> Timestamp {
290 credential.required_until
291 }
292
293 async fn sign_request(
294 &self,
295 _ctx: &Context,
296 req: &mut http::request::Parts,
297 credential: Option<&Self::Credential>,
298 expires_in: Option<Duration>,
299 ) -> Result<()> {
300 let credential = credential.expect("credential must be present");
301 if !credential.is_valid_at(self.required_valid_until(credential, expires_in)) {
302 return Err(Error::credential_invalid(
303 "credential is not valid for operation",
304 ));
305 }
306 req.headers.insert(
307 "x-credential-generation",
308 credential.generation.to_string().parse()?,
309 );
310 Ok(())
311 }
312 }
313
314 fn request_parts() -> http::request::Parts {
315 let mut parts = Request::get("https://example.com/original?x=%2F")
316 .version(Version::HTTP_11)
317 .header("x-original", "value")
318 .body(())
319 .expect("request must build")
320 .into_parts()
321 .0;
322 parts.extensions.insert(Extension("caller"));
323 parts
324 }
325
326 #[test]
327 fn failure_leaves_entire_request_head_unchanged() {
328 let signer = Signer::new(
329 Context::new(),
330 StaticProvider,
331 MutatingSigner { fail: true },
332 );
333 let mut parts = request_parts();
334 let original = parts.clone();
335
336 let result = futures::executor::block_on(signer.sign(&mut parts, None));
337
338 assert!(result.is_err());
339 assert_eq!(parts.method, original.method);
340 assert_eq!(parts.uri, original.uri);
341 assert_eq!(parts.version, original.version);
342 assert_eq!(parts.headers, original.headers);
343 assert_eq!(
344 parts.extensions.get::<Extension>(),
345 original.extensions.get::<Extension>()
346 );
347 }
348
349 #[test]
350 fn success_commits_only_uri_and_headers() {
351 let signer = Signer::new(
352 Context::new(),
353 StaticProvider,
354 MutatingSigner { fail: false },
355 );
356 let mut parts = request_parts();
357 let original = parts.clone();
358
359 futures::executor::block_on(signer.sign(&mut parts, None)).expect("signing must succeed");
360
361 assert_eq!(parts.method, original.method);
362 assert_eq!(parts.version, original.version);
363 assert_eq!(
364 parts.extensions.get::<Extension>(),
365 original.extensions.get::<Extension>()
366 );
367 assert_eq!(
368 parts.uri,
369 "https://signed.example.com/result?auth=1"
370 .parse::<http::Uri>()
371 .expect("URI must parse")
372 );
373 assert_eq!(
374 parts.headers.get("authorization"),
375 Some(&HeaderValue::from_static("signed"))
376 );
377 assert!(!parts.headers.contains_key("x-original"));
378 }
379
380 #[test]
381 fn refreshes_cached_credential_for_operation_requirement() {
382 let base = Timestamp::from_second(1_000).expect("timestamp must be valid");
383 let cached = ExpiringCredential {
384 generation: 1,
385 fresh: true,
386 expires_at: base + Duration::from_secs(20),
387 required_until: base + Duration::from_secs(30),
388 };
389 let refreshed = ExpiringCredential {
390 generation: 2,
391 fresh: true,
392 expires_at: base + Duration::from_secs(20),
393 required_until: base + Duration::from_secs(10),
394 };
395 let (provider, calls) = SequenceProvider::new([Ok(Some(refreshed))]);
396 let signer = Signer::new(Context::new(), provider, OperationSigner);
397 *signer.credential.lock().expect("lock poisoned") = Some(cached);
398
399 let mut parts = request_parts();
400 futures::executor::block_on(signer.sign(&mut parts, None))
401 .expect("refreshed credential must satisfy the recomputed requirement");
402
403 assert_eq!(calls.load(Ordering::SeqCst), 1);
404 assert_eq!(
405 parts.headers.get("x-credential-generation"),
406 Some(&HeaderValue::from_static("2"))
407 );
408 }
409
410 #[test]
411 fn uses_refreshed_credential_that_is_usable_but_not_fresh() {
412 let base = Timestamp::from_second(2_000).expect("timestamp must be valid");
413 let credential = ExpiringCredential {
414 generation: 1,
415 fresh: false,
416 expires_at: base + Duration::from_secs(30),
417 required_until: base + Duration::from_secs(10),
418 };
419 let (provider, calls) =
420 SequenceProvider::new([Ok(Some(credential.clone())), Ok(Some(credential))]);
421 let signer = Signer::new(Context::new(), provider, OperationSigner);
422
423 for _ in 0..2 {
424 let mut parts = request_parts();
425 futures::executor::block_on(signer.sign(&mut parts, None))
426 .expect("usable refreshed credential must be accepted");
427 }
428
429 assert_eq!(calls.load(Ordering::SeqCst), 2);
430 }
431
432 #[test]
433 fn refresh_error_does_not_fall_back_and_caller_can_retry() {
434 let base = Timestamp::from_second(3_000).expect("timestamp must be valid");
435 let cached = ExpiringCredential {
436 generation: 1,
437 fresh: false,
438 expires_at: base + Duration::from_secs(30),
439 required_until: base + Duration::from_secs(10),
440 };
441 let refreshed = ExpiringCredential {
442 generation: 2,
443 fresh: true,
444 expires_at: base + Duration::from_secs(30),
445 required_until: base + Duration::from_secs(10),
446 };
447 let (provider, calls) = SequenceProvider::new([
448 Err(Error::unexpected("injected refresh failure")),
449 Ok(Some(refreshed)),
450 ]);
451 let signer = Signer::new(Context::new(), provider, OperationSigner);
452 *signer.credential.lock().expect("lock poisoned") = Some(cached);
453
454 let mut parts = request_parts();
455 let original = parts.clone();
456 let err = futures::executor::block_on(signer.sign(&mut parts, None))
457 .expect_err("refresh error must be returned");
458 assert_eq!(err.kind(), ErrorKind::Unexpected);
459 assert_eq!(parts.uri, original.uri);
460 assert_eq!(parts.headers, original.headers);
461 assert_eq!(calls.load(Ordering::SeqCst), 1);
462
463 futures::executor::block_on(signer.sign(&mut parts, None))
464 .expect("caller retry must attempt refresh again");
465 assert_eq!(calls.load(Ordering::SeqCst), 2);
466 assert_eq!(
467 parts.headers.get("x-credential-generation"),
468 Some(&HeaderValue::from_static("2"))
469 );
470 }
471
472 #[test]
473 fn missing_refresh_does_not_fall_back_and_caller_can_retry() {
474 let base = Timestamp::from_second(4_000).expect("timestamp must be valid");
475 let cached = ExpiringCredential {
476 generation: 1,
477 fresh: false,
478 expires_at: base + Duration::from_secs(30),
479 required_until: base + Duration::from_secs(10),
480 };
481 let refreshed = ExpiringCredential {
482 generation: 2,
483 fresh: true,
484 expires_at: base + Duration::from_secs(30),
485 required_until: base + Duration::from_secs(10),
486 };
487 let (provider, calls) = SequenceProvider::new([Ok(None), Ok(Some(refreshed))]);
488 let signer = Signer::new(Context::new(), provider, OperationSigner);
489 *signer.credential.lock().expect("lock poisoned") = Some(cached);
490 let mut parts = request_parts();
491 let original = parts.clone();
492
493 let err = futures::executor::block_on(signer.sign(&mut parts, None))
494 .expect_err("missing credential must fail");
495
496 assert_eq!(err.kind(), ErrorKind::CredentialInvalid);
497 assert_eq!(calls.load(Ordering::SeqCst), 1);
498 assert_eq!(parts.uri, original.uri);
499 assert_eq!(parts.headers, original.headers);
500
501 futures::executor::block_on(signer.sign(&mut parts, None))
502 .expect("caller retry must attempt refresh again");
503 assert_eq!(calls.load(Ordering::SeqCst), 2);
504 assert_eq!(
505 parts.headers.get("x-credential-generation"),
506 Some(&HeaderValue::from_static("2"))
507 );
508 }
509
510 #[test]
511 fn debug_is_opaque() {
512 let signer = Signer::new(
513 Context::new(),
514 StaticProvider,
515 MutatingSigner { fail: false },
516 );
517 *signer.credential.lock().expect("lock poisoned") = Some(TestCredential);
518
519 let debug = format!("{signer:?}");
520 assert!(debug.starts_with("Signer"));
521 assert!(!debug.contains("StaticProvider"));
522 assert!(!debug.contains("MutatingSigner"));
523 assert!(!debug.contains("credential:"));
524 }
525}