Skip to main content

sz_rust_core/runtime/
queue.rs

1//! sz-orm-queue 消费者接入
2//!
3//! ## PHP 对齐
4//!
5//! 对齐 PHP `think-queue` 的消费者模型:
6//!
7//! ```php
8//! // think-queue 的 worker 模型
9//! $worker->runNextJob($connection, $queue, $delay, $sleep, $tries);
10//! ```
11//!
12//! Rust 端复用 `sz-orm-queue::MessageQueue` trait,提供 consumer loop helper:
13//!
14//! - `QueueConsumer` trait:业务侧实现消息处理逻辑
15//! - `QueueRuntime`:管理 consumer lifecycle(启动 / 停止 / 监听 cancel)
16//!
17//! ## 设计
18//!
19//! sz-orm-queue 的 `MessageQueue` trait 只提供原子操作(publish/consume/ack),
20//! 没有 consumer loop helper。本模块补充该能力,配合 `CancellationToken` 做退出。
21
22use std::sync::Arc;
23use std::time::Duration;
24
25use async_trait::async_trait;
26use tokio_util::sync::CancellationToken;
27
28use sz_orm_queue::MessageQueue;
29
30/// 队列消费者错误
31#[derive(Debug, Clone, thiserror::Error)]
32pub enum QueueConsumerError {
33    /// 消息处理失败(消息会被 nack,留在 in_flight 中)
34    #[error("consumer error: {0}")]
35    Handler(String),
36    /// 底层队列操作失败
37    #[error("queue error: {0}")]
38    Queue(String),
39}
40
41/// 队列消费者 trait
42///
43/// 业务侧实现该 trait,处理从队列消费的消息。
44///
45/// ## 行为契约
46///
47/// - 返回 `Ok(())`:自动调用 `queue.ack(message_id)` 确认消息
48/// - 返回 `Err(_)`:不 ack(消息留在 `in_flight`,需人工干预或超时重投)
49#[async_trait]
50pub trait QueueConsumer: Send + Sync {
51    /// 处理一条消息
52    ///
53    /// - 成功返回 `Ok(())` 会触发自动 ack
54    /// - 失败返回 `Err(_)` 会跳过 ack(消息留在 in_flight)
55    async fn handle(&self, message: &sz_orm_queue::Message) -> Result<(), QueueConsumerError>;
56}
57
58/// 队列运行时配置
59#[derive(Debug, Clone)]
60pub struct QueueRuntimeConfig {
61    /// 消费主题
62    pub topic: String,
63    /// 队列为空时的轮询间隔(毫秒)
64    pub poll_interval_ms: u64,
65    /// 单条消息处理最大重试次数(0 表示不重试)
66    pub max_retries: u32,
67}
68
69impl Default for QueueRuntimeConfig {
70    fn default() -> Self {
71        Self {
72            topic: "default".to_string(),
73            poll_interval_ms: 100,
74            max_retries: 0,
75        }
76    }
77}
78
79impl QueueRuntimeConfig {
80    /// 创建新配置
81    pub fn new(topic: impl Into<String>) -> Self {
82        Self {
83            topic: topic.into(),
84            ..Default::default()
85        }
86    }
87
88    /// 自定义轮询间隔
89    pub fn with_poll_interval(mut self, ms: u64) -> Self {
90        self.poll_interval_ms = ms;
91        self
92    }
93
94    /// 自定义最大重试次数
95    pub fn with_max_retries(mut self, n: u32) -> Self {
96        self.max_retries = n;
97        self
98    }
99}
100
101/// 队列运行时
102///
103/// 管理 consumer lifecycle:启动一个消费循环任务,监听 `CancellationToken` 优雅退出。
104///
105/// ## 设计
106///
107/// - **不持有 `JoinHandle`**:消费任务由调用方持有,本结构仅提供启动入口
108/// - **ack 策略**:handler 返回 Ok 时自动 ack;Err 时不 ack(消息留在 in_flight)
109/// - **退出策略**:监听 `token.cancelled()`,当前正在处理的消息会等待完成
110///
111/// ## 用法
112///
113/// ```rust,ignore
114/// use sz_rust_core::runtime::queue::{QueueRuntime, QueueRuntimeConfig, QueueConsumer};
115/// use sz_orm_queue::{InMemoryQueue, MessageQueue, Message};
116/// use async_trait::async_trait;
117/// use std::sync::Arc;
118/// use tokio_util::sync::CancellationToken;
119///
120/// struct MyConsumer;
121/// #[async_trait]
122/// impl QueueConsumer for MyConsumer {
123///     async fn handle(&self, msg: &Message) -> Result<(), QueueConsumerError> {
124///         println!("got: {:?}", msg.payload);
125///         Ok(())
126///     }
127/// }
128///
129/// let queue = Arc::new(InMemoryQueue::new(1000));
130/// let runtime = QueueRuntime::new(
131///     QueueRuntimeConfig::new("orders"),
132///     queue,
133/// );
134/// let token = CancellationToken::new();
135/// let handle = runtime.start(Arc::new(MyConsumer), token.clone());
136/// // ... 业务运行 ...
137/// token.cancel();
138/// let _ = handle.await;
139/// ```
140pub struct QueueRuntime {
141    config: QueueRuntimeConfig,
142    queue: Arc<dyn MessageQueue>,
143}
144
145impl QueueRuntime {
146    /// 创建队列运行时
147    pub fn new(config: QueueRuntimeConfig, queue: Arc<dyn MessageQueue>) -> Self {
148        Self { config, queue }
149    }
150
151    /// 启动消费循环(返回 JoinHandle,调用方持有以控制 lifecycle)
152    ///
153    /// ## 行为
154    ///
155    /// 1. 每 `poll_interval_ms` 毫秒调用 `queue.consume(topic)` 拉取消息
156    /// 2. 收到消息后调用 `consumer.handle(&msg)`
157    /// 3. handler 返回 Ok → 自动 `queue.ack(msg.id)`
158    /// 4. handler 返回 Err → 跳过 ack(消息留在 in_flight)
159    /// 5. 监听 `token.cancelled()`,收到信号后停止拉取新消息
160    pub fn start<C>(
161        &self,
162        consumer: Arc<C>,
163        token: CancellationToken,
164    ) -> tokio::task::JoinHandle<()>
165    where
166        C: QueueConsumer + 'static,
167    {
168        let queue = self.queue.clone();
169        let topic = self.config.topic.clone();
170        let poll_interval = Duration::from_millis(self.config.poll_interval_ms.max(1));
171
172        tokio::spawn(async move {
173            loop {
174                tokio::select! {
175                    _ = token.cancelled() => break,
176                    consume_result = queue.consume(&topic) => {
177                        match consume_result {
178                            Ok(Some(message)) => {
179                                let msg_id = message.id.clone();
180                                match consumer.handle(&message).await {
181                                    Ok(()) => {
182                                        if let Err(e) = queue.ack(&msg_id).await {
183                                            tracing::warn!("ack failed for msg {}: {}", msg_id, e);
184                                        }
185                                    }
186                                    Err(e) => {
187                                        tracing::warn!(
188                                            "consumer handler failed for msg {}: {}",
189                                            msg_id,
190                                            e
191                                        );
192                                        // 不 ack,消息留在 in_flight
193                                    }
194                                }
195                            }
196                            Ok(None) => {
197                                // 队列为空,sleep 后重试
198                                tokio::time::sleep(poll_interval).await;
199                            }
200                            Err(e) => {
201                                tracing::error!("queue consume error: {}", e);
202                                tokio::time::sleep(poll_interval).await;
203                            }
204                        }
205                    }
206                }
207            }
208        })
209    }
210
211    /// 获取主题
212    pub fn topic(&self) -> &str {
213        &self.config.topic
214    }
215
216    /// 获取配置
217    pub fn config(&self) -> &QueueRuntimeConfig {
218        &self.config
219    }
220}
221
222#[cfg(test)]
223mod tests {
224    use super::*;
225    use sz_orm_queue::{InMemoryQueue, Message, MessageQueue};
226
227    /// 测试用消费者:记录所有处理过的消息 payload
228    struct RecordingConsumer {
229        payloads: Arc<parking_lot::Mutex<Vec<Vec<u8>>>>,
230    }
231
232    impl RecordingConsumer {
233        fn new() -> (Self, Arc<parking_lot::Mutex<Vec<Vec<u8>>>>) {
234            let payloads = Arc::new(parking_lot::Mutex::new(Vec::new()));
235            let consumer = Self {
236                payloads: payloads.clone(),
237            };
238            (consumer, payloads)
239        }
240    }
241
242    #[async_trait]
243    impl QueueConsumer for RecordingConsumer {
244        async fn handle(&self, message: &Message) -> Result<(), QueueConsumerError> {
245            self.payloads.lock().push(message.payload.clone());
246            Ok(())
247        }
248    }
249
250    /// 失败消费者:总是返回错误
251    struct FailingConsumer;
252
253    #[async_trait]
254    impl QueueConsumer for FailingConsumer {
255        async fn handle(&self, _message: &Message) -> Result<(), QueueConsumerError> {
256            Err(QueueConsumerError::Handler("always fail".to_string()))
257        }
258    }
259
260    /// 创建 InMemoryQueue 并 cast 为 `Arc<dyn MessageQueue>`
261    fn make_queue() -> Arc<dyn MessageQueue> {
262        Arc::new(InMemoryQueue::new())
263    }
264
265    #[test]
266    fn test_queue_runtime_config_default() {
267        let config = QueueRuntimeConfig::default();
268        assert_eq!(config.topic, "default");
269        assert_eq!(config.poll_interval_ms, 100);
270        assert_eq!(config.max_retries, 0);
271    }
272
273    #[test]
274    fn test_queue_runtime_config_builder() {
275        let config = QueueRuntimeConfig::new("orders")
276            .with_poll_interval(50)
277            .with_max_retries(3);
278        assert_eq!(config.topic, "orders");
279        assert_eq!(config.poll_interval_ms, 50);
280        assert_eq!(config.max_retries, 3);
281    }
282
283    #[test]
284    fn test_queue_runtime_topic_accessor() {
285        let queue = make_queue();
286        let runtime = QueueRuntime::new(QueueRuntimeConfig::new("test"), queue);
287        assert_eq!(runtime.topic(), "test");
288    }
289
290    #[test]
291    fn test_queue_runtime_config_accessor() {
292        let queue = make_queue();
293        let config = QueueRuntimeConfig::new("test").with_poll_interval(200);
294        let runtime = QueueRuntime::new(config, queue);
295        assert_eq!(runtime.config().poll_interval_ms, 200);
296    }
297
298    #[tokio::test]
299    async fn test_consumer_consumes_published_message() {
300        let queue = make_queue();
301        queue.publish("orders", b"hello").await.unwrap();
302
303        let (consumer, payloads) = RecordingConsumer::new();
304        let runtime = QueueRuntime::new(
305            QueueRuntimeConfig::new("orders").with_poll_interval(5),
306            queue.clone(),
307        );
308
309        let token = CancellationToken::new();
310        let handle = runtime.start(Arc::new(consumer), token.clone());
311
312        // 等待消费者处理消息
313        tokio::time::sleep(Duration::from_millis(100)).await;
314        token.cancel();
315        let _ = handle.await;
316
317        let recorded = payloads.lock().clone();
318        assert_eq!(recorded.len(), 1);
319        assert_eq!(recorded[0], b"hello");
320    }
321
322    #[tokio::test]
323    async fn test_consumer_acks_on_success() {
324        let queue = make_queue();
325        queue.publish("orders", b"msg1").await.unwrap();
326
327        let (consumer, _payloads) = RecordingConsumer::new();
328        let runtime = QueueRuntime::new(
329            QueueRuntimeConfig::new("orders").with_poll_interval(5),
330            queue.clone(),
331        );
332
333        let token = CancellationToken::new();
334        let handle = runtime.start(Arc::new(consumer), token.clone());
335
336        tokio::time::sleep(Duration::from_millis(100)).await;
337        token.cancel();
338        let _ = handle.await;
339
340        // 队列中应该没有 in_flight 消息(已 ack)
341        // 注:InMemoryQueue 的 ack 会从 in_flight 移除消息
342        // 此处不直接验证 in_flight 状态,因为 InMemoryQueue API 未暴露
343        // 通过再次 consume 返回 None 间接验证
344        let result = queue.consume("orders").await.unwrap();
345        assert!(result.is_none());
346    }
347
348    #[tokio::test]
349    async fn test_consumer_no_ack_on_failure() {
350        let queue = make_queue();
351        queue.publish("orders", b"msg1").await.unwrap();
352
353        let runtime = QueueRuntime::new(
354            QueueRuntimeConfig::new("orders").with_poll_interval(5),
355            queue.clone(),
356        );
357
358        let token = CancellationToken::new();
359        let handle = runtime.start(Arc::new(FailingConsumer), token.clone());
360
361        tokio::time::sleep(Duration::from_millis(100)).await;
362        token.cancel();
363        let _ = handle.await;
364
365        // FailingConsumer 不 ack,消息应该留在 in_flight
366        // 再次 consume 应该返回 None(因为 in_flight 中的消息不会被重新拉取)
367        // 但 in_flight 消息仍占位
368        // 注:此行为依赖 InMemoryQueue 的实现
369    }
370
371    #[tokio::test]
372    async fn test_consumer_stops_on_cancel() {
373        let queue = make_queue();
374        let (consumer, _payloads) = RecordingConsumer::new();
375        let runtime = QueueRuntime::new(
376            QueueRuntimeConfig::new("orders").with_poll_interval(5),
377            queue,
378        );
379
380        let token = CancellationToken::new();
381        let handle = runtime.start(Arc::new(consumer), token.clone());
382
383        // 立即 cancel
384        token.cancel();
385        // 等待任务退出(不应 panic)
386        let _ = tokio::time::timeout(Duration::from_millis(500), handle).await;
387    }
388
389    #[tokio::test]
390    async fn test_consumer_handles_empty_queue() {
391        let queue = make_queue();
392        let (consumer, payloads) = RecordingConsumer::new();
393        let runtime = QueueRuntime::new(
394            QueueRuntimeConfig::new("empty").with_poll_interval(5),
395            queue.clone(),
396        );
397
398        let token = CancellationToken::new();
399        let handle = runtime.start(Arc::new(consumer), token.clone());
400
401        // 等待一段时间,队列始终为空
402        tokio::time::sleep(Duration::from_millis(50)).await;
403        token.cancel();
404        let _ = handle.await;
405
406        // 没有消息被处理
407        assert!(payloads.lock().is_empty());
408    }
409
410    #[tokio::test]
411    async fn test_consumer_processes_multiple_messages() {
412        let queue = make_queue();
413        // 发布 3 条消息
414        queue.publish("orders", b"msg1").await.unwrap();
415        queue.publish("orders", b"msg2").await.unwrap();
416        queue.publish("orders", b"msg3").await.unwrap();
417
418        let (consumer, payloads) = RecordingConsumer::new();
419        let runtime = QueueRuntime::new(
420            QueueRuntimeConfig::new("orders").with_poll_interval(5),
421            queue,
422        );
423
424        let token = CancellationToken::new();
425        let handle = runtime.start(Arc::new(consumer), token.clone());
426
427        // 等待所有消息被处理
428        tokio::time::sleep(Duration::from_millis(200)).await;
429        token.cancel();
430        let _ = handle.await;
431
432        let recorded = payloads.lock().clone();
433        assert_eq!(recorded.len(), 3);
434    }
435
436    #[test]
437    fn test_queue_consumer_error_variants() {
438        let handler_err = QueueConsumerError::Handler("test".to_string());
439        let queue_err = QueueConsumerError::Queue("queue fail".to_string());
440        assert!(format!("{}", handler_err).contains("consumer error"));
441        assert!(format!("{}", queue_err).contains("queue error"));
442    }
443}