Skip to main content

rskit_worker/pool/
runtime.rs

1//! The worker pool runtime: submission, the runner loop, and task execution.
2
3use std::sync::Arc;
4use std::time::Duration;
5
6use tokio::sync::{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::event::Event;
14use crate::handler::Handler;
15use crate::task::TaskHandle;
16
17use super::config::{OverflowPolicy, PoolConfig, PoolStats};
18use super::queue::{PushRejectError, QueueReceiver, SubmitQueue};
19
20struct Envelope<I, O: Clone + Send + 'static> {
21    id: Uuid,
22    input: I,
23    events_bcast: broadcast::Sender<Event<O>>,
24    result_tx: oneshot::Sender<AppResult<O>>,
25    cancel: CancellationToken,
26    event_buffer: usize,
27}
28/// A bounded async worker pool.
29pub struct Pool<I, O>
30where
31    I: Send + 'static,
32    O: Send + Clone + 'static,
33{
34    name: String,
35    queue: SubmitQueue<Envelope<I, O>>,
36    semaphore: Arc<Semaphore>,
37    capacity: usize,
38    event_buffer: usize,
39    overflow_policy: OverflowPolicy,
40    grace_period: Duration,
41    shutdown: CancellationToken,
42    runner: Option<JoinHandle<()>>,
43}
44
45impl<I, O> Pool<I, O>
46where
47    I: Send + 'static,
48    O: Send + Clone + 'static,
49{
50    /// Create a new pool backed by `handler`.
51    ///
52    /// `config.size` is clamped to a minimum of 1 (a zero-sized pool can never execute tasks because no permits would be available);
53    /// a `tracing` warn is emitted when this clamp engages.
54    pub fn new(handler: Arc<dyn Handler<I, O>>, config: PoolConfig) -> Self {
55        let size = if config.size == 0 {
56            tracing::warn!(
57                pool = %config.name,
58                "PoolConfig::size was 0, clamping to 1; a zero-sized pool can never execute tasks"
59            );
60            1
61        } else {
62            config.size
63        };
64        let semaphore = Arc::new(Semaphore::new(size));
65        let (queue, receiver) = SubmitQueue::<Envelope<I, O>>::new(config.queue_size);
66        let shutdown = CancellationToken::new();
67
68        let runner = tokio::spawn(runner_loop(
69            config.name.clone(),
70            handler,
71            receiver,
72            semaphore.clone(),
73            shutdown.clone(),
74        ));
75
76        Pool {
77            name: config.name,
78            queue,
79            semaphore,
80            capacity: size,
81            event_buffer: config.event_buffer,
82            overflow_policy: config.overflow_policy,
83            grace_period: config.grace_period,
84            shutdown,
85            runner: Some(runner),
86        }
87    }
88
89    /// Submit a task; returns a [`TaskHandle`] immediately.
90    pub async fn submit(&self, input: I) -> AppResult<TaskHandle<O>> {
91        let id = Uuid::new_v4();
92        let (bcast_tx, bcast_rx) = broadcast::channel::<Event<O>>(self.event_buffer.max(1));
93        let (result_tx, result_rx) = oneshot::channel::<AppResult<O>>();
94        let cancel = CancellationToken::new();
95
96        let handle = TaskHandle::new(id, bcast_rx, result_rx, cancel.clone());
97        let envelope = Envelope {
98            id,
99            input,
100            events_bcast: bcast_tx,
101            result_tx,
102            cancel,
103            event_buffer: self.event_buffer.max(1),
104        };
105
106        match self.overflow_policy {
107            OverflowPolicy::Block => {
108                self.queue.push_block(envelope).await.map_err(|_| {
109                    AppError::new(
110                        ErrorCode::ServiceUnavailable,
111                        format!("pool '{}' is shut down", self.name),
112                    )
113                })?;
114            }
115            OverflowPolicy::Reject => {
116                self.queue.push_reject(envelope).map_err(|err| match err {
117                    PushRejectError::Closed(_) => AppError::new(
118                        ErrorCode::ServiceUnavailable,
119                        format!("pool '{}' is shut down", self.name),
120                    ),
121                    PushRejectError::Full(_) => AppError::rate_limited()
122                        .with_detail("pool", self.name.clone())
123                        .with_detail("overflow_policy", "reject"),
124                })?;
125            }
126            OverflowPolicy::DropOldest => {
127                let dropped = self.queue.push_drop_oldest(envelope).map_err(|_| {
128                    AppError::new(
129                        ErrorCode::ServiceUnavailable,
130                        format!("pool '{}' is shut down", self.name),
131                    )
132                })?;
133                if let Some(dropped) = dropped {
134                    notify_dropped_task(dropped, &self.name);
135                }
136            }
137        }
138
139        Ok(handle)
140    }
141
142    /// Snapshot of pool activity.
143    pub fn stats(&self) -> PoolStats {
144        let running = self
145            .capacity
146            .saturating_sub(self.semaphore.available_permits());
147        PoolStats {
148            name: self.name.clone(),
149            running,
150            capacity: self.capacity,
151        }
152    }
153
154    /// Number of permits currently available for task execution.
155    #[must_use]
156    pub fn available_permits(&self) -> usize {
157        self.semaphore.available_permits()
158    }
159
160    /// Stop accepting work and ask the runner loop to exit.
161    pub fn close(&self) {
162        self.shutdown.cancel();
163        self.queue.close();
164    }
165
166    /// Cancel all in-flight tasks and shut down the runner loop.
167    pub async fn shutdown(mut self) -> AppResult<()> {
168        self.close();
169        if let Some(runner) = self.runner.take() {
170            let mut runner = runner;
171            let wait = tokio::time::timeout(self.grace_period, &mut runner).await;
172            match wait {
173                Ok(joined) => joined.map_err(|err| {
174                    AppError::new(
175                        ErrorCode::Internal,
176                        format!("pool '{}' runner failed during shutdown: {err}", self.name),
177                    )
178                })?,
179                Err(_) => {
180                    tracing::warn!(
181                        pool = %self.name,
182                        grace_period_ms = self.grace_period.as_millis(),
183                        "shutdown grace period elapsed; aborting runner"
184                    );
185                    self.shutdown.cancel();
186                    runner.abort();
187                    let _ = runner.await;
188                }
189            }
190        }
191        Ok(())
192    }
193}
194
195impl<I, O> Drop for Pool<I, O>
196where
197    I: Send + 'static,
198    O: Send + Clone + 'static,
199{
200    fn drop(&mut self) {
201        self.close();
202        if let Some(runner) = self.runner.take() {
203            runner.abort();
204        }
205    }
206}
207
208fn notify_dropped_task<I, O>(envelope: Envelope<I, O>, pool_name: &str)
209where
210    O: Clone + Send + 'static,
211{
212    let error = AppError::rate_limited()
213        .with_detail("pool", pool_name.to_string())
214        .with_detail("overflow_policy", "drop_oldest");
215    let _ = envelope.events_bcast.send(Event::error(
216        envelope.id,
217        format!("{pool_name}/queue"),
218        error.message().to_string(),
219    ));
220    let _ = envelope.result_tx.send(Err(error));
221}
222
223/// Complete a dequeued envelope with a `ServiceUnavailable` error when the pool is shutting down before the task could be dispatched.
224/// Without this the envelope's `result_tx` would simply be dropped,
225/// leaving any awaiting `TaskHandle::result()` to surface the resulting `RecvError` as a generic channel-closed error rather than a meaningful "pool is shutting down".
226fn fail_envelope_shutdown<I, O>(envelope: Envelope<I, O>, pool_name: &str)
227where
228    O: Clone + Send + 'static,
229{
230    let error = AppError::new(
231        ErrorCode::ServiceUnavailable,
232        format!("pool '{pool_name}' is shutting down"),
233    );
234    let _ = envelope.events_bcast.send(Event::error(
235        envelope.id,
236        format!("{pool_name}/shutdown"),
237        error.message().to_string(),
238    ));
239    let _ = envelope.result_tx.send(Err(error));
240}
241
242async fn runner_loop<I, O>(
243    pool_name: String,
244    handler: Arc<dyn Handler<I, O>>,
245    receiver: QueueReceiver<Envelope<I, O>>,
246    semaphore: Arc<Semaphore>,
247    shutdown: CancellationToken,
248) where
249    I: Send + 'static,
250    O: Send + Clone + 'static,
251{
252    let mut join_set: JoinSet<()> = JoinSet::new();
253
254    loop {
255        let envelope = tokio::select! {
256            biased;
257
258            _ = shutdown.cancelled() => {
259                tracing::info!(pool = %pool_name, "shutdown requested, draining");
260                break;
261            }
262
263            Some(res) = join_set.join_next() => {
264                if let Err(e) = res
265                    && e.is_panic() {
266                        tracing::error!(pool = %pool_name, "task panicked: {:?}", e);
267                    }
268                continue;
269            }
270
271            envelope = receiver.recv() => {
272                match envelope {
273                    Some(e) => e,
274                    None => break,
275                }
276            }
277        };
278
279        let permit = tokio::select! {
280            biased;
281
282            _ = shutdown.cancelled() => {
283                tracing::info!(pool = %pool_name, "shutdown requested while waiting for permit; failing dequeued task");
284                fail_envelope_shutdown(envelope, &pool_name);
285                break;
286            }
287
288            permit = semaphore.clone().acquire_owned() => {
289                match permit {
290                    Ok(p) => p,
291                    Err(_) => {
292                        fail_envelope_shutdown(envelope, &pool_name);
293                        break;
294                    }
295                }
296            }
297        };
298
299        let handler = handler.clone();
300        let pool = pool_name.clone();
301        join_set.spawn(async move {
302            let _permit = permit;
303            run_task(pool, handler, envelope).await;
304        });
305
306        // Reap completed tasks without blocking.
307        while let Some(res) = join_set.try_join_next() {
308            if let Err(e) = res
309                && e.is_panic()
310            {
311                tracing::error!(pool = %pool_name, "task panicked: {:?}", e);
312            }
313        }
314    }
315
316    while let Some(res) = join_set.join_next().await {
317        if let Err(e) = res
318            && e.is_panic()
319        {
320            tracing::error!(pool = %pool_name, "panic during drain: {:?}", e);
321        }
322    }
323
324    tracing::info!(pool = %pool_name, "pool runner exited");
325}
326
327async fn run_task<I, O>(pool_name: String, handler: Arc<dyn Handler<I, O>>, env: Envelope<I, O>)
328where
329    I: Send + 'static,
330    O: Send + Clone + 'static,
331{
332    let task_id = env.id;
333    let worker_id = format!("{pool_name}/{task_id}");
334
335    let (emit_tx, mut emit_rx) = mpsc::channel::<Event<O>>(env.event_buffer);
336    let bcast_tx = env.events_bcast.clone();
337
338    tokio::spawn(async move {
339        while let Some(ev) = emit_rx.recv().await {
340            let _ = bcast_tx.send(ev);
341        }
342    });
343
344    tracing::debug!(pool = %pool_name, task_id = %task_id, "task started");
345    let result = handler.handle(env.input, emit_tx, env.cancel).await;
346
347    match &result {
348        Ok(_) => tracing::debug!(pool = %pool_name, task_id = %task_id, "task succeeded"),
349        Err(e) => tracing::warn!(pool = %pool_name, task_id = %task_id, error = %e, "task failed"),
350    }
351
352    let final_event = match &result {
353        Ok(v) => Event::result(task_id, &worker_id, v.clone()),
354        Err(e) => Event::error(task_id, &worker_id, e.to_string()),
355    };
356    let _ = env.events_bcast.send(final_event);
357    let _ = env.result_tx.send(result);
358}
359
360#[cfg(test)]
361mod tests {
362    use tokio::sync::mpsc;
363
364    use super::*;
365
366    struct EchoHandler;
367
368    #[async_trait::async_trait]
369    impl Handler<u32, u32> for EchoHandler {
370        async fn handle(
371            &self,
372            task: u32,
373            _emit: mpsc::Sender<Event<u32>>,
374            _cancel: CancellationToken,
375        ) -> AppResult<u32> {
376            Ok(task)
377        }
378    }
379    #[tokio::test]
380    async fn closed_pool_submit_reports_service_unavailable_for_each_policy() {
381        for overflow_policy in [
382            OverflowPolicy::Block,
383            OverflowPolicy::Reject,
384            OverflowPolicy::DropOldest,
385        ] {
386            let pool = Pool::new(
387                Arc::new(EchoHandler),
388                PoolConfig::new("closed")
389                    .with_queue_size(1)
390                    .with_overflow_policy(overflow_policy),
391            );
392            pool.close();
393
394            let error = match pool.submit(1).await {
395                Ok(handle) => handle.result().await.unwrap_err(),
396                Err(error) => error,
397            };
398
399            assert_eq!(error.code(), ErrorCode::ServiceUnavailable);
400        }
401    }
402
403    #[tokio::test]
404    async fn pool_stats_and_successful_result_are_reported() {
405        let pool = Pool::new(
406            Arc::new(EchoHandler),
407            PoolConfig::new("echo")
408                .with_size(0)
409                .with_grace_period(Duration::from_millis(50)),
410        );
411
412        let stats = pool.stats();
413        assert_eq!(stats.name, "echo");
414        assert_eq!(stats.capacity, 1);
415        assert!(pool.available_permits() <= 1);
416
417        let handle = pool.submit(7).await.unwrap();
418        let mut events = handle.events();
419        assert_eq!(handle.result().await.unwrap(), 7);
420        let event = events.try_recv().unwrap();
421        assert_eq!(event.data, Some(7));
422
423        pool.shutdown().await.unwrap();
424    }
425}