use std::sync::Arc;
use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
use std::time::Duration;
use async_trait::async_trait;
use amalgam::{
Backplane, Cache, CacheEvent, CacheRegistry, CircuitComponent, Clock,
DefaultEntryOptionsProvider, DistributedCache, DistributedLocker, DistributedSerializer,
EntryOptions, InMemoryDistributedCache, InMemoryDistributedLocker, InProcessBackplane,
JsonSerializer, ManualClock, MaybeValue, Plugin, RecoveryConfig, Result, Tag,
};
struct CountingPlugin {
sets: Arc<AtomicUsize>,
hits: Arc<AtomicUsize>,
}
impl Plugin for CountingPlugin {
fn name(&self) -> &str {
"counting"
}
fn on_event(&self, event: &CacheEvent) {
match event {
CacheEvent::Set { .. } => {
self.sets.fetch_add(1, Ordering::SeqCst);
}
CacheEvent::Hit { stale: false, .. } => {
self.hits.fetch_add(1, Ordering::SeqCst);
}
_ => {}
}
}
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn plugin_receives_set_and_hit_events() {
let sets = Arc::new(AtomicUsize::new(0));
let hits = Arc::new(AtomicUsize::new(0));
let plugin = Arc::new(CountingPlugin {
sets: sets.clone(),
hits: hits.clone(),
});
let clock = Arc::new(ManualClock::default());
let dyn_clock: Arc<dyn Clock> = clock.clone();
let cache: Cache<i32> = Cache::builder().clock(dyn_clock).plugin(plugin).build();
cache.set("k", 1).await;
let v = cache
.get_or_set("k", |ctx| async move { Ok(ctx.value(999)) })
.await
.unwrap();
assert_eq!(v, 1, "fresh L1 value short-circuits the factory");
assert!(
sets.load(Ordering::SeqCst) >= 1,
"plugin observed at least one Set event"
);
assert!(
hits.load(Ordering::SeqCst) >= 1,
"plugin observed at least one fresh Hit event"
);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn distributed_locker_enforces_cross_instance_single_flight() {
let clock = Arc::new(ManualClock::default());
let dyn_clock: Arc<dyn Clock> = clock.clone();
let l2: Arc<dyn DistributedCache> = Arc::new(InMemoryDistributedCache::new(dyn_clock.clone()));
let locker: Arc<dyn DistributedLocker> =
Arc::new(InMemoryDistributedLocker::new(dyn_clock.clone()));
let serializer: Arc<dyn DistributedSerializer<String>> = Arc::new(JsonSerializer);
let opts = EntryOptions::new(Duration::from_secs(60));
let build = || -> Cache<String> {
Cache::builder()
.clock(dyn_clock.clone())
.distributed(l2.clone())
.serializer(serializer.clone())
.distributed_locker(locker.clone())
.default_options(opts.clone())
.build()
};
let cache1 = build();
let cache2 = build();
let calls = Arc::new(AtomicUsize::new(0));
let slow_factory = |calls: Arc<AtomicUsize>| {
move |ctx: amalgam::FactoryContext<String>| async move {
calls.fetch_add(1, Ordering::SeqCst);
tokio::time::sleep(Duration::from_millis(50)).await;
Ok(ctx.value("from-factory".to_owned()))
}
};
let h1 = {
let cache = cache1.clone();
let f = slow_factory(calls.clone());
tokio::spawn(async move { cache.get_or_set("k", f).await })
};
let h2 = {
let cache = cache2.clone();
let f = slow_factory(calls.clone());
tokio::spawn(async move { cache.get_or_set("k", f).await })
};
let v1 = h1.await.unwrap().unwrap();
let v2 = h2.await.unwrap().unwrap();
assert_eq!(v1, "from-factory");
assert_eq!(v2, "from-factory", "loser served the winner's L2 value");
assert_eq!(
calls.load(Ordering::SeqCst),
1,
"the distributed lock collapsed both flights into a single factory run"
);
}
struct FlakyL2 {
inner: Arc<InMemoryDistributedCache>,
down: Arc<AtomicBool>,
}
#[async_trait]
impl DistributedCache for FlakyL2 {
async fn get(&self, key: &str) -> Result<Option<Vec<u8>>> {
if self.down.load(Ordering::SeqCst) {
return Err(amalgam::Error::Distributed("down".into()));
}
self.inner.get(key).await
}
async fn set(&self, key: &str, value: Vec<u8>, ttl: Option<Duration>) -> Result<()> {
if self.down.load(Ordering::SeqCst) {
return Err(amalgam::Error::Distributed("down".into()));
}
self.inner.set(key, value, ttl).await
}
async fn remove(&self, key: &str) -> Result<()> {
self.inner.remove(key).await
}
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn circuit_breaker_opens_then_auto_recovery_replays_write() {
let clock = Arc::new(ManualClock::default());
let dyn_clock: Arc<dyn Clock> = clock.clone();
let inner_l2 = Arc::new(InMemoryDistributedCache::new(dyn_clock.clone()));
let down = Arc::new(AtomicBool::new(true)); let flaky: Arc<dyn DistributedCache> = Arc::new(FlakyL2 {
inner: inner_l2.clone(),
down: down.clone(),
});
let serializer: Arc<dyn DistributedSerializer<String>> = Arc::new(JsonSerializer);
let cache: Cache<String> = Cache::builder()
.clock(dyn_clock.clone())
.distributed(flaky)
.serializer(serializer)
.distributed_circuit_breaker(Duration::from_secs(30))
.auto_recovery(RecoveryConfig {
enabled: true,
delay: Duration::from_millis(100),
max_items: None,
max_retries: None,
})
.build();
let mut events = cache.events().subscribe();
cache.set("k", "v1".to_owned()).await;
const L2_KEY: &str = "v1:k";
assert!(
inner_l2.get(L2_KEY).await.unwrap().is_none(),
"the failed write left nothing in L2"
);
let mut saw_open = false;
for _ in 0..16 {
match events.try_recv() {
Ok(CacheEvent::CircuitBreakerChange {
component: CircuitComponent::Distributed,
closed: false,
}) => {
saw_open = true;
break;
}
Ok(_) => {}
Err(_) => break,
}
}
assert!(
saw_open,
"tripping the L2 breaker emitted CircuitBreakerChange{{ closed: false }}"
);
down.store(false, Ordering::SeqCst);
let mut recovered = false;
for _ in 0..20 {
tokio::time::sleep(Duration::from_millis(50)).await; if inner_l2.get(L2_KEY).await.unwrap().is_some() {
recovered = true;
break;
}
}
assert!(
recovered,
"auto-recovery replayed the write to the inner L2 once it came back up"
);
drop(cache);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn remove_by_tag_propagates_across_nodes() {
let clock = Arc::new(ManualClock::default());
let dyn_clock: Arc<dyn Clock> = clock.clone();
let backplane: Arc<dyn Backplane> = Arc::new(InProcessBackplane::default());
let long = || EntryOptions::new(Duration::from_secs(100));
let tagged = || -> Box<[Tag]> { Box::from([Tag::new("group").unwrap()]) };
let build = |id: &str| -> Cache<String> {
Cache::builder()
.clock(dyn_clock.clone())
.backplane(backplane.clone())
.instance_id(id)
.build()
};
let node_a = build("node-a");
let node_b = build("node-b");
let calls = Arc::new(AtomicUsize::new(0));
{
let calls = calls.clone();
node_a
.get_or_set_full(
"k",
move |ctx| async move {
calls.fetch_add(1, Ordering::SeqCst);
Ok(ctx.value("v1".to_owned()))
},
Some(long()),
tagged(),
MaybeValue::none(),
)
.await
.unwrap();
}
assert_eq!(calls.load(Ordering::SeqCst), 1);
clock.advance(Duration::from_secs(1));
node_b.remove_by_tag("group").await;
tokio::time::sleep(Duration::from_millis(80)).await;
clock.advance(Duration::from_secs(1));
let v = {
let calls = calls.clone();
node_a
.get_or_set_full(
"k",
move |ctx| async move {
calls.fetch_add(1, Ordering::SeqCst);
Ok(ctx.value("v2".to_owned()))
},
Some(long()),
tagged(),
MaybeValue::none(),
)
.await
.unwrap()
};
assert_eq!(
v, "v2",
"the propagated tag marker invalidated node A's entry"
);
assert_eq!(
calls.load(Ordering::SeqCst),
2,
"node A re-ran the factory after the cross-node tag invalidation"
);
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn registry_resolves_named_caches_independently() {
let clock = Arc::new(ManualClock::default());
let dyn_clock: Arc<dyn Clock> = clock.clone();
let registry: CacheRegistry<i32> = CacheRegistry::new();
let make = || -> Cache<i32> { Cache::builder().clock(dyn_clock.clone()).build() };
registry.register("alpha", make());
registry.register("beta", make());
assert_eq!(registry.len(), 2);
let alpha = registry.get("alpha").expect("alpha is registered");
let beta = registry.get("beta").expect("beta is registered");
assert!(registry.get("missing").is_none());
alpha.set("k", 1).await;
beta.set("k", 2).await;
assert_eq!(alpha.try_get("k", None).await.value(), Some(&1));
assert_eq!(beta.try_get("k", None).await.value(), Some(&2));
let builds = Arc::new(AtomicUsize::new(0));
let gamma1 = registry.get_or_create("gamma", || {
builds.fetch_add(1, Ordering::SeqCst);
make()
});
gamma1.set("k", 7).await;
let gamma2 = registry.get_or_create("gamma", || {
builds.fetch_add(1, Ordering::SeqCst);
make()
});
assert_eq!(
builds.load(Ordering::SeqCst),
1,
"gamma was built only once"
);
assert_eq!(
gamma2.try_get("k", None).await.value(),
Some(&7),
"the second get_or_create returned the same gamma cache"
);
}
struct ShortPrefixProvider;
impl DefaultEntryOptionsProvider for ShortPrefixProvider {
fn options_for(&self, key: &str) -> Option<EntryOptions> {
if key.starts_with("short:") {
Some(EntryOptions::new(Duration::from_secs(2)))
} else {
None
}
}
}
#[tokio::test(flavor = "multi_thread", worker_threads = 2)]
async fn default_options_provider_applies_per_key_duration() {
let clock = Arc::new(ManualClock::default());
let dyn_clock: Arc<dyn Clock> = clock.clone();
let cache: Cache<i32> = Cache::builder()
.clock(dyn_clock)
.default_options(EntryOptions::new(Duration::from_secs(100))) .default_options_provider(Arc::new(ShortPrefixProvider))
.build();
let short_calls = Arc::new(AtomicUsize::new(0));
let normal_calls = Arc::new(AtomicUsize::new(0));
let prime = |cache: &Cache<i32>, key: &'static str, counter: Arc<AtomicUsize>, val: i32| {
let cache = cache.clone();
async move {
cache
.get_or_set(key, move |ctx| async move {
counter.fetch_add(1, Ordering::SeqCst);
Ok(ctx.value(val))
})
.await
.unwrap()
}
};
assert_eq!(prime(&cache, "short:x", short_calls.clone(), 1).await, 1);
assert_eq!(prime(&cache, "normal", normal_calls.clone(), 1).await, 1);
assert_eq!(short_calls.load(Ordering::SeqCst), 1);
assert_eq!(normal_calls.load(Ordering::SeqCst), 1);
clock.advance(Duration::from_secs(3));
assert_eq!(prime(&cache, "short:x", short_calls.clone(), 2).await, 2);
assert_eq!(
short_calls.load(Ordering::SeqCst),
2,
"short-prefixed key expired per the provider's 2s duration"
);
assert_eq!(prime(&cache, "normal", normal_calls.clone(), 99).await, 1);
assert_eq!(
normal_calls.load(Ordering::SeqCst),
1,
"normal key used the cache default and is still fresh"
);
}