use std::num::NonZeroUsize;
use std::sync::Arc;
use std::time::Duration;
use cedar_policy::{Entities, EntityUid};
use axess_cache::ClockTtlCache;
use axess_clock::{Clock, SystemClock};
use crate::authz::error::AuthzError;
use crate::authz::provider::RequestEntityProvider;
use crate::session::AuthSession;
const DEFAULT_CAPACITY: usize = 10_000;
const DEFAULT_TTL_SECS: u64 = 60;
#[derive(Hash, Eq, PartialEq, Clone)]
struct EntityCacheKey {
principal: EntityUid,
tenant: Option<String>,
resource: EntityUid,
action: EntityUid,
}
pub struct EntityCache<P>
where
P: RequestEntityProvider,
{
inner: P,
cache: ClockTtlCache<EntityCacheKey, Arc<Entities>>,
}
impl<P> EntityCache<P>
where
P: RequestEntityProvider,
{
pub fn new(inner: P) -> Self {
Self::with_options(
inner,
DEFAULT_CAPACITY,
Duration::from_secs(DEFAULT_TTL_SECS),
Arc::new(SystemClock) as Arc<dyn Clock>,
)
}
pub fn with_options(inner: P, capacity: usize, ttl: Duration, clock: Arc<dyn Clock>) -> Self {
let cap = NonZeroUsize::new(capacity.max(1)).expect("capacity is at least 1");
Self {
inner,
cache: ClockTtlCache::new(cap, ttl, clock),
}
}
pub fn with_capacity(self, capacity: usize) -> Self {
let ttl = Duration::from_secs(DEFAULT_TTL_SECS);
Self::with_options(self.inner, capacity, ttl, Arc::new(SystemClock))
}
pub fn with_ttl(self, ttl: Duration) -> Self {
Self::with_options(self.inner, DEFAULT_CAPACITY, ttl, Arc::new(SystemClock))
}
pub fn with_clock(self, clock: Arc<dyn Clock>) -> Self {
Self::with_options(
self.inner,
DEFAULT_CAPACITY,
Duration::from_secs(DEFAULT_TTL_SECS),
clock,
)
}
pub fn invalidate(
&self,
principal: &EntityUid,
tenant: Option<&str>,
resource: &EntityUid,
action: &EntityUid,
) {
let key = EntityCacheKey {
principal: principal.clone(),
tenant: tenant.map(str::to_string),
resource: resource.clone(),
action: action.clone(),
};
self.cache.invalidate(&key);
}
pub fn invalidate_all(&self) {
self.cache.invalidate_all();
}
pub fn invalidate_principal(&self, principal: &EntityUid) -> usize {
self.cache.invalidate_by(|k| &k.principal == principal)
}
pub fn invalidate_tenant(&self, tenant: &str) -> usize {
self.cache
.invalidate_by(|k| k.tenant.as_deref() == Some(tenant))
}
pub fn inner(&self) -> &P {
&self.inner
}
pub fn stats(&self) -> axess_cache::CacheStats {
self.cache.stats()
}
pub fn reset_stats(&self) {
self.cache.reset_stats();
}
pub fn flush_metrics(&self, metrics: &dyn crate::metrics::AuthnMetrics) {
let snapshot = self.stats();
for _ in 0..snapshot.hits {
metrics.authz_cache_hit();
}
for _ in 0..snapshot.misses {
metrics.authz_cache_miss();
}
for _ in 0..snapshot.capacity_evictions {
metrics.authz_cache_eviction();
}
for _ in 0..snapshot.invalidations {
metrics.authz_cache_invalidation();
}
self.reset_stats();
}
}
impl<P> super::invalidator::CacheInvalidator for EntityCache<P>
where
P: RequestEntityProvider + 'static,
{
type Error = std::convert::Infallible;
async fn invalidate_principal(&self, principal: &EntityUid) -> Result<(), Self::Error> {
let _ = EntityCache::invalidate_principal(self, principal);
Ok(())
}
async fn invalidate_tenant(&self, tenant: &str) -> Result<(), Self::Error> {
let _ = EntityCache::invalidate_tenant(self, tenant);
Ok(())
}
async fn invalidate_all(&self) -> Result<(), Self::Error> {
EntityCache::invalidate_all(self);
Ok(())
}
}
impl<P> RequestEntityProvider for EntityCache<P>
where
P: RequestEntityProvider,
{
fn entities_for<'a>(
&'a self,
session: &'a AuthSession,
principal: &'a EntityUid,
resource: &'a EntityUid,
action: &'a EntityUid,
) -> std::pin::Pin<
Box<dyn std::future::Future<Output = Result<Entities, AuthzError>> + Send + 'a>,
> {
Box::pin(async move {
let tenant = session.tenant_id().await.map(|t| t.to_string().to_string());
let key = EntityCacheKey {
principal: principal.clone(),
tenant,
resource: resource.clone(),
action: action.clone(),
};
let arc = self
.cache
.get_or_try_insert_with(key, || async {
let entities = self
.inner
.entities_for(session, principal, resource, action)
.await?;
Ok::<Arc<Entities>, AuthzError>(Arc::new(entities))
})
.await?;
Ok((*arc).clone())
})
}
}
#[cfg(test)]
mod tests;