tokio_prompt_orchestrator/
worker_pool.rs1use 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#[derive(Debug, Clone)]
42pub struct WorkerConfig {
43 pub min_workers: usize,
45 pub max_workers: usize,
47 pub idle_timeout_secs: u64,
50 pub queue_depth_per_worker: usize,
53}
54
55#[allow(dead_code)]
62pub struct WorkItem<T> {
63 pub id: u64,
65 pub payload: T,
67 pub priority: u8,
69 pub enqueued_at: Instant,
71}
72
73#[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#[derive(Debug, Clone)]
87pub struct PoolStats {
88 pub worker_count: usize,
90 pub queue_depth: usize,
92 pub active_work: usize,
94 pub total_submitted: u64,
96 pub total_completed: u64,
98 pub avg_latency_ms: f64,
100}
101
102type Handler<T, R> = Arc<dyn Fn(T) -> BoxFuture<'static, R> + Send + Sync>;
105
106pub struct WorkerPool<T: Send + 'static, R: Send + 'static> {
112 config: WorkerConfig,
113 handler: Handler<T, R>,
114 tx: mpsc::Sender<InnerItem<T, R>>,
116 rx: Arc<Mutex<mpsc::Receiver<InnerItem<T, R>>>>,
118 worker_count: Arc<AtomicUsize>,
120 active_work: Arc<AtomicUsize>,
122 total_submitted: Arc<AtomicU64>,
124 total_completed: Arc<AtomicU64>,
126 total_latency_ms: Arc<AtomicU64>,
128 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 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 for _ in 0..pool.config.min_workers {
166 pool.scale_up_inner();
167 }
168
169 pool
170 }
171
172 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 self.maybe_scale_up();
191
192 let _ = self.tx.send(inner).await;
196
197 reply_rx
198 }
199
200 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 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 if *shutdown_rx.borrow() {
234 break;
235 }
236
237 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 break;
247 }
248 };
249
250 match item_opt {
251 None => {
252 let current = worker_count.load(Ordering::Relaxed);
254 if current > min_workers {
255 let prev = worker_count.fetch_sub(1, Ordering::Relaxed);
257 if prev <= min_workers {
258 worker_count.fetch_add(1, Ordering::Relaxed);
260 continue;
261 }
262 break;
263 }
264 }
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 let _ = inner.reply.send(result);
278 }
279 }
280 }
281
282 worker_count.fetch_sub(1, Ordering::Relaxed);
283 });
284 }
285
286 pub fn scale_down(&self) {
291 let _ = self.worker_count.load(Ordering::Relaxed);
296 }
297
298 pub async fn drain_and_shutdown(&self) {
302 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 let _ = self.shutdown_tx.send(true);
314
315 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 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#[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(); 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}