use std::future::Future;
use std::num::NonZeroUsize;
use std::sync::Arc;
use std::time::Duration;
use chrono::{DateTime, Utc};
use axess_cache::ClockTtlCache;
use axess_clock::{Clock, SystemClock};
use crate::authn::ids::{DeviceId, TenantId, UserId};
use crate::device::store::{DeviceStore, SweepCounts};
use crate::device::types::{Device, DeviceTrustLevel, FingerprintHash};
const DEFAULT_CAPACITY: usize = 10_000;
const DEFAULT_TTL_SECS: u64 = 60;
type CacheKey = (TenantId, DeviceId);
pub struct CachedDeviceStore<S>
where
S: DeviceStore,
{
inner: S,
cache: Arc<ClockTtlCache<CacheKey, Device>>,
}
impl<S> Clone for CachedDeviceStore<S>
where
S: DeviceStore,
{
fn clone(&self) -> Self {
Self {
inner: self.inner.clone(),
cache: self.cache.clone(),
}
}
}
impl<S> CachedDeviceStore<S>
where
S: DeviceStore,
{
pub fn new(inner: S) -> Self {
Self::with_options(
inner,
DEFAULT_CAPACITY,
Duration::from_secs(DEFAULT_TTL_SECS),
Arc::new(SystemClock),
)
}
pub fn with_options(inner: S, capacity: usize, ttl: Duration, clock: Arc<dyn Clock>) -> Self {
let capacity = NonZeroUsize::new(capacity.max(1)).expect("capacity >= 1");
let cache = Arc::new(ClockTtlCache::new(capacity, ttl, clock));
Self { inner, cache }
}
pub fn with_capacity(mut self, capacity: usize) -> Self {
let cap = NonZeroUsize::new(capacity.max(1)).expect("capacity >= 1");
let ttl = Duration::from_secs(DEFAULT_TTL_SECS);
self.cache = Arc::new(ClockTtlCache::new(
cap,
ttl,
Arc::new(SystemClock) as Arc<dyn Clock>,
));
self
}
pub fn with_ttl(self, ttl: Duration) -> Self {
let cap = self.cache.capacity();
let cache = Arc::new(ClockTtlCache::new(
cap,
ttl,
Arc::new(SystemClock) as Arc<dyn Clock>,
));
Self {
inner: self.inner,
cache,
}
}
pub fn with_clock(self, clock: Arc<dyn Clock>) -> Self {
let cap = self.cache.capacity();
let ttl = Duration::from_secs(DEFAULT_TTL_SECS);
let cache = Arc::new(ClockTtlCache::new(cap, ttl, clock));
Self {
inner: self.inner,
cache,
}
}
pub fn stats(&self) -> axess_cache::CacheStats {
self.cache.stats()
}
pub fn invalidate_all(&self) {
self.cache.invalidate_all();
}
pub fn invalidate_tenant(&self, tenant_id: &TenantId) {
let target = *tenant_id;
self.cache.invalidate_by(|k| k.0 == target);
}
}
impl<S> DeviceStore for CachedDeviceStore<S>
where
S: DeviceStore,
{
type Error = S::Error;
fn load(
&self,
tenant_id: &TenantId,
id: &DeviceId,
) -> impl Future<Output = Result<Option<Device>, Self::Error>> + Send {
let key = (*tenant_id, *id);
let cache = self.cache.clone();
let inner = self.inner.clone();
let tenant = *tenant_id;
let device = *id;
async move {
if let Some(d) = cache.get(&key) {
return Ok(Some(d));
}
let result = inner.load(&tenant, &device).await?;
if let Some(ref d) = result {
cache.insert(key, d.clone());
}
Ok(result)
}
}
fn find_by_fingerprint(
&self,
tenant_id: &TenantId,
hash: &FingerprintHash,
) -> impl Future<Output = Result<Option<Device>, Self::Error>> + Send {
let cache = self.cache.clone();
let inner = self.inner.clone();
let tenant = *tenant_id;
let hash = *hash;
async move {
let result = inner.find_by_fingerprint(&tenant, &hash).await?;
if let Some(ref d) = result {
cache.insert((tenant, d.id), d.clone());
}
Ok(result)
}
}
fn find_for_user(
&self,
tenant_id: &TenantId,
user_id: &UserId,
limit: usize,
) -> impl Future<Output = Result<Vec<Device>, Self::Error>> + Send {
self.inner.find_for_user(tenant_id, user_id, limit)
}
fn find_by_refresh_family(
&self,
tenant_id: &TenantId,
family_id: &str,
) -> impl Future<Output = Result<Vec<Device>, Self::Error>> + Send {
self.inner.find_by_refresh_family(tenant_id, family_id)
}
fn save(&self, device: &Device) -> impl Future<Output = Result<(), Self::Error>> + Send {
let key = (device.tenant_id, device.id);
let cache = self.cache.clone();
let inner = self.inner.clone();
let device = device.clone();
async move {
cache.invalidate(&key);
inner.save(&device).await?;
cache.insert(key, device);
Ok(())
}
}
fn record_sighting(
&self,
tenant_id: &TenantId,
id: &DeviceId,
now: DateTime<Utc>,
) -> impl Future<Output = Result<(), Self::Error>> + Send {
self.inner.record_sighting(tenant_id, id, now)
}
fn set_trust_level(
&self,
tenant_id: &TenantId,
id: &DeviceId,
level: DeviceTrustLevel,
now: DateTime<Utc>,
) -> impl Future<Output = Result<(), Self::Error>> + Send {
let key = (*tenant_id, *id);
let cache = self.cache.clone();
let inner = self.inner.clone();
let tenant = *tenant_id;
let device = *id;
async move {
cache.invalidate(&key);
inner.set_trust_level(&tenant, &device, level, now).await
}
}
fn delete(
&self,
tenant_id: &TenantId,
id: &DeviceId,
) -> impl Future<Output = Result<(), Self::Error>> + Send {
let key = (*tenant_id, *id);
let cache = self.cache.clone();
let inner = self.inner.clone();
let tenant = *tenant_id;
let device = *id;
async move {
cache.invalidate(&key);
inner.delete(&tenant, &device).await
}
}
fn sweep(
&self,
tenant_id: &TenantId,
now: DateTime<Utc>,
) -> impl Future<Output = Result<SweepCounts, Self::Error>> + Send {
self.inner.sweep(tenant_id, now)
}
}
#[cfg(test)]
mod tests;