use std::collections::HashMap;
use std::sync::atomic::{AtomicU64, Ordering};
use std::sync::Arc;
use std::time::Duration;
use crate::time::now_unix_ms;
use async_trait::async_trait;
use futures::StreamExt;
use serde::{Deserialize, Serialize};
use crate::kv::KvStore;
use crate::{PutMeta, Storage};
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct ClaimedMessage {
pub id: String,
pub topic: String,
pub payload: Vec<u8>,
pub attempts: u32,
}
#[derive(Debug, Clone, thiserror::Error)]
pub enum MessagingError {
#[error("messaging backend error: {0}")]
Backend(String),
#[error("messaging decode error: {0}")]
Decode(String),
}
impl MessagingError {
fn backend<E: std::fmt::Display>(err: E) -> Self {
Self::Backend(err.to_string())
}
}
#[async_trait]
pub trait Messaging: Send + Sync {
async fn publish(&self, topic: &str, payload: &[u8]) -> Result<(), MessagingError>;
async fn claim(
&self,
topic: &str,
lease: Duration,
max_batch: usize,
max_attempts: u32,
) -> Result<Vec<ClaimedMessage>, MessagingError>;
async fn ack(&self, msg: &ClaimedMessage) -> Result<(), MessagingError>;
async fn nack(&self, msg: &ClaimedMessage) -> Result<(), MessagingError>;
async fn backlog(&self, _topic: &str) -> Result<usize, MessagingError> {
Ok(0)
}
async fn dead_letter_count(&self, _topic: &str) -> Result<usize, MessagingError> {
Ok(0)
}
async fn purge_dead_letters(&self, _topic: &str) -> Result<usize, MessagingError> {
Ok(0)
}
async fn redrive_dead_letters(&self, _topic: &str) -> Result<usize, MessagingError> {
Ok(0)
}
fn subscribe(
&self,
_topic: &str,
_after: Option<&str>,
) -> futures::stream::BoxStream<'static, StreamEvent> {
futures::stream::empty().boxed()
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub struct StreamEvent {
pub id: String,
pub payload: Vec<u8>,
}
#[derive(Debug, Clone, Serialize, Deserialize, PartialEq, Eq)]
pub struct Record {
#[serde(default = "crate::schema_version")]
pub version: u32,
pub attempts: u32,
pub lease_until_ms: u64,
}
impl Record {
pub fn fresh() -> Self {
Self {
version: crate::SCHEMA_VERSION,
attempts: 0,
lease_until_ms: 0,
}
}
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum ClaimAction {
Lease {
id: String,
record: Record,
},
DeadLetter {
id: String,
record: Record,
},
}
pub fn plan_claim(
mut records: Vec<(String, Record)>,
now_ms: u64,
lease_ms: u64,
max_batch: usize,
max_attempts: u32,
) -> Vec<ClaimAction> {
records.sort_by(|a, b| a.0.cmp(&b.0));
let mut actions = Vec::new();
let mut leased = 0;
for (id, mut record) in records {
if leased >= max_batch {
break;
}
if record.lease_until_ms > now_ms {
continue; }
if record.attempts >= max_attempts {
actions.push(ClaimAction::DeadLetter { id, record });
continue;
}
record.attempts += 1;
record.lease_until_ms = now_ms + lease_ms;
actions.push(ClaimAction::Lease { id, record });
leased += 1;
}
actions
}
pub fn meta_key(topic: &str, id: &str) -> String {
format!("mq/{topic}/{id}")
}
pub fn meta_prefix(topic: &str) -> String {
format!("mq/{topic}/")
}
pub fn payload_key(topic: &str, id: &str) -> String {
format!("mqp/{topic}/{id}")
}
pub fn dead_key(topic: &str, id: &str) -> String {
format!("mqdead/{topic}/{id}")
}
pub fn dead_prefix(topic: &str) -> String {
format!("mqdead/{topic}/")
}
pub fn is_direct_child(key: &str, prefix: &str) -> bool {
key.len() > prefix.len() && !key[prefix.len()..].contains('/')
}
#[derive(Default)]
struct TopicHub {
recent: std::collections::VecDeque<StreamEvent>,
subscribers: Vec<futures::channel::mpsc::Sender<StreamEvent>>,
}
const STREAM_RING: usize = 64;
#[derive(Default)]
pub struct StreamHubs {
live: std::sync::Mutex<HashMap<String, TopicHub>>,
}
impl StreamHubs {
pub fn new() -> Self {
Self::default()
}
pub fn broadcast(&self, topic: &str, id: &str, payload: &[u8]) {
let event = StreamEvent {
id: id.to_string(),
payload: payload.to_vec(),
};
let mut live = self.live.lock().unwrap();
let Some(hub) = live.get_mut(topic) else {
return; };
hub.subscribers
.retain_mut(|tx| match tx.try_send(event.clone()) {
Ok(()) => true,
Err(err) => !err.is_disconnected(), });
hub.recent.push_back(event);
while hub.recent.len() > STREAM_RING {
hub.recent.pop_front();
}
if hub.subscribers.is_empty() {
live.remove(topic);
}
}
pub fn subscribe(
&self,
topic: &str,
after: Option<&str>,
) -> futures::stream::BoxStream<'static, StreamEvent> {
let (tx, rx) = futures::channel::mpsc::channel(64);
let replay: Vec<StreamEvent> = {
let mut live = self.live.lock().unwrap();
let hub = live.entry(topic.to_string()).or_default();
let replay = match after {
Some(after) => hub
.recent
.iter()
.filter(|event| event.id.as_str() > after)
.cloned()
.collect(),
None => Vec::new(),
};
hub.subscribers.push(tx);
replay
};
if replay.is_empty() {
rx.boxed()
} else {
futures::stream::iter(replay).chain(rx).boxed()
}
}
}
pub struct LogMessaging {
storage: Arc<dyn Storage>,
kv: Arc<dyn KvStore>,
claim_lock: futures::lock::Mutex<()>,
seq: AtomicU64,
hubs: StreamHubs,
}
impl LogMessaging {
pub fn new(storage: Arc<dyn Storage>, kv: Arc<dyn KvStore>) -> Self {
Self {
storage,
kv,
claim_lock: futures::lock::Mutex::new(()),
seq: AtomicU64::new(0),
hubs: StreamHubs::new(),
}
}
async fn read_payload(&self, topic: &str, id: &str) -> Result<Vec<u8>, MessagingError> {
let object = self
.storage
.get(&payload_key(topic, id))
.await
.map_err(MessagingError::backend)?;
let mut body = object.body;
let mut buf = Vec::new();
while let Some(chunk) = body.next().await {
buf.extend_from_slice(&chunk.map_err(MessagingError::backend)?);
}
Ok(buf)
}
async fn count_direct(&self, prefix: &str) -> Result<usize, MessagingError> {
let keys = self
.kv
.list_prefix(prefix)
.await
.map_err(MessagingError::backend)?;
Ok(keys.iter().filter(|k| is_direct_child(k, prefix)).count())
}
}
#[async_trait]
impl Messaging for LogMessaging {
async fn publish(&self, topic: &str, payload: &[u8]) -> Result<(), MessagingError> {
let id = format!(
"{:013}-{:016x}",
now_unix_ms(),
self.seq.fetch_add(1, Ordering::Relaxed)
);
let bytes = bytes::Bytes::copy_from_slice(payload);
let body = futures::stream::once(async move { Ok(bytes) }).boxed();
self.storage
.put(&payload_key(topic, &id), body, PutMeta::default())
.await
.map_err(MessagingError::backend)?;
let json = serde_json::to_vec(&Record::fresh()).map_err(MessagingError::backend)?;
self.kv
.put(&meta_key(topic, &id), json)
.await
.map_err(MessagingError::backend)?;
self.hubs.broadcast(topic, &id, payload);
Ok(())
}
async fn claim(
&self,
topic: &str,
lease: Duration,
max_batch: usize,
max_attempts: u32,
) -> Result<Vec<ClaimedMessage>, MessagingError> {
let _guard = self.claim_lock.lock().await;
let now = now_unix_ms();
let prefix = meta_prefix(topic);
let keys = self
.kv
.list_prefix(&prefix)
.await
.map_err(MessagingError::backend)?;
let mut records = Vec::new();
for key in keys {
if !is_direct_child(&key, &prefix) {
continue; }
let Some(raw) = self.kv.get(&key).await.map_err(MessagingError::backend)? else {
continue; };
let record: Record =
serde_json::from_slice(&raw).map_err(|e| MessagingError::Decode(e.to_string()))?;
records.push((key[prefix.len()..].to_string(), record));
}
let actions = plan_claim(
records,
now,
lease.as_millis() as u64,
max_batch,
max_attempts,
);
let mut claimed = Vec::new();
for action in actions {
match action {
ClaimAction::Lease { id, record } => {
let json = serde_json::to_vec(&record).map_err(MessagingError::backend)?;
self.kv
.put(&meta_key(topic, &id), json)
.await
.map_err(MessagingError::backend)?;
let payload = self.read_payload(topic, &id).await?;
claimed.push(ClaimedMessage {
id,
topic: topic.to_string(),
payload,
attempts: record.attempts,
});
}
ClaimAction::DeadLetter { id, record } => {
let json = serde_json::to_vec(&record).map_err(MessagingError::backend)?;
self.kv
.put(&dead_key(topic, &id), json)
.await
.map_err(MessagingError::backend)?;
self.kv
.delete(&meta_key(topic, &id))
.await
.map_err(MessagingError::backend)?;
}
}
}
Ok(claimed)
}
async fn ack(&self, msg: &ClaimedMessage) -> Result<(), MessagingError> {
self.kv
.delete(&meta_key(&msg.topic, &msg.id))
.await
.map_err(MessagingError::backend)?;
self.storage
.delete(&payload_key(&msg.topic, &msg.id))
.await
.map_err(MessagingError::backend)?;
Ok(())
}
async fn backlog(&self, topic: &str) -> Result<usize, MessagingError> {
self.count_direct(&meta_prefix(topic)).await
}
async fn dead_letter_count(&self, topic: &str) -> Result<usize, MessagingError> {
self.count_direct(&dead_prefix(topic)).await
}
async fn nack(&self, msg: &ClaimedMessage) -> Result<(), MessagingError> {
let key = meta_key(&msg.topic, &msg.id);
let Some(raw) = self.kv.get(&key).await.map_err(MessagingError::backend)? else {
return Ok(()); };
let mut record: Record =
serde_json::from_slice(&raw).map_err(|e| MessagingError::Decode(e.to_string()))?;
record.lease_until_ms = 0; let json = serde_json::to_vec(&record).map_err(MessagingError::backend)?;
self.kv
.put(&key, json)
.await
.map_err(MessagingError::backend)?;
Ok(())
}
async fn purge_dead_letters(&self, topic: &str) -> Result<usize, MessagingError> {
let prefix = dead_prefix(topic);
let keys = self
.kv
.list_prefix(&prefix)
.await
.map_err(MessagingError::backend)?;
let mut purged = 0;
for key in keys {
if !is_direct_child(&key, &prefix) {
continue; }
let id = &key[prefix.len()..];
self.storage
.delete(&payload_key(topic, id))
.await
.map_err(MessagingError::backend)?;
self.kv
.delete(&key)
.await
.map_err(MessagingError::backend)?;
purged += 1;
}
Ok(purged)
}
async fn redrive_dead_letters(&self, topic: &str) -> Result<usize, MessagingError> {
let prefix = dead_prefix(topic);
let keys = self
.kv
.list_prefix(&prefix)
.await
.map_err(MessagingError::backend)?;
let mut redriven = 0;
for key in keys {
if !is_direct_child(&key, &prefix) {
continue;
}
let id = &key[prefix.len()..];
let json = serde_json::to_vec(&Record::fresh()).map_err(MessagingError::backend)?;
self.kv
.put(&meta_key(topic, id), json)
.await
.map_err(MessagingError::backend)?;
self.kv
.delete(&key)
.await
.map_err(MessagingError::backend)?;
redriven += 1;
}
Ok(redriven)
}
fn subscribe(
&self,
topic: &str,
after: Option<&str>,
) -> futures::stream::BoxStream<'static, StreamEvent> {
self.hubs.subscribe(topic, after)
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::kv::MemoryKv;
use crate::{ByteStream, GetObject, ObjectMeta, StorageError};
use std::collections::HashMap;
use std::sync::Mutex;
#[derive(Default)]
struct MemStorage {
objects: Mutex<HashMap<String, Vec<u8>>>,
}
#[async_trait]
impl Storage for MemStorage {
async fn get(&self, key: &str) -> Result<GetObject, StorageError> {
let bytes = self
.objects
.lock()
.unwrap()
.get(key)
.cloned()
.ok_or_else(|| StorageError::NotFound(key.to_string()))?;
let size = bytes.len() as u64;
let body: ByteStream =
futures::stream::once(async move { Ok(bytes::Bytes::from(bytes)) }).boxed();
Ok(GetObject {
meta: ObjectMeta {
key: key.to_string(),
size: Some(size),
..Default::default()
},
body,
})
}
async fn get_range(
&self,
key: &str,
_: u64,
_: Option<u64>,
) -> Result<GetObject, StorageError> {
self.get(key).await
}
async fn put(
&self,
key: &str,
mut body: ByteStream,
_: PutMeta,
) -> Result<ObjectMeta, StorageError> {
let mut buf = Vec::new();
while let Some(chunk) = body.next().await {
buf.extend_from_slice(&chunk?);
}
let size = buf.len() as u64;
self.objects.lock().unwrap().insert(key.to_string(), buf);
Ok(ObjectMeta {
key: key.to_string(),
size: Some(size),
..Default::default()
})
}
async fn head(&self, key: &str) -> Result<ObjectMeta, StorageError> {
let map = self.objects.lock().unwrap();
let bytes = map
.get(key)
.ok_or_else(|| StorageError::NotFound(key.to_string()))?;
Ok(ObjectMeta {
key: key.to_string(),
size: Some(bytes.len() as u64),
..Default::default()
})
}
async fn delete(&self, key: &str) -> Result<(), StorageError> {
self.objects.lock().unwrap().remove(key);
Ok(())
}
async fn list(&self, prefix: &str) -> Result<Vec<ObjectMeta>, StorageError> {
Ok(self
.objects
.lock()
.unwrap()
.keys()
.filter(|k| k.starts_with(prefix))
.map(|k| ObjectMeta {
key: k.clone(),
..Default::default()
})
.collect())
}
}
fn mq() -> LogMessaging {
LogMessaging::new(Arc::new(MemStorage::default()), Arc::new(MemoryKv::new()))
}
const LEASE: Duration = Duration::from_secs(30);
#[tokio::test]
async fn publish_claim_ack_roundtrip_and_fifo() {
let mq = mq();
mq.publish("orders/created", b"a").await.unwrap();
mq.publish("orders/created", b"b").await.unwrap();
let batch = mq.claim("orders/created", LEASE, 10, 5).await.unwrap();
assert_eq!(batch.len(), 2);
assert_eq!(batch[0].payload, b"a");
assert_eq!(batch[1].payload, b"b");
assert_eq!(batch[0].attempts, 1);
assert!(mq
.claim("orders/created", LEASE, 10, 5)
.await
.unwrap()
.is_empty());
for m in &batch {
mq.ack(m).await.unwrap();
}
assert!(mq
.claim("orders/created", LEASE, 10, 5)
.await
.unwrap()
.is_empty());
}
#[tokio::test]
async fn topic_scoping_excludes_subtopics() {
let mq = mq();
mq.publish("orders", b"top").await.unwrap();
mq.publish("orders/created", b"sub").await.unwrap();
let batch = mq.claim("orders", LEASE, 10, 5).await.unwrap();
assert_eq!(batch.len(), 1);
assert_eq!(batch[0].payload, b"top");
}
#[tokio::test]
async fn lease_expiry_redelivers() {
let mq = mq();
mq.publish("t", b"x").await.unwrap();
let first = mq.claim("t", Duration::ZERO, 10, 5).await.unwrap();
assert_eq!(first.len(), 1);
assert_eq!(first[0].attempts, 1);
let second = mq.claim("t", LEASE, 10, 5).await.unwrap();
assert_eq!(second.len(), 1);
assert_eq!(second[0].attempts, 2); }
#[tokio::test]
async fn nack_makes_claimable_again() {
let mq = mq();
mq.publish("t", b"x").await.unwrap();
let m = mq.claim("t", LEASE, 10, 5).await.unwrap().pop().unwrap();
mq.nack(&m).await.unwrap();
let again = mq.claim("t", LEASE, 10, 5).await.unwrap();
assert_eq!(again.len(), 1);
assert_eq!(again[0].attempts, 2);
}
#[tokio::test]
async fn subscribe_receives_live_broadcast() {
use futures::StreamExt;
let mq = mq();
let mut sub = mq.subscribe("events", None);
mq.publish("events", b"hello").await.unwrap();
mq.publish("events", b"world").await.unwrap();
assert_eq!(sub.next().await.unwrap().payload, b"hello");
assert_eq!(sub.next().await.unwrap().payload, b"world");
mq.publish("other", b"nope").await.unwrap();
mq.publish("events", b"again").await.unwrap();
assert_eq!(sub.next().await.unwrap().payload, b"again");
}
#[tokio::test]
async fn last_event_id_replays_recent_then_goes_live() {
use futures::StreamExt;
let mq = mq();
let mut keepalive = mq.subscribe("events", None);
mq.publish("events", b"one").await.unwrap();
mq.publish("events", b"two").await.unwrap();
mq.publish("events", b"three").await.unwrap();
let first = keepalive.next().await.unwrap();
assert_eq!(first.payload, b"one");
let mut resumed = mq.subscribe("events", Some(&first.id));
assert_eq!(resumed.next().await.unwrap().payload, b"two");
assert_eq!(resumed.next().await.unwrap().payload, b"three");
mq.publish("events", b"four").await.unwrap();
assert_eq!(resumed.next().await.unwrap().payload, b"four");
}
#[tokio::test]
async fn dropped_subscriber_is_pruned_without_error() {
let mq = mq();
{
let _sub = mq.subscribe("events", None);
} mq.publish("events", b"x").await.unwrap();
}
#[tokio::test]
async fn dead_letters_after_max_attempts() {
let mq = mq();
mq.publish("t", b"x").await.unwrap();
for expected in 1..=2 {
let m = mq.claim("t", Duration::ZERO, 10, 2).await.unwrap();
assert_eq!(m.len(), 1, "attempt {expected}");
assert_eq!(m[0].attempts, expected);
}
let exhausted = mq.claim("t", Duration::ZERO, 10, 2).await.unwrap();
assert!(
exhausted.is_empty(),
"should dead-letter, not deliver a 3rd time"
);
assert_eq!(mq.dead_letter_count("t").await.unwrap(), 1);
}
#[tokio::test]
async fn purge_dead_letters_clears_records_and_payloads() {
let storage: Arc<dyn Storage> = Arc::new(MemStorage::default());
let kv: Arc<dyn KvStore> = Arc::new(MemoryKv::new());
let mq = LogMessaging::new(storage.clone(), kv);
mq.publish("t", b"x").await.unwrap();
let id = mq.claim("t", Duration::ZERO, 10, 1).await.unwrap()[0]
.id
.clone();
assert!(mq
.claim("t", Duration::ZERO, 10, 1)
.await
.unwrap()
.is_empty());
assert_eq!(mq.dead_letter_count("t").await.unwrap(), 1);
let purged = mq.purge_dead_letters("t").await.unwrap();
assert_eq!(purged, 1);
assert_eq!(mq.dead_letter_count("t").await.unwrap(), 0);
assert!(
storage.head(&payload_key("t", &id)).await.is_err(),
"purge frees the dead-lettered payload"
);
}
#[tokio::test]
async fn redrive_dead_letters_requeues_with_fresh_attempts() {
let mq = mq();
mq.publish("t", b"x").await.unwrap();
assert_eq!(mq.claim("t", Duration::ZERO, 10, 1).await.unwrap().len(), 1);
assert!(mq
.claim("t", Duration::ZERO, 10, 1)
.await
.unwrap()
.is_empty());
assert_eq!(mq.dead_letter_count("t").await.unwrap(), 1);
let redriven = mq.redrive_dead_letters("t").await.unwrap();
assert_eq!(redriven, 1);
assert_eq!(mq.dead_letter_count("t").await.unwrap(), 0);
assert_eq!(mq.backlog("t").await.unwrap(), 1);
let again = mq.claim("t", LEASE, 10, 5).await.unwrap();
assert_eq!(again.len(), 1);
assert_eq!(again[0].payload, b"x");
assert_eq!(again[0].attempts, 1, "fresh attempts after redrive");
}
#[tokio::test]
async fn survives_restart_over_shared_backends() {
let storage: Arc<dyn Storage> = Arc::new(MemStorage::default());
let kv: Arc<dyn KvStore> = Arc::new(MemoryKv::new());
{
let mq = LogMessaging::new(storage.clone(), kv.clone());
mq.publish("orders", b"a").await.unwrap();
mq.publish("orders", b"b").await.unwrap();
let batch = mq.claim("orders", Duration::ZERO, 10, 5).await.unwrap();
assert_eq!(batch.len(), 2);
mq.ack(&batch[0]).await.unwrap(); }
let mq = LogMessaging::new(storage, kv);
let batch = mq.claim("orders", LEASE, 10, 5).await.unwrap();
assert_eq!(batch.len(), 1, "only the un-acked message survives");
assert_eq!(batch[0].payload, b"b");
assert_eq!(batch[0].attempts, 2, "redelivery re-charges the attempt");
}
}