Skip to main content

rskit_worker/
pool.rs

1use std::collections::VecDeque;
2use std::sync::Arc;
3use std::time::Duration;
4
5use parking_lot::Mutex;
6use tokio::sync::{Notify, Semaphore, broadcast, mpsc, oneshot};
7use tokio::task::{JoinHandle, JoinSet};
8use tokio_util::sync::CancellationToken;
9use uuid::Uuid;
10
11use rskit_errors::{AppError, AppResult, ErrorCode};
12
13use crate::dispatch::DispatchStrategy;
14use crate::event::Event;
15use crate::handler::Handler;
16use crate::task::TaskHandle;
17
18/// Statistics snapshot for the pool.
19#[derive(Debug, Clone)]
20pub struct PoolStats {
21    /// Human-readable name of the pool.
22    pub name: String,
23    /// Number of tasks currently executing.
24    pub running: usize,
25    /// Maximum concurrent tasks the pool allows.
26    pub capacity: usize,
27}
28
29/// Overflow behavior applied when the submission queue is full.
30#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
31#[non_exhaustive]
32pub enum OverflowPolicy {
33    /// Wait until queue capacity becomes available.
34    #[default]
35    Block,
36    /// Reject the new submission immediately.
37    Reject,
38    /// Drop the oldest queued task and enqueue the new submission.
39    DropOldest,
40}
41
42/// Configuration for a [`Pool`].
43pub struct PoolConfig {
44    /// Human-readable name used in tracing.
45    pub name: String,
46    /// Maximum concurrent tasks (semaphore permits).
47    pub size: usize,
48    /// Capacity of the internal submit queue.
49    pub queue_size: usize,
50    /// Broadcast channel capacity for events per task.
51    pub event_buffer: usize,
52    /// Grace period given to in-flight tasks on shutdown.
53    pub grace_period: Duration,
54    /// Dispatch strategy (reserved for future multi-queue extensions).
55    pub dispatch: DispatchStrategy,
56    /// Queue overflow behavior.
57    pub overflow_policy: OverflowPolicy,
58}
59
60impl Default for PoolConfig {
61    fn default() -> Self {
62        Self {
63            name: "pool".into(),
64            size: available_parallelism(),
65            queue_size: 256,
66            event_buffer: 64,
67            grace_period: Duration::from_secs(30),
68            dispatch: DispatchStrategy::RoundRobin,
69            overflow_policy: OverflowPolicy::Block,
70        }
71    }
72}
73
74impl PoolConfig {
75    /// Create a named pool configuration with sensible defaults.
76    #[must_use]
77    pub fn new(name: impl Into<String>) -> Self {
78        Self {
79            name: name.into(),
80            ..Default::default()
81        }
82    }
83
84    /// Set the maximum number of concurrent tasks. Values below 1 are clamped
85    /// to 1 inside `Pool::new` (with a tracing warning), since a zero-sized
86    /// pool can never execute tasks.
87    #[must_use]
88    pub fn with_size(mut self, size: usize) -> Self {
89        self.size = size;
90        self
91    }
92
93    /// Set the capacity of the internal submit queue.
94    #[must_use]
95    pub fn with_queue_size(mut self, queue_size: usize) -> Self {
96        self.queue_size = queue_size;
97        self
98    }
99
100    /// Set the grace period given to in-flight tasks during shutdown.
101    #[must_use]
102    pub fn with_grace_period(mut self, d: Duration) -> Self {
103        self.grace_period = d;
104        self
105    }
106
107    /// Set the queue overflow policy.
108    #[must_use]
109    pub fn with_overflow_policy(mut self, overflow_policy: OverflowPolicy) -> Self {
110        self.overflow_policy = overflow_policy;
111        self
112    }
113}
114
115fn available_parallelism() -> usize {
116    std::thread::available_parallelism()
117        .map(|n| n.get())
118        .unwrap_or(4)
119}
120
121struct Envelope<I, O: Clone + Send + 'static> {
122    id: Uuid,
123    input: I,
124    events_bcast: broadcast::Sender<Event<O>>,
125    result_tx: oneshot::Sender<AppResult<O>>,
126    cancel: CancellationToken,
127    event_buffer: usize,
128}
129
130struct QueueInner<T> {
131    items: VecDeque<T>,
132    capacity: usize,
133    closed: bool,
134}
135
136struct QueueState<T> {
137    inner: Mutex<QueueInner<T>>,
138    not_empty: Notify,
139    not_full: Notify,
140}
141
142struct SubmitQueue<T> {
143    state: Arc<QueueState<T>>,
144}
145
146enum PushRejectError<T> {
147    Closed(T),
148    Full(T),
149}
150
151struct QueueReceiver<T> {
152    state: Arc<QueueState<T>>,
153}
154
155impl<T> SubmitQueue<T> {
156    fn new(capacity: usize) -> (Self, QueueReceiver<T>) {
157        let state = Arc::new(QueueState {
158            inner: Mutex::new(QueueInner {
159                items: VecDeque::with_capacity(capacity.max(1)),
160                capacity: capacity.max(1),
161                closed: false,
162            }),
163            not_empty: Notify::new(),
164            not_full: Notify::new(),
165        });
166        (
167            Self {
168                state: Arc::clone(&state),
169            },
170            QueueReceiver { state },
171        )
172    }
173
174    async fn push_block(&self, item: T) -> Result<(), T> {
175        let mut item = Some(item);
176        loop {
177            let notified = {
178                let mut inner = self.state.inner.lock();
179                if inner.closed {
180                    return Err(item.take().unwrap_or_else(|| unreachable!("item present")));
181                }
182                if inner.items.len() < inner.capacity {
183                    inner
184                        .items
185                        .push_back(item.take().unwrap_or_else(|| unreachable!("item present")));
186                    self.state.not_empty.notify_one();
187                    return Ok(());
188                }
189                self.state.not_full.notified()
190            };
191            notified.await;
192        }
193    }
194
195    fn push_reject(&self, item: T) -> Result<(), PushRejectError<T>> {
196        let mut inner = self.state.inner.lock();
197        if inner.closed {
198            return Err(PushRejectError::Closed(item));
199        }
200        if inner.items.len() >= inner.capacity {
201            return Err(PushRejectError::Full(item));
202        }
203        inner.items.push_back(item);
204        self.state.not_empty.notify_one();
205        Ok(())
206    }
207
208    fn push_drop_oldest(&self, item: T) -> Result<Option<T>, T> {
209        let mut inner = self.state.inner.lock();
210        if inner.closed {
211            return Err(item);
212        }
213        let dropped = if inner.items.len() >= inner.capacity {
214            inner.items.pop_front()
215        } else {
216            None
217        };
218        inner.items.push_back(item);
219        self.state.not_empty.notify_one();
220        Ok(dropped)
221    }
222
223    fn close(&self) {
224        let mut inner = self.state.inner.lock();
225        inner.closed = true;
226        self.state.not_empty.notify_waiters();
227        self.state.not_full.notify_waiters();
228    }
229}
230
231impl<T> Clone for SubmitQueue<T> {
232    fn clone(&self) -> Self {
233        Self {
234            state: Arc::clone(&self.state),
235        }
236    }
237}
238
239impl<T> QueueReceiver<T> {
240    async fn recv(&self) -> Option<T> {
241        loop {
242            let notified = {
243                let mut inner = self.state.inner.lock();
244                if let Some(item) = inner.items.pop_front() {
245                    self.state.not_full.notify_one();
246                    return Some(item);
247                }
248                if inner.closed {
249                    return None;
250                }
251                self.state.not_empty.notified()
252            };
253            notified.await;
254        }
255    }
256}
257
258/// A bounded async worker pool.
259pub struct Pool<I, O>
260where
261    I: Send + 'static,
262    O: Send + Clone + 'static,
263{
264    name: String,
265    queue: SubmitQueue<Envelope<I, O>>,
266    semaphore: Arc<Semaphore>,
267    capacity: usize,
268    event_buffer: usize,
269    overflow_policy: OverflowPolicy,
270    grace_period: Duration,
271    shutdown: CancellationToken,
272    runner: Option<JoinHandle<()>>,
273}
274
275impl<I, O> Pool<I, O>
276where
277    I: Send + 'static,
278    O: Send + Clone + 'static,
279{
280    /// Create a new pool backed by `handler`.
281    ///
282    /// `config.size` is clamped to a minimum of 1 (a zero-sized pool can never
283    /// execute tasks because no permits would be available); a `tracing` warn
284    /// is emitted when this clamp engages.
285    pub fn new(handler: Arc<dyn Handler<I, O>>, config: PoolConfig) -> Self {
286        let size = if config.size == 0 {
287            tracing::warn!(
288                pool = %config.name,
289                "PoolConfig::size was 0, clamping to 1; a zero-sized pool can never execute tasks"
290            );
291            1
292        } else {
293            config.size
294        };
295        let semaphore = Arc::new(Semaphore::new(size));
296        let (queue, receiver) = SubmitQueue::<Envelope<I, O>>::new(config.queue_size);
297        let shutdown = CancellationToken::new();
298
299        let runner = tokio::spawn(runner_loop(
300            config.name.clone(),
301            handler,
302            receiver,
303            semaphore.clone(),
304            shutdown.clone(),
305        ));
306
307        Pool {
308            name: config.name,
309            queue,
310            semaphore,
311            capacity: size,
312            event_buffer: config.event_buffer,
313            overflow_policy: config.overflow_policy,
314            grace_period: config.grace_period,
315            shutdown,
316            runner: Some(runner),
317        }
318    }
319
320    /// Submit a task; returns a [`TaskHandle`] immediately.
321    pub async fn submit(&self, input: I) -> AppResult<TaskHandle<O>> {
322        let id = Uuid::new_v4();
323        let (bcast_tx, bcast_rx) = broadcast::channel::<Event<O>>(self.event_buffer.max(1));
324        let (result_tx, result_rx) = oneshot::channel::<AppResult<O>>();
325        let cancel = CancellationToken::new();
326
327        let handle = TaskHandle::new(id, bcast_rx, result_rx, cancel.clone());
328        let envelope = Envelope {
329            id,
330            input,
331            events_bcast: bcast_tx,
332            result_tx,
333            cancel,
334            event_buffer: self.event_buffer.max(1),
335        };
336
337        match self.overflow_policy {
338            OverflowPolicy::Block => {
339                self.queue.push_block(envelope).await.map_err(|_| {
340                    AppError::new(
341                        ErrorCode::ServiceUnavailable,
342                        format!("pool '{}' is shut down", self.name),
343                    )
344                })?;
345            }
346            OverflowPolicy::Reject => {
347                self.queue.push_reject(envelope).map_err(|err| match err {
348                    PushRejectError::Closed(_) => AppError::new(
349                        ErrorCode::ServiceUnavailable,
350                        format!("pool '{}' is shut down", self.name),
351                    ),
352                    PushRejectError::Full(_) => AppError::rate_limited()
353                        .with_detail("pool", self.name.clone())
354                        .with_detail("overflow_policy", "reject"),
355                })?;
356            }
357            OverflowPolicy::DropOldest => {
358                let dropped = self.queue.push_drop_oldest(envelope).map_err(|_| {
359                    AppError::new(
360                        ErrorCode::ServiceUnavailable,
361                        format!("pool '{}' is shut down", self.name),
362                    )
363                })?;
364                if let Some(dropped) = dropped {
365                    notify_dropped_task(dropped, &self.name);
366                }
367            }
368        }
369
370        Ok(handle)
371    }
372
373    /// Snapshot of pool activity.
374    pub fn stats(&self) -> PoolStats {
375        let running = self
376            .capacity
377            .saturating_sub(self.semaphore.available_permits());
378        PoolStats {
379            name: self.name.clone(),
380            running,
381            capacity: self.capacity,
382        }
383    }
384
385    /// Number of permits currently available for task execution.
386    #[must_use]
387    pub fn available_permits(&self) -> usize {
388        self.semaphore.available_permits()
389    }
390
391    /// Stop accepting work and ask the runner loop to exit.
392    pub fn close(&self) {
393        self.shutdown.cancel();
394        self.queue.close();
395    }
396
397    /// Cancel all in-flight tasks and shut down the runner loop.
398    pub async fn shutdown(mut self) -> AppResult<()> {
399        self.close();
400        if let Some(runner) = self.runner.take() {
401            let mut runner = runner;
402            let wait = tokio::time::timeout(self.grace_period, &mut runner).await;
403            match wait {
404                Ok(joined) => joined.map_err(|err| {
405                    AppError::new(
406                        ErrorCode::Internal,
407                        format!("pool '{}' runner failed during shutdown: {err}", self.name),
408                    )
409                })?,
410                Err(_) => {
411                    tracing::warn!(
412                        pool = %self.name,
413                        grace_period_ms = self.grace_period.as_millis(),
414                        "shutdown grace period elapsed; aborting runner"
415                    );
416                    self.shutdown.cancel();
417                    runner.abort();
418                    let _ = runner.await;
419                }
420            }
421        }
422        Ok(())
423    }
424}
425
426impl<I, O> Drop for Pool<I, O>
427where
428    I: Send + 'static,
429    O: Send + Clone + 'static,
430{
431    fn drop(&mut self) {
432        self.close();
433        if let Some(runner) = self.runner.take() {
434            runner.abort();
435        }
436    }
437}
438
439fn notify_dropped_task<I, O>(envelope: Envelope<I, O>, pool_name: &str)
440where
441    O: Clone + Send + 'static,
442{
443    let error = AppError::rate_limited()
444        .with_detail("pool", pool_name.to_string())
445        .with_detail("overflow_policy", "drop_oldest");
446    let _ = envelope.events_bcast.send(Event::error(
447        envelope.id,
448        format!("{pool_name}/queue"),
449        error.message().to_string(),
450    ));
451    let _ = envelope.result_tx.send(Err(error));
452}
453
454/// Complete a dequeued envelope with a `ServiceUnavailable` error when the
455/// pool is shutting down before the task could be dispatched. Without this
456/// the envelope's `result_tx` would simply be dropped, leaving any awaiting
457/// `TaskHandle::result()` to surface the resulting `RecvError` as a generic
458/// channel-closed error rather than a meaningful "pool is shutting down".
459fn fail_envelope_shutdown<I, O>(envelope: Envelope<I, O>, pool_name: &str)
460where
461    O: Clone + Send + 'static,
462{
463    let error = AppError::new(
464        ErrorCode::ServiceUnavailable,
465        format!("pool '{pool_name}' is shutting down"),
466    );
467    let _ = envelope.events_bcast.send(Event::error(
468        envelope.id,
469        format!("{pool_name}/shutdown"),
470        error.message().to_string(),
471    ));
472    let _ = envelope.result_tx.send(Err(error));
473}
474
475async fn runner_loop<I, O>(
476    pool_name: String,
477    handler: Arc<dyn Handler<I, O>>,
478    receiver: QueueReceiver<Envelope<I, O>>,
479    semaphore: Arc<Semaphore>,
480    shutdown: CancellationToken,
481) where
482    I: Send + 'static,
483    O: Send + Clone + 'static,
484{
485    let mut join_set: JoinSet<()> = JoinSet::new();
486
487    loop {
488        let envelope = tokio::select! {
489            biased;
490
491            _ = shutdown.cancelled() => {
492                tracing::info!(pool = %pool_name, "shutdown requested, draining");
493                break;
494            }
495
496            Some(res) = join_set.join_next() => {
497                if let Err(e) = res
498                    && e.is_panic() {
499                        tracing::error!(pool = %pool_name, "task panicked: {:?}", e);
500                    }
501                continue;
502            }
503
504            envelope = receiver.recv() => {
505                match envelope {
506                    Some(e) => e,
507                    None => break,
508                }
509            }
510        };
511
512        let permit = tokio::select! {
513            biased;
514
515            _ = shutdown.cancelled() => {
516                tracing::info!(pool = %pool_name, "shutdown requested while waiting for permit; failing dequeued task");
517                fail_envelope_shutdown(envelope, &pool_name);
518                break;
519            }
520
521            permit = semaphore.clone().acquire_owned() => {
522                match permit {
523                    Ok(p) => p,
524                    Err(_) => {
525                        fail_envelope_shutdown(envelope, &pool_name);
526                        break;
527                    }
528                }
529            }
530        };
531
532        let handler = handler.clone();
533        let pool = pool_name.clone();
534        join_set.spawn(async move {
535            let _permit = permit;
536            run_task(pool, handler, envelope).await;
537        });
538
539        // Reap completed tasks without blocking.
540        while let Some(res) = join_set.try_join_next() {
541            if let Err(e) = res
542                && e.is_panic()
543            {
544                tracing::error!(pool = %pool_name, "task panicked: {:?}", e);
545            }
546        }
547    }
548
549    while let Some(res) = join_set.join_next().await {
550        if let Err(e) = res
551            && e.is_panic()
552        {
553            tracing::error!(pool = %pool_name, "panic during drain: {:?}", e);
554        }
555    }
556
557    tracing::info!(pool = %pool_name, "pool runner exited");
558}
559
560async fn run_task<I, O>(pool_name: String, handler: Arc<dyn Handler<I, O>>, env: Envelope<I, O>)
561where
562    I: Send + 'static,
563    O: Send + Clone + 'static,
564{
565    let task_id = env.id;
566    let worker_id = format!("{pool_name}/{task_id}");
567
568    let (emit_tx, mut emit_rx) = mpsc::channel::<Event<O>>(env.event_buffer);
569    let bcast_tx = env.events_bcast.clone();
570
571    tokio::spawn(async move {
572        while let Some(ev) = emit_rx.recv().await {
573            let _ = bcast_tx.send(ev);
574        }
575    });
576
577    tracing::debug!(pool = %pool_name, task_id = %task_id, "task started");
578    let result = handler.handle(env.input, emit_tx, env.cancel).await;
579
580    match &result {
581        Ok(_) => tracing::debug!(pool = %pool_name, task_id = %task_id, "task succeeded"),
582        Err(e) => tracing::warn!(pool = %pool_name, task_id = %task_id, error = %e, "task failed"),
583    }
584
585    let final_event = match &result {
586        Ok(v) => Event::result(task_id, &worker_id, v.clone()),
587        Err(e) => Event::error(task_id, &worker_id, e.to_string()),
588    };
589    let _ = env.events_bcast.send(final_event);
590    let _ = env.result_tx.send(result);
591}