1use 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}
28pub 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 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 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 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 #[must_use]
156 pub fn available_permits(&self) -> usize {
157 self.semaphore.available_permits()
158 }
159
160 pub fn close(&self) {
162 self.shutdown.cancel();
163 self.queue.close();
164 }
165
166 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
223fn 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 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}