use super::InvalidationBus;
use crate::backend::CacheBackend;
use crate::backend::interface::{BackendKind, CacheSetItem};
use crate::error::OxCacheResult;
use async_trait::async_trait;
use std::collections::HashMap;
use std::sync::Arc;
use std::time::Duration;
pub struct InvalidatingBackend {
inner: Arc<dyn CacheBackend>,
bus: Arc<InvalidationBus>,
expire_broadcast: bool,
}
impl InvalidatingBackend {
pub fn new(inner: Arc<dyn CacheBackend>, bus: Arc<InvalidationBus>) -> Self {
Self {
inner,
bus,
expire_broadcast: false,
}
}
pub fn with_expire_broadcast(mut self) -> Self {
self.expire_broadcast = true;
self
}
pub fn inner(&self) -> &Arc<dyn CacheBackend> {
&self.inner
}
}
#[async_trait]
impl crate::backend::CacheReader for InvalidatingBackend {
async fn get(&self, key: &str) -> OxCacheResult<Option<Vec<u8>>> {
self.inner.get(key).await
}
async fn exists(&self, key: &str) -> OxCacheResult<bool> {
self.inner.exists(key).await
}
async fn ttl(&self, key: &str) -> OxCacheResult<Option<Duration>> {
self.inner.ttl(key).await
}
async fn len(&self) -> OxCacheResult<u64> {
self.inner.len().await
}
async fn capacity(&self) -> OxCacheResult<u64> {
self.inner.capacity().await
}
async fn stats(&self) -> OxCacheResult<HashMap<String, String>> {
self.inner.stats().await
}
async fn keys(&self, pattern: &str) -> OxCacheResult<Vec<String>> {
self.inner.keys(pattern).await
}
}
#[async_trait]
impl crate::backend::CacheWriter for InvalidatingBackend {
async fn set(
&self,
key: Arc<str>,
value: Arc<Vec<u8>>,
ttl: Option<Duration>,
) -> OxCacheResult<()> {
self.inner.set(key.clone(), value, ttl).await?;
let _ = self.bus.invalidate_key(&key).await;
Ok(())
}
async fn delete(&self, key: &str) -> OxCacheResult<()> {
self.inner.delete(key).await?;
let _ = self.bus.invalidate_key(key).await;
Ok(())
}
async fn clear(&self) -> OxCacheResult<()> {
self.inner.clear().await?;
let _ = self.bus.invalidate_namespace("*").await;
Ok(())
}
async fn expire(&self, key: &str, ttl: Duration) -> OxCacheResult<bool> {
let result = self.inner.expire(key, ttl).await?;
if self.expire_broadcast {
let _ = self.bus.invalidate_key(key).await;
}
Ok(result)
}
async fn set_many(&self, items: &[CacheSetItem]) -> OxCacheResult<()> {
self.inner.set_many(items).await?;
for (key, _, _) in items {
let _ = self.bus.invalidate_key(key).await;
}
Ok(())
}
async fn delete_many(&self, keys: &[String]) -> OxCacheResult<()> {
self.inner.delete_many(keys).await?;
for key in keys {
let _ = self.bus.invalidate_key(key).await;
}
Ok(())
}
}
#[async_trait]
impl crate::backend::CacheConnector for InvalidatingBackend {
async fn health_check(&self) -> OxCacheResult<()> {
self.inner.health_check().await
}
async fn shutdown(&self) {
self.inner.shutdown().await;
}
fn backend_kind(&self) -> BackendKind {
self.inner.backend_kind()
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::backend::interface::{CacheConnector, CacheReader, CacheWriter};
use crate::backend::{CacheBackend, MockBackend};
use crate::features::invalidation::{
DEFAULT_CHANNEL, InMemoryPubSubTransport, InvalidationConfig,
};
async fn setup() -> (
Arc<InvalidatingBackend>,
Arc<dyn CacheBackend>,
Arc<dyn CacheBackend>,
) {
let transport = Arc::new(InMemoryPubSubTransport::new());
let bus_a = Arc::new(InvalidationBus::new(
transport.clone(),
InvalidationConfig::new("instance-a").with_channel(DEFAULT_CHANNEL),
));
let bus_b = Arc::new(InvalidationBus::new(
transport.clone(),
InvalidationConfig::new("instance-b").with_channel(DEFAULT_CHANNEL),
));
let inner_a: Arc<dyn CacheBackend> = Arc::new(MockBackend::new("mock-a", 100, false));
let inner_b: Arc<dyn CacheBackend> = Arc::new(MockBackend::new("mock-b", 100, false));
let decorated = Arc::new(InvalidatingBackend::new(inner_a.clone(), bus_a));
let _handle = bus_b.spawn_listener(inner_b.clone()).await.unwrap();
tokio::time::sleep(Duration::from_millis(50)).await;
(decorated, inner_a, inner_b)
}
#[tokio::test]
async fn set_publishes_invalidation_to_other_instances() {
let (decorated, _inner_a, inner_b) = setup().await;
inner_b
.set(Arc::from("user:1"), Arc::new(b"stale".to_vec()), None)
.await
.unwrap();
assert!(inner_b.exists("user:1").await.unwrap());
decorated
.set(Arc::from("user:1"), Arc::new(b"fresh".to_vec()), None)
.await
.unwrap();
for _ in 0..50 {
if !inner_b.exists("user:1").await.unwrap() {
break;
}
tokio::time::sleep(Duration::from_millis(10)).await;
}
assert!(
!inner_b.exists("user:1").await.unwrap(),
"装饰器 set 后 B 实例的旧条目应被失效"
);
}
#[tokio::test]
async fn delete_publishes_invalidation_to_other_instances() {
let (decorated, _inner_a, inner_b) = setup().await;
inner_b
.set(Arc::from("user:2"), Arc::new(b"stale".to_vec()), None)
.await
.unwrap();
decorated.delete("user:2").await.unwrap();
for _ in 0..50 {
if !inner_b.exists("user:2").await.unwrap() {
break;
}
tokio::time::sleep(Duration::from_millis(10)).await;
}
assert!(!inner_b.exists("user:2").await.unwrap());
}
#[tokio::test]
async fn read_path_is_passthrough() {
let (decorated, inner_a, _inner_b) = setup().await;
inner_a
.set(Arc::from("k"), Arc::new(b"v".to_vec()), None)
.await
.unwrap();
assert_eq!(decorated.get("k").await.unwrap(), Some(b"v".to_vec()));
assert!(decorated.exists("k").await.unwrap());
assert_eq!(
decorated.backend_kind(),
inner_a.backend_kind(),
"装饰器透传 backend_kind"
);
}
#[tokio::test]
async fn publish_failure_does_not_fail_write() {
let transport = Arc::new(InMemoryPubSubTransport::new());
let bus = Arc::new(InvalidationBus::new(
transport,
InvalidationConfig::new("solo"),
));
let inner: Arc<dyn CacheBackend> = Arc::new(MockBackend::new("mock", 100, false));
let decorated = InvalidatingBackend::new(inner.clone(), bus);
decorated
.set(Arc::from("k"), Arc::new(b"v".to_vec()), None)
.await
.unwrap();
assert_eq!(inner.get("k").await.unwrap(), Some(b"v".to_vec()));
}
async fn setup_two_instances() -> (
Arc<InvalidatingBackend>,
Arc<InvalidatingBackend>,
Arc<dyn CacheBackend>,
crate::features::invalidation::ListenerHandle,
) {
let transport = Arc::new(InMemoryPubSubTransport::new());
let bus_a = Arc::new(InvalidationBus::new(
transport.clone(),
InvalidationConfig::new("instance-a").with_channel(DEFAULT_CHANNEL),
));
let bus_b = Arc::new(InvalidationBus::new(
transport.clone(),
InvalidationConfig::new("instance-b").with_channel(DEFAULT_CHANNEL),
));
let inner_a: Arc<dyn CacheBackend> = Arc::new(MockBackend::new("mock-a", 100, false));
let remote: Arc<dyn CacheBackend> = Arc::new(MockBackend::new("remote", 100, false));
let handle = bus_b.spawn_listener(remote.clone()).await.unwrap();
tokio::time::sleep(Duration::from_millis(50)).await;
let plain = Arc::new(InvalidatingBackend::new(inner_a.clone(), bus_a.clone()));
let broadcasting =
Arc::new(InvalidatingBackend::new(inner_a, bus_a).with_expire_broadcast());
(plain, broadcasting, remote, handle)
}
#[tokio::test]
async fn expire_broadcast_disabled_by_default_keeps_passthrough() {
let (plain, _broadcasting, remote, _handle) = setup_two_instances().await;
plain
.inner()
.set(Arc::from("user:1"), Arc::new(b"v".to_vec()), None)
.await
.unwrap();
remote
.set(Arc::from("user:1"), Arc::new(b"stale".to_vec()), None)
.await
.unwrap();
assert!(
plain
.expire("user:1", Duration::from_secs(30))
.await
.unwrap()
);
for _ in 0..15 {
assert!(
remote.exists("user:1").await.unwrap(),
"默认关闭时 expire 不得广播失效(与既有行为一致)"
);
tokio::time::sleep(Duration::from_millis(10)).await;
}
}
#[tokio::test]
async fn expire_broadcast_enabled_notifies_other_instances() {
let (_plain, broadcasting, remote, _handle) = setup_two_instances().await;
broadcasting
.inner()
.set(Arc::from("user:2"), Arc::new(b"v".to_vec()), None)
.await
.unwrap();
remote
.set(Arc::from("user:2"), Arc::new(b"stale".to_vec()), None)
.await
.unwrap();
assert!(
broadcasting
.expire("user:2", Duration::from_secs(30))
.await
.unwrap()
);
for _ in 0..50 {
if !remote.exists("user:2").await.unwrap() {
return;
}
tokio::time::sleep(Duration::from_millis(10)).await;
}
panic!("开关开启后 expire 应广播失效远端条目");
}
}