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 max_count: Option<u16>,
42 pub idle_timeout: Option<Duration>,
43}
44
45#[async_trait::async_trait]
46pub 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 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}