use std::sync::Arc;
use std::time::Duration;
use async_trait::async_trait;
use tokio_util::sync::CancellationToken;
use crate::orm::MessageQueue;
#[derive(Debug, Clone, thiserror::Error)]
pub enum QueueConsumerError {
#[error("consumer error: {0}")]
Handler(String),
#[error("queue error: {0}")]
Queue(String),
}
#[async_trait]
pub trait QueueConsumer: Send + Sync {
async fn handle(&self, message: &sz_orm_queue::Message) -> Result<(), QueueConsumerError>;
}
#[derive(Debug, Clone)]
pub struct QueueRuntimeConfig {
pub topic: String,
pub poll_interval_ms: u64,
pub max_retries: u32,
}
impl Default for QueueRuntimeConfig {
fn default() -> Self {
Self {
topic: "default".to_string(),
poll_interval_ms: 100,
max_retries: 0,
}
}
}
impl QueueRuntimeConfig {
pub fn new(topic: impl Into<String>) -> Self {
Self {
topic: topic.into(),
..Default::default()
}
}
pub fn with_poll_interval(mut self, ms: u64) -> Self {
self.poll_interval_ms = ms;
self
}
pub fn with_max_retries(mut self, n: u32) -> Self {
self.max_retries = n;
self
}
}
pub struct QueueRuntime {
config: QueueRuntimeConfig,
queue: Arc<dyn MessageQueue>,
}
impl QueueRuntime {
pub fn new(config: QueueRuntimeConfig, queue: Arc<dyn MessageQueue>) -> Self {
Self { config, queue }
}
pub fn start<C>(
&self,
consumer: Arc<C>,
token: CancellationToken,
) -> tokio::task::JoinHandle<()>
where
C: QueueConsumer + 'static,
{
let queue = self.queue.clone();
let topic = self.config.topic.clone();
let poll_interval = Duration::from_millis(self.config.poll_interval_ms.max(1));
tokio::spawn(async move {
loop {
tokio::select! {
_ = token.cancelled() => break,
consume_result = queue.consume(&topic) => {
match consume_result {
Ok(Some(message)) => {
let msg_id = message.id.clone();
match consumer.handle(&message).await {
Ok(()) => {
if let Err(e) = queue.ack(&msg_id).await {
tracing::warn!("ack failed for msg {}: {}", msg_id, e);
}
}
Err(e) => {
tracing::warn!(
"consumer handler failed for msg {}: {}",
msg_id,
e
);
}
}
}
Ok(None) => {
tokio::time::sleep(poll_interval).await;
}
Err(e) => {
tracing::error!("queue consume error: {}", e);
tokio::time::sleep(poll_interval).await;
}
}
}
}
}
})
}
pub fn topic(&self) -> &str {
&self.config.topic
}
pub fn config(&self) -> &QueueRuntimeConfig {
&self.config
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::orm::{InMemoryQueue, Message, MessageQueue};
struct RecordingConsumer {
payloads: Arc<parking_lot::Mutex<Vec<Vec<u8>>>>,
}
impl RecordingConsumer {
fn new() -> (Self, Arc<parking_lot::Mutex<Vec<Vec<u8>>>>) {
let payloads = Arc::new(parking_lot::Mutex::new(Vec::new()));
let consumer = Self {
payloads: payloads.clone(),
};
(consumer, payloads)
}
}
#[async_trait]
impl QueueConsumer for RecordingConsumer {
async fn handle(&self, message: &Message) -> Result<(), QueueConsumerError> {
self.payloads.lock().push(message.payload.clone());
Ok(())
}
}
struct FailingConsumer;
#[async_trait]
impl QueueConsumer for FailingConsumer {
async fn handle(&self, _message: &Message) -> Result<(), QueueConsumerError> {
Err(QueueConsumerError::Handler("always fail".to_string()))
}
}
fn make_queue() -> Arc<dyn MessageQueue> {
Arc::new(InMemoryQueue::new())
}
#[test]
fn test_queue_runtime_config_default() {
let config = QueueRuntimeConfig::default();
assert_eq!(config.topic, "default");
assert_eq!(config.poll_interval_ms, 100);
assert_eq!(config.max_retries, 0);
}
#[test]
fn test_queue_runtime_config_builder() {
let config = QueueRuntimeConfig::new("orders")
.with_poll_interval(50)
.with_max_retries(3);
assert_eq!(config.topic, "orders");
assert_eq!(config.poll_interval_ms, 50);
assert_eq!(config.max_retries, 3);
}
#[test]
fn test_queue_runtime_topic_accessor() {
let queue = make_queue();
let runtime = QueueRuntime::new(QueueRuntimeConfig::new("test"), queue);
assert_eq!(runtime.topic(), "test");
}
#[test]
fn test_queue_runtime_config_accessor() {
let queue = make_queue();
let config = QueueRuntimeConfig::new("test").with_poll_interval(200);
let runtime = QueueRuntime::new(config, queue);
assert_eq!(runtime.config().poll_interval_ms, 200);
}
#[tokio::test]
async fn test_consumer_consumes_published_message() {
let queue = make_queue();
queue.publish("orders", b"hello").await.unwrap();
let (consumer, payloads) = RecordingConsumer::new();
let runtime = QueueRuntime::new(
QueueRuntimeConfig::new("orders").with_poll_interval(5),
queue.clone(),
);
let token = CancellationToken::new();
let handle = runtime.start(Arc::new(consumer), token.clone());
tokio::time::sleep(Duration::from_millis(100)).await;
token.cancel();
let _ = handle.await;
let recorded = payloads.lock().clone();
assert_eq!(recorded.len(), 1);
assert_eq!(recorded[0], b"hello");
}
#[tokio::test]
async fn test_consumer_acks_on_success() {
let queue = make_queue();
queue.publish("orders", b"msg1").await.unwrap();
let (consumer, _payloads) = RecordingConsumer::new();
let runtime = QueueRuntime::new(
QueueRuntimeConfig::new("orders").with_poll_interval(5),
queue.clone(),
);
let token = CancellationToken::new();
let handle = runtime.start(Arc::new(consumer), token.clone());
tokio::time::sleep(Duration::from_millis(100)).await;
token.cancel();
let _ = handle.await;
let result = queue.consume("orders").await.unwrap();
assert!(result.is_none());
}
#[tokio::test]
async fn test_consumer_no_ack_on_failure() {
let queue = make_queue();
queue.publish("orders", b"msg1").await.unwrap();
let runtime = QueueRuntime::new(
QueueRuntimeConfig::new("orders").with_poll_interval(5),
queue.clone(),
);
let token = CancellationToken::new();
let handle = runtime.start(Arc::new(FailingConsumer), token.clone());
tokio::time::sleep(Duration::from_millis(100)).await;
token.cancel();
let _ = handle.await;
}
#[tokio::test]
async fn test_consumer_stops_on_cancel() {
let queue = make_queue();
let (consumer, _payloads) = RecordingConsumer::new();
let runtime = QueueRuntime::new(
QueueRuntimeConfig::new("orders").with_poll_interval(5),
queue,
);
let token = CancellationToken::new();
let handle = runtime.start(Arc::new(consumer), token.clone());
token.cancel();
let _ = tokio::time::timeout(Duration::from_millis(500), handle).await;
}
#[tokio::test]
async fn test_consumer_handles_empty_queue() {
let queue = make_queue();
let (consumer, payloads) = RecordingConsumer::new();
let runtime = QueueRuntime::new(
QueueRuntimeConfig::new("empty").with_poll_interval(5),
queue.clone(),
);
let token = CancellationToken::new();
let handle = runtime.start(Arc::new(consumer), token.clone());
tokio::time::sleep(Duration::from_millis(50)).await;
token.cancel();
let _ = handle.await;
assert!(payloads.lock().is_empty());
}
#[tokio::test]
async fn test_consumer_processes_multiple_messages() {
let queue = make_queue();
queue.publish("orders", b"msg1").await.unwrap();
queue.publish("orders", b"msg2").await.unwrap();
queue.publish("orders", b"msg3").await.unwrap();
let (consumer, payloads) = RecordingConsumer::new();
let runtime = QueueRuntime::new(
QueueRuntimeConfig::new("orders").with_poll_interval(5),
queue,
);
let token = CancellationToken::new();
let handle = runtime.start(Arc::new(consumer), token.clone());
tokio::time::sleep(Duration::from_millis(200)).await;
token.cancel();
let _ = handle.await;
let recorded = payloads.lock().clone();
assert_eq!(recorded.len(), 3);
}
#[test]
fn test_queue_consumer_error_variants() {
let handler_err = QueueConsumerError::Handler("test".to_string());
let queue_err = QueueConsumerError::Queue("queue fail".to_string());
assert!(format!("{}", handler_err).contains("consumer error"));
assert!(format!("{}", queue_err).contains("queue error"));
}
}