Skip to main content

sfo_pool/
keyed_worker_pool.rs

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)]
16/// Configuration for [`KeyedWorkerPool`].
17///
18/// All options default to `None`. Use the `with_*` methods to configure the pool.
19pub struct KeyedWorkerPoolConfig {
20    idle_timeout: Option<Duration>,
21    /// Maximum number of workers whose primary key is the same.
22    ///
23    /// `None` leaves key counts unlimited.
24    max_count_per_key: Option<u16>,
25    /// Maximum number of idle workers retained for each primary key.
26    ///
27    /// This limit is independent for every key. `None` adds no per-key idle
28    /// limit; `Some(0)` disables idle caching for every key.
29    max_idle_count_per_key: Option<u16>,
30}
31
32impl KeyedWorkerPoolConfig {
33    /// Sets the maximum duration for which an idle worker is retained.
34    ///
35    /// `None` disables timeout-based cleanup.
36    pub fn with_idle_timeout(mut self, idle_timeout: Option<Duration>) -> Self {
37        self.idle_timeout = idle_timeout;
38        self
39    }
40
41    /// Sets the worker-count limit for each primary key.
42    ///
43    /// `None` leaves per-key counts unlimited. `Some(0)` is invalid.
44    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    /// Sets the idle-worker cache limit for each primary key.
50    ///
51    /// `None` adds no per-key idle limit. `Some(0)` disables idle caching.
52    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]
59/// A keyed worker managed by [`KeyedWorkerPool`].
60///
61/// Methods on this trait may be called while the pool's internal state lock is held.
62/// Implementations must be non-blocking and must not re-enter APIs on the same pool.
63pub trait KeyedWorker<K: WorkerKey>: Send + 'static {
64    fn is_work(&self) -> bool;
65    /// Returns whether this worker can currently serve its requested primary key.
66    /// Explicit keyed acquisition only reuses workers from that key's primary bucket.
67    /// A worker that is no longer valid for its cached primary key is discarded.
68    fn supports(&self, key: K) -> bool;
69    /// Returns the worker's primary key used for accounting and replacement.
70    /// The pool validates and caches this value when the worker is created.
71    /// If it differs from the cached value when the worker is returned, the worker is discarded.
72    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    /// Creates a usable worker for `key`.
175    ///
176    /// Returning `Ok` asserts that the worker is ready for use. Its primary
177    /// key must be `key`.
178    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    // Idle workers are grouped by primary key; each bucket is ordered from LRU to MRU.
203    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    /// Creates a keyed worker pool with explicit per-key configuration.
371    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                // The first return above evicts immediately from the zero-cap pool.
2444            }
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}