1use crate::{
2 pool_cleared_error, pool_clearing_error, pool_invalid_config_error, PoolError, PoolResult,
3};
4use notify_future::Notify;
5use std::collections::{HashMap, VecDeque};
6use std::hash::Hash;
7use std::ops::{Deref, DerefMut};
8use std::sync::{Arc, Mutex};
9use std::time::{Duration, Instant};
10
11pub trait WorkerKey: Send + 'static + Clone + Hash + Eq + PartialEq {}
12
13impl<T: Send + 'static + Clone + Hash + Eq + PartialEq> WorkerKey for T {}
14
15#[derive(Debug, Clone, Default)]
16pub struct KeyedWorkerPoolConfig {
20 idle_timeout: Option<Duration>,
21 max_count_per_key: Option<u16>,
25 max_idle_count_per_key: Option<u16>,
30}
31
32impl KeyedWorkerPoolConfig {
33 pub fn with_idle_timeout(mut self, idle_timeout: Option<Duration>) -> Self {
37 self.idle_timeout = idle_timeout;
38 self
39 }
40
41 pub fn with_max_count_per_key(mut self, max_count: Option<u16>) -> Self {
45 self.max_count_per_key = max_count;
46 self
47 }
48
49 pub fn with_max_idle_count_per_key(mut self, max_idle_count: Option<u16>) -> Self {
53 self.max_idle_count_per_key = max_idle_count;
54 self
55 }
56}
57
58#[async_trait::async_trait]
59pub trait KeyedWorker<K: WorkerKey>: Send + 'static {
64 fn is_work(&self) -> bool;
65 fn supports(&self, key: K) -> bool;
69 fn primary_key(&self) -> K;
73}
74
75pub struct KeyedWorkerGuard<K: WorkerKey, W: KeyedWorker<K>, F: KeyedWorkerFactory<K, W>> {
76 pool_ref: KeyedWorkerPoolRef<K, W, F>,
77 worker: Option<W>,
78 primary_key: K,
79}
80
81impl<K: WorkerKey, W: KeyedWorker<K>, F: KeyedWorkerFactory<K, W>> KeyedWorkerGuard<K, W, F> {
82 fn new(worker: W, pool_ref: KeyedWorkerPoolRef<K, W, F>, primary_key: K) -> Self {
83 KeyedWorkerGuard {
84 pool_ref,
85 worker: Some(worker),
86 primary_key,
87 }
88 }
89}
90
91impl<K: WorkerKey, W: KeyedWorker<K>, F: KeyedWorkerFactory<K, W>> DerefMut
92 for KeyedWorkerGuard<K, W, F>
93{
94 fn deref_mut(&mut self) -> &mut Self::Target {
95 self.worker.as_mut().unwrap()
96 }
97}
98
99impl<K: WorkerKey, W: KeyedWorker<K>, F: KeyedWorkerFactory<K, W>> Deref
100 for KeyedWorkerGuard<K, W, F>
101{
102 type Target = W;
103
104 fn deref(&self) -> &Self::Target {
105 self.worker.as_ref().unwrap()
106 }
107}
108
109impl<K: WorkerKey, W: KeyedWorker<K>, F: KeyedWorkerFactory<K, W>> Drop
110 for KeyedWorkerGuard<K, W, F>
111{
112 fn drop(&mut self) {
113 if let Some(worker) = self.worker.take() {
114 self.pool_ref.release(worker, self.primary_key.clone());
115 }
116 }
117}
118
119struct KeyedWorkerReservation<K: WorkerKey, W: KeyedWorker<K>, F: KeyedWorkerFactory<K, W>> {
120 pool_ref: KeyedWorkerPoolRef<K, W, F>,
121 requested_key: K,
122 active: bool,
123}
124
125enum ReservationCompletion {
126 Complete,
127 Clearing,
128}
129
130impl<K: WorkerKey, W: KeyedWorker<K>, F: KeyedWorkerFactory<K, W>> KeyedWorkerReservation<K, W, F> {
131 fn new(pool_ref: KeyedWorkerPoolRef<K, W, F>, key: K) -> Self {
132 Self {
133 pool_ref,
134 requested_key: key,
135 active: true,
136 }
137 }
138
139 fn complete(mut self, worker_key: K) -> ReservationCompletion {
140 let (completion, clear_waiters) = {
141 let mut state = self.pool_ref.state.lock().unwrap();
142 state.dec_pending_count_for_key(self.requested_key.clone());
143 if state.clearing {
144 state.current_count -= 1;
145 (
146 ReservationCompletion::Clearing,
147 state.take_clear_waiters_if_done(),
148 )
149 } else {
150 state.inc_worker_count_for_key(worker_key);
151 (ReservationCompletion::Complete, Vec::new())
152 }
153 };
154 self.active = false;
155 for waiter in clear_waiters {
156 waiter.notify(());
157 }
158 completion
159 }
160}
161
162impl<K: WorkerKey, W: KeyedWorker<K>, F: KeyedWorkerFactory<K, W>> Drop
163 for KeyedWorkerReservation<K, W, F>
164{
165 fn drop(&mut self) {
166 if self.active {
167 self.pool_ref.rollback_reservation(&self.requested_key);
168 }
169 }
170}
171
172#[async_trait::async_trait]
173pub trait KeyedWorkerFactory<K: WorkerKey, W: KeyedWorker<K>>: Send + Sync + 'static {
174 async fn create(&self, key: K) -> PoolResult<W>;
179}
180
181struct WaitingItem<K: WorkerKey, W: KeyedWorker<K>, F: KeyedWorkerFactory<K, W>> {
182 future: Notify<KeyedWorkerWaitResult<K, W, F>>,
183 key: K,
184}
185
186struct IdleWorker<K, W> {
187 worker: W,
188 primary_key: K,
189 idle_since: Instant,
190}
191
192enum KeyedWorkerWaitResult<K: WorkerKey, W: KeyedWorker<K>, F: KeyedWorkerFactory<K, W>> {
193 Worker(KeyedWorkerGuard<K, W, F>),
194 Retry,
195 Error(PoolError),
196}
197
198struct WorkerPoolState<K: WorkerKey, W: KeyedWorker<K>, F: KeyedWorkerFactory<K, W>> {
199 current_count: usize,
200 worker_count_by_key: HashMap<K, usize>,
201 pending_count_by_key: HashMap<K, usize>,
202 worker_list: HashMap<K, VecDeque<IdleWorker<K, W>>>,
204 waiting_list: Vec<WaitingItem<K, W, F>>,
205 clearing: bool,
206 clear_waiting_list: Vec<Notify<()>>,
207}
208
209impl<K: WorkerKey, W: KeyedWorker<K>, F: KeyedWorkerFactory<K, W>> WorkerPoolState<K, W, F> {
210 fn push_idle_worker(&mut self, worker: W, primary_key: K) {
211 let idle_worker = IdleWorker {
212 worker,
213 primary_key: primary_key.clone(),
214 idle_since: Instant::now(),
215 };
216 self.worker_list
217 .entry(primary_key)
218 .or_default()
219 .push_back(idle_worker);
220 }
221
222 fn remove_idle_worker(&mut self, primary_key: &K, index: usize) -> IdleWorker<K, W> {
223 let (idle_worker, remove_bucket) = {
224 let workers = self.worker_list.get_mut(primary_key).unwrap();
225 let idle_worker = workers.remove(index).unwrap();
226 (idle_worker, workers.is_empty())
227 };
228 if remove_bucket {
229 self.worker_list.remove(primary_key);
230 }
231 idle_worker
232 }
233
234 fn take_idle_worker_for_key(&mut self, key: &K) -> Option<IdleWorker<K, W>> {
235 let index = self
236 .worker_list
237 .get(key)
238 .map(VecDeque::len)
239 .and_then(|len| len.checked_sub(1))?;
240 Some(self.remove_idle_worker(key, index))
241 }
242
243 fn drain_idle_workers(&mut self) -> Vec<IdleWorker<K, W>> {
244 std::mem::take(&mut self.worker_list)
245 .into_values()
246 .flatten()
247 .collect()
248 }
249
250 fn inc_worker_count_for_key(&mut self, key: K) {
251 let count = self.worker_count_by_key.entry(key).or_insert(0);
252 *count += 1;
253 }
254
255 fn dec_worker_count_for_key(&mut self, key: K) {
256 let mut should_remove = false;
257 if let Some(count) = self.worker_count_by_key.get_mut(&key) {
258 debug_assert!(*count > 0);
259 *count -= 1;
260 should_remove = *count == 0;
261 }
262 if should_remove {
263 self.worker_count_by_key.remove(&key);
264 }
265 }
266
267 fn inc_pending_count_for_key(&mut self, key: K) {
268 let count = self.pending_count_by_key.entry(key).or_insert(0);
269 *count += 1;
270 }
271
272 fn dec_pending_count_for_key(&mut self, key: K) {
273 let mut should_remove = false;
274 if let Some(count) = self.pending_count_by_key.get_mut(&key) {
275 debug_assert!(*count > 0);
276 *count -= 1;
277 should_remove = *count == 0;
278 }
279 if should_remove {
280 self.pending_count_by_key.remove(&key);
281 }
282 }
283
284 fn reserved_count_for_key(&self, key: &K) -> usize {
285 self.worker_count_by_key.get(key).copied().unwrap_or(0)
286 + self.pending_count_by_key.get(key).copied().unwrap_or(0)
287 }
288
289 fn take_clear_waiters_if_done(&mut self) -> Vec<Notify<()>> {
290 if self.clearing && self.current_count == 0 {
291 self.clearing = false;
292 self.clear_waiting_list.drain(..).collect()
293 } else {
294 Vec::new()
295 }
296 }
297
298 fn find_matching_waiter_index_for_worker(&self, worker: &W, primary_key: &K) -> Option<usize> {
299 self.waiting_list.iter().position(|waiting| {
300 if waiting.future.is_canceled() {
301 return false;
302 }
303 waiting.key == *primary_key && worker.supports(waiting.key.clone())
304 })
305 }
306
307 fn remove_canceled_waiters(&mut self) {
308 self.waiting_list
309 .retain(|waiting| !waiting.future.is_canceled());
310 }
311
312 fn drain_waiters(&mut self) -> Vec<Notify<KeyedWorkerWaitResult<K, W, F>>> {
313 self.waiting_list
314 .drain(..)
315 .map(|waiting| waiting.future)
316 .collect()
317 }
318}
319
320pub struct KeyedWorkerPool<K: WorkerKey, W: KeyedWorker<K>, F: KeyedWorkerFactory<K, W>> {
321 factory: Arc<F>,
322 config: KeyedWorkerPoolConfig,
323 state: Mutex<WorkerPoolState<K, W, F>>,
324}
325pub type KeyedWorkerPoolRef<K, W, F> = Arc<KeyedWorkerPool<K, W, F>>;
326
327#[cfg(test)]
328#[test]
329fn test_keyed_worker_pool_config_default_idle_limits() {
330 let config = KeyedWorkerPoolConfig::default();
331 assert_eq!(config.max_idle_count_per_key, None);
332}
333
334#[cfg(test)]
335#[test]
336fn test_keyed_worker_pool_config_builder() {
337 let timeout = Duration::from_secs(1);
338 let config = KeyedWorkerPoolConfig::default()
339 .with_idle_timeout(Some(timeout))
340 .with_max_count_per_key(Some(2))
341 .with_max_idle_count_per_key(Some(1));
342 assert_eq!(config.idle_timeout, Some(timeout));
343 assert_eq!(config.max_count_per_key, Some(2));
344 assert_eq!(config.max_idle_count_per_key, Some(1));
345}
346
347impl<K: WorkerKey, W: KeyedWorker<K>, F: KeyedWorkerFactory<K, W>> KeyedWorkerPool<K, W, F> {
348 fn key_limit_reached(&self, state: &WorkerPoolState<K, W, F>, key: &K) -> bool {
349 self.config
350 .max_count_per_key
351 .map(|max_count| state.reserved_count_for_key(key) >= usize::from(max_count))
352 .unwrap_or(false)
353 }
354
355 fn validate_created_worker(requested_key: &K, worker: &W) -> PoolResult<K> {
356 let worker_key = worker.primary_key();
357 if !worker.supports(worker_key.clone()) {
358 return Err(pool_invalid_config_error(
359 "worker primary key is not valid for itself",
360 ));
361 }
362 if worker_key != requested_key.clone() {
363 return Err(pool_invalid_config_error(
364 "factory returned worker with mismatched key",
365 ));
366 }
367 Ok(worker_key)
368 }
369
370 pub fn new(factory: F, config: KeyedWorkerPoolConfig) -> KeyedWorkerPoolRef<K, W, F> {
372 Arc::new(KeyedWorkerPool {
373 factory: Arc::new(factory),
374 config,
375 state: Mutex::new(WorkerPoolState {
376 current_count: 0,
377 worker_count_by_key: HashMap::new(),
378 pending_count_by_key: HashMap::new(),
379 worker_list: HashMap::new(),
380 waiting_list: Vec::new(),
381 clearing: false,
382 clear_waiting_list: Vec::new(),
383 }),
384 })
385 }
386
387 fn take_expired_idle_workers(
388 state: &mut WorkerPoolState<K, W, F>,
389 idle_timeout: Option<std::time::Duration>,
390 ) -> Vec<IdleWorker<K, W>> {
391 let Some(idle_timeout) = idle_timeout else {
392 return Vec::new();
393 };
394 let mut removed_workers = Vec::new();
395 let now = Instant::now();
396 for workers in state.worker_list.values_mut() {
397 while workers
398 .front()
399 .map(|idle_worker| now.duration_since(idle_worker.idle_since) >= idle_timeout)
400 .unwrap_or(false)
401 {
402 removed_workers.push(workers.pop_front().unwrap());
403 }
404 }
405 state.worker_list.retain(|_, workers| !workers.is_empty());
406 for idle_worker in &removed_workers {
407 state.current_count -= 1;
408 state.dec_worker_count_for_key(idle_worker.primary_key.clone());
409 }
410 removed_workers
411 }
412
413 fn take_expired_idle_workers_for_key(
414 state: &mut WorkerPoolState<K, W, F>,
415 key: &K,
416 idle_timeout: Option<Duration>,
417 ) -> Vec<IdleWorker<K, W>> {
418 let Some(idle_timeout) = idle_timeout else {
419 return Vec::new();
420 };
421 let mut removed_workers = Vec::new();
422 let now = Instant::now();
423 let remove_bucket = if let Some(workers) = state.worker_list.get_mut(key) {
424 while workers
425 .front()
426 .map(|idle_worker| now.duration_since(idle_worker.idle_since) >= idle_timeout)
427 .unwrap_or(false)
428 {
429 removed_workers.push(workers.pop_front().unwrap());
430 }
431 workers.is_empty()
432 } else {
433 false
434 };
435 if remove_bucket {
436 state.worker_list.remove(key);
437 }
438 for idle_worker in &removed_workers {
439 state.current_count -= 1;
440 state.dec_worker_count_for_key(idle_worker.primary_key.clone());
441 }
442 removed_workers
443 }
444
445 pub fn cleanup_idle_worker(&self) -> usize {
446 let (removed_workers, clear_waiters) = {
447 let mut state = self.state.lock().unwrap();
448 let removed_workers =
449 Self::take_expired_idle_workers(&mut state, self.config.idle_timeout);
450 let clear_waiters = state.take_clear_waiters_if_done();
451 (removed_workers, clear_waiters)
452 };
453 for waiter in clear_waiters {
454 waiter.notify(());
455 }
456 let removed_count = removed_workers.len();
457 drop(removed_workers);
458 removed_count
459 }
460
461 pub async fn get_worker(
462 self: &KeyedWorkerPoolRef<K, W, F>,
463 key: K,
464 ) -> PoolResult<KeyedWorkerGuard<K, W, F>> {
465 loop {
466 if self.config.max_count_per_key == Some(0) {
467 return Err(pool_invalid_config_error("pool max_count_per_key is zero"));
468 }
469
470 let (worker, wait, should_create, removed_workers) = {
471 let mut state = self.state.lock().unwrap();
472 if state.clearing {
473 return Err(pool_clearing_error());
474 }
475 state.remove_canceled_waiters();
476
477 let mut removed_workers = Self::take_expired_idle_workers_for_key(
478 &mut state,
479 &key,
480 self.config.idle_timeout,
481 );
482
483 let worker = loop {
484 let Some(idle_worker) = state.take_idle_worker_for_key(&key) else {
485 break None;
486 };
487 let expired = self
488 .config
489 .idle_timeout
490 .map(|idle_timeout| idle_worker.idle_since.elapsed() >= idle_timeout)
491 .unwrap_or(false);
492 if expired
493 || !idle_worker.worker.is_work()
494 || idle_worker.worker.primary_key() != idle_worker.primary_key
495 || !idle_worker.worker.supports(idle_worker.primary_key.clone())
496 {
497 state.current_count -= 1;
498 state.dec_worker_count_for_key(idle_worker.primary_key.clone());
499 removed_workers.push(idle_worker);
500 continue;
501 }
502 break Some((idle_worker.worker, idle_worker.primary_key));
503 };
504
505 if worker.is_some() {
506 (worker, None, false, removed_workers)
507 } else if self.key_limit_reached(&state, &key) {
508 let (notify, waiter) = Notify::new();
509 state.waiting_list.push(WaitingItem {
510 future: notify,
511 key: key.clone(),
512 });
513 (None, Some(waiter), false, removed_workers)
514 } else {
515 state.current_count += 1;
516 state.inc_pending_count_for_key(key.clone());
517 (None, None, true, removed_workers)
518 }
519 };
520
521 let reservation =
522 should_create.then(|| KeyedWorkerReservation::new(self.clone(), key.clone()));
523 drop(removed_workers);
524
525 if let Some((worker, primary_key)) = worker {
526 return Ok(KeyedWorkerGuard::new(worker, self.clone(), primary_key));
527 }
528
529 if let Some(wait) = wait {
530 match wait.await {
531 KeyedWorkerWaitResult::Worker(worker) => return Ok(worker),
532 KeyedWorkerWaitResult::Retry => continue,
533 KeyedWorkerWaitResult::Error(err) => return Err(err),
534 }
535 }
536
537 let reservation = reservation.unwrap();
538 let (worker, primary_key) = match self.factory.create(key.clone()).await {
539 Ok(worker) => {
540 let primary_key = Self::validate_created_worker(&key, &worker)?;
541 (worker, primary_key)
542 }
543 Err(err) => return Err(err),
544 };
545 match reservation.complete(primary_key.clone()) {
546 ReservationCompletion::Complete => {}
547 ReservationCompletion::Clearing => return Err(pool_cleared_error()),
548 }
549 return Ok(KeyedWorkerGuard::new(worker, self.clone(), primary_key));
550 }
551 }
552
553 pub async fn clear_all_worker(&self) {
554 let (waiter, waiting_list, clear_waiters, idle_workers) = {
555 let mut state = self.state.lock().unwrap();
556 let idle_workers = if !state.clearing {
557 state.clearing = true;
558 let idle_workers = state.drain_idle_workers();
559 let cur_worker_count = idle_workers.len();
560 state.current_count -= cur_worker_count;
561 for idle_worker in &idle_workers {
562 state.dec_worker_count_for_key(idle_worker.primary_key.clone());
563 }
564 idle_workers
565 } else {
566 Vec::new()
567 };
568
569 let waiting_list = state.waiting_list.drain(..).collect::<Vec<_>>();
570 if state.current_count == 0 {
571 let clear_waiters = state.take_clear_waiters_if_done();
572 (None, waiting_list, clear_waiters, idle_workers)
573 } else {
574 let (notify, waiter) = Notify::new();
575 state.clear_waiting_list.push(notify);
576 (Some(waiter), waiting_list, Vec::new(), idle_workers)
577 }
578 };
579 for waiting in waiting_list {
580 waiting
581 .future
582 .notify(KeyedWorkerWaitResult::Error(pool_cleared_error()));
583 }
584 for waiter in clear_waiters {
585 waiter.notify(());
586 }
587 drop(idle_workers);
588 if let Some(waiter) = waiter {
589 waiter.await;
590 }
591 }
592
593 fn notify_retry_waiters(waiters: Vec<Notify<KeyedWorkerWaitResult<K, W, F>>>) {
594 for waiter in waiters {
595 waiter.notify(KeyedWorkerWaitResult::Retry);
596 }
597 }
598
599 fn rollback_reservation(&self, key: &K) {
600 let (retry_waiters, clear_waiters) = {
601 let mut state = self.state.lock().unwrap();
602 state.current_count -= 1;
603 state.dec_pending_count_for_key(key.clone());
604 let retry_waiters = state.drain_waiters();
605 let clear_waiters = state.take_clear_waiters_if_done();
606 (retry_waiters, clear_waiters)
607 };
608 Self::notify_retry_waiters(retry_waiters);
609 for waiter in clear_waiters {
610 waiter.notify(());
611 }
612 }
613
614 fn release(self: &KeyedWorkerPoolRef<K, W, F>, work: W, primary_key: K) {
615 enum ReleaseAction<K: WorkerKey, W: KeyedWorker<K>, F: KeyedWorkerFactory<K, W>> {
616 None,
617 Notify(
618 Notify<KeyedWorkerWaitResult<K, W, F>>,
619 KeyedWorkerGuard<K, W, F>,
620 ),
621 Retry(Vec<Notify<KeyedWorkerWaitResult<K, W, F>>>),
622 }
623
624 let primary_key_valid =
625 work.primary_key() == primary_key && work.supports(primary_key.clone());
626 let mut clear_waiters = Vec::new();
627 let mut removed_workers = Vec::new();
628 let action = {
629 let mut state = self.state.lock().unwrap();
630 state.remove_canceled_waiters();
631 if state.clearing {
632 state.current_count -= 1;
633 state.dec_worker_count_for_key(primary_key);
634 clear_waiters = state.take_clear_waiters_if_done();
635 ReleaseAction::None
636 } else if !primary_key_valid {
637 state.current_count -= 1;
638 state.dec_worker_count_for_key(primary_key);
639 let waiters = state.drain_waiters();
640 if !waiters.is_empty() {
641 ReleaseAction::Retry(waiters)
642 } else {
643 ReleaseAction::None
644 }
645 } else if work.is_work() {
646 if let Some(index) =
647 state.find_matching_waiter_index_for_worker(&work, &primary_key)
648 {
649 let waiting_item = state.waiting_list.remove(index);
650 ReleaseAction::Notify(
651 waiting_item.future,
652 KeyedWorkerGuard::new(work, self.clone(), primary_key),
653 )
654 } else {
655 state.push_idle_worker(work, primary_key.clone());
656 if let Some(max_idle_count_per_key) = self.config.max_idle_count_per_key {
657 while state
658 .worker_list
659 .get(&primary_key)
660 .map(VecDeque::len)
661 .unwrap_or(0)
662 > usize::from(max_idle_count_per_key)
663 {
664 let idle_worker = state.remove_idle_worker(&primary_key, 0);
665 state.current_count -= 1;
666 state.dec_worker_count_for_key(idle_worker.primary_key.clone());
667 removed_workers.push(idle_worker);
668 }
669 }
670 ReleaseAction::None
671 }
672 } else {
673 state.dec_worker_count_for_key(primary_key);
674 state.current_count -= 1;
675 let waiters = state.drain_waiters();
676 if !waiters.is_empty() {
677 ReleaseAction::Retry(waiters)
678 } else {
679 clear_waiters = state.take_clear_waiters_if_done();
680 ReleaseAction::None
681 }
682 }
683 };
684
685 for waiter in clear_waiters {
686 waiter.notify(());
687 }
688 drop(removed_workers);
689
690 match action {
691 ReleaseAction::None => {}
692 ReleaseAction::Notify(waiting, worker) => {
693 waiting.notify(KeyedWorkerWaitResult::Worker(worker));
694 }
695 ReleaseAction::Retry(waiters) => {
696 Self::notify_retry_waiters(waiters);
697 }
698 }
699 }
700}
701
702#[cfg(test)]
703mod idle_limit_tests {
704 use super::*;
705 use std::sync::atomic::{AtomicUsize, Ordering};
706
707 #[derive(Clone, Debug, Eq, Hash, PartialEq)]
708 enum Key {
709 A,
710 B,
711 }
712
713 struct TestWorker {
714 id: usize,
715 key: Key,
716 }
717
718 impl KeyedWorker<Key> for TestWorker {
719 fn is_work(&self) -> bool {
720 true
721 }
722
723 fn supports(&self, key: Key) -> bool {
724 self.key == key || (self.key == Key::A && key == Key::B)
725 }
726
727 fn primary_key(&self) -> Key {
728 self.key.clone()
729 }
730 }
731
732 struct TestFactory(AtomicUsize);
733
734 #[async_trait::async_trait]
735 impl KeyedWorkerFactory<Key, TestWorker> for TestFactory {
736 async fn create(&self, key: Key) -> PoolResult<TestWorker> {
737 Ok(TestWorker {
738 id: self.0.fetch_add(1, Ordering::SeqCst),
739 key,
740 })
741 }
742 }
743
744 fn new_pool(per_key_idle: u16) -> KeyedWorkerPoolRef<Key, TestWorker, TestFactory> {
745 KeyedWorkerPool::new(
746 TestFactory(AtomicUsize::new(0)),
747 KeyedWorkerPoolConfig {
748 max_count_per_key: None,
749 max_idle_count_per_key: Some(per_key_idle),
750 idle_timeout: None,
751 },
752 )
753 }
754
755 #[tokio::test]
756 async fn per_key_idle_limits_are_isolated_and_evict_key_lru() {
757 let pool = new_pool(1);
758 let a0 = pool.get_worker(Key::A).await.unwrap();
759 let a1 = pool.get_worker(Key::A).await.unwrap();
760 let b2 = pool.get_worker(Key::B).await.unwrap();
761 drop(a0);
762 drop(b2);
763 drop(a1);
764
765 let a = pool.get_worker(Key::A).await.unwrap();
766 let b = pool.get_worker(Key::B).await.unwrap();
767 assert_eq!(a.id, 1);
768 assert_eq!(b.id, 2);
769 }
770
771 #[tokio::test]
772 async fn acquisition_cleans_expired_worker_behind_bucket_mru() {
773 let pool = KeyedWorkerPool::new(
774 TestFactory(AtomicUsize::new(0)),
775 KeyedWorkerPoolConfig::default().with_idle_timeout(Some(Duration::from_millis(20))),
776 );
777 let a0 = pool.get_worker(Key::A).await.unwrap();
778 let a1 = pool.get_worker(Key::A).await.unwrap();
779
780 drop(a0);
781 tokio::time::sleep(Duration::from_millis(30)).await;
782 drop(a1);
783
784 let a = pool.get_worker(Key::A).await.unwrap();
785 assert_eq!(a.id, 1);
786 assert_eq!(pool.cleanup_idle_worker(), 0);
787 }
788
789 #[tokio::test]
790 async fn keyed_acquisition_only_cleans_and_reuses_requested_bucket() {
791 let pool = KeyedWorkerPool::new(
792 TestFactory(AtomicUsize::new(0)),
793 KeyedWorkerPoolConfig::default().with_idle_timeout(Some(Duration::from_millis(20))),
794 );
795 let a = pool.get_worker(Key::A).await.unwrap();
796 drop(a);
797 tokio::time::sleep(Duration::from_millis(30)).await;
798
799 let b = pool.get_worker(Key::B).await.unwrap();
800 assert_eq!(b.id, 1);
801 assert_eq!(b.primary_key(), Key::B);
802 assert_eq!(pool.cleanup_idle_worker(), 1);
803 }
804
805 #[tokio::test]
806 async fn returned_worker_only_wakes_waiter_for_its_primary_key() {
807 let pool = KeyedWorkerPool::new(
808 TestFactory(AtomicUsize::new(0)),
809 KeyedWorkerPoolConfig::default().with_max_count_per_key(Some(1)),
810 );
811 let a = pool.get_worker(Key::A).await.unwrap();
812 let b = pool.get_worker(Key::B).await.unwrap();
813 let pool_ref = pool.clone();
814 let waiting_b = tokio::spawn(async move { pool_ref.get_worker(Key::B).await });
815 tokio::time::sleep(Duration::from_millis(20)).await;
816
817 drop(a);
818 tokio::task::yield_now().await;
819 assert!(!waiting_b.is_finished());
820
821 drop(b);
822 let b = tokio::time::timeout(Duration::from_secs(1), waiting_b)
823 .await
824 .unwrap()
825 .unwrap()
826 .unwrap();
827 assert_eq!(b.primary_key(), Key::B);
828 }
829
830 #[tokio::test]
831 async fn zero_per_key_idle_limit_disables_idle_cache() {
832 let pool = new_pool(0);
833 let a = pool.get_worker(Key::A).await.unwrap();
834 assert_eq!(a.id, 0);
835 drop(a);
836 let a = pool.get_worker(Key::A).await.unwrap();
837 assert_eq!(a.id, 1);
838 }
839}
840
841#[cfg(test)]
842fn new_keyed_worker_pool<K: WorkerKey, W: KeyedWorker<K>, F: KeyedWorkerFactory<K, W>>(
843 max_count: u16,
844 factory: F,
845) -> KeyedWorkerPoolRef<K, W, F> {
846 KeyedWorkerPool::new(
847 factory,
848 KeyedWorkerPoolConfig {
849 max_count_per_key: Some(max_count),
850 ..Default::default()
851 },
852 )
853}
854
855#[tokio::test]
856async fn test_pool() {
857 struct TestWorker {
858 work: bool,
859 key: TestWorkerKey,
860 }
861
862 #[derive(Clone, Debug, Eq, PartialEq, Hash)]
863 enum TestWorkerKey {
864 A,
865 B,
866 }
867 #[async_trait::async_trait]
868 impl KeyedWorker<TestWorkerKey> for TestWorker {
869 fn is_work(&self) -> bool {
870 self.work
871 }
872
873 fn supports(&self, key: TestWorkerKey) -> bool {
874 self.key == key
875 }
876
877 fn primary_key(&self) -> TestWorkerKey {
878 self.key.clone()
879 }
880 }
881
882 struct TestWorkerFactory;
883
884 #[async_trait::async_trait]
885 impl KeyedWorkerFactory<TestWorkerKey, TestWorker> for TestWorkerFactory {
886 async fn create(&self, key: TestWorkerKey) -> PoolResult<TestWorker> {
887 Ok(TestWorker { work: true, key })
888 }
889 }
890
891 let pool = new_keyed_worker_pool(2, TestWorkerFactory);
892
893 let worker_a1 = pool.get_worker(TestWorkerKey::A).await.unwrap();
894 let worker_a2 = pool.get_worker(TestWorkerKey::A).await.unwrap();
895 let worker_b1 = pool.get_worker(TestWorkerKey::B).await.unwrap();
896 let worker_b2 = pool.get_worker(TestWorkerKey::B).await.unwrap();
897
898 let pool_ref = pool.clone();
899 let keyed_waiter = tokio::spawn(async move { pool_ref.get_worker(TestWorkerKey::B).await });
900 tokio::time::sleep(std::time::Duration::from_millis(20)).await;
901 assert!(!keyed_waiter.is_finished());
902
903 drop(worker_b1);
904 let worker_b = tokio::time::timeout(std::time::Duration::from_secs(1), keyed_waiter)
905 .await
906 .unwrap()
907 .unwrap()
908 .unwrap();
909 drop(worker_a1);
910 drop(worker_a2);
911 drop(worker_b2);
912 drop(worker_b);
913
914 let worker_b1 = pool.get_worker(TestWorkerKey::B).await.unwrap();
915 let worker_b2 = pool.get_worker(TestWorkerKey::B).await.unwrap();
916 let worker1 = pool.get_worker(TestWorkerKey::A).await.unwrap();
917 let worker2 = pool.get_worker(TestWorkerKey::A).await.unwrap();
918
919 let pool_ref = pool.clone();
920 let keyed_a_waiter = tokio::spawn(async move { pool_ref.get_worker(TestWorkerKey::A).await });
921 let pool_ref = pool.clone();
922 let keyed_waiter = tokio::spawn(async move { pool_ref.get_worker(TestWorkerKey::B).await });
923 tokio::time::sleep(std::time::Duration::from_millis(20)).await;
924 assert!(!keyed_a_waiter.is_finished());
925 assert!(!keyed_waiter.is_finished());
926
927 let pool_ref = pool.clone();
928 let clear_task = tokio::spawn(async move {
929 pool_ref.clear_all_worker().await;
930 });
931 tokio::time::sleep(std::time::Duration::from_millis(20)).await;
932
933 assert!(keyed_a_waiter.await.unwrap().is_err());
934 assert!(keyed_waiter.await.unwrap().is_err());
935
936 drop(worker1);
937 drop(worker2);
938 drop(worker_b1);
939 drop(worker_b2);
940
941 tokio::time::timeout(std::time::Duration::from_secs(1), clear_task)
942 .await
943 .unwrap()
944 .unwrap();
945}
946
947#[tokio::test]
948async fn test_clear_all_worker_waits_for_inflight_create() {
949 use std::sync::atomic::{AtomicUsize, Ordering};
950 use std::sync::Arc;
951
952 #[derive(Clone, Debug, Eq, PartialEq, Hash)]
953 enum TestWorkerKey {
954 A,
955 }
956
957 struct TestWorker {
958 key: TestWorkerKey,
959 }
960
961 #[async_trait::async_trait]
962 impl KeyedWorker<TestWorkerKey> for TestWorker {
963 fn is_work(&self) -> bool {
964 true
965 }
966
967 fn supports(&self, key: TestWorkerKey) -> bool {
968 self.key == key
969 }
970
971 fn primary_key(&self) -> TestWorkerKey {
972 self.key.clone()
973 }
974 }
975
976 struct TestWorkerFactory {
977 create_count: Arc<AtomicUsize>,
978 }
979
980 #[async_trait::async_trait]
981 impl KeyedWorkerFactory<TestWorkerKey, TestWorker> for TestWorkerFactory {
982 async fn create(&self, key: TestWorkerKey) -> PoolResult<TestWorker> {
983 self.create_count.fetch_add(1, Ordering::SeqCst);
984 tokio::time::sleep(std::time::Duration::from_millis(100)).await;
985 Ok(TestWorker { key })
986 }
987 }
988
989 let create_count = Arc::new(AtomicUsize::new(0));
990 let pool = new_keyed_worker_pool(
991 1,
992 TestWorkerFactory {
993 create_count: create_count.clone(),
994 },
995 );
996
997 let pool_ref = pool.clone();
998 let worker_task = tokio::spawn(async move { pool_ref.get_worker(TestWorkerKey::A).await });
999 tokio::time::sleep(std::time::Duration::from_millis(100)).await;
1000
1001 pool.clear_all_worker().await;
1002
1003 let worker = worker_task.await.unwrap();
1004 assert!(worker.is_err());
1005 assert_eq!(create_count.load(Ordering::SeqCst), 1);
1006}
1007
1008#[tokio::test]
1009async fn test_concurrent_clear_all_worker() {
1010 #[derive(Clone, Debug, Eq, PartialEq, Hash)]
1011 enum TestWorkerKey {
1012 A,
1013 }
1014
1015 struct TestWorker {
1016 key: TestWorkerKey,
1017 }
1018
1019 #[async_trait::async_trait]
1020 impl KeyedWorker<TestWorkerKey> for TestWorker {
1021 fn is_work(&self) -> bool {
1022 true
1023 }
1024
1025 fn supports(&self, key: TestWorkerKey) -> bool {
1026 self.key == key
1027 }
1028
1029 fn primary_key(&self) -> TestWorkerKey {
1030 self.key.clone()
1031 }
1032 }
1033
1034 struct TestWorkerFactory;
1035
1036 #[async_trait::async_trait]
1037 impl KeyedWorkerFactory<TestWorkerKey, TestWorker> for TestWorkerFactory {
1038 async fn create(&self, key: TestWorkerKey) -> PoolResult<TestWorker> {
1039 Ok(TestWorker { key })
1040 }
1041 }
1042
1043 let pool = new_keyed_worker_pool(1, TestWorkerFactory);
1044 let worker = pool.get_worker(TestWorkerKey::A).await.unwrap();
1045
1046 let pool_ref = pool.clone();
1047 let clear_task1 = tokio::spawn(async move {
1048 pool_ref.clear_all_worker().await;
1049 });
1050
1051 let pool_ref = pool.clone();
1052 let clear_task2 = tokio::spawn(async move {
1053 pool_ref.clear_all_worker().await;
1054 });
1055
1056 tokio::time::sleep(std::time::Duration::from_millis(100)).await;
1057 drop(worker);
1058
1059 tokio::time::timeout(std::time::Duration::from_secs(1), async {
1060 clear_task1.await.unwrap();
1061 clear_task2.await.unwrap();
1062 })
1063 .await
1064 .unwrap();
1065}
1066
1067#[tokio::test]
1068async fn test_zero_max_count_per_key_returns_error() {
1069 #[derive(Clone, Debug, Eq, PartialEq, Hash)]
1070 enum TestWorkerKey {
1071 A,
1072 }
1073
1074 struct TestWorker {
1075 key: TestWorkerKey,
1076 }
1077
1078 #[async_trait::async_trait]
1079 impl KeyedWorker<TestWorkerKey> for TestWorker {
1080 fn is_work(&self) -> bool {
1081 true
1082 }
1083
1084 fn supports(&self, key: TestWorkerKey) -> bool {
1085 self.key == key
1086 }
1087
1088 fn primary_key(&self) -> TestWorkerKey {
1089 self.key.clone()
1090 }
1091 }
1092
1093 struct TestWorkerFactory;
1094
1095 #[async_trait::async_trait]
1096 impl KeyedWorkerFactory<TestWorkerKey, TestWorker> for TestWorkerFactory {
1097 async fn create(&self, key: TestWorkerKey) -> PoolResult<TestWorker> {
1098 Ok(TestWorker { key })
1099 }
1100 }
1101
1102 let pool = new_keyed_worker_pool(0, TestWorkerFactory);
1103 let worker = pool.get_worker(TestWorkerKey::A).await;
1104 assert!(worker.is_err());
1105 assert_eq!(
1106 worker.err().unwrap().code(),
1107 crate::PoolErrorCode::InvalidConfig
1108 );
1109}
1110
1111#[tokio::test]
1112async fn test_keyed_pool_default_config_has_no_per_key_count_limit() {
1113 #[derive(Clone, Debug, Eq, PartialEq, Hash)]
1114 struct TestWorkerKey;
1115
1116 struct TestWorker;
1117
1118 impl KeyedWorker<TestWorkerKey> for TestWorker {
1119 fn is_work(&self) -> bool {
1120 true
1121 }
1122
1123 fn supports(&self, _key: TestWorkerKey) -> bool {
1124 true
1125 }
1126
1127 fn primary_key(&self) -> TestWorkerKey {
1128 TestWorkerKey
1129 }
1130 }
1131
1132 struct TestWorkerFactory;
1133
1134 #[async_trait::async_trait]
1135 impl KeyedWorkerFactory<TestWorkerKey, TestWorker> for TestWorkerFactory {
1136 async fn create(&self, _key: TestWorkerKey) -> PoolResult<TestWorker> {
1137 Ok(TestWorker)
1138 }
1139 }
1140
1141 let pool = KeyedWorkerPool::new(TestWorkerFactory, Default::default());
1142 let worker1 = pool.get_worker(TestWorkerKey).await.unwrap();
1143 let worker2 = pool.get_worker(TestWorkerKey).await.unwrap();
1144 drop((worker1, worker2));
1145}
1146
1147#[tokio::test]
1148async fn test_keyed_pool_waits_when_key_already_has_worker() {
1149 #[derive(Clone, Debug, Eq, PartialEq, Hash)]
1150 enum TestWorkerKey {
1151 B,
1152 }
1153
1154 struct TestWorker {
1155 key: TestWorkerKey,
1156 }
1157
1158 #[async_trait::async_trait]
1159 impl KeyedWorker<TestWorkerKey> for TestWorker {
1160 fn is_work(&self) -> bool {
1161 true
1162 }
1163
1164 fn supports(&self, key: TestWorkerKey) -> bool {
1165 self.key == key
1166 }
1167
1168 fn primary_key(&self) -> TestWorkerKey {
1169 self.key.clone()
1170 }
1171 }
1172
1173 struct TestWorkerFactory;
1174
1175 #[async_trait::async_trait]
1176 impl KeyedWorkerFactory<TestWorkerKey, TestWorker> for TestWorkerFactory {
1177 async fn create(&self, key: TestWorkerKey) -> PoolResult<TestWorker> {
1178 Ok(TestWorker { key })
1179 }
1180 }
1181
1182 let pool = new_keyed_worker_pool(1, TestWorkerFactory);
1183 let _worker = pool.get_worker(TestWorkerKey::B).await.unwrap();
1184
1185 let pool_ref = pool.clone();
1186 let result = tokio::time::timeout(std::time::Duration::from_millis(100), async move {
1187 pool_ref.get_worker(TestWorkerKey::B).await
1188 })
1189 .await;
1190
1191 assert!(result.is_err());
1192}
1193
1194#[tokio::test]
1195async fn test_different_key_is_not_blocked_by_per_key_limit() {
1196 use std::sync::atomic::{AtomicUsize, Ordering};
1197 use std::sync::Arc;
1198
1199 #[derive(Clone, Debug, Eq, PartialEq, Hash)]
1200 enum TestWorkerKey {
1201 A,
1202 B,
1203 }
1204
1205 struct TestWorker {
1206 id: usize,
1207 key: TestWorkerKey,
1208 }
1209
1210 #[async_trait::async_trait]
1211 impl KeyedWorker<TestWorkerKey> for TestWorker {
1212 fn is_work(&self) -> bool {
1213 true
1214 }
1215
1216 fn supports(&self, key: TestWorkerKey) -> bool {
1217 self.key == key
1218 }
1219
1220 fn primary_key(&self) -> TestWorkerKey {
1221 self.key.clone()
1222 }
1223 }
1224
1225 struct TestWorkerFactory {
1226 create_count: Arc<AtomicUsize>,
1227 }
1228
1229 #[async_trait::async_trait]
1230 impl KeyedWorkerFactory<TestWorkerKey, TestWorker> for TestWorkerFactory {
1231 async fn create(&self, key: TestWorkerKey) -> PoolResult<TestWorker> {
1232 let id = self.create_count.fetch_add(1, Ordering::SeqCst);
1233 Ok(TestWorker { id, key })
1234 }
1235 }
1236
1237 let create_count = Arc::new(AtomicUsize::new(0));
1238 let pool = new_keyed_worker_pool(
1239 1,
1240 TestWorkerFactory {
1241 create_count: create_count.clone(),
1242 },
1243 );
1244
1245 let worker_a = pool.get_worker(TestWorkerKey::A).await.unwrap();
1246 let worker_b = pool.get_worker(TestWorkerKey::B).await.unwrap();
1247
1248 assert_eq!(worker_a.id, 0);
1249 assert_eq!(worker_b.id, 1);
1250 assert_eq!(worker_b.primary_key(), TestWorkerKey::B);
1251 assert_eq!(create_count.load(Ordering::SeqCst), 2);
1252}
1253
1254#[tokio::test]
1255async fn test_keyed_create_failure_fails_same_key_waiters() {
1256 #[derive(Clone, Debug, Eq, PartialEq, Hash)]
1257 enum TestWorkerKey {
1258 B,
1259 }
1260
1261 struct TestWorker {
1262 key: TestWorkerKey,
1263 }
1264
1265 #[async_trait::async_trait]
1266 impl KeyedWorker<TestWorkerKey> for TestWorker {
1267 fn is_work(&self) -> bool {
1268 true
1269 }
1270
1271 fn supports(&self, key: TestWorkerKey) -> bool {
1272 self.key == key
1273 }
1274
1275 fn primary_key(&self) -> TestWorkerKey {
1276 self.key.clone()
1277 }
1278 }
1279
1280 struct TestWorkerFactory;
1281
1282 #[async_trait::async_trait]
1283 impl KeyedWorkerFactory<TestWorkerKey, TestWorker> for TestWorkerFactory {
1284 async fn create(&self, _key: TestWorkerKey) -> PoolResult<TestWorker> {
1285 tokio::time::sleep(std::time::Duration::from_millis(50)).await;
1286 Err(crate::pool_invalid_config_error("create failed"))
1287 }
1288 }
1289
1290 let pool = new_keyed_worker_pool(1, TestWorkerFactory);
1291
1292 let pool_ref = pool.clone();
1293 let worker1 = tokio::spawn(async move { pool_ref.get_worker(TestWorkerKey::B).await });
1294 tokio::time::sleep(std::time::Duration::from_millis(20)).await;
1295
1296 let pool_ref = pool.clone();
1297 let worker2 = tokio::spawn(async move { pool_ref.get_worker(TestWorkerKey::B).await });
1298
1299 let (worker1, worker2) = tokio::time::timeout(std::time::Duration::from_secs(1), async {
1300 (worker1.await.unwrap(), worker2.await.unwrap())
1301 })
1302 .await
1303 .unwrap();
1304
1305 assert_eq!(
1306 worker1.err().unwrap().code(),
1307 crate::PoolErrorCode::InvalidConfig
1308 );
1309 assert_eq!(
1310 worker2.err().unwrap().code(),
1311 crate::PoolErrorCode::InvalidConfig
1312 );
1313}
1314
1315#[tokio::test]
1316async fn test_keyed_create_failure_wakes_waiter_to_create() {
1317 use std::sync::atomic::{AtomicUsize, Ordering};
1318 use std::sync::Arc;
1319
1320 #[derive(Clone, Debug, Eq, PartialEq, Hash)]
1321 enum TestWorkerKey {
1322 A,
1323 }
1324
1325 struct TestWorker {
1326 id: usize,
1327 key: TestWorkerKey,
1328 }
1329
1330 #[async_trait::async_trait]
1331 impl KeyedWorker<TestWorkerKey> for TestWorker {
1332 fn is_work(&self) -> bool {
1333 true
1334 }
1335
1336 fn supports(&self, key: TestWorkerKey) -> bool {
1337 self.key == key
1338 }
1339
1340 fn primary_key(&self) -> TestWorkerKey {
1341 self.key.clone()
1342 }
1343 }
1344
1345 struct TestWorkerFactory {
1346 create_count: Arc<AtomicUsize>,
1347 }
1348
1349 #[async_trait::async_trait]
1350 impl KeyedWorkerFactory<TestWorkerKey, TestWorker> for TestWorkerFactory {
1351 async fn create(&self, key: TestWorkerKey) -> PoolResult<TestWorker> {
1352 let id = self.create_count.fetch_add(1, Ordering::SeqCst);
1353 tokio::time::sleep(std::time::Duration::from_millis(50)).await;
1354 if id == 0 && key == TestWorkerKey::A {
1355 Err(crate::pool_invalid_config_error("create failed"))
1356 } else {
1357 Ok(TestWorker { id, key })
1358 }
1359 }
1360 }
1361
1362 let create_count = Arc::new(AtomicUsize::new(0));
1363 let pool = new_keyed_worker_pool(
1364 1,
1365 TestWorkerFactory {
1366 create_count: create_count.clone(),
1367 },
1368 );
1369
1370 let pool_ref = pool.clone();
1371 let keyed = tokio::spawn(async move { pool_ref.get_worker(TestWorkerKey::A).await });
1372 tokio::time::sleep(std::time::Duration::from_millis(10)).await;
1373
1374 let pool_ref = pool.clone();
1375 let waiter = tokio::spawn(async move { pool_ref.get_worker(TestWorkerKey::A).await });
1376
1377 let (keyed, waiter) = tokio::time::timeout(std::time::Duration::from_secs(1), async {
1378 (keyed.await.unwrap(), waiter.await.unwrap())
1379 })
1380 .await
1381 .unwrap();
1382
1383 assert_eq!(
1384 keyed.err().unwrap().code(),
1385 crate::PoolErrorCode::InvalidConfig
1386 );
1387 let waiter = waiter.unwrap();
1388 assert_eq!(waiter.id, 1);
1389 assert_eq!(create_count.load(Ordering::SeqCst), 2);
1390}
1391
1392#[tokio::test]
1393async fn test_keyed_retry_notification_skips_canceled_waiter() {
1394 #[derive(Clone, Debug, Eq, PartialEq, Hash)]
1395 enum TestWorkerKey {
1396 A,
1397 }
1398
1399 struct TestWorker {
1400 key: TestWorkerKey,
1401 }
1402
1403 #[async_trait::async_trait]
1404 impl KeyedWorker<TestWorkerKey> for TestWorker {
1405 fn is_work(&self) -> bool {
1406 true
1407 }
1408
1409 fn supports(&self, key: TestWorkerKey) -> bool {
1410 self.key == key
1411 }
1412
1413 fn primary_key(&self) -> TestWorkerKey {
1414 self.key.clone()
1415 }
1416 }
1417
1418 struct TestWorkerFactory;
1419
1420 #[async_trait::async_trait]
1421 impl KeyedWorkerFactory<TestWorkerKey, TestWorker> for TestWorkerFactory {
1422 async fn create(&self, key: TestWorkerKey) -> PoolResult<TestWorker> {
1423 Ok(TestWorker { key })
1424 }
1425 }
1426
1427 let (canceled_notify, canceled_waiter) = Notify::new();
1428 let _key = TestWorkerKey::A;
1429 drop(canceled_waiter);
1430 let (notify, waiter) = Notify::new();
1431
1432 KeyedWorkerPool::<TestWorkerKey, TestWorker, TestWorkerFactory>::notify_retry_waiters(vec![
1433 canceled_notify,
1434 notify,
1435 ]);
1436
1437 let result = tokio::time::timeout(std::time::Duration::from_secs(1), waiter)
1438 .await
1439 .unwrap();
1440 assert!(matches!(result, KeyedWorkerWaitResult::Retry));
1441}
1442
1443#[tokio::test]
1444async fn test_keyed_request_replaces_non_matching_idle_worker() {
1445 use std::sync::atomic::{AtomicUsize, Ordering};
1446 use std::sync::Arc;
1447
1448 #[derive(Clone, Debug, Eq, PartialEq, Hash)]
1449 enum TestWorkerKey {
1450 A,
1451 B,
1452 }
1453
1454 struct TestWorker {
1455 id: usize,
1456 key: TestWorkerKey,
1457 }
1458
1459 #[async_trait::async_trait]
1460 impl KeyedWorker<TestWorkerKey> for TestWorker {
1461 fn is_work(&self) -> bool {
1462 true
1463 }
1464
1465 fn supports(&self, key: TestWorkerKey) -> bool {
1466 self.key == key
1467 }
1468
1469 fn primary_key(&self) -> TestWorkerKey {
1470 self.key.clone()
1471 }
1472 }
1473
1474 struct TestWorkerFactory {
1475 create_count: Arc<AtomicUsize>,
1476 }
1477
1478 #[async_trait::async_trait]
1479 impl KeyedWorkerFactory<TestWorkerKey, TestWorker> for TestWorkerFactory {
1480 async fn create(&self, key: TestWorkerKey) -> PoolResult<TestWorker> {
1481 let id = self.create_count.fetch_add(1, Ordering::SeqCst);
1482 Ok(TestWorker { id, key })
1483 }
1484 }
1485
1486 let create_count = Arc::new(AtomicUsize::new(0));
1487 let pool = new_keyed_worker_pool(
1488 1,
1489 TestWorkerFactory {
1490 create_count: create_count.clone(),
1491 },
1492 );
1493
1494 {
1495 let worker = pool.get_worker(TestWorkerKey::A).await.unwrap();
1496 assert_eq!(worker.id, 0);
1497 }
1498
1499 let worker = pool.get_worker(TestWorkerKey::B).await.unwrap();
1500 assert_eq!(worker.id, 1);
1501 assert_eq!(worker.primary_key(), TestWorkerKey::B);
1502 assert_eq!(create_count.load(Ordering::SeqCst), 2);
1503}
1504
1505#[tokio::test]
1506async fn test_other_key_does_not_wait_for_returned_non_matching_worker() {
1507 use std::sync::atomic::{AtomicUsize, Ordering};
1508 use std::sync::Arc;
1509
1510 #[derive(Clone, Debug, Eq, PartialEq, Hash)]
1511 enum TestWorkerKey {
1512 A,
1513 B,
1514 }
1515
1516 struct TestWorker {
1517 id: usize,
1518 key: TestWorkerKey,
1519 }
1520
1521 #[async_trait::async_trait]
1522 impl KeyedWorker<TestWorkerKey> for TestWorker {
1523 fn is_work(&self) -> bool {
1524 true
1525 }
1526
1527 fn supports(&self, key: TestWorkerKey) -> bool {
1528 self.key == key
1529 }
1530
1531 fn primary_key(&self) -> TestWorkerKey {
1532 self.key.clone()
1533 }
1534 }
1535
1536 struct TestWorkerFactory {
1537 create_count: Arc<AtomicUsize>,
1538 }
1539
1540 #[async_trait::async_trait]
1541 impl KeyedWorkerFactory<TestWorkerKey, TestWorker> for TestWorkerFactory {
1542 async fn create(&self, key: TestWorkerKey) -> PoolResult<TestWorker> {
1543 let id = self.create_count.fetch_add(1, Ordering::SeqCst);
1544 Ok(TestWorker { id, key })
1545 }
1546 }
1547
1548 let create_count = Arc::new(AtomicUsize::new(0));
1549 let pool = new_keyed_worker_pool(
1550 2,
1551 TestWorkerFactory {
1552 create_count: create_count.clone(),
1553 },
1554 );
1555 let worker_a = pool.get_worker(TestWorkerKey::A).await.unwrap();
1556 let _worker_b = pool.get_worker(TestWorkerKey::B).await.unwrap();
1557
1558 let pool_ref = pool.clone();
1559 let waiter = tokio::spawn(async move { pool_ref.get_worker(TestWorkerKey::B).await });
1560 tokio::time::sleep(std::time::Duration::from_millis(20)).await;
1561 assert!(waiter.is_finished());
1562
1563 let worker = tokio::time::timeout(std::time::Duration::from_secs(1), waiter)
1564 .await
1565 .unwrap()
1566 .unwrap()
1567 .unwrap();
1568 assert_eq!(worker.id, 2);
1569 assert_eq!(worker.primary_key(), TestWorkerKey::B);
1570 assert_eq!(create_count.load(Ordering::SeqCst), 3);
1571 drop(worker_a);
1572}
1573
1574#[tokio::test]
1575async fn test_other_key_does_not_wait_for_unwork_non_matching_worker() {
1576 use std::sync::atomic::{AtomicBool, AtomicUsize, Ordering};
1577 use std::sync::Arc;
1578
1579 #[derive(Clone, Debug, Eq, PartialEq, Hash)]
1580 enum TestWorkerKey {
1581 A,
1582 B,
1583 }
1584
1585 struct TestWorker {
1586 id: usize,
1587 work: AtomicBool,
1588 key: TestWorkerKey,
1589 }
1590
1591 #[async_trait::async_trait]
1592 impl KeyedWorker<TestWorkerKey> for TestWorker {
1593 fn is_work(&self) -> bool {
1594 self.work.load(Ordering::SeqCst)
1595 }
1596
1597 fn supports(&self, key: TestWorkerKey) -> bool {
1598 self.key == key
1599 }
1600
1601 fn primary_key(&self) -> TestWorkerKey {
1602 self.key.clone()
1603 }
1604 }
1605
1606 struct TestWorkerFactory {
1607 create_count: Arc<AtomicUsize>,
1608 }
1609
1610 #[async_trait::async_trait]
1611 impl KeyedWorkerFactory<TestWorkerKey, TestWorker> for TestWorkerFactory {
1612 async fn create(&self, key: TestWorkerKey) -> PoolResult<TestWorker> {
1613 let id = self.create_count.fetch_add(1, Ordering::SeqCst);
1614 Ok(TestWorker {
1615 id,
1616 work: AtomicBool::new(true),
1617 key,
1618 })
1619 }
1620 }
1621
1622 let create_count = Arc::new(AtomicUsize::new(0));
1623 let pool = new_keyed_worker_pool(
1624 2,
1625 TestWorkerFactory {
1626 create_count: create_count.clone(),
1627 },
1628 );
1629 let worker_a = pool.get_worker(TestWorkerKey::A).await.unwrap();
1630 let _worker_b = pool.get_worker(TestWorkerKey::B).await.unwrap();
1631
1632 let pool_ref = pool.clone();
1633 let waiter = tokio::spawn(async move { pool_ref.get_worker(TestWorkerKey::B).await });
1634 tokio::time::sleep(std::time::Duration::from_millis(20)).await;
1635 assert!(waiter.is_finished());
1636
1637 worker_a.work.store(false, Ordering::SeqCst);
1638 let worker = tokio::time::timeout(std::time::Duration::from_secs(1), waiter)
1639 .await
1640 .unwrap()
1641 .unwrap()
1642 .unwrap();
1643 assert_eq!(worker.id, 2);
1644 assert_eq!(worker.primary_key(), TestWorkerKey::B);
1645 assert_eq!(create_count.load(Ordering::SeqCst), 3);
1646 drop(worker_a);
1647}
1648
1649#[tokio::test]
1650async fn test_factory_must_return_matching_key() {
1651 use std::sync::atomic::{AtomicUsize, Ordering};
1652 use std::sync::Arc;
1653
1654 #[derive(Clone, Debug, Eq, PartialEq, Hash)]
1655 enum TestWorkerKey {
1656 A,
1657 B,
1658 }
1659
1660 struct TestWorker {
1661 key: TestWorkerKey,
1662 }
1663
1664 #[async_trait::async_trait]
1665 impl KeyedWorker<TestWorkerKey> for TestWorker {
1666 fn is_work(&self) -> bool {
1667 true
1668 }
1669
1670 fn supports(&self, key: TestWorkerKey) -> bool {
1671 self.key == key
1672 }
1673
1674 fn primary_key(&self) -> TestWorkerKey {
1675 self.key.clone()
1676 }
1677 }
1678
1679 struct TestWorkerFactory {
1680 create_count: Arc<AtomicUsize>,
1681 }
1682
1683 #[async_trait::async_trait]
1684 impl KeyedWorkerFactory<TestWorkerKey, TestWorker> for TestWorkerFactory {
1685 async fn create(&self, key: TestWorkerKey) -> PoolResult<TestWorker> {
1686 let count = self.create_count.fetch_add(1, Ordering::SeqCst);
1687 let key = if count == 0 { TestWorkerKey::A } else { key };
1688 Ok(TestWorker { key })
1689 }
1690 }
1691
1692 let create_count = Arc::new(AtomicUsize::new(0));
1693 let pool = new_keyed_worker_pool(
1694 1,
1695 TestWorkerFactory {
1696 create_count: create_count.clone(),
1697 },
1698 );
1699 let worker = pool.get_worker(TestWorkerKey::B).await;
1700 assert!(worker.is_err());
1701 assert_eq!(
1702 worker.err().unwrap().code(),
1703 crate::PoolErrorCode::InvalidConfig
1704 );
1705
1706 let worker = pool.get_worker(TestWorkerKey::B).await;
1707 assert!(worker.is_ok());
1708 assert_eq!(create_count.load(Ordering::SeqCst), 2);
1709}
1710
1711#[tokio::test(flavor = "multi_thread")]
1712async fn test_keyed_waiter_keeps_queue_priority_over_later_waiter() {
1713 use std::sync::mpsc;
1714
1715 #[derive(Clone, Debug, Eq, PartialEq, Hash)]
1716 enum TestWorkerKey {
1717 B,
1718 }
1719
1720 struct TestWorker {
1721 key: TestWorkerKey,
1722 }
1723
1724 #[async_trait::async_trait]
1725 impl KeyedWorker<TestWorkerKey> for TestWorker {
1726 fn is_work(&self) -> bool {
1727 true
1728 }
1729
1730 fn supports(&self, key: TestWorkerKey) -> bool {
1731 self.key == key
1732 }
1733
1734 fn primary_key(&self) -> TestWorkerKey {
1735 self.key.clone()
1736 }
1737 }
1738
1739 struct TestWorkerFactory;
1740
1741 #[async_trait::async_trait]
1742 impl KeyedWorkerFactory<TestWorkerKey, TestWorker> for TestWorkerFactory {
1743 async fn create(&self, key: TestWorkerKey) -> PoolResult<TestWorker> {
1744 Ok(TestWorker { key })
1745 }
1746 }
1747
1748 let pool = new_keyed_worker_pool(1, TestWorkerFactory);
1749 let worker = pool.get_worker(TestWorkerKey::B).await.unwrap();
1750
1751 let (tx, rx) = mpsc::channel();
1752
1753 let pool_ref = pool.clone();
1754 let tx_keyed = tx.clone();
1755 let keyed_task = tokio::spawn(async move {
1756 let _worker = pool_ref.get_worker(TestWorkerKey::B).await.unwrap();
1757 tx_keyed.send("keyed").unwrap();
1758 tokio::time::sleep(std::time::Duration::from_millis(50)).await;
1759 });
1760
1761 tokio::time::sleep(std::time::Duration::from_millis(20)).await;
1762
1763 let pool_ref = pool.clone();
1764 let later_task = tokio::spawn(async move {
1765 let _worker = pool_ref.get_worker(TestWorkerKey::B).await.unwrap();
1766 tx.send("later").unwrap();
1767 });
1768
1769 tokio::time::sleep(std::time::Duration::from_millis(20)).await;
1770 drop(worker);
1771
1772 let first = rx.recv_timeout(std::time::Duration::from_secs(2)).unwrap();
1773 assert_eq!(first, "keyed");
1774
1775 keyed_task.await.unwrap();
1776 later_task.await.unwrap();
1777}
1778
1779#[tokio::test]
1780async fn test_factory_worker_must_be_valid_for_its_primary_key() {
1781 #[derive(Clone, Debug, Eq, PartialEq, Hash)]
1782 enum TestWorkerKey {
1783 A,
1784 B,
1785 }
1786
1787 struct TestWorker {
1788 key: TestWorkerKey,
1789 }
1790
1791 #[async_trait::async_trait]
1792 impl KeyedWorker<TestWorkerKey> for TestWorker {
1793 fn is_work(&self) -> bool {
1794 true
1795 }
1796
1797 fn supports(&self, key: TestWorkerKey) -> bool {
1798 key == TestWorkerKey::B
1799 }
1800
1801 fn primary_key(&self) -> TestWorkerKey {
1802 self.key.clone()
1803 }
1804 }
1805
1806 struct TestWorkerFactory;
1807
1808 #[async_trait::async_trait]
1809 impl KeyedWorkerFactory<TestWorkerKey, TestWorker> for TestWorkerFactory {
1810 async fn create(&self, _key: TestWorkerKey) -> PoolResult<TestWorker> {
1811 Ok(TestWorker {
1812 key: TestWorkerKey::A,
1813 })
1814 }
1815 }
1816
1817 let pool = new_keyed_worker_pool(1, TestWorkerFactory);
1818 let worker = pool.get_worker(TestWorkerKey::A).await;
1819 assert!(worker.is_err());
1820 assert_eq!(
1821 worker.err().unwrap().code(),
1822 crate::PoolErrorCode::InvalidConfig
1823 );
1824}
1825
1826#[tokio::test]
1827async fn test_keyed_idle_worker_timeout_releases_worker() {
1828 use std::sync::atomic::{AtomicUsize, Ordering};
1829 use std::sync::Arc;
1830
1831 #[derive(Clone, Debug, Eq, PartialEq, Hash)]
1832 enum TestWorkerKey {
1833 A,
1834 B,
1835 }
1836
1837 struct TestWorker {
1838 id: usize,
1839 key: TestWorkerKey,
1840 }
1841
1842 #[async_trait::async_trait]
1843 impl KeyedWorker<TestWorkerKey> for TestWorker {
1844 fn is_work(&self) -> bool {
1845 true
1846 }
1847
1848 fn supports(&self, key: TestWorkerKey) -> bool {
1849 self.key == key
1850 }
1851
1852 fn primary_key(&self) -> TestWorkerKey {
1853 self.key.clone()
1854 }
1855 }
1856
1857 struct TestWorkerFactory {
1858 create_count: Arc<AtomicUsize>,
1859 }
1860
1861 #[async_trait::async_trait]
1862 impl KeyedWorkerFactory<TestWorkerKey, TestWorker> for TestWorkerFactory {
1863 async fn create(&self, key: TestWorkerKey) -> PoolResult<TestWorker> {
1864 let id = self.create_count.fetch_add(1, Ordering::SeqCst);
1865 Ok(TestWorker { id, key })
1866 }
1867 }
1868
1869 let create_count = Arc::new(AtomicUsize::new(0));
1870 let pool = KeyedWorkerPool::new(
1871 TestWorkerFactory {
1872 create_count: create_count.clone(),
1873 },
1874 KeyedWorkerPoolConfig {
1875 max_count_per_key: Some(1),
1876 idle_timeout: Some(std::time::Duration::from_millis(30)),
1877 ..Default::default()
1878 },
1879 );
1880
1881 {
1882 let worker = pool.get_worker(TestWorkerKey::B).await.unwrap();
1883 assert_eq!(worker.id, 0);
1884 assert_eq!(worker.primary_key(), TestWorkerKey::B);
1885 }
1886
1887 tokio::time::sleep(std::time::Duration::from_millis(80)).await;
1888
1889 let worker = pool.get_worker(TestWorkerKey::A).await.unwrap();
1890 assert_eq!(worker.id, 1);
1891 assert_eq!(worker.primary_key(), TestWorkerKey::A);
1892 assert_eq!(create_count.load(Ordering::SeqCst), 2);
1893}
1894
1895#[tokio::test]
1896async fn test_get_keyed_worker_uses_most_recent_matching_idle_worker() {
1897 use std::sync::atomic::{AtomicUsize, Ordering};
1898 use std::sync::Arc;
1899
1900 #[derive(Clone, Debug, Eq, PartialEq, Hash)]
1901 enum TestWorkerKey {
1902 A,
1903 }
1904
1905 struct TestWorker {
1906 id: usize,
1907 key: TestWorkerKey,
1908 }
1909
1910 #[async_trait::async_trait]
1911 impl KeyedWorker<TestWorkerKey> for TestWorker {
1912 fn is_work(&self) -> bool {
1913 true
1914 }
1915
1916 fn supports(&self, key: TestWorkerKey) -> bool {
1917 self.key == key
1918 }
1919
1920 fn primary_key(&self) -> TestWorkerKey {
1921 self.key.clone()
1922 }
1923 }
1924
1925 struct TestWorkerFactory {
1926 create_count: Arc<AtomicUsize>,
1927 }
1928
1929 #[async_trait::async_trait]
1930 impl KeyedWorkerFactory<TestWorkerKey, TestWorker> for TestWorkerFactory {
1931 async fn create(&self, key: TestWorkerKey) -> PoolResult<TestWorker> {
1932 let id = self.create_count.fetch_add(1, Ordering::SeqCst);
1933 Ok(TestWorker { id, key })
1934 }
1935 }
1936
1937 let create_count = Arc::new(AtomicUsize::new(0));
1938 let pool = new_keyed_worker_pool(
1939 2,
1940 TestWorkerFactory {
1941 create_count: create_count.clone(),
1942 },
1943 );
1944
1945 let worker1 = pool.get_worker(TestWorkerKey::A).await.unwrap();
1946 let worker2 = pool.get_worker(TestWorkerKey::A).await.unwrap();
1947 assert_eq!(worker1.id, 0);
1948 assert_eq!(worker2.id, 1);
1949
1950 drop(worker1);
1951 drop(worker2);
1952
1953 let worker = pool.get_worker(TestWorkerKey::A).await.unwrap();
1954 assert_eq!(worker.id, 1);
1955 assert_eq!(create_count.load(Ordering::SeqCst), 2);
1956}
1957
1958#[tokio::test]
1959async fn test_keyed_cleanup_idle_worker_can_be_triggered_externally() {
1960 use std::sync::atomic::{AtomicUsize, Ordering};
1961 use std::sync::Arc;
1962
1963 #[derive(Clone, Debug, Eq, PartialEq, Hash)]
1964 enum TestWorkerKey {
1965 A,
1966 B,
1967 }
1968
1969 struct TestWorker {
1970 id: usize,
1971 key: TestWorkerKey,
1972 }
1973
1974 #[async_trait::async_trait]
1975 impl KeyedWorker<TestWorkerKey> for TestWorker {
1976 fn is_work(&self) -> bool {
1977 true
1978 }
1979
1980 fn supports(&self, key: TestWorkerKey) -> bool {
1981 self.key == key
1982 }
1983
1984 fn primary_key(&self) -> TestWorkerKey {
1985 self.key.clone()
1986 }
1987 }
1988
1989 struct TestWorkerFactory {
1990 create_count: Arc<AtomicUsize>,
1991 }
1992
1993 #[async_trait::async_trait]
1994 impl KeyedWorkerFactory<TestWorkerKey, TestWorker> for TestWorkerFactory {
1995 async fn create(&self, key: TestWorkerKey) -> PoolResult<TestWorker> {
1996 let id = self.create_count.fetch_add(1, Ordering::SeqCst);
1997 Ok(TestWorker { id, key })
1998 }
1999 }
2000
2001 let create_count = Arc::new(AtomicUsize::new(0));
2002 let pool = KeyedWorkerPool::new(
2003 TestWorkerFactory {
2004 create_count: create_count.clone(),
2005 },
2006 KeyedWorkerPoolConfig {
2007 max_count_per_key: Some(1),
2008 idle_timeout: Some(std::time::Duration::from_millis(30)),
2009 ..Default::default()
2010 },
2011 );
2012
2013 {
2014 let worker = pool.get_worker(TestWorkerKey::B).await.unwrap();
2015 assert_eq!(worker.id, 0);
2016 assert_eq!(worker.primary_key(), TestWorkerKey::B);
2017 }
2018
2019 tokio::time::sleep(std::time::Duration::from_millis(80)).await;
2020
2021 assert_eq!(pool.cleanup_idle_worker(), 1);
2022
2023 let worker = pool.get_worker(TestWorkerKey::A).await.unwrap();
2024 assert_eq!(worker.id, 1);
2025 assert_eq!(worker.primary_key(), TestWorkerKey::A);
2026 assert_eq!(create_count.load(Ordering::SeqCst), 2);
2027}
2028
2029#[tokio::test]
2030async fn test_canceled_keyed_create_rolls_back_reservation() {
2031 use std::sync::atomic::{AtomicBool, Ordering};
2032
2033 #[derive(Clone, Debug, Eq, PartialEq, Hash)]
2034 enum TestWorkerKey {
2035 A,
2036 }
2037
2038 struct TestWorker;
2039
2040 #[async_trait::async_trait]
2041 impl KeyedWorker<TestWorkerKey> for TestWorker {
2042 fn is_work(&self) -> bool {
2043 true
2044 }
2045
2046 fn supports(&self, _key: TestWorkerKey) -> bool {
2047 true
2048 }
2049
2050 fn primary_key(&self) -> TestWorkerKey {
2051 TestWorkerKey::A
2052 }
2053 }
2054
2055 struct TestWorkerFactory {
2056 create_started: Arc<AtomicBool>,
2057 allow_create: Arc<AtomicBool>,
2058 }
2059
2060 #[async_trait::async_trait]
2061 impl KeyedWorkerFactory<TestWorkerKey, TestWorker> for TestWorkerFactory {
2062 async fn create(&self, _key: TestWorkerKey) -> PoolResult<TestWorker> {
2063 self.create_started.store(true, Ordering::SeqCst);
2064 while !self.allow_create.load(Ordering::SeqCst) {
2065 tokio::task::yield_now().await;
2066 }
2067 Ok(TestWorker)
2068 }
2069 }
2070
2071 let create_started = Arc::new(AtomicBool::new(false));
2072 let allow_create = Arc::new(AtomicBool::new(false));
2073 let pool = new_keyed_worker_pool(
2074 1,
2075 TestWorkerFactory {
2076 create_started: create_started.clone(),
2077 allow_create: allow_create.clone(),
2078 },
2079 );
2080
2081 let pool_ref = pool.clone();
2082 let create_task = tokio::spawn(async move { pool_ref.get_worker(TestWorkerKey::A).await });
2083 while !create_started.load(Ordering::SeqCst) {
2084 tokio::task::yield_now().await;
2085 }
2086 create_task.abort();
2087 assert!(matches!(create_task.await, Err(err) if err.is_cancelled()));
2088
2089 allow_create.store(true, Ordering::SeqCst);
2090 let worker = tokio::time::timeout(
2091 std::time::Duration::from_secs(1),
2092 pool.get_worker(TestWorkerKey::A),
2093 )
2094 .await
2095 .unwrap()
2096 .unwrap();
2097 drop(worker);
2098
2099 tokio::time::timeout(std::time::Duration::from_secs(1), pool.clear_all_worker())
2100 .await
2101 .unwrap();
2102}
2103
2104#[tokio::test]
2105async fn test_mutating_worker_key_removes_returned_worker() {
2106 #[derive(Clone, Debug, Eq, PartialEq, Hash)]
2107 enum TestWorkerKey {
2108 A,
2109 B,
2110 }
2111
2112 struct TestWorker {
2113 work: bool,
2114 key: TestWorkerKey,
2115 }
2116
2117 #[async_trait::async_trait]
2118 impl KeyedWorker<TestWorkerKey> for TestWorker {
2119 fn is_work(&self) -> bool {
2120 self.work
2121 }
2122
2123 fn supports(&self, key: TestWorkerKey) -> bool {
2124 self.key == key
2125 }
2126
2127 fn primary_key(&self) -> TestWorkerKey {
2128 self.key.clone()
2129 }
2130 }
2131
2132 struct TestWorkerFactory;
2133
2134 #[async_trait::async_trait]
2135 impl KeyedWorkerFactory<TestWorkerKey, TestWorker> for TestWorkerFactory {
2136 async fn create(&self, key: TestWorkerKey) -> PoolResult<TestWorker> {
2137 Ok(TestWorker { work: true, key })
2138 }
2139 }
2140
2141 let pool = KeyedWorkerPool::new(
2142 TestWorkerFactory,
2143 KeyedWorkerPoolConfig {
2144 idle_timeout: None,
2145 max_count_per_key: Some(1),
2146 ..Default::default()
2147 },
2148 );
2149 let mut worker = pool.get_worker(TestWorkerKey::A).await.unwrap();
2150 worker.key = TestWorkerKey::B;
2151 drop(worker);
2152
2153 {
2154 let state = pool.state.lock().unwrap();
2155 assert_eq!(state.current_count, 0);
2156 assert!(state.worker_count_by_key.is_empty());
2157 }
2158
2159 let worker = tokio::time::timeout(
2160 std::time::Duration::from_secs(1),
2161 pool.get_worker(TestWorkerKey::A),
2162 )
2163 .await
2164 .unwrap()
2165 .unwrap();
2166 assert_eq!(worker.primary_key(), TestWorkerKey::A);
2167}
2168
2169#[cfg(test)]
2170mod affected_path_tests {
2171 use super::*;
2172 use std::collections::VecDeque;
2173 use std::sync::atomic::{AtomicBool, Ordering};
2174 use std::sync::mpsc;
2175
2176 #[derive(Clone, Debug, Eq, PartialEq, Hash)]
2177 enum Key {
2178 A,
2179 B,
2180 C,
2181 }
2182
2183 struct BlockingWorker {
2184 key: Key,
2185 }
2186
2187 #[async_trait::async_trait]
2188 impl KeyedWorker<Key> for BlockingWorker {
2189 fn is_work(&self) -> bool {
2190 true
2191 }
2192
2193 fn supports(&self, key: Key) -> bool {
2194 self.key == key
2195 }
2196
2197 fn primary_key(&self) -> Key {
2198 self.key.clone()
2199 }
2200 }
2201
2202 struct BlockingFactory {
2203 create_started: Arc<AtomicBool>,
2204 allow_create: Arc<AtomicBool>,
2205 }
2206
2207 #[async_trait::async_trait]
2208 impl KeyedWorkerFactory<Key, BlockingWorker> for BlockingFactory {
2209 async fn create(&self, key: Key) -> PoolResult<BlockingWorker> {
2210 self.create_started.store(true, Ordering::SeqCst);
2211 while !self.allow_create.load(Ordering::SeqCst) {
2212 tokio::task::yield_now().await;
2213 }
2214 Ok(BlockingWorker { key })
2215 }
2216 }
2217
2218 async fn wait_for_create(create_started: &AtomicBool) {
2219 while !create_started.load(Ordering::SeqCst) {
2220 tokio::task::yield_now().await;
2221 }
2222 }
2223
2224 fn new_blocking_pool(
2225 max_count: u16,
2226 ) -> (
2227 KeyedWorkerPoolRef<Key, BlockingWorker, BlockingFactory>,
2228 Arc<AtomicBool>,
2229 Arc<AtomicBool>,
2230 ) {
2231 let create_started = Arc::new(AtomicBool::new(false));
2232 let allow_create = Arc::new(AtomicBool::new(false));
2233 let pool = new_keyed_worker_pool(
2234 max_count,
2235 BlockingFactory {
2236 create_started: create_started.clone(),
2237 allow_create: allow_create.clone(),
2238 },
2239 );
2240 (pool, create_started, allow_create)
2241 }
2242
2243 #[tokio::test]
2244 async fn test_canceled_create_rolls_back_keyed_pool_reservation() {
2245 let (pool, create_started, allow_create) = new_blocking_pool(1);
2246 let pool_ref = pool.clone();
2247 let create_task = tokio::spawn(async move { pool_ref.get_worker(Key::A).await });
2248 wait_for_create(&create_started).await;
2249 create_task.abort();
2250 assert!(matches!(create_task.await, Err(err) if err.is_cancelled()));
2251
2252 allow_create.store(true, Ordering::SeqCst);
2253 let worker = tokio::time::timeout(Duration::from_secs(1), pool.get_worker(Key::A))
2254 .await
2255 .unwrap()
2256 .unwrap();
2257 drop(worker);
2258 tokio::time::timeout(Duration::from_secs(1), pool.clear_all_worker())
2259 .await
2260 .unwrap();
2261 }
2262
2263 #[tokio::test]
2264 async fn test_canceled_replacement_create_rolls_back_keyed_pool_reservation() {
2265 let (pool, create_started, allow_create) = new_blocking_pool(1);
2266 allow_create.store(true, Ordering::SeqCst);
2267 let worker = pool.get_worker(Key::A).await.unwrap();
2268 drop(worker);
2269
2270 create_started.store(false, Ordering::SeqCst);
2271 allow_create.store(false, Ordering::SeqCst);
2272 let pool_ref = pool.clone();
2273 let create_task = tokio::spawn(async move { pool_ref.get_worker(Key::B).await });
2274 wait_for_create(&create_started).await;
2275 create_task.abort();
2276 assert!(matches!(create_task.await, Err(err) if err.is_cancelled()));
2277
2278 allow_create.store(true, Ordering::SeqCst);
2279 let worker = tokio::time::timeout(Duration::from_secs(1), pool.get_worker(Key::B))
2280 .await
2281 .unwrap()
2282 .unwrap();
2283 drop(worker);
2284 tokio::time::timeout(Duration::from_secs(1), pool.clear_all_worker())
2285 .await
2286 .unwrap();
2287 }
2288
2289 #[tokio::test]
2290 async fn test_canceled_overcommit_create_rolls_back_keyed_pool_reservation() {
2291 let (pool, create_started, allow_create) = new_blocking_pool(1);
2292 allow_create.store(true, Ordering::SeqCst);
2293 let worker_a = pool.get_worker(Key::A).await.unwrap();
2294
2295 create_started.store(false, Ordering::SeqCst);
2296 allow_create.store(false, Ordering::SeqCst);
2297 let pool_ref = pool.clone();
2298 let create_task = tokio::spawn(async move { pool_ref.get_worker(Key::B).await });
2299 wait_for_create(&create_started).await;
2300 create_task.abort();
2301 assert!(matches!(create_task.await, Err(err) if err.is_cancelled()));
2302
2303 {
2304 let state = pool.state.lock().unwrap();
2305 assert_eq!(state.current_count, 1);
2306 assert_eq!(state.reserved_count_for_key(&Key::B), 0);
2307 }
2308
2309 allow_create.store(true, Ordering::SeqCst);
2310 let worker_b = tokio::time::timeout(Duration::from_secs(1), pool.get_worker(Key::B))
2311 .await
2312 .unwrap()
2313 .unwrap();
2314 drop(worker_b);
2315 drop(worker_a);
2316 tokio::time::timeout(Duration::from_secs(1), pool.clear_all_worker())
2317 .await
2318 .unwrap();
2319 }
2320
2321 type DropCallback = Box<dyn FnOnce() + Send>;
2322
2323 struct DropProbeWorker {
2324 working: Arc<AtomicBool>,
2325 key: Key,
2326 on_drop: Option<DropCallback>,
2327 }
2328
2329 #[async_trait::async_trait]
2330 impl KeyedWorker<Key> for DropProbeWorker {
2331 fn is_work(&self) -> bool {
2332 self.working.load(Ordering::SeqCst)
2333 }
2334
2335 fn supports(&self, key: Key) -> bool {
2336 self.key == key
2337 }
2338
2339 fn primary_key(&self) -> Key {
2340 self.key.clone()
2341 }
2342 }
2343
2344 impl Drop for DropProbeWorker {
2345 fn drop(&mut self) {
2346 if let Some(on_drop) = self.on_drop.take() {
2347 on_drop();
2348 }
2349 }
2350 }
2351
2352 struct DropProbeSpec {
2353 working: Arc<AtomicBool>,
2354 on_drop: Option<DropCallback>,
2355 }
2356
2357 struct DropProbeFactory {
2358 specs: Arc<Mutex<VecDeque<DropProbeSpec>>>,
2359 }
2360
2361 #[async_trait::async_trait]
2362 impl KeyedWorkerFactory<Key, DropProbeWorker> for DropProbeFactory {
2363 async fn create(&self, key: Key) -> PoolResult<DropProbeWorker> {
2364 let spec = self.specs.lock().unwrap().pop_front().unwrap();
2365 Ok(DropProbeWorker {
2366 working: spec.working,
2367 key,
2368 on_drop: spec.on_drop,
2369 })
2370 }
2371 }
2372
2373 type DropProbePool = KeyedWorkerPoolRef<Key, DropProbeWorker, DropProbeFactory>;
2374
2375 fn new_drop_probe_pool(
2376 idle_timeout: Option<Duration>,
2377 max_idle_count: Option<u16>,
2378 ) -> (DropProbePool, Arc<Mutex<VecDeque<DropProbeSpec>>>) {
2379 let specs = Arc::new(Mutex::new(VecDeque::new()));
2380 let pool = KeyedWorkerPool::new(
2381 DropProbeFactory {
2382 specs: specs.clone(),
2383 },
2384 KeyedWorkerPoolConfig {
2385 max_count_per_key: Some(1),
2386 max_idle_count_per_key: max_idle_count,
2387 idle_timeout,
2388 },
2389 );
2390 (pool, specs)
2391 }
2392
2393 fn drop_lock_probe(
2394 pool: &DropProbePool,
2395 working: Arc<AtomicBool>,
2396 ) -> (DropProbeSpec, mpsc::Receiver<bool>) {
2397 let (tx, rx) = mpsc::channel();
2398 let pool_ref = Arc::downgrade(pool);
2399 let on_drop = Box::new(move || {
2400 let pool_ref = pool_ref.upgrade().unwrap();
2401 tx.send(pool_ref.state.try_lock().is_ok()).unwrap();
2402 });
2403 (
2404 DropProbeSpec {
2405 working,
2406 on_drop: Some(on_drop),
2407 },
2408 rx,
2409 )
2410 }
2411
2412 fn plain_drop_probe_spec() -> DropProbeSpec {
2413 DropProbeSpec {
2414 working: Arc::new(AtomicBool::new(true)),
2415 on_drop: None,
2416 }
2417 }
2418
2419 #[derive(Copy, Clone)]
2420 enum IdleDropPath {
2421 Cleanup,
2422 IdleLimit,
2423 KeyedInvalidScan,
2424 Clear,
2425 }
2426
2427 async fn assert_idle_drop_path_runs_outside_lock(path: IdleDropPath) {
2428 let idle_timeout = matches!(path, IdleDropPath::Cleanup).then_some(Duration::ZERO);
2429 let max_idle_count = matches!(path, IdleDropPath::IdleLimit).then_some(0);
2430 let (pool, specs) = new_drop_probe_pool(idle_timeout, max_idle_count);
2431 let working = Arc::new(AtomicBool::new(true));
2432 let (spec, drop_result) = drop_lock_probe(&pool, working.clone());
2433 specs.lock().unwrap().push_back(spec);
2434
2435 let worker = pool.get_worker(Key::A).await.unwrap();
2436 drop(worker);
2437
2438 match path {
2439 IdleDropPath::Cleanup => {
2440 assert_eq!(pool.cleanup_idle_worker(), 1);
2441 }
2442 IdleDropPath::IdleLimit => {
2443 }
2445 IdleDropPath::KeyedInvalidScan => {
2446 working.store(false, Ordering::SeqCst);
2447 specs.lock().unwrap().push_back(plain_drop_probe_spec());
2448 let worker = pool.get_worker(Key::A).await.unwrap();
2449 drop(worker);
2450 }
2451 IdleDropPath::Clear => pool.clear_all_worker().await,
2452 }
2453
2454 assert!(drop_result.recv_timeout(Duration::from_secs(1)).unwrap());
2455 }
2456
2457 #[tokio::test]
2458 async fn test_all_keyed_idle_drop_paths_run_outside_state_lock() {
2459 for path in [
2460 IdleDropPath::Cleanup,
2461 IdleDropPath::IdleLimit,
2462 IdleDropPath::KeyedInvalidScan,
2463 IdleDropPath::Clear,
2464 ] {
2465 assert_idle_drop_path_runs_outside_lock(path).await;
2466 }
2467 }
2468
2469 struct MutableWorker {
2470 working: Arc<AtomicBool>,
2471 valid: Arc<AtomicBool>,
2472 key: Key,
2473 }
2474
2475 #[async_trait::async_trait]
2476 impl KeyedWorker<Key> for MutableWorker {
2477 fn is_work(&self) -> bool {
2478 self.working.load(Ordering::SeqCst)
2479 }
2480
2481 fn supports(&self, key: Key) -> bool {
2482 self.valid.load(Ordering::SeqCst) && self.key == key
2483 }
2484
2485 fn primary_key(&self) -> Key {
2486 self.key.clone()
2487 }
2488 }
2489
2490 struct MutableWorkerFactory;
2491
2492 #[async_trait::async_trait]
2493 impl KeyedWorkerFactory<Key, MutableWorker> for MutableWorkerFactory {
2494 async fn create(&self, key: Key) -> PoolResult<MutableWorker> {
2495 Ok(MutableWorker {
2496 working: Arc::new(AtomicBool::new(true)),
2497 valid: Arc::new(AtomicBool::new(true)),
2498 key,
2499 })
2500 }
2501 }
2502
2503 type MutablePool = KeyedWorkerPoolRef<Key, MutableWorker, MutableWorkerFactory>;
2504
2505 fn new_mutable_pool(max_count: u16, idle_timeout: Option<Duration>) -> MutablePool {
2506 KeyedWorkerPool::new(
2507 MutableWorkerFactory,
2508 KeyedWorkerPoolConfig {
2509 max_count_per_key: Some(max_count),
2510 idle_timeout,
2511 ..Default::default()
2512 },
2513 )
2514 }
2515
2516 fn new_limited_mutable_pool(_max_count: u16, max_count_per_key: u16) -> MutablePool {
2517 KeyedWorkerPool::new(
2518 MutableWorkerFactory,
2519 KeyedWorkerPoolConfig {
2520 idle_timeout: None,
2521 max_count_per_key: Some(max_count_per_key),
2522 ..Default::default()
2523 },
2524 )
2525 }
2526
2527 fn assert_accounting_empty(pool: &MutablePool) {
2528 let state = pool.state.lock().unwrap();
2529 assert_eq!(state.current_count, 0);
2530 assert!(state.worker_count_by_key.is_empty());
2531 assert!(state.pending_count_by_key.is_empty());
2532 }
2533
2534 fn assert_only_key(pool: &MutablePool, key: Key, count: usize) {
2535 let state = pool.state.lock().unwrap();
2536 assert_eq!(state.current_count, count);
2537 assert_eq!(state.worker_count_by_key.len(), 1);
2538 assert_eq!(state.worker_count_by_key.get(&key).copied(), Some(count));
2539 }
2540
2541 #[derive(Copy, Clone)]
2542 enum IdleAccountingPath {
2543 Cleanup,
2544 Clear,
2545 KeyedInvalidScan,
2546 }
2547
2548 async fn assert_idle_accounting_path(path: IdleAccountingPath) {
2549 let idle_timeout = matches!(path, IdleAccountingPath::Cleanup).then_some(Duration::ZERO);
2550 let pool = new_mutable_pool(1, idle_timeout);
2551 let worker = pool.get_worker(Key::A).await.unwrap();
2552 let working = worker.working.clone();
2553 drop(worker);
2554
2555 match path {
2556 IdleAccountingPath::Cleanup => {
2557 assert_eq!(pool.cleanup_idle_worker(), 1);
2558 assert_accounting_empty(&pool);
2559 }
2560 IdleAccountingPath::Clear => {
2561 pool.clear_all_worker().await;
2562 assert_accounting_empty(&pool);
2563 }
2564 IdleAccountingPath::KeyedInvalidScan => {
2565 working.store(false, Ordering::SeqCst);
2566 let worker = pool.get_worker(Key::A).await.unwrap();
2567 assert_only_key(&pool, Key::A, 1);
2568 worker.working.store(false, Ordering::SeqCst);
2569 drop(worker);
2570 assert_accounting_empty(&pool);
2571 }
2572 }
2573 }
2574
2575 #[tokio::test]
2576 async fn test_all_idle_accounting_paths() {
2577 for path in [
2578 IdleAccountingPath::Cleanup,
2579 IdleAccountingPath::Clear,
2580 IdleAccountingPath::KeyedInvalidScan,
2581 ] {
2582 assert_idle_accounting_path(path).await;
2583 }
2584 }
2585
2586 async fn wait_for_waiter(pool: &MutablePool) {
2587 loop {
2588 if !pool.state.lock().unwrap().waiting_list.is_empty() {
2589 return;
2590 }
2591 tokio::task::yield_now().await;
2592 }
2593 }
2594
2595 #[tokio::test]
2596 async fn test_canceled_waiter_is_removed_when_worker_returns() {
2597 let pool = new_limited_mutable_pool(1, 1);
2598 let worker_a = pool.get_worker(Key::A).await.unwrap();
2599
2600 let pool_ref = pool.clone();
2601 let waiter = tokio::spawn(async move { pool_ref.get_worker(Key::A).await });
2602 wait_for_waiter(&pool).await;
2603 waiter.abort();
2604 assert!(matches!(waiter.await, Err(err) if err.is_cancelled()));
2605
2606 drop(worker_a);
2607
2608 let state = pool.state.lock().unwrap();
2609 assert!(state.waiting_list.is_empty());
2610 assert_eq!(
2611 state.worker_list.values().map(VecDeque::len).sum::<usize>(),
2612 1
2613 );
2614 }
2615
2616 #[tokio::test]
2617 async fn test_worker_invalid_for_primary_key_wakes_waiter() {
2618 let pool = new_limited_mutable_pool(1, 1);
2619 let worker_a = pool.get_worker(Key::A).await.unwrap();
2620 let valid = worker_a.valid.clone();
2621
2622 let pool_ref = pool.clone();
2623 let waiting_a = tokio::spawn(async move { pool_ref.get_worker(Key::A).await });
2624 wait_for_waiter(&pool).await;
2625
2626 valid.store(false, Ordering::SeqCst);
2627 drop(worker_a);
2628
2629 let replacement_a = tokio::time::timeout(Duration::from_secs(1), waiting_a)
2630 .await
2631 .unwrap()
2632 .unwrap()
2633 .unwrap();
2634 assert_eq!(replacement_a.primary_key(), Key::A);
2635 }
2636
2637 #[tokio::test]
2638 async fn test_idle_worker_invalid_for_primary_key_is_replaced() {
2639 let pool = new_limited_mutable_pool(1, 1);
2640 let worker_a = pool.get_worker(Key::A).await.unwrap();
2641 let valid = worker_a.valid.clone();
2642 drop(worker_a);
2643
2644 valid.store(false, Ordering::SeqCst);
2645
2646 let replacement_a = tokio::time::timeout(Duration::from_secs(1), pool.get_worker(Key::A))
2647 .await
2648 .unwrap()
2649 .unwrap();
2650 assert_eq!(replacement_a.primary_key(), Key::A);
2651 }
2652
2653 #[tokio::test]
2654 async fn test_capped_waiter_does_not_replace_other_key_worker() {
2655 let pool = new_limited_mutable_pool(2, 1);
2656 let worker_a = pool.get_worker(Key::A).await.unwrap();
2657 let worker_b = pool.get_worker(Key::B).await.unwrap();
2658
2659 let pool_ref = pool.clone();
2660 let waiting_a = tokio::spawn(async move { pool_ref.get_worker(Key::A).await });
2661 wait_for_waiter(&pool).await;
2662
2663 drop(worker_b);
2664 tokio::task::yield_now().await;
2665
2666 {
2667 let state = pool.state.lock().unwrap();
2668 assert_eq!(state.current_count, 2);
2669 assert_eq!(
2670 state.worker_list.values().map(VecDeque::len).sum::<usize>(),
2671 1
2672 );
2673 assert_eq!(
2674 state
2675 .worker_list
2676 .get(&Key::B)
2677 .and_then(VecDeque::front)
2678 .unwrap()
2679 .primary_key,
2680 Key::B
2681 );
2682 }
2683 assert!(!waiting_a.is_finished());
2684
2685 drop(worker_a);
2686 let replacement_a = tokio::time::timeout(Duration::from_secs(1), waiting_a)
2687 .await
2688 .unwrap()
2689 .unwrap()
2690 .unwrap();
2691 assert_eq!(replacement_a.primary_key(), Key::A);
2692 }
2693
2694 #[tokio::test]
2695 async fn test_changed_key_worker_is_removed_and_waiter_retries() {
2696 let pool = new_mutable_pool(1, None);
2697 let mut worker = pool.get_worker(Key::A).await.unwrap();
2698 worker.key = Key::B;
2699
2700 let pool_ref = pool.clone();
2701 let waiter = tokio::spawn(async move { pool_ref.get_worker(Key::A).await });
2702 wait_for_waiter(&pool).await;
2703 drop(worker);
2704
2705 let worker = waiter.await.unwrap().unwrap();
2706 assert_eq!(worker.primary_key(), Key::A);
2707 worker.working.store(false, Ordering::SeqCst);
2708 drop(worker);
2709 assert_accounting_empty(&pool);
2710 }
2711
2712 #[tokio::test]
2713 async fn test_cached_key_is_used_while_clearing() {
2714 let pool = new_mutable_pool(1, None);
2715 let mut worker = pool.get_worker(Key::A).await.unwrap();
2716 worker.key = Key::B;
2717
2718 let pool_ref = pool.clone();
2719 let clear_task = tokio::spawn(async move { pool_ref.clear_all_worker().await });
2720 loop {
2721 if pool.state.lock().unwrap().clearing {
2722 break;
2723 }
2724 tokio::task::yield_now().await;
2725 }
2726 drop(worker);
2727 clear_task.await.unwrap();
2728 assert_accounting_empty(&pool);
2729 }
2730
2731 #[tokio::test]
2732 async fn test_changed_key_worker_does_not_wake_other_key_waiter() {
2733 let pool = new_mutable_pool(1, None);
2734 let mut worker_a = pool.get_worker(Key::A).await.unwrap();
2735 let worker_b = pool.get_worker(Key::B).await.unwrap();
2736 worker_a.key = Key::C;
2737
2738 let pool_ref = pool.clone();
2739 let waiter = tokio::spawn(async move { pool_ref.get_worker(Key::B).await });
2740 wait_for_waiter(&pool).await;
2741 drop(worker_a);
2742 tokio::task::yield_now().await;
2743 assert!(!waiter.is_finished());
2744
2745 drop(worker_b);
2746 let replacement_b = waiter.await.unwrap().unwrap();
2747 assert_only_key(&pool, Key::B, 1);
2748 replacement_b.working.store(false, Ordering::SeqCst);
2749 drop(replacement_b);
2750 assert_accounting_empty(&pool);
2751 }
2752
2753 #[tokio::test]
2754 async fn test_changed_key_worker_is_removed_independently() {
2755 let pool = new_mutable_pool(1, None);
2756 let worker_a = pool.get_worker(Key::A).await.unwrap();
2757 let mut worker_b = pool.get_worker(Key::B).await.unwrap();
2758 worker_b.key = Key::C;
2759 drop(worker_b);
2760
2761 assert_only_key(&pool, Key::A, 1);
2762 worker_a.working.store(false, Ordering::SeqCst);
2763 drop(worker_a);
2764 assert_accounting_empty(&pool);
2765 }
2766
2767 #[tokio::test]
2768 async fn test_max_count_per_key_blocks_only_that_key() {
2769 let pool = new_limited_mutable_pool(4, 2);
2770 let worker_a1 = pool.get_worker(Key::A).await.unwrap();
2771 let worker_a2 = pool.get_worker(Key::A).await.unwrap();
2772
2773 let pool_ref = pool.clone();
2774 let waiting_a = tokio::spawn(async move { pool_ref.get_worker(Key::A).await });
2775 wait_for_waiter(&pool).await;
2776 assert!(!waiting_a.is_finished());
2777
2778 let worker_b = pool.get_worker(Key::B).await.unwrap();
2779 {
2780 let state = pool.state.lock().unwrap();
2781 assert_eq!(state.current_count, 3);
2782 assert_eq!(state.reserved_count_for_key(&Key::A), 2);
2783 assert_eq!(state.reserved_count_for_key(&Key::B), 1);
2784 }
2785
2786 drop(worker_a1);
2787 let replacement_a = tokio::time::timeout(Duration::from_secs(1), waiting_a)
2788 .await
2789 .unwrap()
2790 .unwrap()
2791 .unwrap();
2792 assert_eq!(replacement_a.primary_key(), Key::A);
2793
2794 drop(worker_a2);
2795 drop(replacement_a);
2796 drop(worker_b);
2797 }
2798
2799 #[tokio::test]
2800 async fn test_zero_max_count_per_key_returns_error() {
2801 let pool = new_limited_mutable_pool(1, 0);
2802
2803 let keyed_error = pool.get_worker(Key::A).await.err().unwrap();
2804 assert_eq!(keyed_error.code(), crate::PoolErrorCode::InvalidConfig);
2805 }
2806}