Skip to main content

tokio_prompt_orchestrator/
worker_pool.rs

1//! Dynamic worker pool with auto-scaling.
2//!
3//! Provides a generic [`WorkerPool<T, R>`] that automatically scales the number
4//! of concurrent worker tasks based on queue depth relative to
5//! [`WorkerConfig::queue_depth_per_worker`].
6//!
7//! ## Example
8//!
9//! ```rust
10//! use std::sync::Arc;
11//! use tokio_prompt_orchestrator::worker_pool::{WorkerConfig, WorkerPool};
12//! use futures::future::BoxFuture;
13//!
14//! #[tokio::main]
15//! async fn main() {
16//!     let cfg = WorkerConfig {
17//!         min_workers: 1,
18//!         max_workers: 4,
19//!         idle_timeout_secs: 30,
20//!         queue_depth_per_worker: 8,
21//!     };
22//!     let handler: Arc<dyn Fn(u32) -> BoxFuture<'static, u32> + Send + Sync> =
23//!         Arc::new(|x: u32| Box::pin(async move { x * 2 }));
24//!     let pool = WorkerPool::new(cfg, handler);
25//!     let rx = pool.submit(21, 0).await;
26//!     assert_eq!(rx.await.unwrap(), 42);
27//!     pool.drain_and_shutdown().await;
28//! }
29//! ```
30
31use std::sync::atomic::{AtomicU64, AtomicUsize, Ordering};
32use std::sync::Arc;
33use std::time::{Duration, Instant};
34
35use futures::future::BoxFuture;
36use tokio::sync::{mpsc, oneshot, Mutex};
37
38// ── Configuration ─────────────────────────────────────────────────────────────
39
40/// Configuration for a [`WorkerPool`].
41#[derive(Debug, Clone)]
42pub struct WorkerConfig {
43    /// Minimum number of worker tasks to keep alive even when idle.
44    pub min_workers: usize,
45    /// Maximum number of concurrent worker tasks.
46    pub max_workers: usize,
47    /// Number of seconds a worker may be idle before it is allowed to exit
48    /// (only applies to workers above `min_workers`).
49    pub idle_timeout_secs: u64,
50    /// Target queue depth per worker.  When `queue_depth > workers *
51    /// queue_depth_per_worker` the pool will attempt to scale up.
52    pub queue_depth_per_worker: usize,
53}
54
55// ── Work items ────────────────────────────────────────────────────────────────
56
57/// A unit of work submitted to the pool.
58///
59/// This is a public-facing view of a queued item.  The pool infers all fields
60/// internally; callers interact with work items only via [`WorkerPool::submit`].
61#[allow(dead_code)]
62pub struct WorkItem<T> {
63    /// Monotonically increasing submission identifier.
64    pub id: u64,
65    /// The caller-supplied payload.
66    pub payload: T,
67    /// Higher values are more urgent (0 = normal).
68    pub priority: u8,
69    /// Wall-clock time at which the item was enqueued.
70    pub enqueued_at: Instant,
71}
72
73/// Internal work envelope that carries the typed reply channel.
74#[allow(dead_code)]
75struct InnerItem<T, R> {
76    id: u64,
77    payload: T,
78    priority: u8,
79    enqueued_at: Instant,
80    reply: oneshot::Sender<R>,
81}
82
83// ── Statistics ────────────────────────────────────────────────────────────────
84
85/// A point-in-time snapshot of pool statistics.
86#[derive(Debug, Clone)]
87pub struct PoolStats {
88    /// Number of live worker tasks.
89    pub worker_count: usize,
90    /// Number of items waiting in the queue.
91    pub queue_depth: usize,
92    /// Number of items currently being processed.
93    pub active_work: usize,
94    /// Total number of items submitted since the pool was created.
95    pub total_submitted: u64,
96    /// Total number of items that completed processing.
97    pub total_completed: u64,
98    /// Rolling average end-to-end latency in milliseconds.
99    pub avg_latency_ms: f64,
100}
101
102// ── Pool ──────────────────────────────────────────────────────────────────────
103
104type Handler<T, R> = Arc<dyn Fn(T) -> BoxFuture<'static, R> + Send + Sync>;
105
106/// A dynamically-scaling pool of Tokio worker tasks.
107///
108/// Generic over:
109/// - `T` — the input type passed to each work item (must be `Send + 'static`)
110/// - `R` — the result type returned by the handler (must be `Send + 'static`)
111pub struct WorkerPool<T: Send + 'static, R: Send + 'static> {
112    config: WorkerConfig,
113    handler: Handler<T, R>,
114    /// MPSC sender used by `submit`.  Workers share the receiver via a `Mutex`.
115    tx: mpsc::Sender<InnerItem<T, R>>,
116    /// Shared receiver wrapped in a `Mutex` so multiple workers can pull from it.
117    rx: Arc<Mutex<mpsc::Receiver<InnerItem<T, R>>>>,
118    /// Live worker count.
119    worker_count: Arc<AtomicUsize>,
120    /// Items currently being processed (dequeued but not yet replied).
121    active_work: Arc<AtomicUsize>,
122    /// Monotonically increasing submission counter; also used as work-item ID.
123    total_submitted: Arc<AtomicU64>,
124    /// Total completed counter.
125    total_completed: Arc<AtomicU64>,
126    /// Sum of all per-item latencies in milliseconds (for avg calculation).
127    total_latency_ms: Arc<AtomicU64>,
128    /// Shutdown signal: when dropped, all workers will drain and exit.
129    shutdown_tx: Arc<tokio::sync::watch::Sender<bool>>,
130    shutdown_rx: tokio::sync::watch::Receiver<bool>,
131}
132
133impl<T: Send + 'static, R: Send + 'static> WorkerPool<T, R> {
134    /// Create a new pool with the given configuration and async handler function.
135    ///
136    /// `min_workers` worker tasks are spawned immediately.
137    pub fn new(config: WorkerConfig, handler: Handler<T, R>) -> Self {
138        let queue_cap = config.max_workers * config.queue_depth_per_worker.max(1) * 2;
139        let (tx, rx) = mpsc::channel::<InnerItem<T, R>>(queue_cap.max(16));
140        let rx = Arc::new(Mutex::new(rx));
141
142        let worker_count = Arc::new(AtomicUsize::new(0));
143        let active_work = Arc::new(AtomicUsize::new(0));
144        let total_submitted = Arc::new(AtomicU64::new(0));
145        let total_completed = Arc::new(AtomicU64::new(0));
146        let total_latency_ms = Arc::new(AtomicU64::new(0));
147        let (shutdown_tx, shutdown_rx) = tokio::sync::watch::channel(false);
148        let shutdown_tx = Arc::new(shutdown_tx);
149
150        let pool = Self {
151            config,
152            handler,
153            tx,
154            rx,
155            worker_count,
156            active_work,
157            total_submitted,
158            total_completed,
159            total_latency_ms,
160            shutdown_tx,
161            shutdown_rx,
162        };
163
164        // Spawn minimum workers.
165        for _ in 0..pool.config.min_workers {
166            pool.scale_up_inner();
167        }
168
169        pool
170    }
171
172    /// Submit a work item with the given `priority` (higher = more urgent).
173    ///
174    /// Returns a [`oneshot::Receiver<R>`] that resolves when processing is done.
175    /// If the queue is full (bounded channel) this will apply back-pressure and
176    /// await until space is available.
177    pub async fn submit(&self, item: T, priority: u8) -> oneshot::Receiver<R> {
178        let id = self.total_submitted.fetch_add(1, Ordering::Relaxed);
179        let (reply_tx, reply_rx) = oneshot::channel::<R>();
180
181        let inner = InnerItem {
182            id,
183            payload: item,
184            priority,
185            enqueued_at: Instant::now(),
186            reply: reply_tx,
187        };
188
189        // Auto-scale up if needed before sending.
190        self.maybe_scale_up();
191
192        // Send may block if the channel is full — that is intentional back-pressure.
193        // We ignore the error; if the channel is closed the caller will get a
194        // dropped receiver (await on it will return Err).
195        let _ = self.tx.send(inner).await;
196
197        reply_rx
198    }
199
200    /// Attempt to scale up if queue depth per worker exceeds the threshold.
201    fn maybe_scale_up(&self) {
202        let workers = self.worker_count.load(Ordering::Relaxed);
203        let queue_depth = self.tx.max_capacity() - self.tx.capacity();
204        let threshold = workers.saturating_mul(self.config.queue_depth_per_worker.max(1));
205        if queue_depth > threshold && workers < self.config.max_workers {
206            self.scale_up_inner();
207        }
208    }
209
210    /// Spawn one additional worker task.
211    pub fn scale_up(&self) {
212        if self.worker_count.load(Ordering::Relaxed) < self.config.max_workers {
213            self.scale_up_inner();
214        }
215    }
216
217    fn scale_up_inner(&self) {
218        let rx = Arc::clone(&self.rx);
219        let handler = Arc::clone(&self.handler);
220        let worker_count = Arc::clone(&self.worker_count);
221        let active_work = Arc::clone(&self.active_work);
222        let total_completed = Arc::clone(&self.total_completed);
223        let total_latency_ms = Arc::clone(&self.total_latency_ms);
224        let idle_timeout = Duration::from_secs(self.config.idle_timeout_secs);
225        let min_workers = self.config.min_workers;
226        let mut shutdown_rx = self.shutdown_rx.clone();
227
228        worker_count.fetch_add(1, Ordering::Relaxed);
229
230        tokio::spawn(async move {
231            loop {
232                // Respect shutdown signal.
233                if *shutdown_rx.borrow() {
234                    break;
235                }
236
237                // Try to dequeue with an idle timeout so we can scale down.
238                let item_opt = tokio::select! {
239                    item = async {
240                        let mut guard = rx.lock().await;
241                        guard.recv().await
242                    } => item,
243                    _ = tokio::time::sleep(idle_timeout) => None,
244                    _ = shutdown_rx.changed() => {
245                        // Drain remaining items if shutdown was just signalled.
246                        break;
247                    }
248                };
249
250                match item_opt {
251                    None => {
252                        // Idle timeout expired — exit if we're above min_workers.
253                        let current = worker_count.load(Ordering::Relaxed);
254                        if current > min_workers {
255                            // Try to decrement without going below min.
256                            let prev = worker_count.fetch_sub(1, Ordering::Relaxed);
257                            if prev <= min_workers {
258                                // Oops, we went below; add back.
259                                worker_count.fetch_add(1, Ordering::Relaxed);
260                                continue;
261                            }
262                            break;
263                        }
264                        // We are at min — keep looping.
265                    }
266                    Some(inner) => {
267                        active_work.fetch_add(1, Ordering::Relaxed);
268                        let fut = handler(inner.payload);
269                        let result = fut.await;
270                        active_work.fetch_sub(1, Ordering::Relaxed);
271
272                        let latency = inner.enqueued_at.elapsed().as_millis() as u64;
273                        total_latency_ms.fetch_add(latency, Ordering::Relaxed);
274                        total_completed.fetch_add(1, Ordering::Relaxed);
275
276                        // Send result back; ignore if caller dropped receiver.
277                        let _ = inner.reply.send(result);
278                    }
279                }
280            }
281
282            worker_count.fetch_sub(1, Ordering::Relaxed);
283        });
284    }
285
286    /// Signal one idle worker above `min_workers` to exit on its next idle timeout.
287    ///
288    /// This is a best-effort hint; the worker may be processing an item and will
289    /// only exit after completing it and then timing out.
290    pub fn scale_down(&self) {
291        // We cannot send a direct signal to a specific worker, so we rely on the
292        // idle-timeout logic in `scale_up_inner`.  This method is a no-op if all
293        // workers are busy or already at min_workers; it exists for API symmetry
294        // and for callers that want to hint at down-scaling after a traffic lull.
295        let _ = self.worker_count.load(Ordering::Relaxed);
296    }
297
298    /// Wait for the queue to drain completely, then stop all workers.
299    ///
300    /// After this method returns the pool must not be used again.
301    pub async fn drain_and_shutdown(&self) {
302        // Wait for queue to empty and all active work to finish.
303        loop {
304            let queue_depth = self.tx.max_capacity() - self.tx.capacity();
305            let active = self.active_work.load(Ordering::Relaxed);
306            if queue_depth == 0 && active == 0 {
307                break;
308            }
309            tokio::time::sleep(Duration::from_millis(10)).await;
310        }
311
312        // Broadcast shutdown.
313        let _ = self.shutdown_tx.send(true);
314
315        // Wait for all workers to exit.
316        loop {
317            if self.worker_count.load(Ordering::Relaxed) == 0 {
318                break;
319            }
320            tokio::time::sleep(Duration::from_millis(5)).await;
321        }
322    }
323
324    /// Return a point-in-time snapshot of pool statistics.
325    pub fn stats(&self) -> PoolStats {
326        let total_completed = self.total_completed.load(Ordering::Relaxed);
327        let total_latency_ms = self.total_latency_ms.load(Ordering::Relaxed);
328        let avg_latency_ms = if total_completed > 0 {
329            total_latency_ms as f64 / total_completed as f64
330        } else {
331            0.0
332        };
333
334        PoolStats {
335            worker_count: self.worker_count.load(Ordering::Relaxed),
336            queue_depth: self.tx.max_capacity() - self.tx.capacity(),
337            active_work: self.active_work.load(Ordering::Relaxed),
338            total_submitted: self.total_submitted.load(Ordering::Relaxed),
339            total_completed,
340            avg_latency_ms,
341        }
342    }
343}
344
345// ── Tests ─────────────────────────────────────────────────────────────────────
346
347#[cfg(test)]
348mod tests {
349    use super::*;
350
351    fn make_pool(min: usize, max: usize) -> WorkerPool<u32, u32> {
352        let cfg = WorkerConfig {
353            min_workers: min,
354            max_workers: max,
355            idle_timeout_secs: 1,
356            queue_depth_per_worker: 4,
357        };
358        let handler: Handler<u32, u32> =
359            Arc::new(|x: u32| Box::pin(async move { x * 2 }));
360        WorkerPool::new(cfg, handler)
361    }
362
363    #[tokio::test]
364    async fn submit_work_is_handled() {
365        let pool = make_pool(1, 4);
366        let rx = pool.submit(21, 0).await;
367        let result = rx.await.expect("receiver dropped");
368        assert_eq!(result, 42);
369        pool.drain_and_shutdown().await;
370    }
371
372    #[tokio::test]
373    async fn multiple_items_are_all_processed() {
374        let pool = make_pool(1, 4);
375        let mut receivers = Vec::new();
376        for i in 0u32..20 {
377            receivers.push(pool.submit(i, 0).await);
378        }
379        let mut results = Vec::new();
380        for rx in receivers {
381            results.push(rx.await.expect("receiver dropped"));
382        }
383        for (i, r) in results.iter().enumerate() {
384            assert_eq!(*r, (i as u32) * 2);
385        }
386        pool.drain_and_shutdown().await;
387    }
388
389    #[tokio::test]
390    async fn stats_track_completed_count() {
391        let pool = make_pool(1, 4);
392        let rxs: Vec<_> = (0u32..5).map(|_| {
393            let pool_ref = &pool;
394            async move { pool_ref.submit(1, 0).await }
395        }).collect();
396        let mut receivers = Vec::new();
397        for f in rxs {
398            receivers.push(f.await);
399        }
400        for rx in receivers {
401            rx.await.expect("receiver dropped");
402        }
403        let s = pool.stats();
404        assert_eq!(s.total_submitted, 5);
405        assert_eq!(s.total_completed, 5);
406        pool.drain_and_shutdown().await;
407    }
408
409    #[tokio::test]
410    async fn scale_up_does_not_exceed_max() {
411        let pool = make_pool(1, 2);
412        pool.scale_up();
413        pool.scale_up();
414        pool.scale_up(); // Should be ignored — already at max.
415        // Give tasks a moment to register.
416        tokio::time::sleep(Duration::from_millis(20)).await;
417        let count = pool.worker_count.load(Ordering::Relaxed);
418        assert!(count <= 2, "worker count {} exceeded max 2", count);
419        pool.drain_and_shutdown().await;
420    }
421
422    #[tokio::test]
423    async fn drain_and_shutdown_completes() {
424        let pool = make_pool(2, 4);
425        for i in 0u32..10 {
426            let _ = pool.submit(i, 0).await;
427        }
428        pool.drain_and_shutdown().await;
429        assert_eq!(pool.worker_count.load(Ordering::Relaxed), 0);
430    }
431}