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    pub idle_timeout: Option<Duration>,
39}
40
41#[async_trait::async_trait]
42pub trait Worker: Send + 'static {
43    fn is_work(&self) -> bool;
44}
45
46pub struct WorkerGuard<W: Worker, F: WorkerFactory<W>> {
47    pool_ref: WorkerPoolRef<W, F>,
48    worker: Option<W>,
49}
50
51impl<W: Worker, F: WorkerFactory<W>> WorkerGuard<W, F> {
52    fn new(worker: W, pool_ref: WorkerPoolRef<W, F>) -> Self {
53        WorkerGuard {
54            pool_ref,
55            worker: Some(worker),
56        }
57    }
58}
59
60impl<W: Worker, F: WorkerFactory<W>> Deref for WorkerGuard<W, F> {
61    type Target = W;
62
63    fn deref(&self) -> &Self::Target {
64        self.worker.as_ref().unwrap()
65    }
66}
67
68impl<W: Worker, F: WorkerFactory<W>> DerefMut for WorkerGuard<W, F> {
69    fn deref_mut(&mut self) -> &mut Self::Target {
70        self.worker.as_mut().unwrap()
71    }
72}
73
74impl<W: Worker, F: WorkerFactory<W>> Drop for WorkerGuard<W, F> {
75    fn drop(&mut self) {
76        if let Some(worker) = self.worker.take() {
77            self.pool_ref.release(worker);
78        }
79    }
80}
81
82#[async_trait::async_trait]
83pub trait WorkerFactory<W: Worker>: Send + Sync + 'static {
84    async fn create(&self) -> PoolResult<W>;
85}
86
87struct IdleWorker<W: Worker> {
88    worker: W,
89    idle_since: Instant,
90}
91
92enum WorkerWaitResult<W: Worker, F: WorkerFactory<W>> {
93    Worker(WorkerGuard<W, F>),
94    Retry,
95    Error(PoolError),
96}
97
98struct WorkerPoolState<W: Worker, F: WorkerFactory<W>> {
99    current_count: u16,
100    worker_list: VecDeque<IdleWorker<W>>,
101    waiting_list: VecDeque<Notify<WorkerWaitResult<W, F>>>,
102    clearing: bool,
103    clear_waiting_list: Vec<Notify<()>>,
104}
105
106impl<W: Worker, F: WorkerFactory<W>> WorkerPoolState<W, F> {
107    fn take_clear_waiters_if_done(&mut self) -> Vec<Notify<()>> {
108        if self.clearing && self.current_count == 0 {
109            self.clearing = false;
110            self.clear_waiting_list.drain(..).collect()
111        } else {
112            Vec::new()
113        }
114    }
115
116    fn pop_next_waiter(&mut self) -> Option<Notify<WorkerWaitResult<W, F>>> {
117        while let Some(waiter) = self.waiting_list.pop_front() {
118            if !waiter.is_canceled() {
119                return Some(waiter);
120            }
121        }
122        None
123    }
124
125    fn drain_waiters(&mut self) -> Vec<Notify<WorkerWaitResult<W, F>>> {
126        self.waiting_list.drain(..).collect()
127    }
128}
129pub struct WorkerPool<W: Worker, F: WorkerFactory<W>> {
130    factory: Arc<F>,
131    max_count: u16,
132    config: WorkerPoolConfig,
133    state: Mutex<WorkerPoolState<W, F>>,
134}
135pub type WorkerPoolRef<W, F> = Arc<WorkerPool<W, F>>;
136
137impl<W: Worker, F: WorkerFactory<W>> WorkerPool<W, F> {
138    pub fn new(max_count: u16, factory: F) -> WorkerPoolRef<W, F> {
139        Self::new_with_config(max_count, factory, WorkerPoolConfig::default())
140    }
141
142    pub fn new_with_config(
143        max_count: u16,
144        factory: F,
145        config: WorkerPoolConfig,
146    ) -> WorkerPoolRef<W, F> {
147        Arc::new(WorkerPool {
148            factory: Arc::new(factory),
149            max_count,
150            config,
151            state: Mutex::new(WorkerPoolState {
152                current_count: 0,
153                worker_list: VecDeque::with_capacity(max_count as usize),
154                waiting_list: VecDeque::new(),
155                clearing: false,
156                clear_waiting_list: Vec::new(),
157            }),
158        })
159    }
160
161    fn remove_expired_idle_workers(
162        state: &mut WorkerPoolState<W, F>,
163        idle_timeout: Option<Duration>,
164    ) -> u16 {
165        let Some(idle_timeout) = idle_timeout else {
166            return 0;
167        };
168        let mut removed_count = 0;
169        let now = Instant::now();
170        while state
171            .worker_list
172            .front()
173            .map(|idle_worker| now.duration_since(idle_worker.idle_since) >= idle_timeout)
174            .unwrap_or(false)
175        {
176            state.worker_list.pop_front();
177            state.current_count -= 1;
178            removed_count += 1;
179        }
180        removed_count
181    }
182
183    pub fn cleanup_idle_worker(&self) -> u16 {
184        let (removed_count, clear_waiters) = {
185            let mut state = self.state.lock().unwrap();
186            let removed_count =
187                Self::remove_expired_idle_workers(&mut state, self.config.idle_timeout);
188            let clear_waiters = state.take_clear_waiters_if_done();
189            (removed_count, clear_waiters)
190        };
191        for waiter in clear_waiters {
192            waiter.notify(());
193        }
194        removed_count
195    }
196
197    pub async fn get_worker(self: &WorkerPoolRef<W, F>) -> PoolResult<WorkerGuard<W, F>> {
198        loop {
199            if self.max_count == 0 {
200                return Err(pool_invalid_config_error("pool max_count is zero"));
201            }
202
203            let wait = {
204                let mut state = self.state.lock().unwrap();
205                if state.clearing {
206                    return Err(pool_clearing_error());
207                }
208
209                Self::remove_expired_idle_workers(&mut state, self.config.idle_timeout);
210
211                while let Some(idle_worker) = state.worker_list.pop_back() {
212                    let worker = idle_worker.worker;
213                    if !worker.is_work() {
214                        state.current_count -= 1;
215                        continue;
216                    }
217                    return Ok(WorkerGuard::new(worker, self.clone()));
218                }
219
220                if state.current_count < self.max_count {
221                    state.current_count += 1;
222                    None
223                } else {
224                    let (notify, waiter) = Notify::new();
225                    state.waiting_list.push_back(notify);
226                    Some(waiter)
227                }
228            };
229
230            if let Some(wait) = wait {
231                match wait.await {
232                    WorkerWaitResult::Worker(worker) => return Ok(worker),
233                    WorkerWaitResult::Retry => continue,
234                    WorkerWaitResult::Error(err) => return Err(err),
235                }
236            }
237
238            let worker = match self.factory.create().await {
239                Ok(worker) => worker,
240                Err(err) => {
241                    let (retry_waiters, clear_waiters) = {
242                        let mut state = self.state.lock().unwrap();
243                        state.current_count -= 1;
244                        let retry_waiters = state.drain_waiters();
245                        let clear_waiters = state.take_clear_waiters_if_done();
246                        (retry_waiters, clear_waiters)
247                    };
248                    Self::notify_retry_waiters(retry_waiters);
249                    for waiter in clear_waiters {
250                        waiter.notify(());
251                    }
252                    return Err(err);
253                }
254            };
255            let (clearing, clear_waiters) = {
256                let mut state = self.state.lock().unwrap();
257                if state.clearing {
258                    state.current_count -= 1;
259                    (true, state.take_clear_waiters_if_done())
260                } else {
261                    (false, Vec::new())
262                }
263            };
264            for waiter in clear_waiters {
265                waiter.notify(());
266            }
267            if clearing {
268                return Err(pool_cleared_error());
269            }
270            return Ok(WorkerGuard::new(worker, self.clone()));
271        }
272    }
273
274    pub async fn clear_all_worker(&self) {
275        let (waiter, waiting_list, clear_waiters) = {
276            let mut state = self.state.lock().unwrap();
277            if !state.clearing {
278                state.clearing = true;
279                let cur_worker_count = state.worker_list.len();
280                state.worker_list.clear();
281                state.current_count -= cur_worker_count as u16;
282            }
283
284            let waiting_list = state.waiting_list.drain(..).collect::<Vec<_>>();
285            if state.current_count == 0 {
286                let clear_waiters = state.take_clear_waiters_if_done();
287                (None, waiting_list, clear_waiters)
288            } else {
289                let (notify, waiter) = Notify::new();
290                state.clear_waiting_list.push(notify);
291                (Some(waiter), waiting_list, Vec::new())
292            }
293        };
294        for waiting in waiting_list {
295            waiting.notify(WorkerWaitResult::Error(pool_cleared_error()));
296        }
297        for waiter in clear_waiters {
298            waiter.notify(());
299        }
300        if let Some(waiter) = waiter {
301            waiter.await;
302        }
303    }
304
305    fn notify_retry_waiters(waiters: Vec<Notify<WorkerWaitResult<W, F>>>) {
306        for waiter in waiters {
307            waiter.notify(WorkerWaitResult::Retry);
308        }
309    }
310
311    fn release(self: &WorkerPoolRef<W, F>, work: W) {
312        enum ReleaseAction<W: Worker, F: WorkerFactory<W>> {
313            None,
314            Notify(Notify<WorkerWaitResult<W, F>>, WorkerGuard<W, F>),
315            Retry(Vec<Notify<WorkerWaitResult<W, F>>>),
316        }
317
318        let mut clear_waiters = Vec::new();
319        let action = {
320            let mut state = self.state.lock().unwrap();
321            if state.clearing {
322                state.current_count -= 1;
323                clear_waiters = state.take_clear_waiters_if_done();
324                ReleaseAction::None
325            } else if work.is_work() {
326                let future = state.pop_next_waiter();
327                if let Some(future) = future {
328                    ReleaseAction::Notify(future, WorkerGuard::new(work, self.clone()))
329                } else {
330                    state.worker_list.push_back(IdleWorker {
331                        worker: work,
332                        idle_since: Instant::now(),
333                    });
334                    ReleaseAction::None
335                }
336            } else {
337                state.current_count -= 1;
338                let waiters = state.drain_waiters();
339                if !waiters.is_empty() {
340                    ReleaseAction::Retry(waiters)
341                } else {
342                    clear_waiters = state.take_clear_waiters_if_done();
343                    ReleaseAction::None
344                }
345            }
346        };
347
348        for waiter in clear_waiters {
349            waiter.notify(());
350        }
351
352        match action {
353            ReleaseAction::None => {}
354            ReleaseAction::Notify(future, worker) => {
355                future.notify(WorkerWaitResult::Worker(worker));
356            }
357            ReleaseAction::Retry(waiters) => {
358                Self::notify_retry_waiters(waiters);
359            }
360        }
361    }
362}
363
364#[tokio::test]
365async fn test_pool() {
366    struct TestWorker {
367        work: bool,
368    }
369
370    #[async_trait::async_trait]
371    impl Worker for TestWorker {
372        fn is_work(&self) -> bool {
373            self.work
374        }
375    }
376
377    struct TestWorkerFactory;
378
379    #[async_trait::async_trait]
380    impl WorkerFactory<TestWorker> for TestWorkerFactory {
381        async fn create(&self) -> PoolResult<TestWorker> {
382            Ok(TestWorker { work: true })
383        }
384    }
385
386    let pool = WorkerPool::new(2, TestWorkerFactory);
387
388    let worker1 = pool.get_worker().await.unwrap();
389    let worker2 = pool.get_worker().await.unwrap();
390
391    let pool_ref = pool.clone();
392    let waiter = tokio::spawn(async move { pool_ref.get_worker().await });
393    tokio::time::sleep(std::time::Duration::from_millis(20)).await;
394    assert!(!waiter.is_finished());
395
396    drop(worker1);
397    let worker3 = tokio::time::timeout(std::time::Duration::from_secs(1), waiter)
398        .await
399        .unwrap()
400        .unwrap()
401        .unwrap();
402    drop(worker2);
403    drop(worker3);
404
405    let worker1 = pool.get_worker().await.unwrap();
406    let worker2 = pool.get_worker().await.unwrap();
407
408    let pool_ref = pool.clone();
409    let waiter1 = tokio::spawn(async move { pool_ref.get_worker().await });
410    let pool_ref = pool.clone();
411    let waiter2 = tokio::spawn(async move { pool_ref.get_worker().await });
412    tokio::time::sleep(std::time::Duration::from_millis(20)).await;
413    assert!(!waiter1.is_finished());
414    assert!(!waiter2.is_finished());
415
416    let pool_ref = pool.clone();
417    let clear_task = tokio::spawn(async move {
418        pool_ref.clear_all_worker().await;
419    });
420    tokio::time::sleep(std::time::Duration::from_millis(20)).await;
421
422    assert!(waiter1.await.unwrap().is_err());
423    assert!(waiter2.await.unwrap().is_err());
424
425    drop(worker1);
426    drop(worker2);
427
428    tokio::time::timeout(std::time::Duration::from_secs(1), clear_task)
429        .await
430        .unwrap()
431        .unwrap();
432}
433
434#[tokio::test]
435async fn test_clear_all_worker_waits_for_inflight_create() {
436    use std::sync::atomic::{AtomicUsize, Ordering};
437    use std::sync::Arc;
438
439    struct TestWorker;
440
441    #[async_trait::async_trait]
442    impl Worker for TestWorker {
443        fn is_work(&self) -> bool {
444            true
445        }
446    }
447
448    struct TestWorkerFactory {
449        create_count: Arc<AtomicUsize>,
450    }
451
452    #[async_trait::async_trait]
453    impl WorkerFactory<TestWorker> for TestWorkerFactory {
454        async fn create(&self) -> PoolResult<TestWorker> {
455            self.create_count.fetch_add(1, Ordering::SeqCst);
456            tokio::time::sleep(std::time::Duration::from_millis(100)).await;
457            Ok(TestWorker)
458        }
459    }
460
461    let create_count = Arc::new(AtomicUsize::new(0));
462    let pool = WorkerPool::new(
463        1,
464        TestWorkerFactory {
465            create_count: create_count.clone(),
466        },
467    );
468
469    let pool_ref = pool.clone();
470    let worker_task = tokio::spawn(async move { pool_ref.get_worker().await });
471    tokio::time::sleep(std::time::Duration::from_millis(20)).await;
472
473    pool.clear_all_worker().await;
474
475    let worker = worker_task.await.unwrap();
476    assert!(worker.is_err());
477    assert_eq!(create_count.load(Ordering::SeqCst), 1);
478}
479
480#[tokio::test]
481async fn test_concurrent_clear_all_worker() {
482    struct TestWorker;
483
484    #[async_trait::async_trait]
485    impl Worker for TestWorker {
486        fn is_work(&self) -> bool {
487            true
488        }
489    }
490
491    struct TestWorkerFactory;
492
493    #[async_trait::async_trait]
494    impl WorkerFactory<TestWorker> for TestWorkerFactory {
495        async fn create(&self) -> PoolResult<TestWorker> {
496            Ok(TestWorker)
497        }
498    }
499
500    let pool = WorkerPool::new(1, TestWorkerFactory);
501    let worker = pool.get_worker().await.unwrap();
502
503    let pool_ref = pool.clone();
504    let clear_task1 = tokio::spawn(async move {
505        pool_ref.clear_all_worker().await;
506    });
507
508    let pool_ref = pool.clone();
509    let clear_task2 = tokio::spawn(async move {
510        pool_ref.clear_all_worker().await;
511    });
512
513    tokio::time::sleep(std::time::Duration::from_millis(20)).await;
514    drop(worker);
515
516    tokio::time::timeout(std::time::Duration::from_secs(1), async {
517        clear_task1.await.unwrap();
518        clear_task2.await.unwrap();
519    })
520    .await
521    .unwrap();
522}
523
524#[tokio::test]
525async fn test_zero_max_count_returns_error() {
526    struct TestWorker;
527
528    #[async_trait::async_trait]
529    impl Worker for TestWorker {
530        fn is_work(&self) -> bool {
531            true
532        }
533    }
534
535    struct TestWorkerFactory;
536
537    #[async_trait::async_trait]
538    impl WorkerFactory<TestWorker> for TestWorkerFactory {
539        async fn create(&self) -> PoolResult<TestWorker> {
540            Ok(TestWorker)
541        }
542    }
543
544    let pool = WorkerPool::new(0, TestWorkerFactory);
545    let worker = pool.get_worker().await;
546    assert!(worker.is_err());
547    assert_eq!(worker.err().unwrap().code(), PoolErrorCode::InvalidConfig);
548}
549
550#[tokio::test]
551async fn test_create_failure_fails_waiting_workers() {
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            tokio::time::sleep(std::time::Duration::from_millis(50)).await;
567            Err(pool_invalid_config_error("create failed"))
568        }
569    }
570
571    let pool = WorkerPool::new(1, TestWorkerFactory);
572
573    let pool_ref = pool.clone();
574    let worker1 = tokio::spawn(async move { pool_ref.get_worker().await });
575    tokio::time::sleep(std::time::Duration::from_millis(20)).await;
576
577    let pool_ref = pool.clone();
578    let worker2 = tokio::spawn(async move { pool_ref.get_worker().await });
579
580    let (worker1, worker2) = tokio::time::timeout(std::time::Duration::from_secs(1), async {
581        (worker1.await.unwrap(), worker2.await.unwrap())
582    })
583    .await
584    .unwrap();
585
586    assert_eq!(worker1.err().unwrap().code(), PoolErrorCode::InvalidConfig);
587    assert_eq!(worker2.err().unwrap().code(), PoolErrorCode::InvalidConfig);
588}
589
590#[tokio::test]
591async fn test_invalid_worker_drop_outside_runtime_wakes_waiter() {
592    use std::sync::atomic::{AtomicUsize, Ordering};
593    use std::sync::Arc;
594
595    struct TestWorker {
596        id: usize,
597        work: bool,
598    }
599
600    #[async_trait::async_trait]
601    impl Worker for TestWorker {
602        fn is_work(&self) -> bool {
603            self.work
604        }
605    }
606
607    struct TestWorkerFactory {
608        create_count: Arc<AtomicUsize>,
609    }
610
611    #[async_trait::async_trait]
612    impl WorkerFactory<TestWorker> for TestWorkerFactory {
613        async fn create(&self) -> PoolResult<TestWorker> {
614            let id = self.create_count.fetch_add(1, Ordering::SeqCst);
615            Ok(TestWorker { id, work: true })
616        }
617    }
618
619    let create_count = Arc::new(AtomicUsize::new(0));
620    let pool = WorkerPool::new(
621        1,
622        TestWorkerFactory {
623            create_count: create_count.clone(),
624        },
625    );
626
627    let mut worker = pool.get_worker().await.unwrap();
628    assert_eq!(worker.id, 0);
629    worker.work = false;
630
631    let pool_ref = pool.clone();
632    let waiter = tokio::spawn(async move { pool_ref.get_worker().await });
633    tokio::time::sleep(std::time::Duration::from_millis(20)).await;
634    assert!(!waiter.is_finished());
635
636    std::thread::spawn(move || drop(worker)).join().unwrap();
637
638    let worker = tokio::time::timeout(std::time::Duration::from_secs(1), waiter)
639        .await
640        .unwrap()
641        .unwrap()
642        .unwrap();
643    assert_eq!(worker.id, 1);
644    assert_eq!(create_count.load(Ordering::SeqCst), 2);
645}
646
647#[tokio::test]
648async fn test_retry_notification_skips_canceled_waiter() {
649    struct TestWorker;
650
651    #[async_trait::async_trait]
652    impl Worker for TestWorker {
653        fn is_work(&self) -> bool {
654            true
655        }
656    }
657
658    struct TestWorkerFactory;
659
660    #[async_trait::async_trait]
661    impl WorkerFactory<TestWorker> for TestWorkerFactory {
662        async fn create(&self) -> PoolResult<TestWorker> {
663            Ok(TestWorker)
664        }
665    }
666
667    let (canceled_notify, canceled_waiter) = Notify::new();
668    drop(canceled_waiter);
669    let (notify, waiter) = Notify::new();
670
671    WorkerPool::<TestWorker, TestWorkerFactory>::notify_retry_waiters(vec![
672        canceled_notify,
673        notify,
674    ]);
675
676    let result = tokio::time::timeout(std::time::Duration::from_secs(1), waiter)
677        .await
678        .unwrap();
679    assert!(matches!(result, WorkerWaitResult::Retry));
680}
681
682#[tokio::test]
683async fn test_clearing_and_cleared_error_codes() {
684    use std::sync::atomic::{AtomicBool, Ordering};
685    use std::sync::Arc;
686
687    struct TestWorker;
688
689    #[async_trait::async_trait]
690    impl Worker for TestWorker {
691        fn is_work(&self) -> bool {
692            true
693        }
694    }
695
696    struct TestWorkerFactory {
697        should_block: Arc<AtomicBool>,
698    }
699
700    #[async_trait::async_trait]
701    impl WorkerFactory<TestWorker> for TestWorkerFactory {
702        async fn create(&self) -> PoolResult<TestWorker> {
703            while self.should_block.load(Ordering::SeqCst) {
704                tokio::task::yield_now().await;
705            }
706            Ok(TestWorker)
707        }
708    }
709
710    let should_block = Arc::new(AtomicBool::new(true));
711    let pool = WorkerPool::new(
712        1,
713        TestWorkerFactory {
714            should_block: should_block.clone(),
715        },
716    );
717
718    let pool_ref = pool.clone();
719    let inflight = tokio::spawn(async move { pool_ref.get_worker().await });
720    tokio::task::yield_now().await;
721
722    let pool_ref = pool.clone();
723    let clear_task = tokio::spawn(async move {
724        pool_ref.clear_all_worker().await;
725    });
726    tokio::task::yield_now().await;
727
728    let err = pool.get_worker().await.err().unwrap();
729    assert_eq!(err.code(), PoolErrorCode::Clearing);
730
731    should_block.store(false, Ordering::SeqCst);
732    clear_task.await.unwrap();
733
734    let err = inflight.await.unwrap().err().unwrap();
735    assert_eq!(err.code(), PoolErrorCode::Cleared);
736}
737
738#[tokio::test]
739async fn test_idle_worker_timeout_releases_worker() {
740    use std::sync::atomic::{AtomicUsize, Ordering};
741    use std::sync::Arc;
742
743    struct TestWorker {
744        id: usize,
745    }
746
747    #[async_trait::async_trait]
748    impl Worker for TestWorker {
749        fn is_work(&self) -> bool {
750            true
751        }
752    }
753
754    struct TestWorkerFactory {
755        create_count: Arc<AtomicUsize>,
756    }
757
758    #[async_trait::async_trait]
759    impl WorkerFactory<TestWorker> for TestWorkerFactory {
760        async fn create(&self) -> PoolResult<TestWorker> {
761            let id = self.create_count.fetch_add(1, Ordering::SeqCst);
762            Ok(TestWorker { id })
763        }
764    }
765
766    let create_count = Arc::new(AtomicUsize::new(0));
767    let pool = WorkerPool::new_with_config(
768        1,
769        TestWorkerFactory {
770            create_count: create_count.clone(),
771        },
772        WorkerPoolConfig {
773            idle_timeout: Some(std::time::Duration::from_millis(30)),
774            ..WorkerPoolConfig::default()
775        },
776    );
777
778    {
779        let worker = pool.get_worker().await.unwrap();
780        assert_eq!(worker.id, 0);
781    }
782
783    tokio::time::sleep(std::time::Duration::from_millis(80)).await;
784
785    let worker = pool.get_worker().await.unwrap();
786    assert_eq!(worker.id, 1);
787    assert_eq!(create_count.load(Ordering::SeqCst), 2);
788}
789
790#[tokio::test]
791async fn test_idle_worker_reused_before_timeout() {
792    use std::sync::atomic::{AtomicUsize, Ordering};
793    use std::sync::Arc;
794
795    struct TestWorker {
796        id: usize,
797    }
798
799    #[async_trait::async_trait]
800    impl Worker for TestWorker {
801        fn is_work(&self) -> bool {
802            true
803        }
804    }
805
806    struct TestWorkerFactory {
807        create_count: Arc<AtomicUsize>,
808    }
809
810    #[async_trait::async_trait]
811    impl WorkerFactory<TestWorker> for TestWorkerFactory {
812        async fn create(&self) -> PoolResult<TestWorker> {
813            let id = self.create_count.fetch_add(1, Ordering::SeqCst);
814            Ok(TestWorker { id })
815        }
816    }
817
818    let create_count = Arc::new(AtomicUsize::new(0));
819    let pool = WorkerPool::new_with_config(
820        1,
821        TestWorkerFactory {
822            create_count: create_count.clone(),
823        },
824        WorkerPoolConfig {
825            idle_timeout: Some(std::time::Duration::from_secs(1)),
826            ..WorkerPoolConfig::default()
827        },
828    );
829
830    {
831        let worker = pool.get_worker().await.unwrap();
832        assert_eq!(worker.id, 0);
833    }
834
835    tokio::time::sleep(std::time::Duration::from_millis(20)).await;
836
837    let worker = pool.get_worker().await.unwrap();
838    assert_eq!(worker.id, 0);
839    assert_eq!(create_count.load(Ordering::SeqCst), 1);
840}
841
842#[tokio::test]
843async fn test_cleanup_idle_worker_can_be_triggered_externally() {
844    use std::sync::atomic::{AtomicUsize, Ordering};
845    use std::sync::Arc;
846
847    struct TestWorker {
848        id: usize,
849    }
850
851    #[async_trait::async_trait]
852    impl Worker for TestWorker {
853        fn is_work(&self) -> bool {
854            true
855        }
856    }
857
858    struct TestWorkerFactory {
859        create_count: Arc<AtomicUsize>,
860    }
861
862    #[async_trait::async_trait]
863    impl WorkerFactory<TestWorker> for TestWorkerFactory {
864        async fn create(&self) -> PoolResult<TestWorker> {
865            let id = self.create_count.fetch_add(1, Ordering::SeqCst);
866            Ok(TestWorker { id })
867        }
868    }
869
870    let create_count = Arc::new(AtomicUsize::new(0));
871    let pool = WorkerPool::new_with_config(
872        1,
873        TestWorkerFactory {
874            create_count: create_count.clone(),
875        },
876        WorkerPoolConfig {
877            idle_timeout: Some(std::time::Duration::from_millis(30)),
878            ..WorkerPoolConfig::default()
879        },
880    );
881
882    {
883        let worker = pool.get_worker().await.unwrap();
884        assert_eq!(worker.id, 0);
885    }
886
887    tokio::time::sleep(std::time::Duration::from_millis(80)).await;
888
889    assert_eq!(pool.cleanup_idle_worker(), 1);
890
891    let worker = pool.get_worker().await.unwrap();
892    assert_eq!(worker.id, 1);
893    assert_eq!(create_count.load(Ordering::SeqCst), 2);
894}
895
896#[tokio::test]
897async fn test_get_worker_uses_most_recent_idle_worker() {
898    use std::sync::atomic::{AtomicUsize, Ordering};
899    use std::sync::Arc;
900
901    struct TestWorker {
902        id: usize,
903    }
904
905    #[async_trait::async_trait]
906    impl Worker for TestWorker {
907        fn is_work(&self) -> bool {
908            true
909        }
910    }
911
912    struct TestWorkerFactory {
913        create_count: Arc<AtomicUsize>,
914    }
915
916    #[async_trait::async_trait]
917    impl WorkerFactory<TestWorker> for TestWorkerFactory {
918        async fn create(&self) -> PoolResult<TestWorker> {
919            let id = self.create_count.fetch_add(1, Ordering::SeqCst);
920            Ok(TestWorker { id })
921        }
922    }
923
924    let create_count = Arc::new(AtomicUsize::new(0));
925    let pool = WorkerPool::new(
926        2,
927        TestWorkerFactory {
928            create_count: create_count.clone(),
929        },
930    );
931
932    let worker1 = pool.get_worker().await.unwrap();
933    let worker2 = pool.get_worker().await.unwrap();
934    assert_eq!(worker1.id, 0);
935    assert_eq!(worker2.id, 1);
936
937    drop(worker1);
938    drop(worker2);
939
940    let worker = pool.get_worker().await.unwrap();
941    assert_eq!(worker.id, 1);
942    assert_eq!(create_count.load(Ordering::SeqCst), 2);
943}