use super::{DEFAULT_CHANNEL, InvalidationKind, InvalidationMessage, PubSubTransport};
use crate::backend::CacheBackend;
use crate::error::OxCacheResult;
use std::sync::Arc;
use std::sync::atomic::{AtomicBool, Ordering};
use tokio::task::JoinHandle;
#[derive(Debug, Clone)]
pub struct InvalidationConfig {
pub channel: String,
pub instance_id: String,
}
impl InvalidationConfig {
pub fn new(instance_id: impl Into<String>) -> Self {
Self {
channel: DEFAULT_CHANNEL.to_string(),
instance_id: instance_id.into(),
}
}
pub fn with_channel(mut self, channel: impl Into<String>) -> Self {
self.channel = channel.into();
self
}
}
pub struct ListenerHandle {
join: Option<JoinHandle<()>>,
stop: Arc<AtomicBool>,
}
impl ListenerHandle {
pub(crate) fn new(join: JoinHandle<()>, stop: Arc<AtomicBool>) -> Self {
Self {
join: Some(join),
stop,
}
}
pub fn stop(&self) {
self.stop.store(true, Ordering::SeqCst);
}
pub async fn join(mut self) {
if let Some(join) = self.join.take() {
let _ = join.await;
}
}
}
impl Drop for ListenerHandle {
fn drop(&mut self) {
self.stop.store(true, Ordering::SeqCst);
}
}
pub struct InvalidationBus {
transport: Arc<dyn PubSubTransport>,
config: InvalidationConfig,
}
impl std::fmt::Debug for InvalidationBus {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("InvalidationBus")
.field("channel", &self.config.channel)
.field("instance_id", &self.config.instance_id)
.finish()
}
}
impl InvalidationBus {
pub fn new(transport: Arc<dyn PubSubTransport>, config: InvalidationConfig) -> Self {
Self { transport, config }
}
pub fn instance_id(&self) -> &str {
&self.config.instance_id
}
pub async fn invalidate_key(&self, key: &str) -> OxCacheResult<()> {
let msg = InvalidationMessage::key(self.config.instance_id.clone(), key);
self.transport
.publish(&self.config.channel, &msg.encode()?)
.await
}
pub async fn invalidate_namespace(&self, namespace: &str) -> OxCacheResult<()> {
let msg = InvalidationMessage::namespace(self.config.instance_id.clone(), namespace);
self.transport
.publish(&self.config.channel, &msg.encode()?)
.await
}
pub async fn spawn_listener(&self, l1: Arc<dyn CacheBackend>) -> OxCacheResult<ListenerHandle> {
let mut rx = self.transport.subscribe(&self.config.channel).await?;
let instance_id = self.config.instance_id.clone();
let stop = Arc::new(AtomicBool::new(false));
let stop_flag = stop.clone();
let join = tokio::spawn(async move {
loop {
if stop_flag.load(Ordering::SeqCst) {
break;
}
let payload =
match tokio::time::timeout(std::time::Duration::from_millis(100), rx.recv())
.await
{
Ok(Some(payload)) => payload,
Ok(None) => break, Err(_) => continue, };
let msg = match InvalidationMessage::decode(&payload) {
Ok(m) => m,
Err(_) => continue, };
if msg.is_from(&instance_id) {
continue;
}
apply_to_l1(l1.as_ref(), &msg).await;
}
});
Ok(ListenerHandle::new(join, stop))
}
}
async fn apply_to_l1(l1: &dyn CacheBackend, msg: &InvalidationMessage) {
match msg.kind {
InvalidationKind::Key => {
let _ = l1.delete(&msg.target).await;
}
InvalidationKind::Namespace => {
let pattern = if msg.target == "*" {
"*".to_string()
} else {
format!("{}*", msg.target)
};
if let Ok(keys) = l1.keys(&pattern).await {
for key in keys {
let _ = l1.delete(&key).await;
}
}
}
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::backend::MockBackend;
use crate::features::invalidation::InMemoryPubSubTransport;
fn l1() -> Arc<dyn CacheBackend> {
Arc::new(MockBackend::new("mock", 100, false))
}
async fn set_l1(backend: &Arc<dyn CacheBackend>, key: &str, value: &[u8]) {
backend
.set(Arc::from(key), Arc::new(value.to_vec()), None)
.await
.unwrap();
}
#[tokio::test]
async fn cross_instance_invalidation_propagates() {
let transport = Arc::new(InMemoryPubSubTransport::new());
let bus_a = InvalidationBus::new(
transport.clone(),
InvalidationConfig::new("instance-a").with_channel("test-ch"),
);
let bus_b = InvalidationBus::new(
transport.clone(),
InvalidationConfig::new("instance-b").with_channel("test-ch"),
);
let l1_b = l1();
set_l1(&l1_b, "user:1", b"alice").await;
assert!(l1_b.exists("user:1").await.unwrap());
let handle_b = bus_b.spawn_listener(l1_b.clone()).await.unwrap();
tokio::time::sleep(std::time::Duration::from_millis(50)).await;
bus_a.invalidate_key("user:1").await.unwrap();
for _ in 0..50 {
if !l1_b.exists("user:1").await.unwrap() {
break;
}
tokio::time::sleep(std::time::Duration::from_millis(10)).await;
}
assert!(
!l1_b.exists("user:1").await.unwrap(),
"B 的本地 L1 条目应被 A 的广播失效"
);
handle_b.stop();
handle_b.join().await;
}
#[tokio::test]
async fn self_invalidation_is_exempt() {
let transport = Arc::new(InMemoryPubSubTransport::new());
let bus_a = InvalidationBus::new(
transport.clone(),
InvalidationConfig::new("instance-a").with_channel("test-ch"),
);
let l1_a = l1();
set_l1(&l1_a, "user:1", b"alice").await;
let handle_a = bus_a.spawn_listener(l1_a.clone()).await.unwrap();
tokio::time::sleep(std::time::Duration::from_millis(50)).await;
bus_a.invalidate_key("user:1").await.unwrap();
tokio::time::sleep(std::time::Duration::from_millis(200)).await;
assert!(
l1_a.exists("user:1").await.unwrap(),
"自身广播的失效消息应被豁免,A 的本地条目不应被删除"
);
handle_a.stop();
handle_a.join().await;
}
#[tokio::test]
async fn namespace_invalidation_removes_all_matching_keys() {
let transport = Arc::new(InMemoryPubSubTransport::new());
let bus_a = InvalidationBus::new(
transport.clone(),
InvalidationConfig::new("instance-a").with_channel("test-ch"),
);
let bus_b = InvalidationBus::new(
transport.clone(),
InvalidationConfig::new("instance-b").with_channel("test-ch"),
);
let l1_b = l1();
set_l1(&l1_b, "users:1", b"a").await;
set_l1(&l1_b, "users:2", b"b").await;
set_l1(&l1_b, "orders:1", b"c").await;
let handle_b = bus_b.spawn_listener(l1_b.clone()).await.unwrap();
tokio::time::sleep(std::time::Duration::from_millis(50)).await;
bus_a.invalidate_namespace("users:").await.unwrap();
for _ in 0..50 {
if !l1_b.exists("users:1").await.unwrap() && !l1_b.exists("users:2").await.unwrap() {
break;
}
tokio::time::sleep(std::time::Duration::from_millis(10)).await;
}
assert!(!l1_b.exists("users:1").await.unwrap());
assert!(!l1_b.exists("users:2").await.unwrap());
assert!(
l1_b.exists("orders:1").await.unwrap(),
"非匹配前缀的条目不应被失效"
);
handle_b.stop();
handle_b.join().await;
}
#[tokio::test]
async fn malformed_payload_is_ignored() {
let transport = Arc::new(InMemoryPubSubTransport::new());
let bus = InvalidationBus::new(
transport.clone(),
InvalidationConfig::new("instance-a").with_channel("test-ch"),
);
let l1 = l1();
set_l1(&l1, "user:1", b"keep").await;
let handle = bus.spawn_listener(l1.clone()).await.unwrap();
tokio::time::sleep(std::time::Duration::from_millis(50)).await;
transport
.publish("test-ch", "garbage-not-json")
.await
.unwrap();
tokio::time::sleep(std::time::Duration::from_millis(200)).await;
assert!(l1.exists("user:1").await.unwrap());
handle.stop();
handle.join().await;
}
#[test]
fn config_defaults() {
let cfg = InvalidationConfig::new("i");
assert_eq!(cfg.channel, super::super::DEFAULT_CHANNEL);
assert_eq!(cfg.instance_id, "i");
let cfg = cfg.with_channel("custom");
assert_eq!(cfg.channel, "custom");
}
}