Skip to main content

sfo_pool/
worker_pool.rs

1use notify_future::Notify;
2pub use sfo_result::err as pool_err;
3pub use sfo_result::into_err as into_pool_err;
4use std::collections::VecDeque;
5use std::ops::{Deref, DerefMut};
6use std::sync::{Arc, Mutex};
7use std::time::{Duration, Instant};
8
9#[derive(Debug, Copy, Clone, Default, Eq, PartialEq)]
10pub enum PoolErrorCode {
11    #[default]
12    Failed,
13    Clearing,
14    Cleared,
15    InvalidConfig,
16}
17pub type PoolError = sfo_result::Error<PoolErrorCode>;
18pub type PoolResult<T> = sfo_result::Result<T, PoolErrorCode>;
19
20pub(crate) fn pool_error(code: PoolErrorCode, message: &str) -> PoolError {
21    PoolError::new(code, message.to_string())
22}
23
24pub(crate) fn pool_clearing_error() -> PoolError {
25    pool_error(PoolErrorCode::Clearing, "pool is clearing")
26}
27
28pub(crate) fn pool_cleared_error() -> PoolError {
29    pool_error(PoolErrorCode::Cleared, "pool cleared")
30}
31
32pub(crate) fn pool_invalid_config_error(message: &str) -> PoolError {
33    pool_error(PoolErrorCode::InvalidConfig, message)
34}
35
36#[derive(Debug, Clone, Default)]
37/// Configuration for [`WorkerPool`].
38///
39/// All limits default to `None`. Use the `with_*` methods to configure the pool.
40pub struct WorkerPoolConfig {
41    /// Maximum total number of active, idle, and pending workers.
42    ///
43    /// `None` leaves the worker count unlimited.
44    max_count: Option<u16>,
45    /// Maximum number of idle workers retained by the pool.
46    ///
47    /// This does not limit active or pending workers. `None` adds no separate
48    /// idle limit, while `Some(0)` disables idle caching.
49    max_idle_count: Option<u16>,
50    idle_timeout: Option<Duration>,
51}
52
53impl WorkerPoolConfig {
54    /// Sets the maximum total number of workers managed by the pool.
55    ///
56    /// `None` leaves the count unlimited. `Some(0)` is invalid.
57    pub fn with_max_count(mut self, max_count: Option<u16>) -> Self {
58        self.max_count = max_count;
59        self
60    }
61
62    /// Sets the independent idle-worker cache limit.
63    ///
64    /// `None` adds no idle limit. `Some(0)` disables idle caching.
65    pub fn with_max_idle_count(mut self, max_idle_count: Option<u16>) -> Self {
66        self.max_idle_count = max_idle_count;
67        self
68    }
69
70    /// Sets the maximum duration for which an idle worker is retained.
71    ///
72    /// `None` disables timeout-based cleanup.
73    pub fn with_idle_timeout(mut self, idle_timeout: Option<Duration>) -> Self {
74        self.idle_timeout = idle_timeout;
75        self
76    }
77}
78
79#[async_trait::async_trait]
80/// A worker managed by [`WorkerPool`].
81///
82/// Methods on this trait may be called while the pool's internal state lock is held.
83/// Implementations must be non-blocking and must not re-enter APIs on the same pool.
84pub trait Worker: Send + 'static {
85    fn is_work(&self) -> bool;
86}
87
88pub struct WorkerGuard<W: Worker, F: WorkerFactory<W>> {
89    pool_ref: WorkerPoolRef<W, F>,
90    worker: Option<W>,
91}
92
93impl<W: Worker, F: WorkerFactory<W>> WorkerGuard<W, F> {
94    fn new(worker: W, pool_ref: WorkerPoolRef<W, F>) -> Self {
95        WorkerGuard {
96            pool_ref,
97            worker: Some(worker),
98        }
99    }
100}
101
102impl<W: Worker, F: WorkerFactory<W>> Deref for WorkerGuard<W, F> {
103    type Target = W;
104
105    fn deref(&self) -> &Self::Target {
106        self.worker.as_ref().unwrap()
107    }
108}
109
110impl<W: Worker, F: WorkerFactory<W>> DerefMut for WorkerGuard<W, F> {
111    fn deref_mut(&mut self) -> &mut Self::Target {
112        self.worker.as_mut().unwrap()
113    }
114}
115
116impl<W: Worker, F: WorkerFactory<W>> Drop for WorkerGuard<W, F> {
117    fn drop(&mut self) {
118        if let Some(worker) = self.worker.take() {
119            self.pool_ref.release(worker);
120        }
121    }
122}
123
124struct WorkerReservation<W: Worker, F: WorkerFactory<W>> {
125    pool_ref: WorkerPoolRef<W, F>,
126    active: bool,
127}
128
129impl<W: Worker, F: WorkerFactory<W>> WorkerReservation<W, F> {
130    fn new(pool_ref: WorkerPoolRef<W, F>) -> Self {
131        Self {
132            pool_ref,
133            active: true,
134        }
135    }
136
137    fn complete(mut self) -> bool {
138        let (clearing, clear_waiters) = {
139            let mut state = self.pool_ref.state.lock().unwrap();
140            if state.clearing {
141                state.current_count -= 1;
142                (true, state.take_clear_waiters_if_done())
143            } else {
144                (false, Vec::new())
145            }
146        };
147        self.active = false;
148        for waiter in clear_waiters {
149            waiter.notify(());
150        }
151        clearing
152    }
153}
154
155impl<W: Worker, F: WorkerFactory<W>> Drop for WorkerReservation<W, F> {
156    fn drop(&mut self) {
157        if self.active {
158            self.pool_ref.rollback_reservation();
159        }
160    }
161}
162
163#[async_trait::async_trait]
164pub trait WorkerFactory<W: Worker>: Send + Sync + 'static {
165    /// Creates a usable worker.
166    ///
167    /// Returning `Ok` asserts that the worker is ready for use. The pool does not
168    /// call [`Worker::is_work`] before handing a newly created worker to the caller.
169    async fn create(&self) -> PoolResult<W>;
170}
171
172struct IdleWorker<W: Worker> {
173    worker: W,
174    idle_since: Instant,
175}
176
177enum WorkerWaitResult<W: Worker, F: WorkerFactory<W>> {
178    Worker(WorkerGuard<W, F>),
179    Retry,
180    Error(PoolError),
181}
182
183struct WorkerPoolState<W: Worker, F: WorkerFactory<W>> {
184    current_count: usize,
185    worker_list: VecDeque<IdleWorker<W>>,
186    waiting_list: VecDeque<Notify<WorkerWaitResult<W, F>>>,
187    clearing: bool,
188    clear_waiting_list: Vec<Notify<()>>,
189}
190
191impl<W: Worker, F: WorkerFactory<W>> WorkerPoolState<W, F> {
192    fn take_clear_waiters_if_done(&mut self) -> Vec<Notify<()>> {
193        if self.clearing && self.current_count == 0 {
194            self.clearing = false;
195            self.clear_waiting_list.drain(..).collect()
196        } else {
197            Vec::new()
198        }
199    }
200
201    fn pop_next_waiter(&mut self) -> Option<Notify<WorkerWaitResult<W, F>>> {
202        while let Some(waiter) = self.waiting_list.pop_front() {
203            if !waiter.is_canceled() {
204                return Some(waiter);
205            }
206        }
207        None
208    }
209
210    fn drain_waiters(&mut self) -> Vec<Notify<WorkerWaitResult<W, F>>> {
211        self.waiting_list.drain(..).collect()
212    }
213}
214pub struct WorkerPool<W: Worker, F: WorkerFactory<W>> {
215    factory: Arc<F>,
216    config: WorkerPoolConfig,
217    state: Mutex<WorkerPoolState<W, F>>,
218}
219pub type WorkerPoolRef<W, F> = Arc<WorkerPool<W, F>>;
220
221impl<W: Worker, F: WorkerFactory<W>> WorkerPool<W, F> {
222    pub fn new(max_count: u16, factory: F) -> WorkerPoolRef<W, F> {
223        Self::new_with_config(
224            factory,
225            WorkerPoolConfig {
226                max_count: Some(max_count),
227                ..Default::default()
228            },
229        )
230    }
231
232    pub fn new_with_config(factory: F, config: WorkerPoolConfig) -> WorkerPoolRef<W, F> {
233        Arc::new(WorkerPool {
234            factory: Arc::new(factory),
235            config,
236            state: Mutex::new(WorkerPoolState {
237                current_count: 0,
238                worker_list: VecDeque::new(),
239                waiting_list: VecDeque::new(),
240                clearing: false,
241                clear_waiting_list: Vec::new(),
242            }),
243        })
244    }
245
246    fn take_expired_idle_workers(
247        state: &mut WorkerPoolState<W, F>,
248        idle_timeout: Option<Duration>,
249    ) -> Vec<W> {
250        let Some(idle_timeout) = idle_timeout else {
251            return Vec::new();
252        };
253        let mut removed_workers = Vec::new();
254        let now = Instant::now();
255        while state
256            .worker_list
257            .front()
258            .map(|idle_worker| now.duration_since(idle_worker.idle_since) >= idle_timeout)
259            .unwrap_or(false)
260        {
261            let idle_worker = state.worker_list.pop_front().unwrap();
262            state.current_count -= 1;
263            removed_workers.push(idle_worker.worker);
264        }
265        removed_workers
266    }
267
268    pub fn cleanup_idle_worker(&self) -> usize {
269        let (removed_workers, clear_waiters) = {
270            let mut state = self.state.lock().unwrap();
271            let removed_workers =
272                Self::take_expired_idle_workers(&mut state, self.config.idle_timeout);
273            let clear_waiters = state.take_clear_waiters_if_done();
274            (removed_workers, clear_waiters)
275        };
276        for waiter in clear_waiters {
277            waiter.notify(());
278        }
279        let removed_count = removed_workers.len();
280        drop(removed_workers);
281        removed_count
282    }
283
284    pub async fn get_worker(self: &WorkerPoolRef<W, F>) -> PoolResult<WorkerGuard<W, F>> {
285        loop {
286            if self.config.max_count == Some(0) {
287                return Err(pool_invalid_config_error("pool max_count is zero"));
288            }
289
290            let (worker, wait, should_create, removed_workers) = {
291                let mut state = self.state.lock().unwrap();
292                if state.clearing {
293                    return Err(pool_clearing_error());
294                }
295
296                let mut removed_workers =
297                    Self::take_expired_idle_workers(&mut state, self.config.idle_timeout);
298
299                let worker = loop {
300                    let Some(idle_worker) = state.worker_list.pop_back() else {
301                        break None;
302                    };
303                    let worker = idle_worker.worker;
304                    if !worker.is_work() {
305                        state.current_count -= 1;
306                        removed_workers.push(worker);
307                        continue;
308                    }
309                    break Some(worker);
310                };
311
312                if worker.is_some() {
313                    (worker, None, false, removed_workers)
314                } else if self
315                    .config
316                    .max_count
317                    .map(|max_count| state.current_count < usize::from(max_count))
318                    .unwrap_or(true)
319                {
320                    state.current_count += 1;
321                    (None, None, true, removed_workers)
322                } else {
323                    let (notify, waiter) = Notify::new();
324                    state.waiting_list.push_back(notify);
325                    (None, Some(waiter), false, removed_workers)
326                }
327            };
328
329            let reservation = should_create.then(|| WorkerReservation::new(self.clone()));
330            drop(removed_workers);
331
332            if let Some(worker) = worker {
333                return Ok(WorkerGuard::new(worker, self.clone()));
334            }
335
336            if let Some(wait) = wait {
337                match wait.await {
338                    WorkerWaitResult::Worker(worker) => return Ok(worker),
339                    WorkerWaitResult::Retry => continue,
340                    WorkerWaitResult::Error(err) => return Err(err),
341                }
342            }
343
344            let reservation = reservation.unwrap();
345            let worker = match self.factory.create().await {
346                Ok(worker) => worker,
347                Err(err) => return Err(err),
348            };
349            if reservation.complete() {
350                return Err(pool_cleared_error());
351            }
352            return Ok(WorkerGuard::new(worker, self.clone()));
353        }
354    }
355
356    pub async fn clear_all_worker(&self) {
357        let (waiter, waiting_list, clear_waiters, idle_workers) = {
358            let mut state = self.state.lock().unwrap();
359            let idle_workers = if !state.clearing {
360                state.clearing = true;
361                let cur_worker_count = state.worker_list.len();
362                let idle_workers = state
363                    .worker_list
364                    .drain(..)
365                    .map(|idle_worker| idle_worker.worker)
366                    .collect::<Vec<_>>();
367                state.current_count -= cur_worker_count;
368                idle_workers
369            } else {
370                Vec::new()
371            };
372
373            let waiting_list = state.waiting_list.drain(..).collect::<Vec<_>>();
374            if state.current_count == 0 {
375                let clear_waiters = state.take_clear_waiters_if_done();
376                (None, waiting_list, clear_waiters, idle_workers)
377            } else {
378                let (notify, waiter) = Notify::new();
379                state.clear_waiting_list.push(notify);
380                (Some(waiter), waiting_list, Vec::new(), idle_workers)
381            }
382        };
383        for waiting in waiting_list {
384            waiting.notify(WorkerWaitResult::Error(pool_cleared_error()));
385        }
386        for waiter in clear_waiters {
387            waiter.notify(());
388        }
389        drop(idle_workers);
390        if let Some(waiter) = waiter {
391            waiter.await;
392        }
393    }
394
395    fn notify_retry_waiters(waiters: Vec<Notify<WorkerWaitResult<W, F>>>) {
396        for waiter in waiters {
397            waiter.notify(WorkerWaitResult::Retry);
398        }
399    }
400
401    fn rollback_reservation(&self) {
402        let (retry_waiters, clear_waiters) = {
403            let mut state = self.state.lock().unwrap();
404            state.current_count -= 1;
405            let retry_waiters = state.drain_waiters();
406            let clear_waiters = state.take_clear_waiters_if_done();
407            (retry_waiters, clear_waiters)
408        };
409        Self::notify_retry_waiters(retry_waiters);
410        for waiter in clear_waiters {
411            waiter.notify(());
412        }
413    }
414
415    fn release(self: &WorkerPoolRef<W, F>, work: W) {
416        enum ReleaseAction<W: Worker, F: WorkerFactory<W>> {
417            None,
418            Notify(Notify<WorkerWaitResult<W, F>>, WorkerGuard<W, F>),
419            Retry(Vec<Notify<WorkerWaitResult<W, F>>>),
420        }
421
422        let mut clear_waiters = Vec::new();
423        let mut removed_workers = Vec::new();
424        let action = {
425            let mut state = self.state.lock().unwrap();
426            if state.clearing {
427                state.current_count -= 1;
428                clear_waiters = state.take_clear_waiters_if_done();
429                ReleaseAction::None
430            } else if work.is_work() {
431                let future = state.pop_next_waiter();
432                if let Some(future) = future {
433                    ReleaseAction::Notify(future, WorkerGuard::new(work, self.clone()))
434                } else {
435                    state.worker_list.push_back(IdleWorker {
436                        worker: work,
437                        idle_since: Instant::now(),
438                    });
439                    if let Some(max_idle_count) = self.config.max_idle_count {
440                        while state.worker_list.len() > usize::from(max_idle_count) {
441                            let idle_worker = state.worker_list.pop_front().unwrap();
442                            state.current_count -= 1;
443                            removed_workers.push(idle_worker.worker);
444                        }
445                    }
446                    ReleaseAction::None
447                }
448            } else {
449                state.current_count -= 1;
450                let waiters = state.drain_waiters();
451                if !waiters.is_empty() {
452                    ReleaseAction::Retry(waiters)
453                } else {
454                    clear_waiters = state.take_clear_waiters_if_done();
455                    ReleaseAction::None
456                }
457            }
458        };
459
460        for waiter in clear_waiters {
461            waiter.notify(());
462        }
463        drop(removed_workers);
464
465        match action {
466            ReleaseAction::None => {}
467            ReleaseAction::Notify(future, worker) => {
468                future.notify(WorkerWaitResult::Worker(worker));
469            }
470            ReleaseAction::Retry(waiters) => {
471                Self::notify_retry_waiters(waiters);
472            }
473        }
474    }
475}
476
477#[cfg(test)]
478mod idle_limit_tests {
479    use super::*;
480    use std::sync::atomic::{AtomicUsize, Ordering};
481
482    struct TestWorker {
483        id: usize,
484        dropped: Arc<Mutex<Vec<usize>>>,
485    }
486
487    impl Worker for TestWorker {
488        fn is_work(&self) -> bool {
489            true
490        }
491    }
492
493    impl Drop for TestWorker {
494        fn drop(&mut self) {
495            self.dropped.lock().unwrap().push(self.id);
496        }
497    }
498
499    struct TestFactory {
500        next_id: AtomicUsize,
501        dropped: Arc<Mutex<Vec<usize>>>,
502    }
503
504    #[async_trait::async_trait]
505    impl WorkerFactory<TestWorker> for TestFactory {
506        async fn create(&self) -> PoolResult<TestWorker> {
507            Ok(TestWorker {
508                id: self.next_id.fetch_add(1, Ordering::SeqCst),
509                dropped: self.dropped.clone(),
510            })
511        }
512    }
513
514    fn new_pool(max_count: u16, max_idle_count: u16) -> WorkerPoolRef<TestWorker, TestFactory> {
515        let dropped = Arc::new(Mutex::new(Vec::new()));
516        WorkerPool::new_with_config(
517            TestFactory {
518                next_id: AtomicUsize::new(0),
519                dropped,
520            },
521            WorkerPoolConfig {
522                max_count: Some(max_count),
523                max_idle_count: Some(max_idle_count),
524                ..Default::default()
525            },
526        )
527    }
528
529    #[tokio::test]
530    async fn idle_limit_only_caps_idle_workers_and_evicts_lru() {
531        let pool = new_pool(3, 1);
532        let worker0 = pool.get_worker().await.unwrap();
533        let worker1 = pool.get_worker().await.unwrap();
534        let worker2 = pool.get_worker().await.unwrap();
535        assert_eq!(worker2.id, 2);
536
537        drop(worker0);
538        drop(worker1);
539        drop(worker2);
540        assert_eq!(*pool.factory.dropped.lock().unwrap(), vec![0, 1]);
541
542        let newest = pool.get_worker().await.unwrap();
543        assert_eq!(newest.id, 2);
544        let newly_created = pool.get_worker().await.unwrap();
545        assert_eq!(newly_created.id, 3);
546    }
547
548    #[tokio::test]
549    async fn zero_idle_limit_disables_idle_cache() {
550        let pool = new_pool(2, 0);
551        let worker = pool.get_worker().await.unwrap();
552        assert_eq!(worker.id, 0);
553        drop(worker);
554        assert_eq!(*pool.factory.dropped.lock().unwrap(), vec![0]);
555
556        let worker = pool.get_worker().await.unwrap();
557        assert_eq!(worker.id, 1);
558    }
559}
560
561#[tokio::test]
562async fn test_pool() {
563    struct TestWorker {
564        work: bool,
565    }
566
567    #[async_trait::async_trait]
568    impl Worker for TestWorker {
569        fn is_work(&self) -> bool {
570            self.work
571        }
572    }
573
574    struct TestWorkerFactory;
575
576    #[async_trait::async_trait]
577    impl WorkerFactory<TestWorker> for TestWorkerFactory {
578        async fn create(&self) -> PoolResult<TestWorker> {
579            Ok(TestWorker { work: true })
580        }
581    }
582
583    let pool = WorkerPool::new(2, TestWorkerFactory);
584
585    let worker1 = pool.get_worker().await.unwrap();
586    let worker2 = pool.get_worker().await.unwrap();
587
588    let pool_ref = pool.clone();
589    let waiter = tokio::spawn(async move { pool_ref.get_worker().await });
590    tokio::time::sleep(std::time::Duration::from_millis(20)).await;
591    assert!(!waiter.is_finished());
592
593    drop(worker1);
594    let worker3 = tokio::time::timeout(std::time::Duration::from_secs(1), waiter)
595        .await
596        .unwrap()
597        .unwrap()
598        .unwrap();
599    drop(worker2);
600    drop(worker3);
601
602    let worker1 = pool.get_worker().await.unwrap();
603    let worker2 = pool.get_worker().await.unwrap();
604
605    let pool_ref = pool.clone();
606    let waiter1 = tokio::spawn(async move { pool_ref.get_worker().await });
607    let pool_ref = pool.clone();
608    let waiter2 = tokio::spawn(async move { pool_ref.get_worker().await });
609    tokio::time::sleep(std::time::Duration::from_millis(20)).await;
610    assert!(!waiter1.is_finished());
611    assert!(!waiter2.is_finished());
612
613    let pool_ref = pool.clone();
614    let clear_task = tokio::spawn(async move {
615        pool_ref.clear_all_worker().await;
616    });
617    tokio::time::sleep(std::time::Duration::from_millis(20)).await;
618
619    assert!(waiter1.await.unwrap().is_err());
620    assert!(waiter2.await.unwrap().is_err());
621
622    drop(worker1);
623    drop(worker2);
624
625    tokio::time::timeout(std::time::Duration::from_secs(1), clear_task)
626        .await
627        .unwrap()
628        .unwrap();
629}
630
631#[tokio::test]
632async fn test_clear_all_worker_waits_for_inflight_create() {
633    use std::sync::atomic::{AtomicUsize, Ordering};
634    use std::sync::Arc;
635
636    struct TestWorker;
637
638    #[async_trait::async_trait]
639    impl Worker for TestWorker {
640        fn is_work(&self) -> bool {
641            true
642        }
643    }
644
645    struct TestWorkerFactory {
646        create_count: Arc<AtomicUsize>,
647    }
648
649    #[async_trait::async_trait]
650    impl WorkerFactory<TestWorker> for TestWorkerFactory {
651        async fn create(&self) -> PoolResult<TestWorker> {
652            self.create_count.fetch_add(1, Ordering::SeqCst);
653            tokio::time::sleep(std::time::Duration::from_millis(100)).await;
654            Ok(TestWorker)
655        }
656    }
657
658    let create_count = Arc::new(AtomicUsize::new(0));
659    let pool = WorkerPool::new(
660        1,
661        TestWorkerFactory {
662            create_count: create_count.clone(),
663        },
664    );
665
666    let pool_ref = pool.clone();
667    let worker_task = tokio::spawn(async move { pool_ref.get_worker().await });
668    tokio::time::sleep(std::time::Duration::from_millis(20)).await;
669
670    pool.clear_all_worker().await;
671
672    let worker = worker_task.await.unwrap();
673    assert!(worker.is_err());
674    assert_eq!(create_count.load(Ordering::SeqCst), 1);
675}
676
677#[tokio::test]
678async fn test_concurrent_clear_all_worker() {
679    struct TestWorker;
680
681    #[async_trait::async_trait]
682    impl Worker for TestWorker {
683        fn is_work(&self) -> bool {
684            true
685        }
686    }
687
688    struct TestWorkerFactory;
689
690    #[async_trait::async_trait]
691    impl WorkerFactory<TestWorker> for TestWorkerFactory {
692        async fn create(&self) -> PoolResult<TestWorker> {
693            Ok(TestWorker)
694        }
695    }
696
697    let pool = WorkerPool::new(1, TestWorkerFactory);
698    let worker = pool.get_worker().await.unwrap();
699
700    let pool_ref = pool.clone();
701    let clear_task1 = tokio::spawn(async move {
702        pool_ref.clear_all_worker().await;
703    });
704
705    let pool_ref = pool.clone();
706    let clear_task2 = tokio::spawn(async move {
707        pool_ref.clear_all_worker().await;
708    });
709
710    tokio::time::sleep(std::time::Duration::from_millis(20)).await;
711    drop(worker);
712
713    tokio::time::timeout(std::time::Duration::from_secs(1), async {
714        clear_task1.await.unwrap();
715        clear_task2.await.unwrap();
716    })
717    .await
718    .unwrap();
719}
720
721#[tokio::test]
722async fn test_zero_max_count_returns_error() {
723    struct TestWorker;
724
725    #[async_trait::async_trait]
726    impl Worker for TestWorker {
727        fn is_work(&self) -> bool {
728            true
729        }
730    }
731
732    struct TestWorkerFactory;
733
734    #[async_trait::async_trait]
735    impl WorkerFactory<TestWorker> for TestWorkerFactory {
736        async fn create(&self) -> PoolResult<TestWorker> {
737            Ok(TestWorker)
738        }
739    }
740
741    let pool = WorkerPool::new(0, TestWorkerFactory);
742    let worker = pool.get_worker().await;
743    assert!(worker.is_err());
744    assert_eq!(worker.err().unwrap().code(), PoolErrorCode::InvalidConfig);
745}
746
747#[test]
748fn test_worker_pool_config_default_max_count() {
749    let config = WorkerPoolConfig::default();
750    assert_eq!(config.max_count, None);
751    assert_eq!(config.max_idle_count, None);
752}
753
754#[test]
755fn test_worker_pool_config_builder() {
756    let timeout = Duration::from_secs(1);
757    let config = WorkerPoolConfig::default()
758        .with_max_count(Some(2))
759        .with_max_idle_count(Some(1))
760        .with_idle_timeout(Some(timeout));
761    assert_eq!(config.max_count, Some(2));
762    assert_eq!(config.max_idle_count, Some(1));
763    assert_eq!(config.idle_timeout, Some(timeout));
764}
765
766#[tokio::test]
767async fn test_worker_pool_default_config_has_no_max_count() {
768    struct TestWorker;
769
770    impl Worker for TestWorker {
771        fn is_work(&self) -> bool {
772            true
773        }
774    }
775
776    struct TestWorkerFactory;
777
778    #[async_trait::async_trait]
779    impl WorkerFactory<TestWorker> for TestWorkerFactory {
780        async fn create(&self) -> PoolResult<TestWorker> {
781            Ok(TestWorker)
782        }
783    }
784
785    let pool = WorkerPool::new_with_config(TestWorkerFactory, Default::default());
786    let worker1 = pool.get_worker().await.unwrap();
787    let worker2 = pool.get_worker().await.unwrap();
788    drop((worker1, worker2));
789}
790
791#[tokio::test]
792async fn test_create_failure_fails_waiting_workers() {
793    struct TestWorker;
794
795    #[async_trait::async_trait]
796    impl Worker for TestWorker {
797        fn is_work(&self) -> bool {
798            true
799        }
800    }
801
802    struct TestWorkerFactory;
803
804    #[async_trait::async_trait]
805    impl WorkerFactory<TestWorker> for TestWorkerFactory {
806        async fn create(&self) -> PoolResult<TestWorker> {
807            tokio::time::sleep(std::time::Duration::from_millis(50)).await;
808            Err(pool_invalid_config_error("create failed"))
809        }
810    }
811
812    let pool = WorkerPool::new(1, TestWorkerFactory);
813
814    let pool_ref = pool.clone();
815    let worker1 = tokio::spawn(async move { pool_ref.get_worker().await });
816    tokio::time::sleep(std::time::Duration::from_millis(20)).await;
817
818    let pool_ref = pool.clone();
819    let worker2 = tokio::spawn(async move { pool_ref.get_worker().await });
820
821    let (worker1, worker2) = tokio::time::timeout(std::time::Duration::from_secs(1), async {
822        (worker1.await.unwrap(), worker2.await.unwrap())
823    })
824    .await
825    .unwrap();
826
827    assert_eq!(worker1.err().unwrap().code(), PoolErrorCode::InvalidConfig);
828    assert_eq!(worker2.err().unwrap().code(), PoolErrorCode::InvalidConfig);
829}
830
831#[tokio::test]
832async fn test_invalid_worker_drop_outside_runtime_wakes_waiter() {
833    use std::sync::atomic::{AtomicUsize, Ordering};
834    use std::sync::Arc;
835
836    struct TestWorker {
837        id: usize,
838        work: bool,
839    }
840
841    #[async_trait::async_trait]
842    impl Worker for TestWorker {
843        fn is_work(&self) -> bool {
844            self.work
845        }
846    }
847
848    struct TestWorkerFactory {
849        create_count: Arc<AtomicUsize>,
850    }
851
852    #[async_trait::async_trait]
853    impl WorkerFactory<TestWorker> for TestWorkerFactory {
854        async fn create(&self) -> PoolResult<TestWorker> {
855            let id = self.create_count.fetch_add(1, Ordering::SeqCst);
856            Ok(TestWorker { id, work: true })
857        }
858    }
859
860    let create_count = Arc::new(AtomicUsize::new(0));
861    let pool = WorkerPool::new(
862        1,
863        TestWorkerFactory {
864            create_count: create_count.clone(),
865        },
866    );
867
868    let mut worker = pool.get_worker().await.unwrap();
869    assert_eq!(worker.id, 0);
870    worker.work = false;
871
872    let pool_ref = pool.clone();
873    let waiter = tokio::spawn(async move { pool_ref.get_worker().await });
874    tokio::time::sleep(std::time::Duration::from_millis(20)).await;
875    assert!(!waiter.is_finished());
876
877    std::thread::spawn(move || drop(worker)).join().unwrap();
878
879    let worker = tokio::time::timeout(std::time::Duration::from_secs(1), waiter)
880        .await
881        .unwrap()
882        .unwrap()
883        .unwrap();
884    assert_eq!(worker.id, 1);
885    assert_eq!(create_count.load(Ordering::SeqCst), 2);
886}
887
888#[tokio::test]
889async fn test_retry_notification_skips_canceled_waiter() {
890    struct TestWorker;
891
892    #[async_trait::async_trait]
893    impl Worker for TestWorker {
894        fn is_work(&self) -> bool {
895            true
896        }
897    }
898
899    struct TestWorkerFactory;
900
901    #[async_trait::async_trait]
902    impl WorkerFactory<TestWorker> for TestWorkerFactory {
903        async fn create(&self) -> PoolResult<TestWorker> {
904            Ok(TestWorker)
905        }
906    }
907
908    let (canceled_notify, canceled_waiter) = Notify::new();
909    drop(canceled_waiter);
910    let (notify, waiter) = Notify::new();
911
912    WorkerPool::<TestWorker, TestWorkerFactory>::notify_retry_waiters(vec![
913        canceled_notify,
914        notify,
915    ]);
916
917    let result = tokio::time::timeout(std::time::Duration::from_secs(1), waiter)
918        .await
919        .unwrap();
920    assert!(matches!(result, WorkerWaitResult::Retry));
921}
922
923#[tokio::test]
924async fn test_clearing_and_cleared_error_codes() {
925    use std::sync::atomic::{AtomicBool, Ordering};
926    use std::sync::Arc;
927
928    struct TestWorker;
929
930    #[async_trait::async_trait]
931    impl Worker for TestWorker {
932        fn is_work(&self) -> bool {
933            true
934        }
935    }
936
937    struct TestWorkerFactory {
938        should_block: Arc<AtomicBool>,
939    }
940
941    #[async_trait::async_trait]
942    impl WorkerFactory<TestWorker> for TestWorkerFactory {
943        async fn create(&self) -> PoolResult<TestWorker> {
944            while self.should_block.load(Ordering::SeqCst) {
945                tokio::task::yield_now().await;
946            }
947            Ok(TestWorker)
948        }
949    }
950
951    let should_block = Arc::new(AtomicBool::new(true));
952    let pool = WorkerPool::new(
953        1,
954        TestWorkerFactory {
955            should_block: should_block.clone(),
956        },
957    );
958
959    let pool_ref = pool.clone();
960    let inflight = tokio::spawn(async move { pool_ref.get_worker().await });
961    tokio::task::yield_now().await;
962
963    let pool_ref = pool.clone();
964    let clear_task = tokio::spawn(async move {
965        pool_ref.clear_all_worker().await;
966    });
967    tokio::task::yield_now().await;
968
969    let err = pool.get_worker().await.err().unwrap();
970    assert_eq!(err.code(), PoolErrorCode::Clearing);
971
972    should_block.store(false, Ordering::SeqCst);
973    clear_task.await.unwrap();
974
975    let err = inflight.await.unwrap().err().unwrap();
976    assert_eq!(err.code(), PoolErrorCode::Cleared);
977}
978
979#[tokio::test]
980async fn test_idle_worker_timeout_releases_worker() {
981    use std::sync::atomic::{AtomicUsize, Ordering};
982    use std::sync::Arc;
983
984    struct TestWorker {
985        id: usize,
986    }
987
988    #[async_trait::async_trait]
989    impl Worker for TestWorker {
990        fn is_work(&self) -> bool {
991            true
992        }
993    }
994
995    struct TestWorkerFactory {
996        create_count: Arc<AtomicUsize>,
997    }
998
999    #[async_trait::async_trait]
1000    impl WorkerFactory<TestWorker> for TestWorkerFactory {
1001        async fn create(&self) -> PoolResult<TestWorker> {
1002            let id = self.create_count.fetch_add(1, Ordering::SeqCst);
1003            Ok(TestWorker { id })
1004        }
1005    }
1006
1007    let create_count = Arc::new(AtomicUsize::new(0));
1008    let pool = WorkerPool::new_with_config(
1009        TestWorkerFactory {
1010            create_count: create_count.clone(),
1011        },
1012        WorkerPoolConfig {
1013            max_count: Some(1),
1014            idle_timeout: Some(std::time::Duration::from_millis(30)),
1015            ..Default::default()
1016        },
1017    );
1018
1019    {
1020        let worker = pool.get_worker().await.unwrap();
1021        assert_eq!(worker.id, 0);
1022    }
1023
1024    tokio::time::sleep(std::time::Duration::from_millis(80)).await;
1025
1026    let worker = pool.get_worker().await.unwrap();
1027    assert_eq!(worker.id, 1);
1028    assert_eq!(create_count.load(Ordering::SeqCst), 2);
1029}
1030
1031#[tokio::test]
1032async fn test_idle_worker_reused_before_timeout() {
1033    use std::sync::atomic::{AtomicUsize, Ordering};
1034    use std::sync::Arc;
1035
1036    struct TestWorker {
1037        id: usize,
1038    }
1039
1040    #[async_trait::async_trait]
1041    impl Worker for TestWorker {
1042        fn is_work(&self) -> bool {
1043            true
1044        }
1045    }
1046
1047    struct TestWorkerFactory {
1048        create_count: Arc<AtomicUsize>,
1049    }
1050
1051    #[async_trait::async_trait]
1052    impl WorkerFactory<TestWorker> for TestWorkerFactory {
1053        async fn create(&self) -> PoolResult<TestWorker> {
1054            let id = self.create_count.fetch_add(1, Ordering::SeqCst);
1055            Ok(TestWorker { id })
1056        }
1057    }
1058
1059    let create_count = Arc::new(AtomicUsize::new(0));
1060    let pool = WorkerPool::new_with_config(
1061        TestWorkerFactory {
1062            create_count: create_count.clone(),
1063        },
1064        WorkerPoolConfig {
1065            max_count: Some(1),
1066            idle_timeout: Some(std::time::Duration::from_secs(1)),
1067            ..Default::default()
1068        },
1069    );
1070
1071    {
1072        let worker = pool.get_worker().await.unwrap();
1073        assert_eq!(worker.id, 0);
1074    }
1075
1076    tokio::time::sleep(std::time::Duration::from_millis(20)).await;
1077
1078    let worker = pool.get_worker().await.unwrap();
1079    assert_eq!(worker.id, 0);
1080    assert_eq!(create_count.load(Ordering::SeqCst), 1);
1081}
1082
1083#[tokio::test]
1084async fn test_cleanup_idle_worker_can_be_triggered_externally() {
1085    use std::sync::atomic::{AtomicUsize, Ordering};
1086    use std::sync::Arc;
1087
1088    struct TestWorker {
1089        id: usize,
1090    }
1091
1092    #[async_trait::async_trait]
1093    impl Worker for TestWorker {
1094        fn is_work(&self) -> bool {
1095            true
1096        }
1097    }
1098
1099    struct TestWorkerFactory {
1100        create_count: Arc<AtomicUsize>,
1101    }
1102
1103    #[async_trait::async_trait]
1104    impl WorkerFactory<TestWorker> for TestWorkerFactory {
1105        async fn create(&self) -> PoolResult<TestWorker> {
1106            let id = self.create_count.fetch_add(1, Ordering::SeqCst);
1107            Ok(TestWorker { id })
1108        }
1109    }
1110
1111    let create_count = Arc::new(AtomicUsize::new(0));
1112    let pool = WorkerPool::new_with_config(
1113        TestWorkerFactory {
1114            create_count: create_count.clone(),
1115        },
1116        WorkerPoolConfig {
1117            max_count: Some(1),
1118            idle_timeout: Some(std::time::Duration::from_millis(30)),
1119            ..Default::default()
1120        },
1121    );
1122
1123    {
1124        let worker = pool.get_worker().await.unwrap();
1125        assert_eq!(worker.id, 0);
1126    }
1127
1128    tokio::time::sleep(std::time::Duration::from_millis(80)).await;
1129
1130    assert_eq!(pool.cleanup_idle_worker(), 1);
1131
1132    let worker = pool.get_worker().await.unwrap();
1133    assert_eq!(worker.id, 1);
1134    assert_eq!(create_count.load(Ordering::SeqCst), 2);
1135}
1136
1137#[tokio::test]
1138async fn test_get_worker_uses_most_recent_idle_worker() {
1139    use std::sync::atomic::{AtomicUsize, Ordering};
1140    use std::sync::Arc;
1141
1142    struct TestWorker {
1143        id: usize,
1144    }
1145
1146    #[async_trait::async_trait]
1147    impl Worker for TestWorker {
1148        fn is_work(&self) -> bool {
1149            true
1150        }
1151    }
1152
1153    struct TestWorkerFactory {
1154        create_count: Arc<AtomicUsize>,
1155    }
1156
1157    #[async_trait::async_trait]
1158    impl WorkerFactory<TestWorker> for TestWorkerFactory {
1159        async fn create(&self) -> PoolResult<TestWorker> {
1160            let id = self.create_count.fetch_add(1, Ordering::SeqCst);
1161            Ok(TestWorker { id })
1162        }
1163    }
1164
1165    let create_count = Arc::new(AtomicUsize::new(0));
1166    let pool = WorkerPool::new(
1167        2,
1168        TestWorkerFactory {
1169            create_count: create_count.clone(),
1170        },
1171    );
1172
1173    let worker1 = pool.get_worker().await.unwrap();
1174    let worker2 = pool.get_worker().await.unwrap();
1175    assert_eq!(worker1.id, 0);
1176    assert_eq!(worker2.id, 1);
1177
1178    drop(worker1);
1179    drop(worker2);
1180
1181    let worker = pool.get_worker().await.unwrap();
1182    assert_eq!(worker.id, 1);
1183    assert_eq!(create_count.load(Ordering::SeqCst), 2);
1184}
1185
1186#[tokio::test]
1187async fn test_canceled_create_rolls_back_reservation() {
1188    use std::sync::atomic::{AtomicBool, Ordering};
1189
1190    struct TestWorker;
1191
1192    #[async_trait::async_trait]
1193    impl Worker for TestWorker {
1194        fn is_work(&self) -> bool {
1195            true
1196        }
1197    }
1198
1199    struct TestWorkerFactory {
1200        create_started: Arc<AtomicBool>,
1201        allow_create: Arc<AtomicBool>,
1202    }
1203
1204    #[async_trait::async_trait]
1205    impl WorkerFactory<TestWorker> for TestWorkerFactory {
1206        async fn create(&self) -> PoolResult<TestWorker> {
1207            self.create_started.store(true, Ordering::SeqCst);
1208            while !self.allow_create.load(Ordering::SeqCst) {
1209                tokio::task::yield_now().await;
1210            }
1211            Ok(TestWorker)
1212        }
1213    }
1214
1215    let create_started = Arc::new(AtomicBool::new(false));
1216    let allow_create = Arc::new(AtomicBool::new(false));
1217    let pool = WorkerPool::new(
1218        1,
1219        TestWorkerFactory {
1220            create_started: create_started.clone(),
1221            allow_create: allow_create.clone(),
1222        },
1223    );
1224
1225    let pool_ref = pool.clone();
1226    let create_task = tokio::spawn(async move { pool_ref.get_worker().await });
1227    while !create_started.load(Ordering::SeqCst) {
1228        tokio::task::yield_now().await;
1229    }
1230    create_task.abort();
1231    assert!(matches!(create_task.await, Err(err) if err.is_cancelled()));
1232
1233    allow_create.store(true, Ordering::SeqCst);
1234    let worker = tokio::time::timeout(std::time::Duration::from_secs(1), pool.get_worker())
1235        .await
1236        .unwrap()
1237        .unwrap();
1238    drop(worker);
1239
1240    tokio::time::timeout(std::time::Duration::from_secs(1), pool.clear_all_worker())
1241        .await
1242        .unwrap();
1243}
1244
1245#[tokio::test]
1246async fn test_cleanup_drops_idle_worker_outside_state_lock() {
1247    use std::sync::mpsc;
1248
1249    type DropCallback = Box<dyn FnOnce() + Send>;
1250
1251    struct TestWorker {
1252        on_drop: Option<DropCallback>,
1253    }
1254
1255    #[async_trait::async_trait]
1256    impl Worker for TestWorker {
1257        fn is_work(&self) -> bool {
1258            true
1259        }
1260    }
1261
1262    impl Drop for TestWorker {
1263        fn drop(&mut self) {
1264            if let Some(on_drop) = self.on_drop.take() {
1265                on_drop();
1266            }
1267        }
1268    }
1269
1270    struct TestWorkerFactory {
1271        on_drop: Arc<Mutex<Option<DropCallback>>>,
1272    }
1273
1274    #[async_trait::async_trait]
1275    impl WorkerFactory<TestWorker> for TestWorkerFactory {
1276        async fn create(&self) -> PoolResult<TestWorker> {
1277            Ok(TestWorker {
1278                on_drop: self.on_drop.lock().unwrap().take(),
1279            })
1280        }
1281    }
1282
1283    let on_drop = Arc::new(Mutex::new(None));
1284    let pool = WorkerPool::new_with_config(
1285        TestWorkerFactory {
1286            on_drop: on_drop.clone(),
1287        },
1288        WorkerPoolConfig {
1289            max_count: Some(1),
1290            idle_timeout: Some(Duration::ZERO),
1291            ..Default::default()
1292        },
1293    );
1294    let (tx, rx) = mpsc::channel();
1295    let pool_ref = pool.clone();
1296    *on_drop.lock().unwrap() = Some(Box::new(move || {
1297        pool_ref.cleanup_idle_worker();
1298        tx.send(()).unwrap();
1299    }));
1300
1301    let worker = pool.get_worker().await.unwrap();
1302    drop(worker);
1303
1304    let pool_ref = pool.clone();
1305    let cleanup_thread = std::thread::spawn(move || pool_ref.cleanup_idle_worker());
1306    rx.recv_timeout(Duration::from_secs(1)).unwrap();
1307    assert_eq!(cleanup_thread.join().unwrap(), 1);
1308}
1309
1310#[cfg(test)]
1311mod affected_drop_path_tests {
1312    use super::*;
1313    use std::collections::VecDeque;
1314    use std::sync::atomic::{AtomicBool, Ordering};
1315    use std::sync::mpsc;
1316
1317    type DropCallback = Box<dyn FnOnce() + Send>;
1318
1319    struct TestWorker {
1320        working: Arc<AtomicBool>,
1321        on_drop: Option<DropCallback>,
1322    }
1323
1324    #[async_trait::async_trait]
1325    impl Worker for TestWorker {
1326        fn is_work(&self) -> bool {
1327            self.working.load(Ordering::SeqCst)
1328        }
1329    }
1330
1331    impl Drop for TestWorker {
1332        fn drop(&mut self) {
1333            if let Some(on_drop) = self.on_drop.take() {
1334                on_drop();
1335            }
1336        }
1337    }
1338
1339    struct WorkerSpec {
1340        working: Arc<AtomicBool>,
1341        on_drop: Option<DropCallback>,
1342    }
1343
1344    struct TestWorkerFactory {
1345        specs: Arc<Mutex<VecDeque<WorkerSpec>>>,
1346    }
1347
1348    #[async_trait::async_trait]
1349    impl WorkerFactory<TestWorker> for TestWorkerFactory {
1350        async fn create(&self) -> PoolResult<TestWorker> {
1351            let spec = self.specs.lock().unwrap().pop_front().unwrap();
1352            Ok(TestWorker {
1353                working: spec.working,
1354                on_drop: spec.on_drop,
1355            })
1356        }
1357    }
1358
1359    fn new_pool() -> (
1360        WorkerPoolRef<TestWorker, TestWorkerFactory>,
1361        Arc<Mutex<VecDeque<WorkerSpec>>>,
1362    ) {
1363        let specs = Arc::new(Mutex::new(VecDeque::new()));
1364        let pool = WorkerPool::new(
1365            1,
1366            TestWorkerFactory {
1367                specs: specs.clone(),
1368            },
1369        );
1370        (pool, specs)
1371    }
1372
1373    fn lock_check_spec(
1374        pool: &WorkerPoolRef<TestWorker, TestWorkerFactory>,
1375        working: Arc<AtomicBool>,
1376    ) -> (WorkerSpec, mpsc::Receiver<bool>) {
1377        let (tx, rx) = mpsc::channel();
1378        let pool_ref = Arc::downgrade(pool);
1379        let on_drop = Box::new(move || {
1380            let pool_ref = pool_ref.upgrade().unwrap();
1381            tx.send(pool_ref.state.try_lock().is_ok()).unwrap();
1382        });
1383        (
1384            WorkerSpec {
1385                working,
1386                on_drop: Some(on_drop),
1387            },
1388            rx,
1389        )
1390    }
1391
1392    fn plain_spec() -> WorkerSpec {
1393        WorkerSpec {
1394            working: Arc::new(AtomicBool::new(true)),
1395            on_drop: None,
1396        }
1397    }
1398
1399    #[tokio::test]
1400    async fn test_invalid_idle_worker_is_dropped_outside_state_lock() {
1401        let (pool, specs) = new_pool();
1402        let working = Arc::new(AtomicBool::new(true));
1403        let (spec, drop_result) = lock_check_spec(&pool, working.clone());
1404        specs.lock().unwrap().push_back(spec);
1405
1406        let worker = pool.get_worker().await.unwrap();
1407        drop(worker);
1408        working.store(false, Ordering::SeqCst);
1409        specs.lock().unwrap().push_back(plain_spec());
1410
1411        let replacement = pool.get_worker().await.unwrap();
1412        assert!(drop_result.recv_timeout(Duration::from_secs(1)).unwrap());
1413        drop(replacement);
1414    }
1415
1416    #[tokio::test]
1417    async fn test_clear_drops_idle_worker_outside_state_lock() {
1418        let (pool, specs) = new_pool();
1419        let (spec, drop_result) = lock_check_spec(&pool, Arc::new(AtomicBool::new(true)));
1420        specs.lock().unwrap().push_back(spec);
1421
1422        let worker = pool.get_worker().await.unwrap();
1423        drop(worker);
1424        pool.clear_all_worker().await;
1425
1426        assert!(drop_result.recv_timeout(Duration::from_secs(1)).unwrap());
1427    }
1428
1429    #[tokio::test]
1430    async fn test_idle_limit_drops_worker_outside_state_lock() {
1431        let specs = Arc::new(Mutex::new(VecDeque::new()));
1432        let pool = WorkerPool::new_with_config(
1433            TestWorkerFactory {
1434                specs: specs.clone(),
1435            },
1436            WorkerPoolConfig {
1437                max_count: Some(1),
1438                max_idle_count: Some(0),
1439                ..Default::default()
1440            },
1441        );
1442        let (spec, drop_result) = lock_check_spec(&pool, Arc::new(AtomicBool::new(true)));
1443        specs.lock().unwrap().push_back(spec);
1444
1445        let worker = pool.get_worker().await.unwrap();
1446        drop(worker);
1447
1448        assert!(drop_result.recv_timeout(Duration::from_secs(1)).unwrap());
1449    }
1450}