use crate::canonical_message::tracing_support::LazyMessageIds;
use crate::models::RedisStreamsConfig;
use crate::traits::{
BoxFuture, ConsumerError, EndpointStatus, MessageConsumer, MessageDisposition,
MessagePublisher, PublisherError, ReceivedBatch, SentBatch,
};
use crate::CanonicalMessage;
use crate::APP_NAME;
use anyhow::anyhow;
use async_channel::{bounded, Receiver, Sender};
use async_trait::async_trait;
use redis::aio::ConnectionManager;
use redis::streams::{
StreamAutoClaimOptions, StreamAutoClaimReply, StreamReadOptions, StreamReadReply,
};
use redis::{AsyncCommands, IntoConnectionInfo};
use std::any::Any;
use std::collections::{HashMap, VecDeque};
use std::time::{Duration, Instant};
use tracing::trace;
const PAYLOAD_FIELD: &str = "payload";
const MESSAGE_ID_FIELD: &str = "mqb_message_id";
const DEFAULT_BLOCK_MS: u64 = 5000;
const DEFAULT_BUFFER: usize = 128;
const DEFAULT_REDELIVERY_MS: u64 = 60_000;
fn open_client(config: &RedisStreamsConfig) -> anyhow::Result<redis::Client> {
let url = url_with_credentials(config);
let info = url
.as_str()
.into_connection_info()
.map_err(|e| anyhow!("Invalid Redis URL: {}", e))?;
redis::Client::open(info).map_err(|e| anyhow!("Failed to open Redis client: {}", e))
}
fn encode_userinfo(s: &str) -> String {
let mut out = String::with_capacity(s.len());
for b in s.bytes() {
if b.is_ascii_alphanumeric() || matches!(b, b'-' | b'.' | b'_' | b'~') {
out.push(b as char);
} else {
out.push_str(&format!("%{:02X}", b));
}
}
out
}
fn url_with_credentials(config: &RedisStreamsConfig) -> String {
if config.username.is_none() && config.password.is_none() {
return config.url.clone();
}
let Some(scheme_end) = config.url.find("://") else {
return config.url.clone();
};
let (scheme, rest) = config.url.split_at(scheme_end + 3);
let authority_end = rest.find('/').unwrap_or(rest.len());
if rest[..authority_end].contains('@') {
return config.url.clone();
}
let user = config
.username
.as_deref()
.map(encode_userinfo)
.unwrap_or_default();
let userinfo = match &config.password {
Some(password) => format!("{}:{}@", user, encode_userinfo(password)),
None => format!("{}@", user),
};
format!("{}{}{}", scheme, userinfo, rest)
}
pub struct RedisStreamsPublisher {
conn: ConnectionManager,
stream: String,
maxlen: Option<usize>,
approx_trim: bool,
}
impl RedisStreamsPublisher {
pub async fn new(config: &RedisStreamsConfig) -> anyhow::Result<Self> {
let stream = config
.stream
.clone()
.ok_or_else(|| anyhow!("Stream key is required for Redis Streams publisher"))?;
let client = open_client(config)?;
let conn = ConnectionManager::new(client)
.await
.map_err(|e| anyhow!("Failed to connect to Redis: {}", e))?;
Ok(Self {
conn,
stream,
maxlen: config.maxlen,
approx_trim: config.approx_trim.unwrap_or(true),
})
}
}
#[async_trait]
impl MessagePublisher for RedisStreamsPublisher {
async fn send_batch(
&self,
mut messages: Vec<CanonicalMessage>,
) -> Result<SentBatch, PublisherError> {
trace!(stream = %self.stream, count = messages.len(), message_ids = ?LazyMessageIds(&messages), "Publishing batch of Redis Streams messages");
if messages.is_empty() {
return Ok(SentBatch::Ack);
}
let mut pipe = redis::pipe();
for message in &mut messages {
message.strip_source_metadata();
pipe.cmd("XADD").arg(&self.stream);
if let Some(maxlen) = self.maxlen {
pipe.arg("MAXLEN");
if self.approx_trim {
pipe.arg("~");
}
pipe.arg(maxlen);
}
pipe.arg("*")
.arg(PAYLOAD_FIELD)
.arg(message.payload.as_ref())
.arg(MESSAGE_ID_FIELD)
.arg(format!("{:032x}", message.message_id));
for (key, value) in &message.metadata {
if key == PAYLOAD_FIELD || key == MESSAGE_ID_FIELD {
continue;
}
pipe.arg(key).arg(value);
}
}
let mut conn = self.conn.clone();
pipe.query_async::<()>(&mut conn)
.await
.map_err(|e| PublisherError::Retryable(anyhow!("Redis XADD failed: {}", e)))?;
Ok(SentBatch::Ack)
}
async fn status(&self) -> EndpointStatus {
let mut conn = self.conn.clone();
let healthy = redis::cmd("PING")
.query_async::<String>(&mut conn)
.await
.is_ok();
EndpointStatus {
healthy,
target: self.stream.clone(),
error: if healthy {
None
} else {
Some("Redis PING failed".to_string())
},
..Default::default()
}
}
fn as_any(&self) -> &dyn Any {
self
}
}
struct StreamEntry {
id: String,
msg: CanonicalMessage,
}
pub struct RedisStreamsConsumer {
rx: Receiver<Result<StreamEntry, ConsumerError>>,
ack_conn: ConnectionManager,
stream: String,
group: Option<String>,
buffer: VecDeque<StreamEntry>,
exit_on_empty: bool,
}
impl RedisStreamsConsumer {
pub async fn new(config: &RedisStreamsConfig) -> anyhow::Result<Self> {
let stream = config
.stream
.clone()
.ok_or_else(|| anyhow!("Stream key is required for Redis Streams consumer"))?;
let client = open_client(config)?;
let mut read_conn = ConnectionManager::new(client.clone())
.await
.map_err(|e| anyhow!("Failed to connect to Redis: {}", e))?;
let ack_conn = ConnectionManager::new(client.clone())
.await
.map_err(|e| anyhow!("Failed to connect to Redis: {}", e))?;
let block_ms = config.block_ms.unwrap_or(DEFAULT_BLOCK_MS) as usize;
let count = config.internal_buffer_size.unwrap_or(DEFAULT_BUFFER).max(1);
let redelivery_ms = config
.redelivery_timeout_ms
.unwrap_or(DEFAULT_REDELIVERY_MS);
let group = if config.subscriber_mode {
None
} else {
let group = config
.group
.clone()
.unwrap_or_else(|| format!("{}-{}", APP_NAME, stream));
let start_id = if config.read_from_start { "0" } else { "$" };
let created: redis::RedisResult<()> = read_conn
.xgroup_create_mkstream(&stream, &group, start_id)
.await;
if let Err(e) = created {
if e.code() != Some("BUSYGROUP") {
return Err(anyhow!("Failed to create Redis consumer group: {}", e));
}
}
Some(group)
};
let consumer_name = config.consumer_name.clone().unwrap_or_else(|| {
format!("{}-{:032x}", APP_NAME, fast_uuid_v7::gen_id_with_sub_ms_4())
});
let readers = if group.is_some() {
config.reader_connections.unwrap_or(1).max(1)
} else {
1
};
let (tx, rx) =
bounded::<Result<StreamEntry, ConsumerError>>(count.saturating_mul(readers).max(count));
for i in 0..readers {
let read_conn = if i == 0 {
read_conn.clone()
} else {
ConnectionManager::new(client.clone())
.await
.map_err(|e| anyhow!("Failed to connect to Redis: {}", e))?
};
let consumer_name = if readers == 1 {
consumer_name.clone()
} else {
format!("{}-{}", consumer_name, i)
};
spawn_stream_reader(ReaderCtx {
read_conn,
stream: stream.clone(),
group: group.clone(),
consumer_name,
count,
block_ms,
redelivery_ms,
reclaim: i == 0,
tx: tx.clone(),
});
}
drop(tx);
Ok(Self {
rx,
ack_conn,
stream,
group,
buffer: VecDeque::new(),
exit_on_empty: false,
})
}
}
struct ReaderCtx {
read_conn: ConnectionManager,
stream: String,
group: Option<String>,
consumer_name: String,
count: usize,
block_ms: usize,
redelivery_ms: u64,
reclaim: bool,
tx: Sender<Result<StreamEntry, ConsumerError>>,
}
fn spawn_stream_reader(ctx: ReaderCtx) {
let ReaderCtx {
mut read_conn,
stream: task_stream,
group: task_group,
consumer_name,
count,
block_ms,
redelivery_ms,
reclaim,
tx,
} = ctx;
tokio::spawn(async move {
let mut last_id = String::from("$");
let reclaim_enabled = reclaim && task_group.is_some() && redelivery_ms > 0;
let reclaim_interval = Duration::from_millis(redelivery_ms.min(block_ms as u64).max(1000));
let mut last_reclaim: Option<Instant> = None;
let mut reclaim_cursor = String::from("0-0");
loop {
if reclaim_enabled && last_reclaim.is_none_or(|t| t.elapsed() >= reclaim_interval) {
last_reclaim = Some(Instant::now());
if let Some(group) = &task_group {
let claim_opts = StreamAutoClaimOptions::default().count(count);
let claimed: redis::RedisResult<StreamAutoClaimReply> = read_conn
.xautoclaim_options(
&task_stream,
group,
&consumer_name,
redelivery_ms as usize,
&reclaim_cursor,
claim_opts,
)
.await;
match claimed {
Ok(reply) => {
reclaim_cursor = reply.next_stream_id;
for entry in reply.claimed {
let stream_entry = parse_entry(entry.id, entry.map);
if tx.send(Ok(stream_entry)).await.is_err() {
return; }
}
}
Err(e) => {
trace!(stream = %task_stream, error = %e, "Redis XAUTOCLAIM failed")
}
}
}
}
let mut opts = StreamReadOptions::default().count(count).block(block_ms);
let read_ids: Vec<&str> = match &task_group {
Some(group) => {
opts = opts.group(group, &consumer_name);
vec![">"]
}
None => vec![last_id.as_str()],
};
let reply: redis::RedisResult<Option<StreamReadReply>> = read_conn
.xread_options(&[&task_stream], read_ids.as_slice(), &opts)
.await;
match reply {
Ok(Some(reply)) => {
for key in reply.keys {
for entry in key.ids {
if task_group.is_none() {
last_id = entry.id.clone();
}
let stream_entry = parse_entry(entry.id, entry.map);
if tx.send(Ok(stream_entry)).await.is_err() {
return; }
}
}
}
Ok(None) => {}
Err(e) if e.is_timeout() => {}
Err(e) => {
if tx
.send(Err(ConsumerError::Connection(anyhow!(
"Redis XREAD failed: {}",
e
))))
.await
.is_err()
{
return;
}
tokio::time::sleep(std::time::Duration::from_millis(500)).await;
}
}
}
});
}
fn parse_entry(id: String, map: HashMap<String, redis::Value>) -> StreamEntry {
let mut payload = Vec::new();
let mut message_id = None;
let mut metadata = HashMap::new();
for (key, value) in map {
if key == PAYLOAD_FIELD {
payload = redis::from_redis_value::<Vec<u8>>(value).unwrap_or_default();
} else if key == MESSAGE_ID_FIELD {
if let Ok(s) = redis::from_redis_value::<String>(value) {
if let Ok(n) = u128::from_str_radix(&s, 16) {
message_id = Some(n);
}
}
} else {
if crate::canonical_message::is_source_metadata_key(&key) {
continue;
}
if let Ok(s) = redis::from_redis_value::<String>(value) {
metadata.insert(key, s);
}
}
}
let mut msg = CanonicalMessage::new(payload, message_id);
msg.metadata = metadata;
if crate::canonical_message::source_metadata_enabled() {
msg.metadata
.insert("mqb.src.redis_stream_id".to_string(), id.clone());
}
StreamEntry { id, msg }
}
#[async_trait]
impl MessageConsumer for RedisStreamsConsumer {
fn commit_requires_order(&self) -> bool {
false
}
fn set_exit_on_empty(&mut self, exit_on_empty: bool) {
self.exit_on_empty = exit_on_empty;
}
async fn receive_batch(&mut self, max_messages: usize) -> Result<ReceivedBatch, ConsumerError> {
if max_messages == 0 {
return Ok(ReceivedBatch {
messages: Vec::new(),
commit: Box::new(|_| Box::pin(async { Ok(()) })),
});
}
if self.buffer.is_empty() {
let Some(res) = crate::traits::drain_gated(self.exit_on_empty, self.rx.recv()).await
else {
return Ok(ReceivedBatch::empty());
};
let entry = res.map_err(|_| ConsumerError::EndOfStream)??;
self.buffer.push_back(entry);
}
while self.buffer.len() < max_messages {
match self.rx.try_recv() {
Ok(Ok(entry)) => self.buffer.push_back(entry),
Ok(Err(_)) => break,
Err(_) => break,
}
}
let mut messages = Vec::with_capacity(max_messages);
let mut ids = Vec::with_capacity(max_messages);
while messages.len() < max_messages {
if let Some(entry) = self.buffer.pop_front() {
ids.push(entry.id);
messages.push(entry.msg);
} else {
break;
}
}
trace!(stream = %self.stream, count = messages.len(), message_ids = ?LazyMessageIds(&messages), "Received batch of Redis Streams messages");
let group = self.group.clone();
let stream = self.stream.clone();
let mut ack_conn = self.ack_conn.clone();
let commit = Box::new(move |dispositions: Vec<MessageDisposition>| {
Box::pin(async move {
let Some(group) = group else {
return Ok(());
};
let ack_ids: Vec<String> = ids
.into_iter()
.zip(dispositions)
.filter_map(|(id, disposition)| match disposition {
MessageDisposition::Ack | MessageDisposition::Reply(_) => Some(id),
MessageDisposition::Nack => None,
})
.collect();
if !ack_ids.is_empty() {
ack_conn
.xack::<_, _, _, ()>(&stream, &group, &ack_ids)
.await
.map_err(|e| anyhow!("Redis XACK failed: {}", e))?;
}
Ok(())
}) as BoxFuture<'static, anyhow::Result<()>>
});
Ok(ReceivedBatch { messages, commit })
}
async fn status(&self) -> EndpointStatus {
let mut conn = self.ack_conn.clone();
let healthy = redis::cmd("PING")
.query_async::<String>(&mut conn)
.await
.is_ok();
EndpointStatus {
healthy,
target: self.stream.clone(),
pending: Some(self.buffer.len() + self.rx.len()),
error: if healthy {
None
} else {
Some("Redis PING failed".to_string())
},
..Default::default()
}
}
fn as_any(&self) -> &dyn Any {
self
}
}
#[cfg(test)]
mod tests {
use super::*;
fn bulk(s: &str) -> redis::Value {
redis::Value::BulkString(s.as_bytes().to_vec())
}
#[test]
fn parse_entry_reads_payload_and_message_id() {
let mut map = HashMap::new();
map.insert(PAYLOAD_FIELD.to_string(), bulk("hello"));
map.insert(
MESSAGE_ID_FIELD.to_string(),
bulk("0000000000000000000000000000002a"),
);
map.insert("user_key".to_string(), bulk("kept"));
let entry = parse_entry("1-0".to_string(), map);
assert_eq!(entry.msg.payload.as_ref(), b"hello");
assert_eq!(entry.msg.message_id, 0x2a);
assert_eq!(
entry.msg.metadata.get("user_key").map(String::as_str),
Some("kept")
);
assert!(!entry.msg.metadata.contains_key(PAYLOAD_FIELD));
assert!(!entry.msg.metadata.contains_key(MESSAGE_ID_FIELD));
}
#[test]
fn parse_entry_strips_spoofed_source_metadata() {
let mut map = HashMap::new();
map.insert(PAYLOAD_FIELD.to_string(), bulk("body"));
map.insert("mqb.src.kafka_offset".to_string(), bulk("999"));
map.insert("user_key".to_string(), bulk("kept"));
let entry = parse_entry("2-0".to_string(), map);
assert!(!entry.msg.metadata.contains_key("mqb.src.kafka_offset"));
assert_eq!(
entry.msg.metadata.get("user_key").map(String::as_str),
Some("kept")
);
}
#[test]
fn parse_entry_without_message_id_gets_generated_id() {
let mut map = HashMap::new();
map.insert(PAYLOAD_FIELD.to_string(), bulk("body"));
let entry = parse_entry("3-0".to_string(), map);
assert_ne!(entry.msg.message_id, 0);
}
}