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#[derive(Debug, Clone)]
20pub struct PoolStats {
21 pub name: String,
23 pub running: usize,
25 pub capacity: usize,
27}
28
29#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)]
31#[non_exhaustive]
32pub enum OverflowPolicy {
33 #[default]
35 Block,
36 Reject,
38 DropOldest,
40}
41
42pub struct PoolConfig {
44 pub name: String,
46 pub size: usize,
48 pub queue_size: usize,
50 pub event_buffer: usize,
52 pub grace_period: Duration,
54 pub dispatch: DispatchStrategy,
56 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 #[must_use]
77 pub fn new(name: impl Into<String>) -> Self {
78 Self {
79 name: name.into(),
80 ..Default::default()
81 }
82 }
83
84 #[must_use]
88 pub fn with_size(mut self, size: usize) -> Self {
89 self.size = size;
90 self
91 }
92
93 #[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 #[must_use]
102 pub fn with_grace_period(mut self, d: Duration) -> Self {
103 self.grace_period = d;
104 self
105 }
106
107 #[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
258pub 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 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 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 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 #[must_use]
387 pub fn available_permits(&self) -> usize {
388 self.semaphore.available_permits()
389 }
390
391 pub fn close(&self) {
393 self.shutdown.cancel();
394 self.queue.close();
395 }
396
397 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
454fn 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 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}