1use std::num::NonZeroUsize;
12use std::sync::Arc;
13use std::time::Duration;
14
15use arc_swap::ArcSwapOption;
16use jiff::{SignedDuration, Timestamp};
17use tokio::sync::watch;
18use tollgate_auth::{CredentialVerifier, HmacRegistry, Verified};
19use tollgate_store::{Clock, CredentialSet, KeySource, StoreError};
20use zeroize::Zeroizing;
21
22#[derive(Debug, Clone)]
30pub struct KeyManagerConfig {
31 pub refresh_interval: Duration,
39 pub fetch_timeout: Duration,
45 pub pass_timeout: Duration,
52 pub page_limit: NonZeroUsize,
55 pub max_pages: NonZeroUsize,
61 pub max_age: Duration,
72 pub shutdown_timeout: Duration,
76}
77
78impl Default for KeyManagerConfig {
79 fn default() -> Self {
80 Self {
81 refresh_interval: Duration::from_secs(5),
82 fetch_timeout: Duration::from_secs(5),
83 pass_timeout: Duration::from_secs(10),
84 page_limit: tollgate_store::DEFAULT_KEY_PAGE_LIMIT,
85 max_pages: NonZeroUsize::new(1024).unwrap(),
86 max_age: Duration::from_secs(30),
87 shutdown_timeout: Duration::from_secs(5),
88 }
89 }
90}
91
92#[derive(Debug, Clone, Copy, PartialEq, Eq)]
95pub struct KeyManagerConfigError(pub &'static str);
96impl std::fmt::Display for KeyManagerConfigError {
97 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
98 f.write_str(self.0)
99 }
100}
101impl std::error::Error for KeyManagerConfigError {}
102
103impl KeyManagerConfig {
104 pub fn validate(&self) -> Result<(), KeyManagerConfigError> {
114 for duration in [
115 self.refresh_interval,
116 self.fetch_timeout,
117 self.pass_timeout,
118 self.max_age,
119 self.shutdown_timeout,
120 ] {
121 if duration.is_zero()
122 || tokio::time::Instant::now().checked_add(duration).is_none()
123 || SignedDuration::try_from(duration).is_err()
124 {
125 return Err(KeyManagerConfigError(
126 "key manager durations must be positive and representable",
127 ));
128 }
129 }
130 tollgate_store::validate_key_page_limit(self.page_limit)
131 .map_err(|_| KeyManagerConfigError("credential page limit exceeds 4096"))?;
132 if self.fetch_timeout > self.pass_timeout {
133 return Err(KeyManagerConfigError(
134 "key fetch timeout exceeds the pass budget",
135 ));
136 }
137 if self
140 .pass_timeout
141 .checked_mul(2)
142 .and_then(|fetches| self.refresh_interval.checked_add(fetches))
143 .is_none_or(|cycle| cycle >= self.max_age)
144 {
145 return Err(KeyManagerConfigError(
146 "key projection max_age must exceed refresh_interval + 2 * pass_timeout",
147 ));
148 }
149 Ok(())
150 }
151}
152
153struct Projection {
154 revision: u64,
155 registry: HmacRegistry,
156 fetched_at: Timestamp,
157 usable_until: Timestamp,
158 projected_keys: usize,
159}
160
161#[derive(Clone, Default)]
164pub struct KeyVerifier {
165 current: Arc<ArcSwapOption<Projection>>,
166}
167
168impl CredentialVerifier for KeyVerifier {
169 fn verify(&self, credential: &[u8]) -> Option<Verified> {
170 self.current.load().as_ref()?.registry.verify(credential)
171 }
172}
173
174#[derive(Debug, Clone, Copy, PartialEq, Eq)]
176pub enum KeyManagerHealth {
177 Starting,
179 Healthy,
181 Degraded,
184 Stopped,
186 Failed,
189}
190
191#[derive(Debug, Default, Clone, Copy, PartialEq, Eq)]
193pub struct KeyManagerStats {
194 pub attempts: u64,
196 pub refreshes: u64,
198 pub failures: u64,
202 pub timeouts: u64,
204 pub pages: u64,
206 pub revision_conflicts: u64,
210 pub pass_timeouts: u64,
212 pub page_budget_exceeded: u64,
215 pub counter_overflow: bool,
218}
219
220impl KeyManagerStats {
221 fn increment(value: &mut u64, overflow: &mut bool) {
222 if let Some(next) = value.checked_add(1) {
223 *value = next;
224 } else {
225 *overflow = true;
226 }
227 }
228}
229
230#[derive(Debug, Clone, Copy)]
231struct Progress {
232 health: KeyManagerHealth,
233 stats: KeyManagerStats,
234}
235
236#[derive(Debug, Clone, Copy)]
239pub struct KeyManagerReport {
240 pub health: KeyManagerHealth,
243 pub stats: KeyManagerStats,
245 pub ready: bool,
247 pub projected_keys: usize,
249 pub revision: Option<u64>,
252 pub fetched_at: Option<Timestamp>,
254 pub usable_until: Option<Timestamp>,
257}
258
259#[derive(Clone)]
262pub struct KeyManagerMonitor {
263 progress: watch::Receiver<Progress>,
264 verifier: KeyVerifier,
265}
266impl KeyManagerMonitor {
267 pub fn report(&self, now: Timestamp) -> KeyManagerReport {
271 let exited = self.progress.has_changed().is_err();
272 let mut progress = *self.progress.borrow();
273 if exited && progress.health != KeyManagerHealth::Stopped {
274 progress.health = KeyManagerHealth::Failed;
275 }
276 let projection = self.verifier.current.load();
277 let running = matches!(
278 progress.health,
279 KeyManagerHealth::Healthy | KeyManagerHealth::Degraded
280 );
281 KeyManagerReport {
282 health: progress.health,
283 stats: progress.stats,
284 ready: running && projection.as_ref().is_some_and(|p| now < p.usable_until),
285 projected_keys: projection.as_ref().map_or(0, |p| p.projected_keys),
286 revision: projection.as_ref().map(|p| p.revision),
287 fetched_at: projection.as_ref().map(|p| p.fetched_at),
288 usable_until: projection.as_ref().map(|p| p.usable_until),
289 }
290 }
291
292 pub async fn changed(&mut self) -> Result<(), watch::error::RecvError> {
300 self.progress.changed().await
301 }
302}
303
304#[derive(Debug, Clone, Copy)]
306pub struct KeyManagerShutdownReport {
307 pub stats: KeyManagerStats,
309 pub task_failed: bool,
312 pub deadline_expired: bool,
315}
316
317#[must_use = "retain the key manager to own credential refresh"]
325pub struct KeyManager {
326 task: Option<tokio::task::JoinHandle<()>>,
327 stop: watch::Sender<bool>,
328 monitor: KeyManagerMonitor,
329 shutdown_timeout: Duration,
330}
331
332impl KeyManager {
333 pub fn spawn(
344 source: Arc<dyn KeySource>,
345 secret: &[u8],
346 clock: Arc<dyn Clock>,
347 config: KeyManagerConfig,
348 ) -> Result<Self, KeyManagerConfigError> {
349 config.validate()?;
350 if secret.len() < 32 {
351 return Err(KeyManagerConfigError(
352 "credential projection HMAC secret must contain at least 32 bytes",
353 ));
354 }
355 let max_age = SignedDuration::try_from(config.max_age)
356 .map_err(|_| KeyManagerConfigError("key projection max_age overflow"))?;
357 clock.now().checked_add(max_age).map_err(|_| {
358 KeyManagerConfigError("key projection deadline exceeds the timestamp range")
359 })?;
360 let verifier = KeyVerifier::default();
361 let (stop, stopping) = watch::channel(false);
362 let (publisher, progress) = watch::channel(Progress {
363 health: KeyManagerHealth::Starting,
364 stats: KeyManagerStats::default(),
365 });
366 let monitor = KeyManagerMonitor {
367 progress,
368 verifier: verifier.clone(),
369 };
370 let shutdown_timeout = config.shutdown_timeout;
371 let secret = Zeroizing::new(secret.to_vec());
372 let task = tokio::spawn(run(
373 source,
374 secret,
375 clock,
376 config,
377 max_age,
378 stopping,
379 Publisher {
380 verifier,
381 sender: publisher,
382 },
383 ));
384 Ok(Self {
385 task: Some(task),
386 stop,
387 monitor,
388 shutdown_timeout,
389 })
390 }
391
392 pub fn verifier(&self) -> KeyVerifier {
396 self.monitor.verifier.clone()
397 }
398 pub fn monitor(&self) -> KeyManagerMonitor {
400 self.monitor.clone()
401 }
402
403 pub async fn shutdown(mut self) -> KeyManagerShutdownReport {
408 crate::signal(&self.stop, true, "key-manager shutdown");
409 let joined = tokio::time::timeout(
410 self.shutdown_timeout,
411 self.task.as_mut().expect("key manager owns its task"),
412 )
413 .await;
414 let (task_failed, deadline_expired) = match joined {
415 Ok(Ok(())) => (false, false),
416 Ok(Err(error)) => {
417 tracing::error!(%error, "key manager task died");
418 (true, false)
419 }
420 Err(_) => {
421 tracing::error!("key manager shutdown deadline expired");
422 (true, true)
423 }
424 };
425 KeyManagerShutdownReport {
426 stats: self.monitor.progress.borrow().stats,
427 task_failed,
428 deadline_expired,
429 }
430 }
431}
432
433impl Drop for KeyManager {
434 fn drop(&mut self) {
435 if let Some(task) = self.task.take() {
436 task.abort();
437 }
438 }
439}
440
441struct Publisher {
445 verifier: KeyVerifier,
446 sender: watch::Sender<Progress>,
447}
448impl Drop for Publisher {
449 fn drop(&mut self) {
450 self.verifier.current.store(None);
451 }
452}
453
454fn projection(
455 secret: &[u8],
456 keys: CredentialSet,
457 started: Timestamp,
458 until: Timestamp,
459 clock: &dyn Clock,
460 latest_source_time: Timestamp,
461) -> Result<Projection, StoreError> {
462 let now = clock.now().max(latest_source_time);
463 let registry = HmacRegistry::new(secret);
464 registry.install(
467 keys.records()
468 .iter()
469 .filter(|key| key.not_after.is_none_or(|end| now < end))
470 .map(|key| {
471 (
472 key.principal,
473 key.digest,
474 Some(key.not_after.map_or(until, |end| end.min(until))),
475 )
476 }),
477 );
478 if clock.now() >= until {
481 return Err(StoreError(
482 "credential projection expired before publication".into(),
483 ));
484 }
485 let projected_keys = registry.len();
486 Ok(Projection {
487 revision: 0, registry,
489 fetched_at: started,
490 usable_until: until,
491 projected_keys,
492 })
493}
494
495async fn refresh(
496 source: &dyn KeySource,
497 secret: &[u8],
498 clock: &dyn Clock,
499 config: &KeyManagerConfig,
500 max_age: SignedDuration,
501 started: Timestamp,
502 stats: &mut KeyManagerStats,
503) -> Result<Projection, &'static str> {
504 use tokio::time::Instant;
505 let deadline = Instant::now() + config.pass_timeout;
506 let until = started
507 .checked_add(max_age)
508 .map_err(|_| "timestamp-overflow")?;
509 let mut revision = None;
510 let mut after = None;
511 let mut records = Vec::new();
512 let mut latest_source_time = started;
513 let mut requests = 0usize;
514 loop {
515 if Instant::now() >= deadline {
516 KeyManagerStats::increment(&mut stats.pass_timeouts, &mut stats.counter_overflow);
517 return Err("pass-deadline");
518 }
519 if requests == config.max_pages.get() {
520 KeyManagerStats::increment(
521 &mut stats.page_budget_exceeded,
522 &mut stats.counter_overflow,
523 );
524 return Err("page-budget");
525 }
526 requests += 1; let call_deadline = deadline.min(Instant::now() + config.fetch_timeout);
528 let page = match tokio::time::timeout_at(
529 call_deadline,
530 source.active_keys_page(clock.now(), after, config.page_limit),
531 )
532 .await
533 {
534 Ok(Ok(page)) => page,
535 Ok(Err(_)) => return Err("source-read"),
536 Err(_) => {
537 if call_deadline == deadline {
538 KeyManagerStats::increment(
539 &mut stats.pass_timeouts,
540 &mut stats.counter_overflow,
541 );
542 return Err("pass-deadline");
543 }
544 KeyManagerStats::increment(&mut stats.timeouts, &mut stats.counter_overflow);
545 return Err("call-deadline");
546 }
547 };
548 page.validate_request(after, config.page_limit)
549 .map_err(|_| "page-request-mismatch")?;
550 KeyManagerStats::increment(&mut stats.pages, &mut stats.counter_overflow);
551 if revision.is_some_and(|previous| previous != page.revision()) {
552 KeyManagerStats::increment(&mut stats.revision_conflicts, &mut stats.counter_overflow);
555 records.clear();
556 revision = None;
557 after = None;
558 latest_source_time = started;
559 } else {
560 revision = Some(page.revision());
561 latest_source_time = latest_source_time.max(page.as_of());
562 after = page.next_after();
563 records.extend(page.into_records());
564 if after.is_none() {
565 let keys = CredentialSet::try_new(records).map_err(|_| "invalid-identity-set")?;
566 let mut candidate =
567 projection(secret, keys, started, until, clock, latest_source_time)
568 .map_err(|_| "expired-projection")?;
569 if Instant::now() >= deadline {
570 KeyManagerStats::increment(
571 &mut stats.pass_timeouts,
572 &mut stats.counter_overflow,
573 );
574 return Err("pass-deadline");
575 }
576 candidate.revision = revision.expect("a terminal page supplied its revision");
577 return Ok(candidate);
578 }
579 }
580 tokio::task::yield_now().await;
581 }
582}
583
584async fn run(
585 source: Arc<dyn KeySource>,
586 secret: Zeroizing<Vec<u8>>,
587 clock: Arc<dyn Clock>,
588 config: KeyManagerConfig,
589 max_age: SignedDuration,
590 mut stop: watch::Receiver<bool>,
591 publisher: Publisher,
592) {
593 let mut progress = *publisher.sender.borrow();
594 loop {
595 let started = clock.now();
596 KeyManagerStats::increment(
597 &mut progress.stats.attempts,
598 &mut progress.stats.counter_overflow,
599 );
600 publisher.sender.send_replace(progress);
601 let result = tokio::select! {
602 biased;
603 _ = stop.changed() => break,
604 result = refresh(&*source, &secret, &*clock, &config, max_age, started, &mut progress.stats) => result,
605 };
606 let result = result.and_then(|candidate| {
608 if publisher
609 .verifier
610 .current
611 .load()
612 .as_ref()
613 .is_some_and(|previous| candidate.revision < previous.revision)
614 {
615 Err("revision-regression")
616 } else {
617 Ok(candidate)
618 }
619 });
620 match result {
621 Ok(projection) => {
622 publisher.verifier.current.store(Some(Arc::new(projection)));
623 KeyManagerStats::increment(
624 &mut progress.stats.refreshes,
625 &mut progress.stats.counter_overflow,
626 );
627 if progress.health == KeyManagerHealth::Degraded {
628 tracing::info!("credential projection refresh recovered");
629 }
630 progress.health = KeyManagerHealth::Healthy;
631 }
632 Err(error) => {
633 tracing::warn!(
636 timeouts = progress.stats.timeouts,
637 reason = error,
638 "credential projection refresh failed; retaining the previous expiry"
639 );
640 KeyManagerStats::increment(
641 &mut progress.stats.failures,
642 &mut progress.stats.counter_overflow,
643 );
644 progress.health = KeyManagerHealth::Degraded;
645 }
646 }
647 publisher.sender.send_replace(progress);
648 tokio::select! {
649 biased;
650 _ = stop.changed() => break,
651 _ = tokio::time::sleep(config.refresh_interval) => {},
652 }
653 }
654 progress.health = KeyManagerHealth::Stopped;
655 publisher.sender.send_replace(progress);
656}
657
658#[cfg(test)]
659mod tests {
660 use super::*;
661 use proptest::prelude::*;
662 use tollgate_core::KeyId;
663 use tollgate_store::CredentialRecord;
664
665 proptest! {
666 #[test]
667 fn projected_evidence_is_bounded_by_both_source_expiry_and_fetch_start(
668 started in -1_000_000i64..1_000_000,
669 age in 1i64..10_000,
670 elapsed in 0i64..20_000,
671 expiry in proptest::option::of(-1_010_000i64..1_020_000),
672 ) {
673 let stamp = |s| Timestamp::from_second(s).unwrap();
674 let secret = b"fixture-projection-proof-secret-108";
675 let key = HmacRegistry::new(secret).mint(KeyId(1)).unwrap();
676 let keys = CredentialSet::try_new(vec![CredentialRecord {
677 key_id: key.key_id,
678 principal: key.principal,
679 digest: key.digest,
680 not_after: expiry.map(stamp),
681 }]).unwrap();
682 let clock = tollgate_store::ManualClock::new(stamp(started + elapsed));
683 let result = projection(secret, keys, stamp(started), stamp(started + age), &clock, stamp(started));
684 if elapsed >= age {
685 prop_assert!(result.is_err());
686 } else {
687 let table = result.unwrap();
688 let proof = table.registry.verify(&key.secret);
689 if expiry.is_some_and(|end| end <= started + elapsed) {
690 prop_assert!(proof.is_none());
691 prop_assert_eq!(table.projected_keys, 0);
692 } else {
693 let deadline = proof.unwrap().reusable_until.unwrap().as_second();
694 prop_assert_eq!(deadline, expiry.map_or(started + age, |end| end.min(started + age)));
695 prop_assert!(deadline > started + elapsed);
696 prop_assert_eq!(table.projected_keys, 1);
697 }
698 }
699 }
700 }
701
702 #[test]
703 fn diagnostic_counter_overflow_is_visible_without_wrapping() {
704 let mut count = u64::MAX;
705 let mut overflow = false;
706 KeyManagerStats::increment(&mut count, &mut overflow);
707 assert_eq!(count, u64::MAX);
708 assert!(overflow);
709 }
710}