use crate::error::MqError;
use async_trait::async_trait;
use serde::{Deserialize, Serialize};
use std::collections::{HashMap, VecDeque};
use std::sync::Arc;
use std::time::{Duration, Instant};
use tokio::sync::{Notify, RwLock};
#[async_trait]
pub trait MessageQueue: Send + Sync {
async fn publish(&self, topic: &str, message: &[u8]) -> Result<(), MqError>;
async fn consume(&self, topic: &str) -> Result<Option<Message>, MqError>;
async fn ack(&self, message_id: &str) -> Result<(), MqError>;
async fn subscribe(&self, topic: &str) -> Result<(), MqError>;
async fn nack(&self, _message_id: &str) -> Result<(), MqError> {
Err(MqError::NotSupported("nack not supported".to_string()))
}
async fn reject(&self, _message_id: &str) -> Result<(), MqError> {
Err(MqError::NotSupported("reject not supported".to_string()))
}
}
#[derive(Debug, Clone, Serialize, Deserialize)]
pub struct Message {
pub topic: String,
pub payload: Vec<u8>,
pub key: Option<String>,
pub timestamp: i64,
pub headers: HashMap<String, String>,
#[serde(default)]
pub id: String,
#[serde(default)]
pub retry_count: u32,
}
impl Message {
pub fn new(topic: impl Into<String>, payload: Vec<u8>) -> Self {
Self {
topic: topic.into(),
payload,
key: None,
timestamp: current_timestamp(),
headers: HashMap::new(),
id: String::new(),
retry_count: 0,
}
}
pub fn with_key(mut self, key: impl Into<String>) -> Self {
self.key = Some(key.into());
self
}
pub fn text(&self) -> Option<&str> {
std::str::from_utf8(&self.payload).ok()
}
pub fn json<T: serde::de::DeserializeOwned>(&self) -> Option<T> {
serde_json::from_slice(&self.payload).ok()
}
pub fn text_message(topic: impl Into<String>, text: impl Into<String>) -> Self {
Self::new(topic, text.into().into_bytes())
}
pub fn json_message<T: serde::Serialize>(
topic: impl Into<String>,
data: &T,
) -> Result<Self, MqError> {
let payload = serde_json::to_vec(data)?;
Ok(Self::new(topic, payload))
}
}
fn current_timestamp() -> i64 {
use std::time::{SystemTime, UNIX_EPOCH};
match SystemTime::now().duration_since(UNIX_EPOCH) {
Ok(d) => d.as_millis() as i64,
Err(e) => {
eprintln!(
"WARN: current_timestamp: system time before UNIX_EPOCH: {} (duration_secs={})",
e,
e.duration().as_secs()
);
0
}
}
}
pub struct QueueConfig {
pub provider: MqProvider,
pub brokers: Vec<String>,
pub group_id: Option<String>,
pub username: Option<String>,
pub password: Option<String>,
}
impl Default for QueueConfig {
fn default() -> Self {
Self {
provider: MqProvider::Kafka(KafkaConfig::default()),
brokers: vec!["localhost:9092".to_string()],
group_id: None,
username: None,
password: None,
}
}
}
impl QueueConfig {
pub fn new() -> Self {
Self::default()
}
pub fn with_provider(mut self, provider: MqProvider) -> Self {
self.provider = provider;
self
}
pub fn with_brokers(mut self, brokers: Vec<String>) -> Self {
self.brokers = brokers;
self
}
pub fn with_group(mut self, group: impl Into<String>) -> Self {
self.group_id = Some(group.into());
self
}
pub fn with_auth(mut self, username: impl Into<String>, password: impl Into<String>) -> Self {
self.username = Some(username.into());
self.password = Some(password.into());
self
}
}
#[derive(Debug, Clone)]
pub enum MqProvider {
Kafka(KafkaConfig),
RabbitMQ(RabbitConfig),
RocketMQ(RocketConfig),
ActiveMQ(ActiveConfig),
Nats(NatsConfig),
Pulsar(PulsarConfig),
}
#[derive(Debug, Clone, Default)]
pub struct KafkaConfig {
pub client_id: Option<String>,
pub acks: Option<String>,
pub retries: Option<u32>,
}
#[derive(Debug, Clone, Default)]
pub struct RabbitConfig {
pub virtual_host: Option<String>,
}
#[derive(Debug, Clone, Default)]
pub struct RocketConfig {
pub namespace: Option<String>,
}
#[derive(Debug, Clone, Default)]
pub struct ActiveConfig {
pub broker_url: Option<String>,
}
#[derive(Debug, Clone, Default)]
pub struct NatsConfig {
pub name: Option<String>,
}
#[derive(Debug, Clone, Default)]
pub struct PulsarConfig {
pub service_url: Option<String>,
}
#[derive(Debug, Clone)]
pub struct ReconnectPolicy {
pub max_retries: u32,
pub initial_delay_ms: u64,
pub max_delay_ms: u64,
pub multiplier: f64,
}
impl Default for ReconnectPolicy {
fn default() -> Self {
Self {
max_retries: 5,
initial_delay_ms: 100,
max_delay_ms: 10_000,
multiplier: 2.0,
}
}
}
impl ReconnectPolicy {
pub fn new() -> Self {
Self::default()
}
pub fn next_delay(&self, attempt: u32) -> Duration {
let delay_ms = (self.initial_delay_ms as f64) * self.multiplier.powi(attempt as i32);
let delay_ms = delay_ms.max(0.0).min(self.max_delay_ms as f64);
Duration::from_millis(delay_ms as u64)
}
}
#[derive(Debug, Clone, Default)]
pub struct ReconnectState {
pub attempts: u32,
pub last_reconnect: Option<Instant>,
}
#[derive(Debug, Clone)]
pub struct BackpressurePolicy {
pub max_queue_size: usize,
pub on_overflow: OverflowStrategy,
}
impl Default for BackpressurePolicy {
fn default() -> Self {
Self {
max_queue_size: 10_000,
on_overflow: OverflowStrategy::Reject,
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum OverflowStrategy {
Block,
DropOldest,
DropNewest,
Reject,
}
pub struct InMemoryQueue {
inner: Arc<RwLock<InMemoryQueueInner>>,
}
impl Clone for InMemoryQueue {
fn clone(&self) -> Self {
Self {
inner: Arc::clone(&self.inner),
}
}
}
struct InMemoryQueueInner {
queues: HashMap<String, VecDeque<Message>>,
in_flight: HashMap<String, Message>,
subscribers: HashMap<String, usize>,
next_id: u64,
max_messages_per_topic: usize,
dead_letters: HashMap<String, VecDeque<Message>>,
max_retries: u32,
backpressure: Option<BackpressurePolicy>,
notify: HashMap<String, Arc<Notify>>,
}
const DEFAULT_MAX_MESSAGES_PER_TOPIC: usize = 100_000;
const DEFAULT_MAX_RETRIES: u32 = 3;
impl InMemoryQueue {
pub fn new() -> Self {
Self::with_max_messages_per_topic(DEFAULT_MAX_MESSAGES_PER_TOPIC)
}
pub fn with_max_messages_per_topic(max: usize) -> Self {
Self {
inner: Arc::new(RwLock::new(InMemoryQueueInner {
queues: HashMap::new(),
in_flight: HashMap::new(),
subscribers: HashMap::new(),
next_id: 1,
max_messages_per_topic: max,
dead_letters: HashMap::new(),
max_retries: DEFAULT_MAX_RETRIES,
backpressure: None,
notify: HashMap::new(),
})),
}
}
pub fn with_max_retries(max_retries: u32) -> Self {
Self {
inner: Arc::new(RwLock::new(InMemoryQueueInner {
queues: HashMap::new(),
in_flight: HashMap::new(),
subscribers: HashMap::new(),
next_id: 1,
max_messages_per_topic: DEFAULT_MAX_MESSAGES_PER_TOPIC,
dead_letters: HashMap::new(),
max_retries,
backpressure: None,
notify: HashMap::new(),
})),
}
}
pub fn with_backpressure(policy: BackpressurePolicy) -> Self {
let max = policy.max_queue_size;
Self {
inner: Arc::new(RwLock::new(InMemoryQueueInner {
queues: HashMap::new(),
in_flight: HashMap::new(),
subscribers: HashMap::new(),
next_id: 1,
max_messages_per_topic: max,
dead_letters: HashMap::new(),
max_retries: DEFAULT_MAX_RETRIES,
backpressure: Some(policy),
notify: HashMap::new(),
})),
}
}
pub async fn message_count(&self, topic: &str) -> usize {
let inner = self.inner.read().await;
inner.queues.get(topic).map(|q| q.len()).unwrap_or(0)
}
pub async fn subscriber_count(&self, topic: &str) -> usize {
let inner = self.inner.read().await;
*inner.subscribers.get(topic).unwrap_or(&0)
}
pub async fn in_flight_count(&self) -> usize {
let inner = self.inner.read().await;
inner.in_flight.len()
}
pub async fn dead_letter_count(&self, topic: &str) -> usize {
let inner = self.inner.read().await;
inner.dead_letters.get(topic).map(|q| q.len()).unwrap_or(0)
}
pub async fn consume_dead_letter(&self, topic: &str) -> Option<Message> {
let mut inner = self.inner.write().await;
if let Some(dq) = inner.dead_letters.get_mut(topic) {
return dq.pop_front();
}
None
}
pub async fn requeue_dead_letter(&self, message_id: &str) -> Result<(), MqError> {
let mut inner = self.inner.write().await;
for dq in inner.dead_letters.values_mut() {
let mut found_idx = None;
for (idx, m) in dq.iter().enumerate() {
if m.id == message_id {
found_idx = Some(idx);
break;
}
}
if let Some(idx) = found_idx {
let mut msg = dq.remove(idx).expect("checked: idx exists");
msg.retry_count = 0;
inner
.queues
.entry(msg.topic.clone())
.or_insert_with(VecDeque::new)
.push_back(msg);
return Ok(());
}
}
Err(MqError::NotSupported(format!(
"Dead letter not found: {}",
message_id
)))
}
}
impl Default for InMemoryQueue {
fn default() -> Self {
Self::new()
}
}
#[async_trait]
impl MessageQueue for InMemoryQueue {
async fn publish(&self, topic: &str, message: &[u8]) -> Result<(), MqError> {
let strategy = {
let inner = self.inner.read().await;
inner
.backpressure
.as_ref()
.map(|p| p.on_overflow)
.unwrap_or(OverflowStrategy::Reject)
};
match strategy {
OverflowStrategy::Block => self.publish_with_block(topic, message).await,
_ => self.publish_immediate(topic, message, strategy).await,
}
}
async fn consume(&self, topic: &str) -> Result<Option<Message>, MqError> {
let mut inner = self.inner.write().await;
let queue = inner
.queues
.entry(topic.to_string())
.or_insert_with(VecDeque::new);
if let Some(msg) = queue.pop_front() {
inner.in_flight.insert(msg.id.clone(), msg.clone());
if let Some(notify) = inner.notify.get(topic) {
notify.notify_one();
}
Ok(Some(msg))
} else {
Ok(None)
}
}
async fn ack(&self, message_id: &str) -> Result<(), MqError> {
let mut inner = self.inner.write().await;
inner.in_flight.remove(message_id).ok_or_else(|| {
MqError::NotSupported(format!("Message not found for ack: {}", message_id))
})?;
Ok(())
}
async fn subscribe(&self, topic: &str) -> Result<(), MqError> {
let mut inner = self.inner.write().await;
*inner.subscribers.entry(topic.to_string()).or_insert(0) += 1;
Ok(())
}
async fn nack(&self, message_id: &str) -> Result<(), MqError> {
let mut inner = self.inner.write().await;
let mut msg = inner.in_flight.remove(message_id).ok_or_else(|| {
MqError::NotSupported(format!("Message not found for nack: {}", message_id))
})?;
msg.retry_count = msg.retry_count.saturating_add(1);
if msg.retry_count >= inner.max_retries {
inner
.dead_letters
.entry(msg.topic.clone())
.or_insert_with(VecDeque::new)
.push_back(msg);
} else {
inner
.queues
.entry(msg.topic.clone())
.or_insert_with(VecDeque::new)
.push_back(msg);
}
Ok(())
}
async fn reject(&self, message_id: &str) -> Result<(), MqError> {
let mut inner = self.inner.write().await;
let msg = inner.in_flight.remove(message_id).ok_or_else(|| {
MqError::NotSupported(format!("Message not found for reject: {}", message_id))
})?;
inner
.dead_letters
.entry(msg.topic.clone())
.or_insert_with(VecDeque::new)
.push_back(msg);
Ok(())
}
}
impl InMemoryQueue {
async fn publish_immediate(
&self,
topic: &str,
message: &[u8],
strategy: OverflowStrategy,
) -> Result<(), MqError> {
let mut inner = self.inner.write().await;
let current_count = inner.queues.get(topic).map(|q| q.len()).unwrap_or(0);
if current_count >= inner.max_messages_per_topic {
match strategy {
OverflowStrategy::DropOldest => {
let queue = inner
.queues
.entry(topic.to_string())
.or_insert_with(VecDeque::new);
queue.pop_front();
}
OverflowStrategy::DropNewest => {
return Ok(());
}
OverflowStrategy::Reject => {
return Err(MqError::Publish(format!(
"topic '{}' is full: {} >= {} messages (H-3 protection)",
topic, current_count, inner.max_messages_per_topic
)));
}
OverflowStrategy::Block => {
return Err(MqError::Publish(format!(
"topic '{}' overflow with Block strategy: use publish_with_block instead",
topic
)));
}
}
}
let id = format!("msg-{}", inner.next_id);
inner.next_id = inner
.next_id
.checked_add(1)
.ok_or_else(|| MqError::Publish("message id overflow: u64::MAX reached".to_string()))?;
let msg = Message {
id,
retry_count: 0,
..Message::new(topic, message.to_vec())
};
inner
.queues
.entry(topic.to_string())
.or_insert_with(VecDeque::new)
.push_back(msg);
Ok(())
}
async fn publish_with_block(&self, topic: &str, message: &[u8]) -> Result<(), MqError> {
loop {
let notify = {
let mut inner = self.inner.write().await;
let current_count = inner.queues.get(topic).map(|q| q.len()).unwrap_or(0);
if current_count < inner.max_messages_per_topic {
let id = format!("msg-{}", inner.next_id);
inner.next_id = inner.next_id.checked_add(1).ok_or_else(|| {
MqError::Publish("message id overflow: u64::MAX reached".to_string())
})?;
let msg = Message {
id,
retry_count: 0,
..Message::new(topic, message.to_vec())
};
inner
.queues
.entry(topic.to_string())
.or_insert_with(VecDeque::new)
.push_back(msg);
return Ok(());
}
inner
.notify
.entry(topic.to_string())
.or_insert_with(|| Arc::new(Notify::new()))
.clone()
};
notify.notified().await;
}
}
}
pub struct QueueWrapper {
queue: Box<dyn MessageQueue>,
reconnect: Option<ReconnectPolicy>,
}
impl QueueWrapper {
pub fn new(provider: MqProvider) -> Self {
let queue: Box<dyn MessageQueue> = match provider {
MqProvider::Kafka(_) => Box::new(crate::kafka::InMemoryKafkaQueue::new()),
MqProvider::RabbitMQ(_) => Box::new(crate::rabbitmq::InMemoryRabbitmqQueue::new()),
MqProvider::RocketMQ(_) => Box::new(crate::rocketmq::InMemoryRocketmqQueue::new()),
MqProvider::ActiveMQ(_) => Box::new(crate::activemq::InMemoryActivemqQueue::new()),
MqProvider::Nats(_) => Box::new(crate::nats::InMemoryNatsQueue::new()),
MqProvider::Pulsar(_) => Box::new(crate::pulsar::InMemoryPulsarQueue::new()),
};
Self {
queue,
reconnect: None,
}
}
#[cfg(test)]
pub(crate) fn with_queue(queue: Box<dyn MessageQueue>) -> Self {
Self {
queue,
reconnect: None,
}
}
pub fn with_reconnect(mut self, policy: ReconnectPolicy) -> Self {
self.reconnect = Some(policy);
self
}
pub async fn publish(&self, topic: &str, message: &[u8]) -> Result<(), MqError> {
if let Some(policy) = &self.reconnect {
let mut attempts = 0u32;
loop {
match self.queue.publish(topic, message).await {
Ok(()) => return Ok(()),
Err(MqError::Connection(_)) if attempts < policy.max_retries => {
let delay = policy.next_delay(attempts);
tokio::time::sleep(delay).await;
attempts += 1;
}
Err(e) => return Err(e),
}
}
} else {
self.queue.publish(topic, message).await
}
}
pub async fn consume(&self, topic: &str) -> Result<Option<Message>, MqError> {
if let Some(policy) = &self.reconnect {
let mut attempts = 0u32;
loop {
match self.queue.consume(topic).await {
Ok(msg) => return Ok(msg),
Err(MqError::Connection(_)) if attempts < policy.max_retries => {
let delay = policy.next_delay(attempts);
tokio::time::sleep(delay).await;
attempts += 1;
}
Err(e) => return Err(e),
}
}
} else {
self.queue.consume(topic).await
}
}
pub async fn ack(&self, message_id: &str) -> Result<(), MqError> {
self.queue.ack(message_id).await
}
pub async fn subscribe(&self, topic: &str) -> Result<(), MqError> {
self.queue.subscribe(topic).await
}
pub async fn nack(&self, message_id: &str) -> Result<(), MqError> {
self.queue.nack(message_id).await
}
pub async fn reject(&self, message_id: &str) -> Result<(), MqError> {
self.queue.reject(message_id).await
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::sync::atomic::{AtomicU32, Ordering};
#[tokio::test]
async fn test_in_memory_queue_basic() {
let queue = InMemoryQueue::new();
queue.publish("topic1", b"hello").await.unwrap();
let msg = queue
.consume("topic1")
.await
.unwrap()
.expect("msg should exist");
assert_eq!(msg.payload, b"hello");
queue.ack(&msg.id).await.unwrap();
}
#[tokio::test]
async fn test_l2_next_id_overflow_protection() {
let queue = InMemoryQueue::new();
{
let mut inner = queue.inner.write().await;
inner.next_id = u64::MAX;
}
let result = queue.publish("topic1", b"msg").await;
assert!(result.is_err());
match result {
Err(MqError::Publish(msg)) => {
assert!(
msg.contains("overflow"),
"expected overflow error, got: {}",
msg
);
}
_ => panic!("Expected MqError::Publish with overflow message"),
}
}
#[tokio::test]
async fn test_l2_next_id_near_max() {
let queue = InMemoryQueue::new();
{
let mut inner = queue.inner.write().await;
inner.next_id = u64::MAX - 1;
}
let result1 = queue.publish("topic1", b"msg1").await;
assert!(result1.is_ok());
let result2 = queue.publish("topic1", b"msg2").await;
assert!(result2.is_err());
}
#[tokio::test]
async fn test_nack_requeues_message_with_retry_count() {
let queue = InMemoryQueue::new();
queue.publish("topic", b"msg1").await.unwrap();
let msg = queue.consume("topic").await.unwrap().unwrap();
assert_eq!(msg.retry_count, 0);
queue.nack(&msg.id).await.unwrap();
assert_eq!(queue.message_count("topic").await, 1);
assert_eq!(queue.in_flight_count().await, 0);
let msg2 = queue.consume("topic").await.unwrap().unwrap();
assert_eq!(msg2.id, msg.id);
assert_eq!(msg2.retry_count, 1);
}
#[tokio::test]
async fn test_nack_increments_retry_count() {
let queue = InMemoryQueue::with_max_retries(10);
queue.publish("topic", b"data").await.unwrap();
for expected_retry in 0..5u32 {
let msg = queue.consume("topic").await.unwrap().unwrap();
assert_eq!(
msg.retry_count, expected_retry,
"consume should show retry_count before this iteration's nack"
);
queue.nack(&msg.id).await.unwrap();
}
assert_eq!(queue.message_count("topic").await, 1);
assert_eq!(queue.dead_letter_count("topic").await, 0);
}
#[tokio::test]
async fn test_nack_max_retries_sends_to_dlx() {
let queue = InMemoryQueue::with_max_retries(3);
queue.publish("topic", b"payload").await.unwrap();
let msg = queue.consume("topic").await.unwrap().unwrap();
queue.nack(&msg.id).await.unwrap();
assert_eq!(queue.message_count("topic").await, 1);
assert_eq!(queue.dead_letter_count("topic").await, 0);
let msg = queue.consume("topic").await.unwrap().unwrap();
assert_eq!(msg.retry_count, 1);
queue.nack(&msg.id).await.unwrap();
assert_eq!(queue.message_count("topic").await, 1);
assert_eq!(queue.dead_letter_count("topic").await, 0);
let msg = queue.consume("topic").await.unwrap().unwrap();
assert_eq!(msg.retry_count, 2);
queue.nack(&msg.id).await.unwrap();
assert_eq!(queue.message_count("topic").await, 0);
assert_eq!(queue.dead_letter_count("topic").await, 1);
}
#[tokio::test]
async fn test_nack_max_retries_zero_sends_to_dlx_immediately() {
let queue = InMemoryQueue::with_max_retries(0);
queue.publish("topic", b"msg").await.unwrap();
let msg = queue.consume("topic").await.unwrap().unwrap();
assert_eq!(msg.retry_count, 0);
queue.nack(&msg.id).await.unwrap();
assert_eq!(queue.message_count("topic").await, 0);
assert_eq!(queue.dead_letter_count("topic").await, 1);
let dlq_msg = queue
.consume_dead_letter("topic")
.await
.expect("should have dead letter");
assert_eq!(dlq_msg.retry_count, 1);
}
#[tokio::test]
async fn test_nack_unknown_message_id_returns_error() {
let queue = InMemoryQueue::new();
let result = queue.nack("nonexistent-id").await;
assert!(result.is_err());
match result {
Err(MqError::NotSupported(msg)) => {
assert!(msg.contains("not found for nack"));
}
_ => panic!("Expected MqError::NotSupported"),
}
}
#[tokio::test]
async fn test_reject_sends_to_dead_letter_queue() {
let queue = InMemoryQueue::new();
queue.publish("topic", b"bad-msg").await.unwrap();
let msg = queue.consume("topic").await.unwrap().unwrap();
queue.reject(&msg.id).await.unwrap();
assert_eq!(queue.message_count("topic").await, 0);
assert_eq!(queue.in_flight_count().await, 0);
assert_eq!(queue.dead_letter_count("topic").await, 1);
}
#[tokio::test]
async fn test_reject_does_not_increment_retry_count() {
let queue = InMemoryQueue::new();
queue.publish("topic", b"msg").await.unwrap();
let msg = queue.consume("topic").await.unwrap().unwrap();
assert_eq!(msg.retry_count, 0);
queue.reject(&msg.id).await.unwrap();
let dlq_msg = queue
.consume_dead_letter("topic")
.await
.expect("should have dead letter");
assert_eq!(
dlq_msg.retry_count, 0,
"reject should not increment retry_count"
);
}
#[tokio::test]
async fn test_reject_unknown_message_id_returns_error() {
let queue = InMemoryQueue::new();
let result = queue.reject("nonexistent-id").await;
assert!(result.is_err());
}
#[tokio::test]
async fn test_reject_empty_in_flight_returns_error() {
let queue = InMemoryQueue::new();
let result = queue.reject("any-id").await;
assert!(result.is_err());
assert_eq!(queue.dead_letter_count("topic").await, 0);
}
#[tokio::test]
async fn test_dead_letter_count_empty_topic() {
let queue = InMemoryQueue::new();
assert_eq!(queue.dead_letter_count("no-such-topic").await, 0);
}
#[tokio::test]
async fn test_dead_letter_count_after_reject() {
let queue = InMemoryQueue::new();
queue.publish("topic", b"m1").await.unwrap();
queue.publish("topic", b"m2").await.unwrap();
let m1 = queue.consume("topic").await.unwrap().unwrap();
queue.reject(&m1.id).await.unwrap();
assert_eq!(queue.dead_letter_count("topic").await, 1);
let m2 = queue.consume("topic").await.unwrap().unwrap();
queue.reject(&m2.id).await.unwrap();
assert_eq!(queue.dead_letter_count("topic").await, 2);
}
#[tokio::test]
async fn test_consume_dead_letter() {
let queue = InMemoryQueue::new();
queue.publish("topic", b"first").await.unwrap();
queue.publish("topic", b"second").await.unwrap();
let m1 = queue.consume("topic").await.unwrap().unwrap();
queue.reject(&m1.id).await.unwrap();
let m2 = queue.consume("topic").await.unwrap().unwrap();
queue.reject(&m2.id).await.unwrap();
let d1 = queue
.consume_dead_letter("topic")
.await
.expect("should have dead letter");
assert_eq!(d1.payload, b"first");
let d2 = queue
.consume_dead_letter("topic")
.await
.expect("should have dead letter");
assert_eq!(d2.payload, b"second");
assert!(queue.consume_dead_letter("topic").await.is_none());
}
#[tokio::test]
async fn test_consume_dead_letter_empty() {
let queue = InMemoryQueue::new();
assert!(queue.consume_dead_letter("no-such-topic").await.is_none());
}
#[tokio::test]
async fn test_requeue_dead_letter_resets_retry_count() {
let queue = InMemoryQueue::with_max_retries(2);
queue.publish("topic", b"msg").await.unwrap();
let m1 = queue.consume("topic").await.unwrap().unwrap();
queue.nack(&m1.id).await.unwrap();
let m2 = queue.consume("topic").await.unwrap().unwrap();
assert_eq!(m2.retry_count, 1);
queue.nack(&m2.id).await.unwrap();
assert_eq!(queue.dead_letter_count("topic").await, 1);
assert_eq!(queue.message_count("topic").await, 0);
queue.requeue_dead_letter(&m1.id).await.unwrap();
assert_eq!(queue.dead_letter_count("topic").await, 0);
assert_eq!(queue.message_count("topic").await, 1);
let m3 = queue.consume("topic").await.unwrap().unwrap();
assert_eq!(m3.id, m1.id);
assert_eq!(
m3.retry_count, 0,
"retry_count should be reset after requeue"
);
}
#[tokio::test]
async fn test_requeue_dead_letter_not_found() {
let queue = InMemoryQueue::new();
let result = queue.requeue_dead_letter("nonexistent-id").await;
assert!(result.is_err());
match result {
Err(MqError::NotSupported(msg)) => {
assert!(msg.contains("Dead letter not found"));
}
_ => panic!("Expected MqError::NotSupported"),
}
}
#[test]
fn test_reconnect_policy_default_values() {
let policy = ReconnectPolicy::default();
assert_eq!(policy.max_retries, 5);
assert_eq!(policy.initial_delay_ms, 100);
assert_eq!(policy.max_delay_ms, 10_000);
assert!((policy.multiplier - 2.0).abs() < f64::EPSILON);
}
#[test]
fn test_reconnect_policy_next_delay_exponential() {
let policy = ReconnectPolicy {
max_retries: 5,
initial_delay_ms: 100,
max_delay_ms: 10_000,
multiplier: 2.0,
};
assert_eq!(policy.next_delay(0), Duration::from_millis(100));
assert_eq!(policy.next_delay(1), Duration::from_millis(200));
assert_eq!(policy.next_delay(2), Duration::from_millis(400));
assert_eq!(policy.next_delay(3), Duration::from_millis(800));
}
#[test]
fn test_reconnect_policy_next_delay_capped_at_max() {
let policy = ReconnectPolicy {
max_retries: 10,
initial_delay_ms: 100,
max_delay_ms: 1000,
multiplier: 2.0,
};
assert_eq!(policy.next_delay(4), Duration::from_millis(1000));
assert_eq!(policy.next_delay(10), Duration::from_millis(1000));
}
#[test]
fn test_reconnect_policy_zero_attempt() {
let policy = ReconnectPolicy {
max_retries: 3,
initial_delay_ms: 500,
max_delay_ms: 10_000,
multiplier: 3.0,
};
assert_eq!(policy.next_delay(0), Duration::from_millis(500));
}
struct FailingQueue {
call_count: AtomicU32,
succeed_on_attempt: u32,
}
impl FailingQueue {
fn new(succeed_on_attempt: u32) -> Self {
Self {
call_count: AtomicU32::new(0),
succeed_on_attempt,
}
}
}
#[async_trait]
impl MessageQueue for FailingQueue {
async fn publish(&self, _topic: &str, _message: &[u8]) -> Result<(), MqError> {
let attempt = self.call_count.fetch_add(1, Ordering::SeqCst) + 1;
if attempt >= self.succeed_on_attempt {
Ok(())
} else {
Err(MqError::Connection(
"simulated connection error".to_string(),
))
}
}
async fn consume(&self, _topic: &str) -> Result<Option<Message>, MqError> {
Err(MqError::Connection(
"simulated connection error".to_string(),
))
}
async fn ack(&self, _message_id: &str) -> Result<(), MqError> {
Err(MqError::Connection("simulated".to_string()))
}
async fn subscribe(&self, _topic: &str) -> Result<(), MqError> {
Err(MqError::Connection("simulated".to_string()))
}
}
struct PublishErrorQueue;
#[async_trait]
impl MessageQueue for PublishErrorQueue {
async fn publish(&self, _topic: &str, _message: &[u8]) -> Result<(), MqError> {
Err(MqError::Publish("non-connection error".to_string()))
}
async fn consume(&self, _topic: &str) -> Result<Option<Message>, MqError> {
Err(MqError::Publish("non-connection error".to_string()))
}
async fn ack(&self, _message_id: &str) -> Result<(), MqError> {
Ok(())
}
async fn subscribe(&self, _topic: &str) -> Result<(), MqError> {
Ok(())
}
}
#[tokio::test]
async fn test_reconnect_retries_on_connection_error() {
let failing = FailingQueue::new(3);
let wrapper = QueueWrapper::with_queue(Box::new(failing)).with_reconnect(ReconnectPolicy {
max_retries: 5,
initial_delay_ms: 1, max_delay_ms: 10,
multiplier: 2.0,
});
let result = wrapper.publish("topic", b"data").await;
assert!(result.is_ok(), "should succeed after retries");
}
#[tokio::test]
async fn test_reconnect_gives_up_after_max_retries() {
let failing = FailingQueue::new(u32::MAX);
let wrapper = QueueWrapper::with_queue(Box::new(failing)).with_reconnect(ReconnectPolicy {
max_retries: 2,
initial_delay_ms: 1,
max_delay_ms: 10,
multiplier: 2.0,
});
let result = wrapper.publish("topic", b"data").await;
assert!(result.is_err());
match result {
Err(MqError::Connection(_)) => {}
_ => panic!("Expected MqError::Connection"),
}
}
#[tokio::test]
async fn test_reconnect_no_retry_on_non_connection_error() {
let wrapper =
QueueWrapper::with_queue(Box::new(PublishErrorQueue)).with_reconnect(ReconnectPolicy {
max_retries: 5,
initial_delay_ms: 1,
max_delay_ms: 10,
multiplier: 2.0,
});
let result = wrapper.publish("topic", b"data").await;
assert!(result.is_err());
match result {
Err(MqError::Publish(msg)) => {
assert!(msg.contains("non-connection error"));
}
_ => panic!("Expected MqError::Publish"),
}
}
#[tokio::test]
async fn test_backpressure_drop_oldest() {
let queue = InMemoryQueue::with_backpressure(BackpressurePolicy {
max_queue_size: 2,
on_overflow: OverflowStrategy::DropOldest,
});
queue.publish("topic", b"m1").await.unwrap();
queue.publish("topic", b"m2").await.unwrap();
queue.publish("topic", b"m3").await.unwrap();
assert_eq!(queue.message_count("topic").await, 2);
let m1 = queue.consume("topic").await.unwrap().unwrap();
assert_eq!(m1.payload, b"m2", "oldest should be dropped");
let m2 = queue.consume("topic").await.unwrap().unwrap();
assert_eq!(m2.payload, b"m3");
}
#[tokio::test]
async fn test_backpressure_drop_newest() {
let queue = InMemoryQueue::with_backpressure(BackpressurePolicy {
max_queue_size: 1,
on_overflow: OverflowStrategy::DropNewest,
});
queue.publish("topic", b"m1").await.unwrap();
let result = queue.publish("topic", b"m2").await;
assert!(result.is_ok());
assert_eq!(queue.message_count("topic").await, 1);
let m = queue.consume("topic").await.unwrap().unwrap();
assert_eq!(m.payload, b"m1", "newest should be dropped");
}
#[tokio::test]
async fn test_backpressure_reject() {
let queue = InMemoryQueue::with_backpressure(BackpressurePolicy {
max_queue_size: 1,
on_overflow: OverflowStrategy::Reject,
});
queue.publish("topic", b"m1").await.unwrap();
let result = queue.publish("topic", b"m2").await;
assert!(result.is_err());
assert_eq!(queue.message_count("topic").await, 1);
}
#[tokio::test]
async fn test_backpressure_block_unblocks_on_consume() {
let queue = InMemoryQueue::with_backpressure(BackpressurePolicy {
max_queue_size: 1,
on_overflow: OverflowStrategy::Block,
});
queue.publish("topic", b"m1").await.unwrap();
let queue_clone = queue.clone();
let handle = tokio::spawn(async move { queue_clone.publish("topic", b"m2").await });
tokio::time::sleep(Duration::from_millis(50)).await;
assert!(!handle.is_finished(), "publish should be blocked");
queue.consume("topic").await.unwrap();
let result = tokio::time::timeout(Duration::from_secs(1), handle)
.await
.expect("publish should complete after consume");
assert!(result.is_ok(), "publish should succeed: {:?}", result);
assert_eq!(queue.message_count("topic").await, 1);
let m = queue.consume("topic").await.unwrap().unwrap();
assert_eq!(m.payload, b"m2");
}
#[tokio::test]
async fn test_backpressure_block_times_out_when_queue_stays_full() {
let queue = InMemoryQueue::with_backpressure(BackpressurePolicy {
max_queue_size: 1,
on_overflow: OverflowStrategy::Block,
});
queue.publish("topic", b"m1").await.unwrap();
let result =
tokio::time::timeout(Duration::from_millis(100), queue.publish("topic", b"m2")).await;
assert!(result.is_err(), "publish should block and time out");
assert_eq!(queue.message_count("topic").await, 1);
}
#[tokio::test]
async fn test_backpressure_block_isolated_per_topic() {
let queue = InMemoryQueue::with_backpressure(BackpressurePolicy {
max_queue_size: 1,
on_overflow: OverflowStrategy::Block,
});
queue.publish("topic-a", b"a1").await.unwrap();
queue.publish("topic-b", b"b1").await.unwrap();
let queue_clone = queue.clone();
let handle = tokio::spawn(async move { queue_clone.publish("topic-a", b"a2").await });
tokio::time::sleep(Duration::from_millis(50)).await;
assert!(!handle.is_finished(), "topic-a publish should be blocked");
queue.consume("topic-b").await.unwrap();
tokio::time::sleep(Duration::from_millis(50)).await;
assert!(
!handle.is_finished(),
"topic-a publish should still be blocked after topic-b consume"
);
queue.consume("topic-a").await.unwrap();
let result = tokio::time::timeout(Duration::from_secs(1), handle)
.await
.expect("publish should complete after topic-a consume");
assert!(result.is_ok());
}
}