Skip to main content

rust_zero_core/
queue.rs

1use crate::{
2    CounterVec, GaugeVec, HistogramOptions, HistogramVec, Metrics, MetricsError, VectorOptions,
3};
4use std::{
5    error::Error,
6    fmt,
7    future::Future,
8    sync::{
9        atomic::{AtomicUsize, Ordering},
10        Arc,
11    },
12    time::{Duration, Instant},
13};
14use tokio::{
15    sync::{broadcast, mpsc, watch, Mutex},
16    task::{JoinError, JoinSet},
17};
18
19/// Creates a bounded, backpressure-aware queue for asynchronous service handoff.
20pub fn bounded<T>(capacity: usize) -> (QueueSender<T>, QueueReceiver<T>) {
21    assert!(capacity > 0, "queue capacity must be greater than zero");
22    let (sender, receiver) = mpsc::channel(capacity);
23    (QueueSender { sender }, QueueReceiver { receiver })
24}
25
26/// Sending half of a bounded service queue.
27#[derive(Clone)]
28pub struct QueueSender<T> {
29    sender: mpsc::Sender<T>,
30}
31
32impl<T> QueueSender<T> {
33    pub async fn send(&self, value: T) -> Result<(), mpsc::error::SendError<T>> {
34        self.sender.send(value).await
35    }
36
37    pub fn try_send(&self, value: T) -> Result<(), mpsc::error::TrySendError<T>> {
38        self.sender.try_send(value)
39    }
40}
41
42/// Receiving half of a bounded service queue.
43pub struct QueueReceiver<T> {
44    receiver: mpsc::Receiver<T>,
45}
46
47impl<T> QueueReceiver<T> {
48    pub async fn recv(&mut self) -> Option<T> {
49        self.receiver.recv().await
50    }
51}
52
53/// Configuration for a supervised in-process consumer queue.
54#[derive(Debug, Clone, PartialEq, Eq)]
55pub struct QueueRuntimeConfig {
56    pub name: String,
57    pub capacity: usize,
58    pub workers: usize,
59    pub event_capacity: usize,
60    pub shutdown_timeout: Duration,
61}
62
63impl QueueRuntimeConfig {
64    pub fn new(name: impl Into<String>, capacity: usize, workers: usize) -> Self {
65        Self {
66            name: name.into(),
67            capacity,
68            workers,
69            event_capacity: 256,
70            shutdown_timeout: Duration::from_secs(30),
71        }
72    }
73
74    pub fn validate(&self) -> Result<(), QueueConfigError> {
75        if self.name.trim().is_empty() {
76            return Err(QueueConfigError::EmptyName);
77        }
78        if self.capacity == 0 {
79            return Err(QueueConfigError::ZeroCapacity);
80        }
81        if self.workers == 0 {
82            return Err(QueueConfigError::ZeroWorkers);
83        }
84        if self.event_capacity == 0 {
85            return Err(QueueConfigError::ZeroEventCapacity);
86        }
87        if self.shutdown_timeout.is_zero() {
88            return Err(QueueConfigError::ZeroShutdownTimeout);
89        }
90        Ok(())
91    }
92}
93
94#[derive(Debug, Clone, PartialEq, Eq)]
95pub enum QueueConfigError {
96    EmptyName,
97    ZeroCapacity,
98    ZeroWorkers,
99    ZeroEventCapacity,
100    ZeroShutdownTimeout,
101}
102
103impl fmt::Display for QueueConfigError {
104    fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
105        formatter.write_str(match self {
106            Self::EmptyName => "queue name must not be empty",
107            Self::ZeroCapacity => "queue capacity must be greater than zero",
108            Self::ZeroWorkers => "queue worker count must be greater than zero",
109            Self::ZeroEventCapacity => "queue event capacity must be greater than zero",
110            Self::ZeroShutdownTimeout => "queue shutdown timeout must be greater than zero",
111        })
112    }
113}
114
115impl Error for QueueConfigError {}
116
117/// Lifecycle and processing notifications emitted by a running queue.
118#[derive(Debug, Clone, PartialEq)]
119pub enum QueueEvent {
120    Queued,
121    Started {
122        worker: usize,
123    },
124    Succeeded {
125        worker: usize,
126        elapsed: Duration,
127    },
128    Failed {
129        worker: usize,
130        elapsed: Duration,
131        message: String,
132    },
133    Paused,
134    Resumed,
135    Shutdown,
136}
137
138/// Prometheus instruments shared by queue producers and consumers.
139#[derive(Clone)]
140pub struct QueueMetrics {
141    messages: CounterVec,
142    inflight: GaugeVec,
143    duration: HistogramVec,
144}
145
146impl QueueMetrics {
147    pub fn new(registry: &Metrics, namespace: impl Into<String>) -> Result<Self, MetricsError> {
148        let namespace = namespace.into();
149        let messages = registry.counter_vec(
150            VectorOptions::new("messages_total", "Queue messages by queue and outcome")
151                .with_namespace(namespace.clone())
152                .with_subsystem("queue")
153                .with_labels(["queue", "outcome"]),
154        )?;
155        let inflight = registry.gauge_vec(
156            VectorOptions::new("inflight", "Queue messages currently being processed")
157                .with_namespace(namespace.clone())
158                .with_subsystem("queue")
159                .with_labels(["queue"]),
160        )?;
161        let duration = registry.histogram_vec(
162            HistogramOptions::new("processing_duration_seconds", "Queue processing latency")
163                .with_vector_options(
164                    VectorOptions::new("processing_duration_seconds", "Queue processing latency")
165                        .with_namespace(namespace)
166                        .with_subsystem("queue")
167                        .with_labels(["queue", "outcome"]),
168                ),
169        )?;
170        Ok(Self {
171            messages,
172            inflight,
173            duration,
174        })
175    }
176
177    fn message(&self, queue: &str, outcome: &str) {
178        let _ = self.messages.inc(&[queue, outcome]);
179    }
180
181    fn begin(&self, queue: &str) {
182        let _ = self.inflight.inc(&[queue]);
183    }
184
185    fn finish(&self, queue: &str, outcome: &str, elapsed: Duration) {
186        let _ = self.inflight.add(-1.0, &[queue]);
187        let _ = self.message_and_duration(queue, outcome, elapsed);
188    }
189
190    fn message_and_duration(
191        &self,
192        queue: &str,
193        outcome: &str,
194        elapsed: Duration,
195    ) -> Result<(), MetricsError> {
196        self.messages.inc(&[queue, outcome])?;
197        self.duration
198            .observe(elapsed.as_secs_f64(), &[queue, outcome])
199    }
200}
201
202#[derive(Debug, Clone, Copy, PartialEq, Eq)]
203enum QueueState {
204    Running,
205    Paused,
206    Shutdown,
207}
208
209/// A clonable producer for a supervised queue.
210pub struct QueueProducer<T> {
211    name: Arc<str>,
212    sender: mpsc::Sender<T>,
213    events: broadcast::Sender<QueueEvent>,
214    metrics: Option<QueueMetrics>,
215}
216
217impl<T> Clone for QueueProducer<T> {
218    fn clone(&self) -> Self {
219        Self {
220            name: Arc::clone(&self.name),
221            sender: self.sender.clone(),
222            events: self.events.clone(),
223            metrics: self.metrics.clone(),
224        }
225    }
226}
227
228impl<T> QueueProducer<T> {
229    pub fn name(&self) -> &str {
230        &self.name
231    }
232
233    pub fn capacity(&self) -> usize {
234        self.sender.capacity()
235    }
236
237    pub fn subscribe(&self) -> broadcast::Receiver<QueueEvent> {
238        self.events.subscribe()
239    }
240
241    pub async fn push(&self, value: T) -> Result<(), mpsc::error::SendError<T>> {
242        self.sender.send(value).await?;
243        self.record_queued();
244        Ok(())
245    }
246
247    pub fn try_push(&self, value: T) -> Result<(), mpsc::error::TrySendError<T>> {
248        self.sender.try_send(value)?;
249        self.record_queued();
250        Ok(())
251    }
252
253    fn record_queued(&self) {
254        let _ = self.events.send(QueueEvent::Queued);
255        if let Some(metrics) = &self.metrics {
256            metrics.message(&self.name, "queued");
257        }
258    }
259}
260
261/// Starts a queue and supervises a configurable pool of consumers.
262pub struct QueueRuntime;
263
264impl QueueRuntime {
265    pub fn start<T, H, Fut, E>(
266        config: QueueRuntimeConfig,
267        handler: H,
268    ) -> Result<(QueueProducer<T>, RunningQueue), QueueConfigError>
269    where
270        T: Send + 'static,
271        H: Fn(T) -> Fut + Send + Sync + 'static,
272        Fut: Future<Output = Result<(), E>> + Send + 'static,
273        E: fmt::Display + Send + 'static,
274    {
275        Self::start_inner(config, None, handler)
276    }
277
278    pub fn start_with_metrics<T, H, Fut, E>(
279        config: QueueRuntimeConfig,
280        metrics: QueueMetrics,
281        handler: H,
282    ) -> Result<(QueueProducer<T>, RunningQueue), QueueConfigError>
283    where
284        T: Send + 'static,
285        H: Fn(T) -> Fut + Send + Sync + 'static,
286        Fut: Future<Output = Result<(), E>> + Send + 'static,
287        E: fmt::Display + Send + 'static,
288    {
289        Self::start_inner(config, Some(metrics), handler)
290    }
291
292    fn start_inner<T, H, Fut, E>(
293        config: QueueRuntimeConfig,
294        metrics: Option<QueueMetrics>,
295        handler: H,
296    ) -> Result<(QueueProducer<T>, RunningQueue), QueueConfigError>
297    where
298        T: Send + 'static,
299        H: Fn(T) -> Fut + Send + Sync + 'static,
300        Fut: Future<Output = Result<(), E>> + Send + 'static,
301        E: fmt::Display + Send + 'static,
302    {
303        config.validate()?;
304        let (sender, receiver) = mpsc::channel(config.capacity);
305        let receiver = Arc::new(Mutex::new(receiver));
306        let handler = Arc::new(handler);
307        let (state_sender, state_receiver) = watch::channel(QueueState::Running);
308        let (events, _) = broadcast::channel(config.event_capacity);
309        let name: Arc<str> = Arc::from(config.name.as_str());
310        let mut tasks = JoinSet::new();
311
312        for worker in 0..config.workers {
313            let receiver = Arc::clone(&receiver);
314            let handler = Arc::clone(&handler);
315            let state = state_receiver.clone();
316            let events = events.clone();
317            let metrics = metrics.clone();
318            let name = Arc::clone(&name);
319            tasks.spawn(async move {
320                worker_loop(worker, name, receiver, state, events, metrics, handler).await
321            });
322        }
323
324        Ok((
325            QueueProducer {
326                name,
327                sender,
328                events: events.clone(),
329                metrics,
330            },
331            RunningQueue {
332                state: state_sender,
333                events,
334                tasks,
335                workers: config.workers,
336                shutdown_timeout: config.shutdown_timeout,
337            },
338        ))
339    }
340}
341
342async fn worker_loop<T, H, Fut, E>(
343    worker: usize,
344    name: Arc<str>,
345    receiver: Arc<Mutex<mpsc::Receiver<T>>>,
346    mut state: watch::Receiver<QueueState>,
347    events: broadcast::Sender<QueueEvent>,
348    metrics: Option<QueueMetrics>,
349    handler: Arc<H>,
350) where
351    T: Send + 'static,
352    H: Fn(T) -> Fut + Send + Sync + 'static,
353    Fut: Future<Output = Result<(), E>> + Send + 'static,
354    E: fmt::Display + Send + 'static,
355{
356    loop {
357        let current_state = *state.borrow();
358        match current_state {
359            QueueState::Shutdown => return,
360            QueueState::Paused => {
361                if state.changed().await.is_err() {
362                    return;
363                }
364                continue;
365            }
366            QueueState::Running => {}
367        }
368
369        let next = async {
370            let mut receiver = receiver.lock().await;
371            receiver.recv().await
372        };
373        let value = tokio::select! {
374            biased;
375            changed = state.changed() => {
376                if changed.is_err() { return; }
377                continue;
378            }
379            value = next => value,
380        };
381        let Some(value) = value else {
382            return;
383        };
384
385        let _ = events.send(QueueEvent::Started { worker });
386        if let Some(metrics) = &metrics {
387            metrics.begin(&name);
388        }
389        let started = Instant::now();
390        match handler(value).await {
391            Ok(()) => {
392                let elapsed = started.elapsed();
393                if let Some(metrics) = &metrics {
394                    metrics.finish(&name, "succeeded", elapsed);
395                }
396                let _ = events.send(QueueEvent::Succeeded { worker, elapsed });
397            }
398            Err(error) => {
399                let elapsed = started.elapsed();
400                if let Some(metrics) = &metrics {
401                    metrics.finish(&name, "failed", elapsed);
402                }
403                let _ = events.send(QueueEvent::Failed {
404                    worker,
405                    elapsed,
406                    message: error.to_string(),
407                });
408            }
409        }
410    }
411}
412
413/// Control and supervision handle for a running queue.
414pub struct RunningQueue {
415    state: watch::Sender<QueueState>,
416    events: broadcast::Sender<QueueEvent>,
417    tasks: JoinSet<()>,
418    workers: usize,
419    shutdown_timeout: Duration,
420}
421
422impl RunningQueue {
423    pub fn subscribe(&self) -> broadcast::Receiver<QueueEvent> {
424        self.events.subscribe()
425    }
426
427    pub fn is_paused(&self) -> bool {
428        *self.state.borrow() == QueueState::Paused
429    }
430
431    pub fn pause(&self) {
432        if *self.state.borrow() == QueueState::Running {
433            let _ = self.state.send(QueueState::Paused);
434            let _ = self.events.send(QueueEvent::Paused);
435        }
436    }
437
438    pub fn resume(&self) {
439        if *self.state.borrow() == QueueState::Paused {
440            let _ = self.state.send(QueueState::Running);
441            let _ = self.events.send(QueueEvent::Resumed);
442        }
443    }
444
445    /// Requests shutdown and waits for all workers to stop within the configured deadline.
446    pub async fn shutdown(mut self) -> Result<(), QueueRuntimeError> {
447        let _ = self.state.send(QueueState::Shutdown);
448        let _ = self.events.send(QueueEvent::Shutdown);
449        self.drain().await
450    }
451
452    /// Waits for workers to exit because all producers were dropped.
453    ///
454    /// A panic in any worker is surfaced and causes the remaining workers to be stopped.
455    pub async fn wait(mut self) -> Result<(), QueueRuntimeError> {
456        while let Some(result) = self.tasks.join_next().await {
457            self.workers = self.workers.saturating_sub(1);
458            if let Err(error) = result {
459                let _ = self.state.send(QueueState::Shutdown);
460                self.tasks.abort_all();
461                return Err(join_error(error));
462            }
463        }
464        Ok(())
465    }
466
467    async fn drain(&mut self) -> Result<(), QueueRuntimeError> {
468        let drain = async {
469            while let Some(result) = self.tasks.join_next().await {
470                self.workers = self.workers.saturating_sub(1);
471                result.map_err(join_error)?;
472            }
473            Ok(())
474        };
475        match tokio::time::timeout(self.shutdown_timeout, drain).await {
476            Ok(result) => result,
477            Err(_) => {
478                self.tasks.abort_all();
479                Err(QueueRuntimeError::ShutdownTimeout {
480                    remaining: self.workers,
481                })
482            }
483        }
484    }
485}
486
487#[derive(Debug, Clone, PartialEq, Eq)]
488pub enum QueueRuntimeError {
489    WorkerPanicked(String),
490    ShutdownTimeout { remaining: usize },
491}
492
493impl fmt::Display for QueueRuntimeError {
494    fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
495        match self {
496            Self::WorkerPanicked(message) => write!(formatter, "queue worker panicked: {message}"),
497            Self::ShutdownTimeout { remaining } => write!(
498                formatter,
499                "queue shutdown timed out with {remaining} worker(s) remaining"
500            ),
501        }
502    }
503}
504
505impl Error for QueueRuntimeError {}
506
507fn join_error(error: JoinError) -> QueueRuntimeError {
508    QueueRuntimeError::WorkerPanicked(error.to_string())
509}
510
511/// Round-robin, non-blocking pusher that fails over when a queue is full or closed.
512pub struct BalancedPusher<T> {
513    producers: Arc<[QueueProducer<T>]>,
514    next: AtomicUsize,
515}
516
517impl<T> BalancedPusher<T> {
518    pub fn new(
519        producers: impl IntoIterator<Item = QueueProducer<T>>,
520    ) -> Result<Self, PusherConfigError> {
521        let producers: Arc<[QueueProducer<T>]> = producers.into_iter().collect::<Vec<_>>().into();
522        if producers.is_empty() {
523            return Err(PusherConfigError::Empty);
524        }
525        Ok(Self {
526            producers,
527            next: AtomicUsize::new(0),
528        })
529    }
530
531    pub fn try_push(&self, mut value: T) -> Result<usize, PushError<T>> {
532        let start = self.next.fetch_add(1, Ordering::Relaxed) % self.producers.len();
533        for offset in 0..self.producers.len() {
534            let index = (start + offset) % self.producers.len();
535            match self.producers[index].try_push(value) {
536                Ok(()) => return Ok(index),
537                Err(mpsc::error::TrySendError::Full(returned))
538                | Err(mpsc::error::TrySendError::Closed(returned)) => value = returned,
539            }
540        }
541        Err(PushError { value })
542    }
543}
544
545/// Pusher that offers a clone of each message to every configured queue.
546pub struct FanoutPusher<T> {
547    producers: Arc<[QueueProducer<T>]>,
548}
549
550impl<T> FanoutPusher<T> {
551    pub fn new(
552        producers: impl IntoIterator<Item = QueueProducer<T>>,
553    ) -> Result<Self, PusherConfigError> {
554        let producers: Arc<[QueueProducer<T>]> = producers.into_iter().collect::<Vec<_>>().into();
555        if producers.is_empty() {
556            return Err(PusherConfigError::Empty);
557        }
558        Ok(Self { producers })
559    }
560}
561
562impl<T: Clone> FanoutPusher<T> {
563    /// Returns the indexes of queues that were full or closed.
564    pub fn try_push(&self, value: T) -> Vec<usize> {
565        self.producers
566            .iter()
567            .enumerate()
568            .filter_map(|(index, producer)| producer.try_push(value.clone()).err().map(|_| index))
569            .collect()
570    }
571}
572
573#[derive(Debug, Clone, PartialEq, Eq)]
574pub enum PusherConfigError {
575    Empty,
576}
577
578impl fmt::Display for PusherConfigError {
579    fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
580        formatter.write_str("at least one queue producer is required")
581    }
582}
583
584impl Error for PusherConfigError {}
585
586#[derive(Debug)]
587pub struct PushError<T> {
588    pub value: T,
589}
590
591impl<T> fmt::Display for PushError<T> {
592    fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
593        formatter.write_str("all queue producers are full or closed")
594    }
595}
596
597impl<T: fmt::Debug> Error for PushError<T> {}
598
599#[cfg(test)]
600mod tests {
601    use super::*;
602    use std::sync::atomic::{AtomicUsize, Ordering};
603    use tokio::sync::Notify;
604
605    #[tokio::test]
606    async fn preserves_fifo_order_and_backpressure() {
607        let (sender, mut receiver) = bounded(1);
608        sender.try_send(1).unwrap();
609        assert!(sender.try_send(2).is_err());
610        assert_eq!(receiver.recv().await, Some(1));
611        sender.send(2).await.unwrap();
612        assert_eq!(receiver.recv().await, Some(2));
613    }
614
615    #[tokio::test]
616    async fn pauses_resumes_reports_failures_and_records_metrics() {
617        let registry = Metrics::new();
618        let metrics = QueueMetrics::new(&registry, "test").unwrap();
619        let processed = Arc::new(AtomicUsize::new(0));
620        let (producer, running) =
621            QueueRuntime::start_with_metrics(QueueRuntimeConfig::new("emails", 4, 2), metrics, {
622                let processed = Arc::clone(&processed);
623                move |value: usize| {
624                    let processed = Arc::clone(&processed);
625                    async move {
626                        processed.fetch_add(1, Ordering::SeqCst);
627                        if value == 2 {
628                            Err("rejected")
629                        } else {
630                            Ok(())
631                        }
632                    }
633                }
634            })
635            .unwrap();
636        let mut events = running.subscribe();
637        running.pause();
638        producer.push(1).await.unwrap();
639        tokio::task::yield_now().await;
640        assert_eq!(processed.load(Ordering::SeqCst), 0);
641        running.resume();
642        producer.push(2).await.unwrap();
643
644        tokio::time::timeout(Duration::from_secs(1), async {
645            while processed.load(Ordering::SeqCst) < 2 {
646                tokio::task::yield_now().await;
647            }
648        })
649        .await
650        .unwrap();
651        let mut saw_failure = false;
652        while let Ok(event) = events.try_recv() {
653            saw_failure |= matches!(event, QueueEvent::Failed { .. });
654        }
655        assert!(saw_failure);
656        running.shutdown().await.unwrap();
657
658        let rendered = registry.render();
659        assert!(rendered.contains("test_queue_messages_total"));
660        assert!(rendered.contains("outcome=\"succeeded\""));
661        assert!(rendered.contains("outcome=\"failed\""));
662    }
663
664    #[tokio::test]
665    async fn shutdown_is_bounded_when_a_handler_does_not_finish() {
666        let blocked = Arc::new(Notify::new());
667        let mut config = QueueRuntimeConfig::new("blocked", 1, 1);
668        config.shutdown_timeout = Duration::from_millis(10);
669        let (producer, running) = QueueRuntime::start(config, {
670            let blocked = Arc::clone(&blocked);
671            move |_: ()| {
672                let blocked = Arc::clone(&blocked);
673                async move {
674                    blocked.notified().await;
675                    Ok::<_, &'static str>(())
676                }
677            }
678        })
679        .unwrap();
680        producer.push(()).await.unwrap();
681        tokio::task::yield_now().await;
682        assert_eq!(
683            running.shutdown().await,
684            Err(QueueRuntimeError::ShutdownTimeout { remaining: 1 })
685        );
686    }
687
688    #[tokio::test]
689    async fn balanced_failover_and_fanout_route_messages() {
690        let (first, first_running) =
691            QueueRuntime::start(QueueRuntimeConfig::new("first", 1, 1), |_: usize| async {
692                Ok::<_, &'static str>(())
693            })
694            .unwrap();
695        let (second, second_running) =
696            QueueRuntime::start(QueueRuntimeConfig::new("second", 1, 1), |_: usize| async {
697                Ok::<_, &'static str>(())
698            })
699            .unwrap();
700
701        let balanced = BalancedPusher::new([first.clone(), second.clone()]).unwrap();
702        assert_eq!(balanced.try_push(1).unwrap(), 0);
703        assert_eq!(balanced.try_push(2).unwrap(), 1);
704
705        let fanout = FanoutPusher::new([first, second]).unwrap();
706        tokio::task::yield_now().await;
707        assert!(fanout.try_push(3).is_empty());
708
709        first_running.shutdown().await.unwrap();
710        second_running.shutdown().await.unwrap();
711    }
712}