use crate::error::{OxCacheError, OxCacheResult};
use async_trait::async_trait;
use std::collections::HashMap;
use std::sync::Mutex;
use tokio::sync::mpsc::UnboundedReceiver;
pub struct SubscriptionReceiver {
rx: UnboundedReceiver<String>,
}
impl SubscriptionReceiver {
pub(crate) fn new(rx: UnboundedReceiver<String>) -> Self {
Self { rx }
}
pub async fn recv(&mut self) -> Option<String> {
self.rx.recv().await
}
pub fn try_recv(&mut self) -> Option<String> {
self.rx.try_recv().ok()
}
}
#[async_trait]
pub trait PubSubTransport: Send + Sync + 'static {
async fn publish(&self, channel: &str, payload: &str) -> OxCacheResult<()>;
async fn subscribe(&self, channel: &str) -> OxCacheResult<SubscriptionReceiver>;
}
#[derive(Default)]
pub struct InMemoryPubSubTransport {
channels: Mutex<HashMap<String, Vec<tokio::sync::mpsc::UnboundedSender<String>>>>,
}
impl InMemoryPubSubTransport {
pub fn new() -> Self {
Self::default()
}
pub fn subscriber_count(&self, channel: &str) -> usize {
self.channels
.lock()
.map(|map| map.get(channel).map(Vec::len).unwrap_or(0))
.unwrap_or(0)
}
}
#[async_trait]
impl PubSubTransport for InMemoryPubSubTransport {
async fn publish(&self, channel: &str, payload: &str) -> OxCacheResult<()> {
if let Ok(mut map) = self.channels.lock()
&& let Some(list) = map.get_mut(channel)
{
let payload = payload.to_string();
list.retain(|tx| tx.send(payload.clone()).is_ok());
}
Ok(())
}
async fn subscribe(&self, channel: &str) -> OxCacheResult<SubscriptionReceiver> {
let (tx, rx) = tokio::sync::mpsc::unbounded_channel();
match self.channels.lock() {
Ok(mut map) => {
map.entry(channel.to_string()).or_default().push(tx);
}
Err(_) => {
return Err(OxCacheError::Operation(
"in-memory pubsub transport lock poisoned".to_string(),
));
}
}
Ok(SubscriptionReceiver::new(rx))
}
}
pub struct RedisPubSubTransport {
client: redis::Client,
publish_conn: tokio::sync::Mutex<redis::aio::ConnectionManager>,
connection_string: String,
}
impl RedisPubSubTransport {
pub async fn new(connection_string: &str) -> OxCacheResult<Self> {
let client = redis::Client::open(connection_string)
.map_err(|e| OxCacheError::Connection(format!("invalid redis url: {e}")))?;
let publish_conn = tokio::time::timeout(
std::time::Duration::from_secs(5),
client.get_connection_manager(),
)
.await
.map_err(|_| OxCacheError::Connection("redis pubsub connect timeout".to_string()))?
.map_err(|e| OxCacheError::Connection(format!("redis pubsub connect failed: {e}")))?;
Ok(Self {
client,
publish_conn: tokio::sync::Mutex::new(publish_conn),
connection_string: connection_string.to_string(),
})
}
pub fn connection_string(&self) -> &str {
&self.connection_string
}
async fn run_subscription(
client: redis::Client,
channel: String,
tx: tokio::sync::mpsc::UnboundedSender<String>,
) {
use futures::stream::StreamExt;
let mut backoff_ms: u64 = 200;
loop {
let subscribed = async {
let mut pubsub = client.get_async_pubsub().await?;
pubsub.subscribe(channel.as_str()).await?;
Ok::<_, redis::RedisError>(pubsub)
};
let mut pubsub = match subscribed.await {
Ok(ps) => {
backoff_ms = 200; ps
}
Err(_) => {
tokio::time::sleep(std::time::Duration::from_millis(backoff_ms)).await;
backoff_ms = (backoff_ms.saturating_mul(2)).min(10_000);
continue;
}
};
let mut stream = pubsub.on_message();
while let Some(msg) = stream.next().await {
let payload: Option<String> = msg.get_payload().ok();
if let Some(payload) = payload
&& tx.send(payload).is_err()
{
return;
}
}
tokio::time::sleep(std::time::Duration::from_millis(backoff_ms)).await;
backoff_ms = (backoff_ms.saturating_mul(2)).min(10_000);
}
}
}
#[async_trait]
impl PubSubTransport for RedisPubSubTransport {
async fn publish(&self, channel: &str, payload: &str) -> OxCacheResult<()> {
let mut conn = self.publish_conn.lock().await;
let _: i64 = redis::cmd("PUBLISH")
.arg(channel)
.arg(payload)
.query_async(&mut *conn)
.await
.map_err(|e| OxCacheError::Connection(format!("redis PUBLISH failed: {e}")))?;
Ok(())
}
async fn subscribe(&self, channel: &str) -> OxCacheResult<SubscriptionReceiver> {
let (tx, rx) = tokio::sync::mpsc::unbounded_channel();
tokio::spawn(Self::run_subscription(
self.client.clone(),
channel.to_string(),
tx,
));
Ok(SubscriptionReceiver::new(rx))
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::Arc;
#[tokio::test]
async fn in_memory_subscribe_then_publish_delivers() {
let transport = Arc::new(InMemoryPubSubTransport::new());
let mut rx = transport.subscribe("ch").await.unwrap();
transport.publish("ch", "hello").await.unwrap();
assert_eq!(rx.recv().await, Some("hello".to_string()));
}
#[tokio::test]
async fn in_memory_publish_before_subscribe_is_dropped() {
let transport = Arc::new(InMemoryPubSubTransport::new());
transport.publish("ch", "early").await.unwrap();
let mut rx = transport.subscribe("ch").await.unwrap();
transport.publish("ch", "late").await.unwrap();
assert_eq!(rx.recv().await, Some("late".to_string()));
assert!(rx.try_recv().is_none());
}
#[tokio::test]
async fn in_memory_broadcast_to_all_subscribers() {
let transport = Arc::new(InMemoryPubSubTransport::new());
let mut rx1 = transport.subscribe("ch").await.unwrap();
let mut rx2 = transport.subscribe("ch").await.unwrap();
assert_eq!(transport.subscriber_count("ch"), 2);
transport.publish("ch", "fanout").await.unwrap();
assert_eq!(rx1.recv().await, Some("fanout".to_string()));
assert_eq!(rx2.recv().await, Some("fanout".to_string()));
}
#[tokio::test]
async fn in_memory_disconnected_senders_are_pruned() {
let transport = Arc::new(InMemoryPubSubTransport::new());
let rx = transport.subscribe("ch").await.unwrap();
assert_eq!(transport.subscriber_count("ch"), 1);
drop(rx);
transport.publish("ch", "after-drop").await.unwrap();
assert_eq!(transport.subscriber_count("ch"), 0, "死 sender 应被剔除");
}
#[tokio::test]
async fn in_memory_channels_are_isolated() {
let transport = Arc::new(InMemoryPubSubTransport::new());
let mut rx_a = transport.subscribe("a").await.unwrap();
let mut rx_b = transport.subscribe("b").await.unwrap();
transport.publish("a", "msg-a").await.unwrap();
assert_eq!(rx_a.recv().await, Some("msg-a".to_string()));
assert!(rx_b.try_recv().is_none());
}
#[test]
fn in_memory_publish_without_subscribers_is_noop() {
let transport = InMemoryPubSubTransport::new();
futures::executor::block_on(async {
transport.publish("ghost", "x").await.unwrap();
});
}
}