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