use std::num::NonZeroUsize;
use std::sync::Mutex;
use std::time::{Duration, Instant};
use lru::LruCache;
use super::did::DidDocument;
#[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,
}
impl DidDocCache {
pub fn new(size: NonZeroUsize, positive_ttl: Duration, negative_ttl: Duration) -> Self {
Self {
inner: Mutex::new(LruCache::new(size)),
positive_ttl,
negative_ttl,
}
}
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 Instant::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 = Instant::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 = Instant::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>>,
}
impl JtiCache {
pub fn new(size: NonZeroUsize) -> Self {
Self {
inner: Mutex::new(LruCache::new(size)),
}
}
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)
&& Instant::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::*;
fn nz(n: usize) -> NonZeroUsize {
NonZeroUsize::new(n).unwrap()
}
#[test]
fn doc_cache_returns_cached_then_expires() {
let cache = DidDocCache::new(nz(10), Duration::from_millis(50), Duration::from_millis(5));
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"),
}
std::thread::sleep(Duration::from_millis(60));
assert!(cache.get("did:plc:a").is_none(), "must expire");
}
#[test]
fn doc_cache_negative_has_shorter_ttl() {
let cache = DidDocCache::new(nz(10), Duration::from_secs(60), Duration::from_millis(5));
cache.insert_err("did:plc:bad".into());
assert!(matches!(cache.get("did:plc:bad"), Some(CachedResolve::Err)));
std::thread::sleep(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 cache = JtiCache::new(nz(1000));
let exp = Instant::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 cache = JtiCache::new(nz(1000));
let exp = Instant::now() + Duration::from_millis(5);
cache
.check_and_record("did:plc:a", "j1", exp)
.expect("first ok");
std::thread::sleep(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 cache = JtiCache::new(nz(1000));
let exp = Instant::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");
}
}