1use crate::time::Timestamp;
19use crate::{
20 Context, Error, GrantCredential, GrantCredentialDyn, ProvideCredential, ProvideCredentialDyn,
21 Result, SigningCredential,
22};
23use std::any::type_name;
24use std::fmt::{Debug, Formatter};
25use std::sync::{Arc, Mutex};
26use std::time::Duration;
27
28#[derive(Clone)]
35pub struct Granter<K: SigningCredential> {
36 ctx: Context,
37 provider: Arc<dyn ProvideCredentialDyn<Credential = K>>,
38 granter: Arc<dyn GrantCredentialDyn<Credential = K>>,
39 credential: Arc<Mutex<Option<K>>>,
40}
41
42impl<K: SigningCredential> Debug for Granter<K> {
43 fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
44 f.debug_struct("Granter")
45 .field("credential_type", &type_name::<K>())
46 .finish_non_exhaustive()
47 }
48}
49
50impl<K: SigningCredential> Granter<K> {
51 pub fn new(
53 ctx: Context,
54 provider: impl ProvideCredential<Credential = K>,
55 granter: impl GrantCredential<Credential = K>,
56 ) -> Self {
57 Self {
58 ctx,
59 provider: Arc::new(provider),
60 granter: Arc::new(granter),
61 credential: Arc::new(Mutex::new(None)),
62 }
63 }
64
65 pub fn with_context(mut self, ctx: Context) -> Self {
67 self.ctx = ctx;
68 self.credential = Arc::new(Mutex::new(None));
69 self
70 }
71
72 pub fn with_credential_provider(
74 mut self,
75 provider: impl ProvideCredential<Credential = K>,
76 ) -> Self {
77 self.provider = Arc::new(provider);
78 self.credential = Arc::new(Mutex::new(None));
79 self
80 }
81
82 pub fn with_credential_granter(
89 mut self,
90 granter: impl GrantCredential<Credential = K>,
91 ) -> Self {
92 self.granter = Arc::new(granter);
93 self
94 }
95
96 pub async fn grant(&self, expires_in: Option<Duration>) -> Result<K> {
105 let credential = self.credential.lock().expect("lock poisoned").clone();
106 let credential = match credential {
107 Some(credential)
108 if credential.is_valid()
109 && credential.is_valid_at(
110 self.granter
111 .required_valid_until_dyn(&credential, expires_in),
112 ) =>
113 {
114 credential
115 }
116 _ => {
117 let credential = self
118 .provider
119 .provide_credential_dyn(&self.ctx)
120 .await?
121 .ok_or_else(|| {
122 Error::credential_invalid("failed to load source credential")
123 .with_context(format!("credential_type: {}", type_name::<K>()))
124 })?;
125
126 let required_until = self
127 .granter
128 .required_valid_until_dyn(&credential, expires_in);
129 if !credential.is_valid_at(required_until) {
130 return Err(Error::credential_invalid(
131 "refreshed source credential expires before the granting deadline",
132 )
133 .with_context(format!("credential_type: {}", type_name::<K>()))
134 .with_context(format!("required_valid_until: {required_until}")));
135 }
136
137 *self.credential.lock().expect("lock poisoned") = Some(credential.clone());
138 credential
139 }
140 };
141
142 let granted = self
143 .granter
144 .grant_credential_dyn(&self.ctx, &credential, expires_in)
145 .await?;
146 let now = Timestamp::now();
147 if !granted.is_valid_at(now) {
148 return Err(
149 Error::credential_invalid("granted credential is not currently usable")
150 .with_context(format!("credential_type: {}", type_name::<K>()))
151 .with_context(format!("validated_at: {now}")),
152 );
153 }
154
155 Ok(granted)
156 }
157}
158
159#[cfg(test)]
160mod tests {
161 use super::*;
162 use crate::time::Timestamp;
163 use crate::{ErrorKind, GrantCredentialDyn, StaticEnv};
164 use std::collections::HashMap;
165 use std::sync::atomic::{AtomicUsize, Ordering};
166
167 #[derive(Clone)]
168 struct TestCredential {
169 generation: usize,
170 fresh: bool,
171 expires_at: Timestamp,
172 secret: Arc<String>,
173 }
174
175 impl Debug for TestCredential {
176 fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
177 f.debug_struct("TestCredential")
178 .field("generation", &self.generation)
179 .field("secret", &self.secret)
180 .finish()
181 }
182 }
183
184 impl SigningCredential for TestCredential {
185 fn is_valid(&self) -> bool {
186 self.fresh && self.is_valid_at(Timestamp::now() + Duration::from_secs(20))
187 }
188
189 fn is_valid_at(&self, timestamp: Timestamp) -> bool {
190 !self.secret.is_empty() && self.expires_at > timestamp
191 }
192 }
193
194 #[derive(Clone)]
195 struct CountingProvider {
196 calls: Arc<AtomicUsize>,
197 secret: Arc<String>,
198 expires_at: Timestamp,
199 fresh: bool,
200 }
201
202 impl Debug for CountingProvider {
203 fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
204 f.debug_struct("CountingProvider")
205 .field("secret", &self.secret)
206 .finish()
207 }
208 }
209
210 impl CountingProvider {
211 fn new(secret: &str, expires_at: Timestamp) -> (Self, Arc<AtomicUsize>) {
212 let calls = Arc::new(AtomicUsize::new(0));
213 (
214 Self {
215 calls: calls.clone(),
216 secret: Arc::new(secret.to_string()),
217 expires_at,
218 fresh: true,
219 },
220 calls,
221 )
222 }
223
224 fn with_fresh(mut self, fresh: bool) -> Self {
225 self.fresh = fresh;
226 self
227 }
228 }
229
230 impl ProvideCredential for CountingProvider {
231 type Credential = TestCredential;
232
233 async fn provide_credential(&self, ctx: &Context) -> Result<Option<Self::Credential>> {
234 let call = self.calls.fetch_add(1, Ordering::SeqCst) + 1;
235 let generation = ctx
236 .env_var("generation")
237 .and_then(|value| value.parse().ok())
238 .unwrap_or(call);
239 Ok(Some(TestCredential {
240 generation,
241 fresh: self.fresh,
242 expires_at: self.expires_at,
243 secret: self.secret.clone(),
244 }))
245 }
246 }
247
248 #[derive(Clone)]
249 struct CountingGranter {
250 calls: Arc<AtomicUsize>,
251 required_until: Timestamp,
252 output_expires_at: Timestamp,
253 secret: Arc<String>,
254 }
255
256 impl Debug for CountingGranter {
257 fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
258 f.debug_struct("CountingGranter")
259 .field("secret", &self.secret)
260 .finish()
261 }
262 }
263
264 impl CountingGranter {
265 fn new(
266 secret: &str,
267 required_until: Timestamp,
268 output_expires_at: Timestamp,
269 ) -> (Self, Arc<AtomicUsize>) {
270 let calls = Arc::new(AtomicUsize::new(0));
271 (
272 Self {
273 calls: calls.clone(),
274 required_until,
275 output_expires_at,
276 secret: Arc::new(secret.to_string()),
277 },
278 calls,
279 )
280 }
281 }
282
283 impl GrantCredential for CountingGranter {
284 type Credential = TestCredential;
285
286 fn required_valid_until(
287 &self,
288 _credential: &Self::Credential,
289 _expires_in: Option<Duration>,
290 ) -> Timestamp {
291 self.required_until
292 }
293
294 async fn grant_credential(
295 &self,
296 _ctx: &Context,
297 credential: &Self::Credential,
298 _expires_in: Option<Duration>,
299 ) -> Result<Self::Credential> {
300 let call = self.calls.fetch_add(1, Ordering::SeqCst) + 1;
301 Ok(TestCredential {
302 generation: credential.generation * 100 + call,
303 fresh: true,
304 expires_at: self.output_expires_at,
305 secret: Arc::new(format!("granted-{call}")),
306 })
307 }
308 }
309
310 #[derive(Clone)]
311 struct ErrorProvider {
312 calls: Arc<AtomicUsize>,
313 }
314
315 impl Debug for ErrorProvider {
316 fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
317 f.debug_struct("ErrorProvider").finish_non_exhaustive()
318 }
319 }
320
321 impl ProvideCredential for ErrorProvider {
322 type Credential = TestCredential;
323
324 async fn provide_credential(&self, _ctx: &Context) -> Result<Option<Self::Credential>> {
325 self.calls.fetch_add(1, Ordering::SeqCst);
326 Err(Error::unexpected("source provider failed"))
327 }
328 }
329
330 #[derive(Clone)]
331 struct FailOnceGranter {
332 calls: Arc<AtomicUsize>,
333 required_until: Timestamp,
334 output_expires_at: Timestamp,
335 }
336
337 impl Debug for FailOnceGranter {
338 fn fmt(&self, f: &mut Formatter<'_>) -> std::fmt::Result {
339 f.debug_struct("FailOnceGranter").finish_non_exhaustive()
340 }
341 }
342
343 impl GrantCredential for FailOnceGranter {
344 type Credential = TestCredential;
345
346 fn required_valid_until(
347 &self,
348 _credential: &Self::Credential,
349 _expires_in: Option<Duration>,
350 ) -> Timestamp {
351 self.required_until
352 }
353
354 async fn grant_credential(
355 &self,
356 _ctx: &Context,
357 credential: &Self::Credential,
358 _expires_in: Option<Duration>,
359 ) -> Result<Self::Credential> {
360 let call = self.calls.fetch_add(1, Ordering::SeqCst) + 1;
361 if call == 1 {
362 return Err(Error::unexpected("grant operation failed"));
363 }
364 Ok(TestCredential {
365 generation: credential.generation * 100 + call,
366 fresh: true,
367 expires_at: self.output_expires_at,
368 secret: Arc::new(format!("granted-{call}")),
369 })
370 }
371 }
372
373 fn future_timestamp(seconds: u64) -> Timestamp {
374 Timestamp::now() + Duration::from_secs(seconds)
375 }
376
377 fn context_with_generation(generation: usize) -> Context {
378 Context::new().with_env(StaticEnv {
379 home_dir: None,
380 envs: HashMap::from([("generation".to_string(), generation.to_string())]),
381 })
382 }
383
384 #[test]
385 fn dyn_bridge_forwards_deadline_and_grant() {
386 let required_until = future_timestamp(60);
387 let output_expires_at = future_timestamp(120);
388 let (operation, calls) =
389 CountingGranter::new("operation-secret", required_until, output_expires_at);
390 let operation: Arc<dyn GrantCredentialDyn<Credential = TestCredential>> =
391 Arc::new(operation);
392 let credential = TestCredential {
393 generation: 4,
394 fresh: true,
395 expires_at: future_timestamp(120),
396 secret: Arc::new("source-secret".to_string()),
397 };
398
399 assert_eq!(
400 operation.required_valid_until(&credential, None),
401 required_until
402 );
403 let granted = futures::executor::block_on(operation.grant_credential(
404 &Context::new(),
405 &credential,
406 None,
407 ))
408 .expect("dyn grant must succeed");
409
410 assert_eq!(granted.generation, 401);
411 assert_eq!(calls.load(Ordering::SeqCst), 1);
412 }
413
414 #[test]
415 fn caches_only_source_and_shares_it_across_clones() {
416 let (provider, provider_calls) =
417 CountingProvider::new("provider-secret", future_timestamp(300));
418 let source_secret = provider.secret.clone();
419 let (operation, operation_calls) = CountingGranter::new(
420 "operation-secret",
421 future_timestamp(30),
422 future_timestamp(120),
423 );
424 let granter = Granter::new(context_with_generation(7), provider, operation);
425
426 let first = futures::executor::block_on(granter.grant(None)).expect("grant must succeed");
427 let second =
428 futures::executor::block_on(granter.clone().grant(None)).expect("grant must succeed");
429
430 assert_eq!(first.generation, 701);
431 assert_eq!(second.generation, 702);
432 assert_eq!(provider_calls.load(Ordering::SeqCst), 1);
433 assert_eq!(operation_calls.load(Ordering::SeqCst), 2);
434 assert!(!Arc::ptr_eq(&first.secret, &source_secret));
435 assert!(!Arc::ptr_eq(&second.secret, &source_secret));
436 }
437
438 #[test]
439 fn replacements_follow_source_cache_isolation_contract() {
440 let (provider, provider_calls) =
441 CountingProvider::new("provider-secret", future_timestamp(300));
442 let (operation, _) = CountingGranter::new(
443 "operation-secret",
444 future_timestamp(30),
445 future_timestamp(120),
446 );
447 let granter = Granter::new(context_with_generation(1), provider, operation);
448 futures::executor::block_on(granter.grant(None)).expect("initial grant must succeed");
449
450 let (replacement_operation, _) = CountingGranter::new(
451 "replacement-operation-secret",
452 future_timestamp(30),
453 future_timestamp(120),
454 );
455 let replacement = granter
456 .clone()
457 .with_credential_granter(replacement_operation);
458 let granted = futures::executor::block_on(replacement.grant(None))
459 .expect("operation replacement must reuse source");
460 assert_eq!(granted.generation / 100, 1);
461 assert_eq!(provider_calls.load(Ordering::SeqCst), 1);
462
463 let isolated_context = granter.clone().with_context(context_with_generation(2));
464 let granted = futures::executor::block_on(isolated_context.grant(None))
465 .expect("context replacement must reload source");
466 assert_eq!(granted.generation / 100, 2);
467 assert_eq!(provider_calls.load(Ordering::SeqCst), 2);
468
469 let (replacement_provider, replacement_provider_calls) =
470 CountingProvider::new("replacement-provider-secret", future_timestamp(300));
471 let isolated_provider = granter.with_credential_provider(replacement_provider);
472 futures::executor::block_on(isolated_provider.grant(None))
473 .expect("provider replacement must reload source");
474 assert_eq!(replacement_provider_calls.load(Ordering::SeqCst), 1);
475 }
476
477 #[test]
478 fn rejects_unusable_source_without_caching_it() {
479 let required_until = future_timestamp(120);
480 let (provider, provider_calls) =
481 CountingProvider::new("source-secret", future_timestamp(60));
482 let (operation, operation_calls) =
483 CountingGranter::new("operation-secret", required_until, future_timestamp(180));
484 let granter = Granter::new(Context::new(), provider, operation);
485
486 for _ in 0..2 {
487 let err = futures::executor::block_on(granter.grant(None))
488 .expect_err("short-lived source must be rejected");
489 assert_eq!(err.kind(), ErrorKind::CredentialInvalid);
490 assert!(!format!("{err:?}").contains("source-secret"));
491 }
492
493 assert_eq!(provider_calls.load(Ordering::SeqCst), 2);
494 assert_eq!(operation_calls.load(Ordering::SeqCst), 0);
495 }
496
497 #[test]
498 fn refreshed_source_needs_exact_validity_but_is_not_reused_when_stale() {
499 let (provider, provider_calls) =
500 CountingProvider::new("source-secret", future_timestamp(300));
501 let provider = provider.with_fresh(false);
502 let (operation, operation_calls) = CountingGranter::new(
503 "operation-secret",
504 future_timestamp(30),
505 future_timestamp(120),
506 );
507 let granter = Granter::new(Context::new(), provider, operation);
508
509 futures::executor::block_on(granter.grant(None))
510 .expect("exact-valid refreshed source must be accepted");
511 futures::executor::block_on(granter.grant(None))
512 .expect("stale cached source must be refreshed again");
513
514 assert_eq!(provider_calls.load(Ordering::SeqCst), 2);
515 assert_eq!(operation_calls.load(Ordering::SeqCst), 2);
516 }
517
518 #[test]
519 fn provider_and_grant_errors_do_not_create_output_cache_state() {
520 let provider_calls = Arc::new(AtomicUsize::new(0));
521 let (operation, operation_calls) = CountingGranter::new(
522 "operation-secret",
523 future_timestamp(30),
524 future_timestamp(120),
525 );
526 let provider_error = Granter::new(
527 Context::new(),
528 ErrorProvider {
529 calls: provider_calls.clone(),
530 },
531 operation,
532 );
533 for _ in 0..2 {
534 futures::executor::block_on(provider_error.grant(None))
535 .expect_err("provider error must be returned");
536 }
537 assert_eq!(provider_calls.load(Ordering::SeqCst), 2);
538 assert_eq!(operation_calls.load(Ordering::SeqCst), 0);
539
540 let (provider, provider_calls) =
541 CountingProvider::new("source-secret", future_timestamp(300));
542 let grant_calls = Arc::new(AtomicUsize::new(0));
543 let granter = Granter::new(
544 Context::new(),
545 provider,
546 FailOnceGranter {
547 calls: grant_calls.clone(),
548 required_until: future_timestamp(30),
549 output_expires_at: future_timestamp(120),
550 },
551 );
552 futures::executor::block_on(granter.grant(None))
553 .expect_err("first grant error must be returned");
554 let output = futures::executor::block_on(granter.grant(None))
555 .expect("second grant must execute again");
556
557 assert_eq!(output.generation, 102);
558 assert_eq!(provider_calls.load(Ordering::SeqCst), 1);
559 assert_eq!(grant_calls.load(Ordering::SeqCst), 2);
560 }
561
562 #[test]
563 fn rejects_output_that_is_expired_after_granting() {
564 let (provider, _) = CountingProvider::new("source-secret", future_timestamp(120));
565 let (operation, _) = CountingGranter::new(
566 "operation-secret",
567 future_timestamp(30),
568 Timestamp::now() - Duration::from_secs(1),
569 );
570 let granter = Granter::new(Context::new(), provider, operation);
571
572 let err = futures::executor::block_on(granter.grant(None))
573 .expect_err("expired output must be rejected");
574 assert_eq!(err.kind(), ErrorKind::CredentialInvalid);
575 assert!(!format!("{err:?}").contains("source-secret"));
576 assert!(!format!("{err:?}").contains("operation-secret"));
577 }
578
579 #[test]
580 fn debug_is_opaque_even_after_source_is_cached() {
581 let (provider, _) = CountingProvider::new("provider-secret", future_timestamp(300));
582 let (operation, _) = CountingGranter::new(
583 "operation-secret",
584 future_timestamp(30),
585 future_timestamp(120),
586 );
587 let granter = Granter::new(Context::new(), provider, operation);
588 futures::executor::block_on(granter.grant(None)).expect("grant must succeed");
589
590 let debug = format!("{granter:?}");
591 assert!(debug.starts_with("Granter"));
592 assert!(!debug.contains("provider-secret"));
593 assert!(!debug.contains("operation-secret"));
594 assert!(!debug.contains("granted-"));
595 }
596}