1use std::{
19 collections::VecDeque,
20 mem::ManuallyDrop,
21 ops::{Deref, DerefMut},
22 sync::{Arc, Mutex},
23};
24
25#[derive(Debug)]
27pub struct ObjectPool<T> {
28 queue: Mutex<VecDeque<T>>,
30
31 capacity: Option<usize>,
34}
35
36impl<T> ObjectPool<T> {
37 pub fn new<A>(args: A, initial_size: usize, capacity: Option<usize>) -> Self
44 where
45 T: AsPooled<A>,
46 A: Clone,
47 {
48 let queue = (0..initial_size).map(|_| T::create(args.clone())).collect();
49
50 Self {
51 queue: Mutex::new(queue),
52 capacity,
53 }
54 }
55
56 pub fn try_new<A>(
63 args: A,
64 initial_size: usize,
65 capacity: Option<usize>,
66 ) -> Result<Self, T::Error>
67 where
68 T: TryAsPooled<A>,
69 A: Clone,
70 {
71 let queue = (0..initial_size)
72 .map(|_| T::try_create(args.clone()))
73 .collect::<Result<_, _>>()?;
74 Ok(Self {
75 queue: Mutex::new(queue),
76 capacity,
77 })
78 }
79
80 pub fn try_get_ref<A>(&self, args: A) -> Result<PooledRef<'_, T>, T::Error>
87 where
88 T: TryAsPooled<A>,
89 {
90 let item = self.try_get_or_create(args)?;
91 Ok(PooledRef {
92 item: ManuallyDrop::new(item),
93 parent: self,
94 })
95 }
96
97 pub fn get_ref<A>(&self, args: A) -> PooledRef<'_, T>
106 where
107 T: AsPooled<A>,
108 {
109 let item = self.get_or_create(args);
110 PooledRef {
111 item: ManuallyDrop::new(item),
112 parent: self,
113 }
114 }
115
116 pub fn try_get<A>(self: &Arc<Self>, args: A) -> Result<PooledArc<T>, T::Error>
123 where
124 T: TryAsPooled<A>,
125 {
126 let item = self.try_get_or_create(args)?;
127 Ok(PooledArc {
128 item: ManuallyDrop::new(item),
129 parent: self.clone(),
130 })
131 }
132
133 pub fn get<A>(self: &Arc<Self>, args: A) -> PooledArc<T>
142 where
143 T: AsPooled<A>,
144 {
145 let item = self.get_or_create(args);
146 PooledArc {
147 item: ManuallyDrop::new(item),
148 parent: self.clone(),
149 }
150 }
151
152 pub fn len(&self) -> usize {
154 self.lock().len()
155 }
156
157 pub fn is_empty(&self) -> bool {
159 self.len() == 0
160 }
161
162 fn try_get_or_create<A>(&self, args: A) -> Result<T, T::Error>
167 where
168 T: TryAsPooled<A>,
169 {
170 let maybe = self.lock().pop_front();
174 if let Some(mut item) = maybe {
175 item.try_modify(args)?;
176 Ok(item)
177 } else {
178 T::try_create(args)
179 }
180 }
181
182 fn get_or_create<A>(&self, args: A) -> T
183 where
184 T: AsPooled<A>,
185 {
186 let maybe = self.lock().pop_front();
190 if let Some(mut item) = maybe {
191 item.modify(args);
192 item
193 } else {
194 T::create(args)
195 }
196 }
197
198 fn lock(&self) -> std::sync::MutexGuard<'_, VecDeque<T>> {
199 match self.queue.lock() {
200 Ok(guard) => guard,
201 Err(poisoned) => {
202 self.queue.clear_poison();
212 poisoned.into_inner()
213 }
214 }
215 }
216}
217
218pub trait TryAsPooled<A>
225where
226 Self: Sized,
227{
228 type Error;
230
231 fn try_create(args: A) -> Result<Self, Self::Error>;
233
234 fn try_modify(&mut self, args: A) -> Result<(), Self::Error>;
243}
244
245pub trait AsPooled<A> {
251 fn create(args: A) -> Self;
253
254 fn modify(&mut self, args: A);
263}
264
265#[derive(Debug, Clone, Copy)]
273pub struct Undef {
274 pub len: usize,
275}
276
277impl Undef {
278 pub fn new(len: usize) -> Self {
280 Self { len }
281 }
282}
283
284impl<T> AsPooled<Undef> for Vec<T>
285where
286 T: Default + Clone,
287{
288 fn create(undef: Undef) -> Self {
289 vec![T::default(); undef.len]
290 }
291
292 fn modify(&mut self, undef: Undef) {
293 self.resize(undef.len, T::default())
294 }
295}
296
297#[derive(Debug)]
301pub struct PooledRef<'a, T> {
302 item: ManuallyDrop<T>,
303 parent: &'a ObjectPool<T>,
304}
305
306impl<T> Drop for PooledRef<'_, T> {
307 fn drop(&mut self) {
308 let mut guard = self.parent.lock();
309 if guard.len() < self.parent.capacity.unwrap_or(usize::MAX) {
310 guard.push_back(unsafe { ManuallyDrop::take(&mut self.item) });
312 } else {
313 std::mem::drop(guard);
315
316 unsafe { ManuallyDrop::drop(&mut self.item) };
318 }
319 }
320}
321
322impl<T> Deref for PooledRef<'_, T> {
323 type Target = T;
324
325 fn deref(&self) -> &Self::Target {
326 &self.item
327 }
328}
329
330impl<T> DerefMut for PooledRef<'_, T> {
331 fn deref_mut(&mut self) -> &mut Self::Target {
332 &mut self.item
333 }
334}
335
336#[derive(Debug)]
342pub struct PooledArc<T> {
343 item: ManuallyDrop<T>,
344 parent: Arc<ObjectPool<T>>,
345}
346
347impl<T> Drop for PooledArc<T> {
348 fn drop(&mut self) {
349 let mut guard = self.parent.lock();
350 if guard.len() < self.parent.capacity.unwrap_or(usize::MAX) {
351 guard.push_back(unsafe { ManuallyDrop::take(&mut self.item) });
353 } else {
354 std::mem::drop(guard);
356
357 unsafe { ManuallyDrop::drop(&mut self.item) };
359 }
360 }
361}
362
363impl<T> Deref for PooledArc<T> {
364 type Target = T;
365
366 fn deref(&self) -> &Self::Target {
367 &self.item
368 }
369}
370
371impl<T> DerefMut for PooledArc<T> {
372 fn deref_mut(&mut self) -> &mut Self::Target {
373 &mut self.item
374 }
375}
376
377#[derive(Debug)]
380pub enum PoolOption<T> {
381 NonPooled(T),
382 Pooled(PooledArc<T>),
383}
384
385impl<T> PoolOption<T> {
386 pub fn non_pooled(item: T) -> Self {
387 PoolOption::NonPooled(item)
388 }
389
390 pub fn try_non_pooled_create<A>(args: A) -> Result<Self, T::Error>
391 where
392 T: TryAsPooled<A>,
393 {
394 Ok(PoolOption::NonPooled(T::try_create(args)?))
395 }
396
397 pub fn non_pooled_create<A>(args: A) -> Self
398 where
399 T: AsPooled<A>,
400 {
401 PoolOption::NonPooled(T::create(args))
402 }
403
404 pub fn pooled<A>(pool: &Arc<ObjectPool<T>>, args: A) -> Self
405 where
406 T: AsPooled<A>,
407 {
408 PoolOption::Pooled(pool.get(args))
409 }
410
411 pub fn try_pooled<A>(pool: &Arc<ObjectPool<T>>, args: A) -> Result<Self, T::Error>
412 where
413 T: TryAsPooled<A>,
414 {
415 Ok(PoolOption::Pooled(pool.try_get(args)?))
416 }
417
418 pub fn is_pooled(&self) -> bool {
419 matches!(self, PoolOption::Pooled(_))
420 }
421
422 pub fn is_non_pooled(&self) -> bool {
423 matches!(self, PoolOption::NonPooled(_))
424 }
425}
426
427impl<T> Deref for PoolOption<T> {
428 type Target = T;
429
430 fn deref(&self) -> &Self::Target {
431 match self {
432 PoolOption::NonPooled(item) => item,
433 PoolOption::Pooled(item) => item,
434 }
435 }
436}
437
438impl<T> DerefMut for PoolOption<T> {
439 fn deref_mut(&mut self) -> &mut Self::Target {
440 match self {
441 PoolOption::NonPooled(item) => item,
442 PoolOption::Pooled(item) => item,
443 }
444 }
445}
446
447#[cfg(test)]
452mod tests {
453 use super::*;
454
455 use crate::assert_contains;
456
457 #[derive(Debug)]
458 struct TestItem {
459 value: Box<u32>,
460 panic_on_drop: bool,
461 }
462
463 impl TestItem {
464 fn new(value: u32) -> Self {
465 Self {
466 value: Box::new(value),
467 panic_on_drop: false,
468 }
469 }
470 }
471
472 impl AsPooled<u32> for TestItem {
473 fn create(value: u32) -> Self {
474 TestItem::new(value)
475 }
476
477 fn modify(&mut self, value: u32) {
478 *self.value = value;
479 self.panic_on_drop = false;
480 }
481 }
482
483 impl TryAsPooled<i32> for TestItem {
484 type Error = ();
485
486 fn try_create(value: i32) -> Result<Self, Self::Error> {
487 match value.try_into() {
488 Ok(v) => Ok(TestItem::new(v)),
489 Err(_) => Err(()),
490 }
491 }
492
493 fn try_modify(&mut self, value: i32) -> Result<(), Self::Error> {
494 match value.try_into() {
495 Ok(v) => {
496 *self.value = v;
497 self.panic_on_drop = false;
498 Ok(())
499 }
500 Err(_) => Err(()),
501 }
502 }
503 }
504
505 impl Drop for TestItem {
506 fn drop(&mut self) {
507 if self.panic_on_drop {
508 panic!("panicking on drop");
509 }
510 }
511 }
512
513 struct TestPanic;
515
516 impl AsPooled<TestPanic> for TestItem {
517 fn create(_: TestPanic) -> Self {
518 panic!("panicking on create")
519 }
520
521 fn modify(&mut self, _: TestPanic) {
522 panic!("panicking on modify")
523 }
524 }
525
526 impl TryAsPooled<TestPanic> for TestItem {
527 type Error = ();
528
529 fn try_create(_: TestPanic) -> Result<Self, Self::Error> {
530 panic!("panicking on try_create")
531 }
532
533 fn try_modify(&mut self, _: TestPanic) -> Result<(), Self::Error> {
534 panic!("panicking on try_modify")
535 }
536 }
537
538 #[test]
539 fn test_pool_basic_tests() {
540 let pool = ObjectPool::<TestItem>::new(42, 2, None);
541 assert_eq!(pool.len(), 2);
542
543 let item1 = pool.get_ref(100);
544 assert_eq!(*item1.value, 100);
545 assert_eq!(pool.len(), 1);
546
547 let item2 = pool.get_ref(200);
548 assert_eq!(*item2.value, 200);
549 assert_eq!(pool.len(), 0);
550
551 let item = pool.get_ref(300);
552 assert_eq!(*item.value, 300);
553 assert_eq!(pool.len(), 0);
554 {
555 let item = pool.get_ref(400);
556 assert_eq!(*item.value, 400);
557 assert_eq!(pool.len(), 0);
558 }
559 assert_eq!(pool.len(), 1); {
561 let item_a = pool.get_ref(500);
562 assert_eq!(*item_a.value, 500);
563 assert_eq!(pool.len(), 0); let item_b = pool.get_ref(600);
565 assert_eq!(*item_b.value, 600);
566 assert_eq!(pool.len(), 0); }
568 assert_eq!(pool.len(), 2); let pool = ObjectPool::<TestItem>::new(42, 1, None);
571 let item = pool.get_ref(100);
572 assert_eq!(*item.value, 100);
573 }
574
575 #[test]
576 fn test_pool_basic_tests_with_try() {
577 let pool_result = ObjectPool::<TestItem>::try_new(-1, 2, Some(100));
579 assert!(
580 pool_result.is_err(),
581 "Pool creation should fail with negative args"
582 );
583
584 let pool = ObjectPool::<TestItem>::try_new(42, 2, None).unwrap();
585 assert_eq!(pool.len(), 2);
586 let item1 = pool.try_get_ref(100).unwrap();
587 assert_eq!(*item1.value, 100);
588 assert_eq!(pool.len(), 1);
589 let item2 = pool.try_get_ref(200).unwrap();
590 assert_eq!(*item2.value, 200);
591 assert_eq!(pool.len(), 0);
592 let item = pool.try_get_ref(300).unwrap();
593 assert_eq!(*item.value, 300);
594 assert_eq!(pool.len(), 0);
595 {
596 let item = pool.try_get_ref(400).unwrap();
597 assert_eq!(*item.value, 400);
598 assert_eq!(pool.len(), 0);
599 }
600 assert_eq!(pool.len(), 1); {
602 let item_a = pool.try_get_ref(500).unwrap();
603 assert_eq!(*item_a.value, 500);
604 assert_eq!(pool.len(), 0); let item_b = pool.try_get_ref(600).unwrap();
606 assert_eq!(*item_b.value, 600);
607 assert_eq!(pool.len(), 0); }
609 assert_eq!(pool.len(), 2); let pool = ObjectPool::<TestItem>::try_new(42, 1, Some(100)).unwrap();
612 let item = pool.try_get_ref(100).unwrap();
613 assert_eq!(*item.value, 100);
614 }
615
616 #[test]
617 fn test_pool_with_arc() {
618 let pool = &Arc::new(ObjectPool::<TestItem>::new(42, 1, None));
619 let item = pool.get(100);
620 assert_eq!(*item.value, 100);
621 assert_eq!(pool.len(), 0);
622
623 let item = pool.get(200);
624 assert_eq!(*item.value, 200);
625 assert_eq!(pool.len(), 0);
626 {
627 let item = pool.get(400);
628 assert_eq!(*item.value, 400);
629 assert_eq!(pool.len(), 0);
630 }
631 assert_eq!(pool.len(), 1); let item = pool.try_get_ref(400).unwrap();
633 assert_eq!(*item.value, 400);
634 assert_eq!(pool.len(), 0); let item = pool.try_get(500).unwrap();
636 assert_eq!(*item.value, 500);
637 }
638
639 #[test]
640 fn test_pool_max_capacity_ref() {
641 let pool = ObjectPool::<TestItem>::new(42, 1, Some(1));
642 assert_eq!(pool.len(), 1);
643 assert!(!pool.is_empty());
644 assert_eq!(pool.len(), pool.capacity.unwrap()); {
646 let item = pool.get_ref(100);
647 assert_eq!(pool.len(), 0); assert!(pool.is_empty());
649 assert!(pool.len() < pool.capacity.unwrap()); assert_eq!(*item.value, 100);
651 }
652 assert_eq!(pool.len(), 1); assert_eq!(pool.len(), pool.capacity.unwrap()); {
655 let item1 = pool.get_ref(100);
656 assert_eq!(pool.len(), 0); let item2 = pool.get_ref(200);
658 assert_eq!(pool.len(), 0); let item3 = pool.get_ref(300);
660 assert_eq!(pool.len(), 0); assert!(*item1.value == 100 && *item2.value == 200 && *item3.value == 300);
662 }
664 assert_eq!(pool.len(), pool.capacity.unwrap()); assert_eq!(pool.len(), 1); }
667
668 #[test]
669 fn test_pool_max_capacity_pooled_item() {
670 let pool = &Arc::new(ObjectPool::<TestItem>::new(42, 1, Some(1)));
671 assert_eq!(pool.len(), 1);
672 assert_eq!(pool.len(), pool.capacity.unwrap()); {
674 let item = pool.get(100);
675 assert_eq!(pool.len(), 0); assert!(pool.len() < pool.capacity.unwrap()); assert_eq!(*item.value, 100);
678 }
679 assert_eq!(pool.len(), 1); assert_eq!(pool.len(), pool.capacity.unwrap()); {
682 let item1 = pool.get(100);
683 assert_eq!(pool.len(), 0); let item2 = pool.get(200);
685 assert_eq!(pool.len(), 0); let item3 = pool.get(300);
687 assert_eq!(pool.len(), 0); assert!(*item1.value == 100 && *item2.value == 200 && *item3.value == 300);
689 }
691 assert_eq!(pool.len(), pool.capacity.unwrap()); assert_eq!(pool.len(), 1); }
694
695 #[test]
696 fn test_pool_options() {
697 let item = PoolOption::non_pooled(TestItem::new(42));
699 assert_eq!(*item.value, 42);
700 assert!(item.is_non_pooled());
701 assert!(!item.is_pooled());
702
703 let item = PoolOption::<TestItem>::non_pooled_create(100);
704 assert_eq!(*item.value, 100);
705 assert!(item.is_non_pooled());
706
707 let pool = Arc::new(ObjectPool::<TestItem>::new(42, 1, None));
709 let item = PoolOption::pooled(&pool, 100);
710 assert_eq!(*item.value, 100);
711 assert!(item.is_pooled());
712 assert!(!item.is_non_pooled());
713
714 let item = PoolOption::try_pooled(&pool, 100).unwrap();
716 assert_eq!(*item.value, 100);
717 assert!(item.is_pooled());
718 let item = PoolOption::<TestItem>::try_non_pooled_create(200).unwrap();
719 assert_eq!(*item.value, 200);
720 assert!(item.is_non_pooled());
721 assert!(!item.is_pooled());
722 let item_result = PoolOption::<TestItem>::try_non_pooled_create(-200);
723 assert!(
724 item_result.is_err(),
725 "Creating non-pooled item with negative args should fail"
726 );
727 }
728
729 #[test]
730 fn test_pool_ref_deref_mut() {
731 let pool = ObjectPool::<TestItem>::new(42, 1, None);
733 let item = pool.get_ref(100);
734
735 assert_eq!(*item.value, 100);
736 let mut item = pool.get_ref(100);
737 *item.value = 200;
738 assert_eq!(*item.value, 200);
739
740 let item_ref: &TestItem = &item;
741 assert_eq!(*item_ref.value, 200);
742
743 let item_ref_mut: &mut TestItem = &mut item;
744 assert_eq!(*item_ref_mut.value, 200);
745
746 *item_ref_mut.value = 300;
747 assert_eq!(*item_ref_mut.value, 300);
748
749 let pool = &Arc::new(ObjectPool::<TestItem>::new(42, 1, None));
751 let mut item = pool.get_ref(100);
752 assert_eq!(*item.value, 100);
753
754 *item.value = 200;
755 assert_eq!(*item.value, 200);
756
757 let pool = Arc::new(ObjectPool::<TestItem>::new(42, 1, None));
759 let mut item = PoolOption::pooled(&pool, 100);
760 assert_eq!(*item.value, 100);
761
762 *item.value = 200;
763 assert_eq!(*item.value, 200);
764
765 let mut item = PoolOption::non_pooled(TestItem::new(42));
766 assert_eq!(*item.value, 42);
767
768 *item.value = 100;
769 assert_eq!(*item.value, 100);
770 }
771
772 fn check_error(err: &dyn std::any::Any, contains: &str) {
791 match err.downcast_ref::<&'static str>() {
792 Some(msg) => assert_contains!(msg, contains),
793 None => panic!("incorrect downcast type"),
794 }
795 }
796
797 #[test]
798 fn test_panic_during_create() {
799 let pool = ObjectPool::<TestItem>::new(0u32, 0, Some(1));
800
801 let err = std::panic::catch_unwind(|| {
803 let _ = pool.get_ref(TestPanic);
804 })
805 .unwrap_err();
806
807 check_error(&*err, "panicking on create");
808
809 assert!(
810 !pool.queue.is_poisoned(),
811 "lock should be released while calling trait implementations"
812 );
813
814 assert_eq!(pool.len(), 0);
815 }
816
817 #[test]
818 fn test_panic_during_try_create() {
819 let pool = ObjectPool::<TestItem>::new(0u32, 0, Some(1));
820
821 let err = std::panic::catch_unwind(|| {
823 let _ = pool.try_get_ref(TestPanic);
824 })
825 .unwrap_err();
826
827 check_error(&*err, "panicking on try_create");
828
829 assert!(
830 !pool.queue.is_poisoned(),
831 "lock should be released while calling trait implementations"
832 );
833
834 assert_eq!(pool.len(), 0);
835 }
836
837 #[test]
838 fn test_panic_during_modify() {
839 let pool = ObjectPool::<TestItem>::new(0u32, 0, Some(1));
840
841 let _ = pool.get_ref(0u32);
843 assert_eq!(pool.len(), 1);
844
845 let err = std::panic::catch_unwind(|| {
846 let _ = pool.get_ref(TestPanic);
847 })
848 .unwrap_err();
849
850 check_error(&*err, "panicking on modify");
851
852 assert!(
853 !pool.queue.is_poisoned(),
854 "lock should be released while calling trait implementations"
855 );
856
857 assert_eq!(
858 pool.len(),
859 0,
860 "we should not return a potentially torn object to the pool"
861 );
862 }
863
864 #[test]
865 fn test_panic_during_try_modify() {
866 let pool = ObjectPool::<TestItem>::new(0u32, 0, Some(1));
867
868 let _ = pool.get_ref(0u32);
870 assert_eq!(pool.len(), 1);
871
872 let err = std::panic::catch_unwind(|| {
873 let _ = pool.try_get_ref(TestPanic);
874 })
875 .unwrap_err();
876
877 check_error(&*err, "panicking on try_modify");
878
879 assert!(
880 !pool.queue.is_poisoned(),
881 "lock should be released while calling trait implementations"
882 );
883
884 assert_eq!(
885 pool.len(),
886 0,
887 "we should not return a potentially torn object to the pool"
888 );
889 }
890
891 #[test]
893 fn test_panic_during_drop_ref() {
894 let pool = ObjectPool::<TestItem>::new(0u32, 0, Some(1));
895
896 let mut a = pool.get_ref(0u32);
899 let _ = pool.get_ref(1u32);
900 assert_eq!(pool.len(), 1);
901
902 a.panic_on_drop = true;
903 let err = std::panic::catch_unwind(move || std::mem::drop(a)).unwrap_err();
904 check_error(&*err, "panicking on drop");
905
906 assert!(
907 !pool.queue.is_poisoned(),
908 "lock should be released while calling object drop"
909 );
910 assert_eq!(pool.len(), 1);
911 }
912
913 #[test]
915 fn test_panic_during_drop_arc() {
916 let pool = Arc::new(ObjectPool::<TestItem>::new(0u32, 0, Some(1)));
917
918 let mut a = pool.get(0u32);
921 let _ = pool.get(1u32);
922 assert_eq!(pool.len(), 1);
923
924 a.panic_on_drop = true;
925 let err = std::panic::catch_unwind(move || std::mem::drop(a)).unwrap_err();
926 check_error(&*err, "panicking on drop");
927
928 assert!(
929 !pool.queue.is_poisoned(),
930 "lock should be released while calling object drop"
931 );
932 assert_eq!(pool.len(), 1);
933 }
934
935 #[test]
937 fn test_panic_recovery() {
938 let pool = ObjectPool::<TestItem>::new(0u32, 1, Some(1));
939
940 let err = std::panic::catch_unwind(|| {
941 let _guard = pool.queue.lock();
942 panic!("yeet");
943 })
944 .unwrap_err();
945
946 check_error(&*err, "yeet");
947
948 assert!(pool.queue.is_poisoned());
949
950 let _ = pool.get_ref(1u32);
951 assert!(!pool.queue.is_poisoned(), "poison should be cleared");
952 }
953
954 #[test]
959 fn test_undef() {
960 let mut x: Vec<f32> = Vec::<f32>::create(Undef::new(10));
961 assert_eq!(x.len(), 10);
962
963 x.modify(Undef::new(0));
964 assert_eq!(x.len(), 0);
965
966 x.modify(Undef::new(20));
967 assert_eq!(x.len(), 20);
968 }
969}