use crate::error::MqError;
use crate::queue::{Message, MessageQueue, RedisConfig, RedisMode};
use async_trait::async_trait;
use futures::StreamExt;
use redis::aio::{MultiplexedConnection, PubSubStream};
use redis::{AsyncCommands, RedisError};
use std::collections::HashMap;
use std::pin::Pin;
use std::sync::Arc;
use tokio::sync::Mutex;
const STREAM_START_FROM_BEGINNING: &str = "0";
const LIST_BLOCK_TIMEOUT_SECS: f64 = 1.0;
const POLL_TIMEOUT_MS: u64 = 100;
pub struct RedisQueueProvider {
config: RedisConfig,
conn: MultiplexedConnection,
pubsub_streams: Arc<Mutex<HashMap<String, Pin<Box<PubSubStream>>>>>,
in_flight: Arc<Mutex<HashMap<String, (String, String)>>>,
}
impl RedisQueueProvider {
pub async fn new(url: impl Into<String>) -> Result<Self, MqError> {
Self::connect(RedisConfig {
url: Some(url.into()),
..RedisConfig::default()
})
.await
}
pub async fn connect(mut config: RedisConfig) -> Result<Self, MqError> {
let url = config
.url
.clone()
.unwrap_or_else(|| "redis://127.0.0.1:6379/0".to_string());
let client = redis::Client::open(url.as_str())
.map_err(|e| MqError::Connection(format!("Redis client open failed: {e}")))?;
let conn = client
.get_multiplexed_async_connection()
.await
.map_err(|e| MqError::Connection(format!("Redis connect failed: {e}")))?;
if config.consumer_group.is_none() {
config.consumer_group = Some("sz-orm-queue-group".to_string());
}
if config.consumer_name.is_none() {
config.consumer_name = Some("consumer-1".to_string());
}
Ok(Self {
config,
conn,
pubsub_streams: Arc::new(Mutex::new(HashMap::new())),
in_flight: Arc::new(Mutex::new(HashMap::new())),
})
}
pub fn mode(&self) -> RedisMode {
self.config.mode
}
fn group(&self) -> &str {
self.config
.consumer_group
.as_deref()
.unwrap_or("sz-orm-queue-group")
}
fn consumer(&self) -> &str {
self.config
.consumer_name
.as_deref()
.unwrap_or("consumer-1")
}
async fn publish_list(&self, topic: &str, message: &[u8]) -> Result<(), MqError> {
let mut conn = self.conn.clone();
let _: i64 = conn
.rpush(topic, message.to_vec())
.await
.map_err(|e| MqError::Publish(format!("Redis RPUSH failed: {e}")))?;
Ok(())
}
async fn consume_list(&self, topic: &str) -> Result<Option<Message>, MqError> {
let mut conn = self.conn.clone();
let result: Option<(String, Vec<u8>)> = conn
.blpop(topic, LIST_BLOCK_TIMEOUT_SECS)
.await
.map_err(|e| MqError::Connection(format!("Redis BLPOP failed: {e}")))?;
Ok(result.map(|(_, payload)| Message {
topic: topic.to_string(),
payload,
key: None,
timestamp: current_timestamp_millis(),
headers: HashMap::new(),
id: format!("redis-list-{}", current_timestamp_millis()),
retry_count: 0,
}))
}
async fn publish_pubsub(&self, topic: &str, message: &[u8]) -> Result<(), MqError> {
let mut conn = self.conn.clone();
let _: i64 = conn
.publish(topic, message.to_vec())
.await
.map_err(|e| MqError::Publish(format!("Redis PUBLISH failed: {e}")))?;
Ok(())
}
async fn consume_pubsub(&self, topic: &str) -> Result<Option<Message>, MqError> {
let need_subscribe = !self.pubsub_streams.lock().await.contains_key(topic);
if need_subscribe {
self.subscribe_pubsub(topic).await?;
}
let mut streams = self.pubsub_streams.lock().await;
let stream = match streams.get_mut(topic) {
Some(s) => s,
None => {
return Err(MqError::Subscribe(
"Redis pubsub stream not available".into(),
))
}
};
match tokio::time::timeout(
std::time::Duration::from_millis(POLL_TIMEOUT_MS),
stream.next(),
)
.await
{
Ok(Some(msg)) => Ok(Some(Message {
topic: msg.get_channel_name().to_string(),
payload: msg.get_payload_bytes().to_vec(),
key: None,
timestamp: current_timestamp_millis(),
headers: HashMap::new(),
id: format!("redis-pubsub-{}", current_timestamp_millis()),
retry_count: 0,
})),
_ => Ok(None),
}
}
async fn subscribe_pubsub(&self, topic: &str) -> Result<(), MqError> {
let url = self
.config
.url
.clone()
.unwrap_or_else(|| "redis://127.0.0.1:6379/0".to_string());
let client = redis::Client::open(url.as_str())
.map_err(|e| MqError::Connection(format!("Redis client open failed: {e}")))?;
let mut pubsub = client
.get_async_pubsub()
.await
.map_err(|e| MqError::Connection(format!("Redis pubsub connect failed: {e}")))?;
pubsub
.subscribe(topic)
.await
.map_err(|e| MqError::Subscribe(format!("Redis SUBSCRIBE failed: {e}")))?;
let stream = pubsub.into_on_message();
self.pubsub_streams
.lock()
.await
.insert(topic.to_string(), Box::pin(stream));
Ok(())
}
async fn publish_stream(&self, topic: &str, message: &[u8]) -> Result<(), MqError> {
let mut conn = self.conn.clone();
let field_value = ("payload", message.to_vec());
let _: String = conn
.xadd(topic, "*", &[field_value])
.await
.map_err(|e| MqError::Publish(format!("Redis XADD failed: {e}")))?;
Ok(())
}
async fn consume_stream(&self, topic: &str) -> Result<Option<Message>, MqError> {
self.ensure_group(topic).await?;
let mut conn = self.conn.clone();
let options = redis::streams::StreamReadOptions::default()
.group(self.group(), self.consumer())
.count(1)
.block(POLL_TIMEOUT_MS as usize);
let keys = [topic];
let ids = [">"];
let reply: redis::streams::StreamReadReply = conn
.xread_options(&keys[..], &ids[..], &options)
.await
.map_err(|e| MqError::Connection(format!("Redis XREADGROUP failed: {e}")))?;
for stream_key in reply.keys {
for entry in stream_key.ids {
let entry_id = entry.id.clone();
let payload: Vec<u8> = entry.get("payload").unwrap_or_default();
let message_id = format!("{topic}::{entry_id}");
self.in_flight
.lock()
.await
.insert(message_id.clone(), (topic.to_string(), self.group().to_string()));
return Ok(Some(Message {
topic: topic.to_string(),
payload,
key: Some(entry_id),
timestamp: current_timestamp_millis(),
headers: HashMap::new(),
id: message_id,
retry_count: 0,
}));
}
}
Ok(None)
}
async fn ensure_group(&self, topic: &str) -> Result<(), MqError> {
let mut conn = self.conn.clone();
let result: Result<(), RedisError> = conn
.xgroup_create(topic, self.group(), STREAM_START_FROM_BEGINNING)
.await;
if let Err(e) = result {
let msg = e.to_string();
if !msg.contains("BUSYGROUP") {
return Err(MqError::Connection(format!(
"Redis XGROUP CREATE failed: {msg}"
)));
}
}
Ok(())
}
}
#[async_trait]
impl MessageQueue for RedisQueueProvider {
async fn publish(&self, topic: &str, message: &[u8]) -> Result<(), MqError> {
match self.config.mode {
RedisMode::List => self.publish_list(topic, message).await,
RedisMode::PubSub => self.publish_pubsub(topic, message).await,
RedisMode::Stream => self.publish_stream(topic, message).await,
}
}
async fn consume(&self, topic: &str) -> Result<Option<Message>, MqError> {
match self.config.mode {
RedisMode::List => self.consume_list(topic).await,
RedisMode::PubSub => self.consume_pubsub(topic).await,
RedisMode::Stream => self.consume_stream(topic).await,
}
}
async fn ack(&self, message_id: &str) -> Result<(), MqError> {
let entry = self.in_flight.lock().await.remove(message_id);
match entry {
None => Ok(()), Some((stream, group)) => {
let entry_id = message_id
.rsplit_once("::")
.map(|(_, id)| id.to_string())
.unwrap_or_else(|| message_id.to_string());
let mut conn = self.conn.clone();
let _: i64 = conn
.xack(stream, group, &[entry_id])
.await
.map_err(|e| MqError::Publish(format!("Redis XACK failed: {e}")))?;
Ok(())
}
}
}
async fn subscribe(&self, topic: &str) -> Result<(), MqError> {
match self.config.mode {
RedisMode::PubSub => self.subscribe_pubsub(topic).await,
RedisMode::Stream => self.ensure_group(topic).await,
RedisMode::List => Ok(()),
}
}
}
fn current_timestamp_millis() -> i64 {
use std::time::{SystemTime, UNIX_EPOCH};
SystemTime::now()
.duration_since(UNIX_EPOCH)
.unwrap_or_default()
.as_millis() as i64
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_redis_config_types_compile() {
let cfg = RedisConfig {
mode: RedisMode::Stream,
..RedisConfig::default()
};
assert_eq!(cfg.mode, RedisMode::Stream);
assert_eq!(cfg.pool_size, 8);
}
#[test]
fn test_redis_config_default_is_list_mode() {
let cfg = RedisConfig::default();
assert_eq!(cfg.mode, RedisMode::List);
assert!(cfg.consumer_group.is_none());
}
#[tokio::test]
#[ignore = "需真实 Redis 服务器"]
async fn test_redis_list_publish_and_consume() {
let queue = RedisQueueProvider::connect(RedisConfig::default())
.await
.unwrap();
queue.publish("sz-orm-list", b"hello-list").await.unwrap();
let msg = queue
.consume("sz-orm-list")
.await
.unwrap()
.expect("应有消息");
assert_eq!(msg.payload, b"hello-list");
queue.ack(&msg.id).await.unwrap();
}
#[tokio::test]
#[ignore = "需真实 Redis 服务器"]
async fn test_redis_pubsub_publish_and_consume() {
let queue = RedisQueueProvider::connect(RedisConfig {
mode: RedisMode::PubSub,
..RedisConfig::default()
})
.await
.unwrap();
queue.subscribe("sz-orm-pubsub").await.unwrap();
tokio::time::sleep(std::time::Duration::from_millis(100)).await;
queue
.publish("sz-orm-pubsub", b"hello-pubsub")
.await
.unwrap();
let msg = queue
.consume("sz-orm-pubsub")
.await
.unwrap()
.expect("应有消息");
assert_eq!(msg.payload, b"hello-pubsub");
}
#[tokio::test]
#[ignore = "需真实 Redis 服务器"]
async fn test_redis_stream_publish_consume_ack() {
let queue = RedisQueueProvider::connect(RedisConfig {
mode: RedisMode::Stream,
..RedisConfig::default()
})
.await
.unwrap();
{
let mut conn = queue.conn.clone();
let _: i64 = conn.del("sz-orm-stream").await.unwrap_or(0);
}
queue
.publish("sz-orm-stream", b"hello-stream")
.await
.unwrap();
let msg = queue
.consume("sz-orm-stream")
.await
.unwrap()
.expect("应有消息");
assert_eq!(msg.payload, b"hello-stream");
queue.ack(&msg.id).await.unwrap();
}
}