use alloc::boxed::Box;
use alloc::vec::Vec;
use core::mem;
use core::time::Duration;
use std::sync::{RwLock, RwLockReadGuard};
use std::time::Instant;
use crate::crypto::TicketProducer;
use crate::error::Error;
pub struct TicketRotator {
pub(crate) generator: fn() -> Result<Box<dyn TicketProducer>, Error>,
lifetime: Duration,
state: RwLock<TicketRotatorState>,
}
impl TicketRotator {
pub fn new(
lifetime: Duration,
generator: fn() -> Result<Box<dyn TicketProducer>, Error>,
) -> Result<Self, Error> {
Ok(Self {
generator,
lifetime,
state: RwLock::new(TicketRotatorState {
current: Some(Generation {
producer: generator()?,
expires_at: Instant::now() + lifetime,
}),
previous: None,
}),
})
}
fn encrypt_at(&self, message: &[u8], now: Instant) -> Option<Vec<u8>> {
let state = self.maybe_roll(now)?;
if let Some(current) = &state.current {
return current.producer.encrypt(message);
}
let Some(prev) = &state.previous else {
return None;
};
if !prev.in_grace_period(now, self.lifetime) {
return None;
}
prev.producer.encrypt(message)
}
fn decrypt_at(&self, ciphertext: &[u8], now: Instant) -> Option<Vec<u8>> {
let state = self.maybe_roll(now)?;
if let Some(current) = &state.current {
if let Some(plain) = current.producer.decrypt(ciphertext) {
return Some(plain);
}
}
let Some(prev) = &state.previous else {
return None;
};
if !prev.in_grace_period(now, self.lifetime) {
return None;
}
prev.producer.decrypt(ciphertext)
}
pub(crate) fn maybe_roll(
&self,
now: Instant,
) -> Option<RwLockReadGuard<'_, TicketRotatorState>> {
{
let read = self.state.read().ok()?;
match &read.current {
Some(current) if now <= current.expires_at => return Some(read),
_ => {}
}
}
let mut write = self.state.write().ok()?;
if let Some(current) = &write.current {
if now <= current.expires_at {
drop(write);
return self.state.read().ok();
}
}
let next = (self.generator)()
.ok()
.map(|producer| Generation {
producer,
expires_at: now + self.lifetime,
});
let prev = mem::replace(&mut write.current, next);
if prev.is_some() {
write.previous = prev;
}
drop(write);
self.state.read().ok()
}
}
impl TicketProducer for TicketRotator {
fn encrypt(&self, message: &[u8]) -> Option<Vec<u8>> {
self.encrypt_at(message, Instant::now())
}
fn decrypt(&self, ciphertext: &[u8]) -> Option<Vec<u8>> {
self.decrypt_at(ciphertext, Instant::now())
}
fn lifetime(&self) -> Duration {
self.lifetime
}
}
impl core::fmt::Debug for TicketRotator {
fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
f.debug_struct("TicketRotator")
.finish_non_exhaustive()
}
}
#[derive(Debug)]
pub(crate) struct TicketRotatorState {
current: Option<Generation>,
previous: Option<Generation>,
}
#[derive(Debug)]
struct Generation {
producer: Box<dyn TicketProducer>,
expires_at: Instant,
}
impl Generation {
fn in_grace_period(&self, now: Instant, lifetime: Duration) -> bool {
now <= self.expires_at + lifetime
}
}
#[cfg(test)]
mod tests {
use core::sync::atomic::{AtomicU8, Ordering};
use core::time::Duration;
use super::*;
#[test]
fn ticketrotator_switching_test() {
let t = TicketRotator::new(Duration::from_secs(1), FakeTicketer::new).unwrap();
let now = Instant::now();
let cipher1 = t.encrypt(b"ticket 1").unwrap();
assert_eq!(t.decrypt(&cipher1).unwrap(), b"ticket 1");
{
t.maybe_roll(now + Duration::from_secs(10));
}
let cipher2 = t.encrypt(b"ticket 2").unwrap();
assert_eq!(t.decrypt(&cipher1).unwrap(), b"ticket 1");
assert_eq!(t.decrypt(&cipher2).unwrap(), b"ticket 2");
{
t.maybe_roll(now + Duration::from_secs(20));
}
let cipher3 = t.encrypt(b"ticket 3").unwrap();
assert!(t.decrypt(&cipher1).is_none());
assert_eq!(t.decrypt(&cipher2).unwrap(), b"ticket 2");
assert_eq!(t.decrypt(&cipher3).unwrap(), b"ticket 3");
}
#[test]
fn ticketrotator_remains_usable_over_temporary_ticketer_creation_failure() {
let mut t = TicketRotator::new(Duration::from_secs(1), FakeTicketer::new).unwrap();
let expiry = t
.state
.read()
.unwrap()
.current
.as_ref()
.unwrap()
.expires_at;
let cipher1 = t.encrypt(b"ticket 1").unwrap();
assert_eq!(t.decrypt(&cipher1).unwrap(), b"ticket 1");
t.generator = fail_generator;
let t1 = expiry;
drop(t.maybe_roll(t1));
assert!(t.encrypt_at(b"ticket 2", t1).is_some());
let t2 = expiry + Duration::from_secs(1);
let cipher3 = t.encrypt_at(b"ticket 3", t2).unwrap();
assert_eq!(t.decrypt_at(&cipher1, t2).unwrap(), b"ticket 1");
assert_eq!(t.decrypt_at(&cipher3, t2).unwrap(), b"ticket 3");
let t3 = expiry + Duration::from_secs(2);
assert_eq!(t.encrypt_at(b"ticket 4", t3), None);
assert_eq!(t.decrypt_at(&cipher3, t3), None);
t.generator = FakeTicketer::new;
let t4 = expiry + Duration::from_secs(3);
drop(t.maybe_roll(t4));
let t5 = expiry + Duration::from_secs(4);
let cipher5 = t.encrypt_at(b"ticket 5", t5).unwrap();
assert!(t.decrypt_at(&cipher1, t5).is_none());
assert!(t.decrypt_at(&cipher3, t5).is_none());
assert_eq!(t.decrypt_at(&cipher5, t5).unwrap(), b"ticket 5");
t.generator = fail_generator;
let mut write = t.state.write().unwrap();
write.current = None;
write.previous = None;
drop(write);
assert!(t.encrypt(b"ticket 6").is_none());
}
#[derive(Debug)]
struct FakeTicketer {
generation: u8,
}
impl FakeTicketer {
#[expect(clippy::new_ret_no_self)]
fn new() -> Result<Box<dyn TicketProducer>, Error> {
Ok(Box::new(Self {
generation: std::dbg!(FAKE_GEN.fetch_add(1, Ordering::SeqCst)),
}))
}
}
impl TicketProducer for FakeTicketer {
fn encrypt(&self, message: &[u8]) -> Option<Vec<u8>> {
let mut v = Vec::with_capacity(1 + message.len());
v.push(self.generation);
v.extend(
message
.iter()
.copied()
.map(|b| b ^ self.generation),
);
Some(v)
}
fn decrypt(&self, ciphertext: &[u8]) -> Option<Vec<u8>> {
if ciphertext.first()? != &self.generation {
return None;
}
Some(
ciphertext[1..]
.iter()
.copied()
.map(|b| b ^ self.generation)
.collect(),
)
}
fn lifetime(&self) -> Duration {
Duration::ZERO }
}
static FAKE_GEN: AtomicU8 = AtomicU8::new(0);
fn fail_generator() -> Result<Box<dyn TicketProducer>, Error> {
Err(Error::FailedToGetRandomBytes)
}
}