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