use super::*;
use axess_clock::testing::MockClock;
use std::collections::HashSet;
use std::str::FromStr;
use std::sync::atomic::{AtomicUsize, Ordering};
struct CountingProvider {
calls: Arc<AtomicUsize>,
}
impl RequestEntityProvider for CountingProvider {
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>,
> {
let _ = (session, action);
let calls = self.calls.clone();
let principal = principal.clone();
let resource = resource.clone();
Box::pin(async move {
calls.fetch_add(1, Ordering::SeqCst);
let p = cedar_policy::Entity::new(
principal,
std::collections::HashMap::new(),
HashSet::new(),
)
.unwrap();
let r = cedar_policy::Entity::new(
resource,
std::collections::HashMap::new(),
HashSet::new(),
)
.unwrap();
Ok(Entities::from_entities(vec![p, r], None).unwrap())
})
}
}
fn guest_session() -> AuthSession {
use crate::session::SessionData;
use crate::session::id::SessionId;
use crate::session::layer::{SessionHandle, SessionInner};
use tokio::sync::RwLock;
let inner = SessionInner {
id: SessionId::new(&axess_rng::SystemRng),
data: SessionData::default(),
modified: false,
regenerate: false,
pre_cycle_id: None,
pending_fingerprint: None,
max_custom_bytes: 64 * 1024,
};
AuthSession(SessionHandle(Arc::new(RwLock::new(inner))))
}
fn principal() -> EntityUid {
EntityUid::from_str("App::User::\"alice\"").unwrap()
}
fn action() -> EntityUid {
EntityUid::from_str("App::Action::\"View\"").unwrap()
}
fn doc(id: &str) -> EntityUid {
EntityUid::from_str(&format!("App::Doc::\"{id}\"")).unwrap()
}
#[tokio::test]
async fn first_call_misses_then_caches() {
let calls = Arc::new(AtomicUsize::new(0));
let cached = EntityCache::new(CountingProvider {
calls: calls.clone(),
});
let s = guest_session();
let p = principal();
let a = action();
let r1 = doc("doc-1");
let _ = cached.entities_for(&s, &p, &r1, &a).await.unwrap();
assert_eq!(calls.load(Ordering::SeqCst), 1);
let _ = cached.entities_for(&s, &p, &r1, &a).await.unwrap();
assert_eq!(
calls.load(Ordering::SeqCst),
1,
"cache hit should not invoke inner"
);
let r2 = doc("doc-2");
let _ = cached.entities_for(&s, &p, &r2, &a).await.unwrap();
assert_eq!(calls.load(Ordering::SeqCst), 2);
}
#[tokio::test]
async fn invalidate_evicts_cached_entry() {
let calls = Arc::new(AtomicUsize::new(0));
let cached = EntityCache::new(CountingProvider {
calls: calls.clone(),
});
let s = guest_session();
let p = principal();
let a = action();
let r = doc("doc-1");
let _ = cached.entities_for(&s, &p, &r, &a).await.unwrap();
let _ = cached.entities_for(&s, &p, &r, &a).await.unwrap();
assert_eq!(calls.load(Ordering::SeqCst), 1);
cached.invalidate(&p, None, &r, &a);
let _ = cached.entities_for(&s, &p, &r, &a).await.unwrap();
assert_eq!(
calls.load(Ordering::SeqCst),
2,
"after invalidate, next call should re-invoke inner"
);
}
#[tokio::test]
async fn concurrent_cold_misses_share_one_inner_call() {
let calls = Arc::new(AtomicUsize::new(0));
let cached = Arc::new(EntityCache::new(CountingProvider {
calls: calls.clone(),
}));
let p = principal();
let a = action();
let r = doc("doc-1");
const N: usize = 8;
let mut handles = Vec::with_capacity(N);
for _ in 0..N {
let cached = cached.clone();
let p = p.clone();
let a = a.clone();
let r = r.clone();
let s = guest_session();
handles.push(tokio::spawn(async move {
cached.entities_for(&s, &p, &r, &a).await.map(|_| ())
}));
}
for h in handles {
h.await.unwrap().unwrap();
}
assert_eq!(
calls.load(Ordering::SeqCst),
1,
"single-flight must collapse N concurrent cold misses into 1 inner call"
);
}
#[tokio::test]
async fn entries_expire_under_injected_clock() {
let clock = Arc::new(MockClock::now());
let calls = Arc::new(AtomicUsize::new(0));
let cached = EntityCache::with_options(
CountingProvider {
calls: calls.clone(),
},
DEFAULT_CAPACITY,
Duration::from_secs(60),
clock.clone() as Arc<dyn Clock>,
);
let s = guest_session();
let p = principal();
let a = action();
let r = doc("doc-1");
let _ = cached.entities_for(&s, &p, &r, &a).await.unwrap();
assert_eq!(calls.load(Ordering::SeqCst), 1);
clock.advance_secs(30);
let _ = cached.entities_for(&s, &p, &r, &a).await.unwrap();
assert_eq!(calls.load(Ordering::SeqCst), 1, "still inside TTL");
clock.advance_secs(31);
let _ = cached.entities_for(&s, &p, &r, &a).await.unwrap();
assert_eq!(
calls.load(Ordering::SeqCst),
2,
"TTL expired under MockClock; must re-fetch from inner"
);
}