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
19pub 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#[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
42pub 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#[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#[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#[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
209pub 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
261pub 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
413pub 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 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 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
511pub 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
545pub 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 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(®istry, "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}