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}