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