use super::{ListenerHandle, PubSubTransport};
use crate::backend::CacheBackend;
use crate::error::OxCacheResult;
use std::sync::Arc;
use std::time::Duration;
#[derive(Debug, Clone)]
pub struct KeyspaceNotificationConfig {
pub channels: Vec<String>,
}
impl Default for KeyspaceNotificationConfig {
fn default() -> Self {
Self {
channels: vec![
"__keyevent@0__:del".to_string(),
"__keyevent@0__:expired".to_string(),
],
}
}
}
impl KeyspaceNotificationConfig {
pub fn with_channels(mut self, channels: Vec<String>) -> Self {
self.channels = channels;
self
}
pub fn for_db(db: i32) -> Self {
Self {
channels: vec![
format!("__keyevent@{db}__:del"),
format!("__keyevent@{db}__:expired"),
],
}
}
}
pub struct KeyspaceNotificationListener;
impl KeyspaceNotificationListener {
pub async fn spawn(
transport: Arc<dyn PubSubTransport>,
config: KeyspaceNotificationConfig,
l1: Arc<dyn CacheBackend>,
) -> OxCacheResult<ListenerHandle> {
let mut merged: Vec<_> = Vec::new();
for channel in &config.channels {
let rx = transport.subscribe(channel).await?;
merged.push(rx);
}
let stop = Arc::new(std::sync::atomic::AtomicBool::new(false));
let stop_flag = stop.clone();
let join = tokio::spawn(async move {
loop {
if stop_flag.load(std::sync::atomic::Ordering::SeqCst) {
break;
}
let mut payload: Option<String> = None;
let mut closed: Vec<usize> = Vec::new();
for (idx, rx) in merged.iter_mut().enumerate() {
match tokio::time::timeout(Duration::from_millis(1), rx.recv()).await {
Ok(Some(msg)) => {
payload = Some(msg);
break;
}
Ok(None) => closed.push(idx),
Err(_) => continue,
}
}
for idx in closed.into_iter().rev() {
merged.remove(idx);
}
if merged.is_empty() {
return;
}
let Some(key) = payload else {
tokio::time::sleep(Duration::from_millis(10)).await;
continue;
};
let _ = l1.delete(&key).await;
}
});
Ok(ListenerHandle::new(join, stop))
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::backend::MockBackend;
use crate::features::invalidation::InMemoryPubSubTransport;
use std::sync::Arc;
fn l1() -> Arc<dyn CacheBackend> {
Arc::new(MockBackend::new("mock", 100, false))
}
async fn set_entry(l1: &Arc<dyn CacheBackend>, key: &str) {
l1.set(Arc::from(key), Arc::new(b"v".to_vec()), None)
.await
.unwrap();
}
#[tokio::test]
async fn external_del_event_invalidates_local_l1() {
let transport = Arc::new(InMemoryPubSubTransport::new());
let l1 = l1();
set_entry(&l1, "user:1").await;
let handle = KeyspaceNotificationListener::spawn(
transport.clone(),
KeyspaceNotificationConfig::default(),
l1.clone(),
)
.await
.unwrap();
tokio::time::sleep(Duration::from_millis(50)).await;
transport
.publish("__keyevent@0__:del", "user:1")
.await
.unwrap();
for _ in 0..50 {
if !l1.exists("user:1").await.unwrap() {
break;
}
tokio::time::sleep(Duration::from_millis(10)).await;
}
assert!(
!l1.exists("user:1").await.unwrap(),
"外部 DEL 库失效本地条目"
);
handle.stop();
handle.join().await;
}
#[tokio::test]
async fn expired_event_invalidates_local_l1() {
let transport = Arc::new(InMemoryPubSubTransport::new());
let l1 = l1();
set_entry(&l1, "session:42").await;
let handle = KeyspaceNotificationListener::spawn(
transport.clone(),
KeyspaceNotificationConfig::default(),
l1.clone(),
)
.await
.unwrap();
tokio::time::sleep(Duration::from_millis(50)).await;
transport
.publish("__keyevent@0__:expired", "session:42")
.await
.unwrap();
for _ in 0..50 {
if !l1.exists("session:42").await.unwrap() {
break;
}
tokio::time::sleep(Duration::from_millis(10)).await;
}
assert!(!l1.exists("session:42").await.unwrap());
handle.stop();
handle.join().await;
}
#[tokio::test]
async fn unsubscribed_channels_are_ignored() {
let transport = Arc::new(InMemoryPubSubTransport::new());
let l1 = l1();
set_entry(&l1, "keep:1").await;
let handle = KeyspaceNotificationListener::spawn(
transport.clone(),
KeyspaceNotificationConfig::for_db(0),
l1.clone(),
)
.await
.unwrap();
tokio::time::sleep(Duration::from_millis(50)).await;
transport
.publish("__keyevent@1__:del", "keep:1")
.await
.unwrap();
transport
.publish("__keyevent@0__:hdel", "keep:1")
.await
.unwrap();
tokio::time::sleep(Duration::from_millis(150)).await;
assert!(l1.exists("keep:1").await.unwrap(), "未订阅频道不得失效");
handle.stop();
handle.join().await;
}
#[test]
fn default_channels_cover_del_and_expired() {
let cfg = KeyspaceNotificationConfig::default();
assert_eq!(cfg.channels.len(), 2);
assert!(cfg.channels[0].ends_with(":del"));
assert!(cfg.channels[1].ends_with(":expired"));
let cfg = KeyspaceNotificationConfig::for_db(3);
assert!(
cfg.channels[0].contains("@3__:"),
"got {:?}",
cfg.channels[0]
);
}
}