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