Skip to main content

get_size2/
tracker.rs

1#[cfg(feature = "alloc")]
2use alloc::boxed::Box;
3#[cfg(all(feature = "alloc", not(feature = "std")))]
4use alloc::collections::BTreeSet;
5#[cfg(feature = "std")]
6use std::collections::HashSet;
7#[cfg(feature = "std")]
8use std::sync::{Arc, Mutex, RwLock};
9
10/// A tracker which makes sure that shared ownership objects are only accounted for once.
11pub trait GetSizeTracker {
12    /// Tracks an arbitrary object located at `addr`.
13    ///
14    /// Returns `true` if the reference, as indexed by the pointed to `addr`, has not yet
15    /// been seen by this tracker. Otherwise it returns `false`.
16    fn track<A>(&mut self, addr: *const A) -> bool;
17}
18
19impl<T: GetSizeTracker> GetSizeTracker for &mut T {
20    fn track<A>(&mut self, addr: *const A) -> bool {
21        GetSizeTracker::track(*self, addr)
22    }
23}
24
25#[cfg(feature = "alloc")]
26impl<T: GetSizeTracker> GetSizeTracker for Box<T> {
27    fn track<A>(&mut self, addr: *const A) -> bool {
28        GetSizeTracker::track(&mut **self, addr)
29    }
30}
31
32#[cfg(feature = "std")]
33impl<T: GetSizeTracker> GetSizeTracker for Mutex<T> {
34    fn track<A>(&mut self, addr: *const A) -> bool {
35        let tracker = self
36            .get_mut()
37            .unwrap_or_else(std::sync::PoisonError::into_inner);
38
39        GetSizeTracker::track(&mut *tracker, addr)
40    }
41}
42
43#[cfg(feature = "std")]
44impl<T: GetSizeTracker> GetSizeTracker for RwLock<T> {
45    fn track<A>(&mut self, addr: *const A) -> bool {
46        let mut tracker = self
47            .write()
48            .unwrap_or_else(std::sync::PoisonError::into_inner);
49
50        GetSizeTracker::track(&mut *tracker, addr)
51    }
52}
53
54#[cfg(feature = "std")]
55impl<T: GetSizeTracker> GetSizeTracker for Arc<Mutex<T>> {
56    fn track<A>(&mut self, addr: *const A) -> bool {
57        let mut tracker = self
58            .lock()
59            .unwrap_or_else(std::sync::PoisonError::into_inner);
60
61        GetSizeTracker::track(&mut *tracker, addr)
62    }
63}
64
65#[cfg(feature = "std")]
66impl<T: GetSizeTracker> GetSizeTracker for Arc<RwLock<T>> {
67    fn track<A>(&mut self, addr: *const A) -> bool {
68        let mut tracker = self
69            .write()
70            .unwrap_or_else(std::sync::PoisonError::into_inner);
71
72        GetSizeTracker::track(&mut *tracker, addr)
73    }
74}
75
76// The set of addresses already seen by a `StandardTracker`. A `HashSet` needs the random state
77// provided by `std`, so `no_std` builds fall back to the `alloc` only `BTreeSet`.
78#[cfg(feature = "std")]
79type SeenAddresses = HashSet<usize>;
80#[cfg(all(feature = "alloc", not(feature = "std")))]
81type SeenAddresses = BTreeSet<usize>;
82
83/// A simple standard tracker which can be used to track shared ownership references.
84#[cfg(feature = "alloc")]
85#[derive(Debug, Default)]
86pub struct StandardTracker {
87    inner: SeenAddresses,
88}
89
90#[cfg(feature = "alloc")]
91impl StandardTracker {
92    #[must_use]
93    pub fn new() -> Self {
94        Self::default()
95    }
96
97    pub fn clear(&mut self) {
98        self.inner.clear();
99    }
100}
101
102#[cfg(feature = "alloc")]
103impl GetSizeTracker for StandardTracker {
104    fn track<A>(&mut self, addr: *const A) -> bool {
105        self.inner.insert(addr.addr())
106    }
107}
108
109/// A pseudo tracker which does not track anything.
110#[derive(Debug, Clone, Copy, Default)]
111pub struct NoTracker {
112    answer: bool,
113}
114
115impl NoTracker {
116    /// Creates a new pseudo tracker, which will always return the given `answer`.
117    #[must_use]
118    pub const fn new(answer: bool) -> Self {
119        Self { answer }
120    }
121
122    /// Get the answer which will always be returned by this pseudo tracker.
123    #[must_use]
124    pub const fn answer(&self) -> bool {
125        self.answer
126    }
127
128    /// Changes the answer which will always be returned by this pseudo tracker.
129    pub const fn set_answer(&mut self, answer: bool) {
130        self.answer = answer;
131    }
132}
133
134impl GetSizeTracker for NoTracker {
135    fn track<A>(&mut self, _addr: *const A) -> bool {
136        self.answer
137    }
138}