use std::num::NonZeroUsize;
use std::sync::Arc;
use std::sync::Mutex;
use std::time::{Duration, Instant};
use lru::LruCache;
use super::did::DidDocument;
pub(crate) trait Clock: Send + Sync {
fn now(&self) -> Instant;
}
pub(crate) struct SystemClock;
impl Clock for SystemClock {
fn now(&self) -> Instant {
Instant::now()
}
}
#[derive(Debug, Clone)]
pub(crate) enum CachedResolve {
Ok(DidDocument),
Err,
}
pub(crate) struct DidDocCache {
inner: Mutex<LruCache<String, (CachedResolve, Instant)>>,
positive_ttl: Duration,
negative_ttl: Duration,
clock: Arc<dyn Clock>,
}
impl DidDocCache {
pub fn new(
size: NonZeroUsize,
positive_ttl: Duration,
negative_ttl: Duration,
clock: Arc<dyn Clock>,
) -> Self {
Self {
inner: Mutex::new(LruCache::new(size)),
positive_ttl,
negative_ttl,
clock,
}
}
pub fn get(&self, did: &str) -> Option<CachedResolve> {
let mut inner = self.inner.lock().expect("cache mutex poisoned");
let (cached, expires_at) = inner.get(did)?;
if self.clock.now() >= *expires_at {
let key = did.to_owned();
inner.pop(&key);
return None;
}
Some(cached.clone())
}
pub fn insert_ok(&self, did: String, doc: DidDocument) {
let exp = self.clock.now() + self.positive_ttl;
self.inner
.lock()
.expect("cache mutex poisoned")
.put(did, (CachedResolve::Ok(doc), exp));
}
pub fn insert_err(&self, did: String) {
let exp = self.clock.now() + self.negative_ttl;
self.inner
.lock()
.expect("cache mutex poisoned")
.put(did, (CachedResolve::Err, exp));
}
}
pub(crate) struct JtiCache {
inner: Mutex<LruCache<(String, String), Instant>>,
clock: Arc<dyn Clock>,
}
impl JtiCache {
pub fn new(size: NonZeroUsize, clock: Arc<dyn Clock>) -> Self {
Self {
inner: Mutex::new(LruCache::new(size)),
clock,
}
}
pub fn check_and_record(
&self,
iss: &str,
jti: &str,
expires_at: Instant,
) -> Result<(), Replay> {
let mut inner = self.inner.lock().expect("cache mutex poisoned");
let key = (iss.to_owned(), jti.to_owned());
if let Some(existing_exp) = inner.get(&key)
&& self.clock.now() < *existing_exp
{
return Err(Replay);
}
inner.put(key, expires_at);
Ok(())
}
}
#[derive(Debug, thiserror::Error)]
#[error("replay detected")]
pub struct Replay;
#[cfg(test)]
mod tests {
use super::*;
use std::sync::atomic::{AtomicU64, Ordering};
fn nz(n: usize) -> NonZeroUsize {
NonZeroUsize::new(n).unwrap()
}
struct MockClock {
base: Instant,
offset_ms: AtomicU64,
}
impl MockClock {
fn new() -> Self {
Self {
base: Instant::now(),
offset_ms: AtomicU64::new(0),
}
}
fn advance(&self, by: Duration) {
self.offset_ms
.fetch_add(by.as_millis() as u64, Ordering::Relaxed);
}
}
impl Clock for MockClock {
fn now(&self) -> Instant {
self.base + Duration::from_millis(self.offset_ms.load(Ordering::Relaxed))
}
}
#[test]
fn doc_cache_returns_cached_then_expires() {
let clock = Arc::new(MockClock::new());
let cache = DidDocCache::new(
nz(10),
Duration::from_millis(50),
Duration::from_millis(5),
clock.clone(),
);
let doc = DidDocument {
id: "did:plc:a".into(),
verification_method: vec![],
};
cache.insert_ok("did:plc:a".into(), doc.clone());
match cache.get("did:plc:a").expect("hit") {
CachedResolve::Ok(got) => assert_eq!(got.id, "did:plc:a"),
_ => panic!("expected Ok"),
}
clock.advance(Duration::from_millis(60));
assert!(cache.get("did:plc:a").is_none(), "must expire");
}
#[test]
fn doc_cache_negative_has_shorter_ttl() {
let clock = Arc::new(MockClock::new());
let cache = DidDocCache::new(
nz(10),
Duration::from_secs(60),
Duration::from_millis(5),
clock.clone(),
);
cache.insert_err("did:plc:bad".into());
assert!(matches!(cache.get("did:plc:bad"), Some(CachedResolve::Err)));
clock.advance(Duration::from_millis(15));
assert!(cache.get("did:plc:bad").is_none(), "neg entry must expire");
}
#[test]
fn jti_cache_second_use_is_replay() {
let clock = Arc::new(MockClock::new());
let cache = JtiCache::new(nz(1000), clock.clone());
let exp = clock.now() + Duration::from_secs(60);
cache
.check_and_record("did:plc:a", "j1", exp)
.expect("first ok");
let second = cache.check_and_record("did:plc:a", "j1", exp);
assert!(second.is_err(), "second use must be replay");
}
#[test]
fn jti_cache_expiry_permits_reuse() {
let clock = Arc::new(MockClock::new());
let cache = JtiCache::new(nz(1000), clock.clone());
let exp = clock.now() + Duration::from_millis(5);
cache
.check_and_record("did:plc:a", "j1", exp)
.expect("first ok");
clock.advance(Duration::from_millis(20));
cache
.check_and_record("did:plc:a", "j1", exp)
.expect("reuse ok post-expiry");
}
#[test]
fn jti_cache_distinguishes_iss() {
let clock = Arc::new(MockClock::new());
let cache = JtiCache::new(nz(1000), clock.clone());
let exp = clock.now() + Duration::from_secs(60);
cache.check_and_record("did:plc:a", "same", exp).unwrap();
cache
.check_and_record("did:plc:b", "same", exp)
.expect("different iss, same jti, not replay");
}
}