1use std::{
4 cell::UnsafeCell,
5 ops::{Deref, DerefMut},
6 sync::{Arc, Weak, atomic::AtomicPtr},
7 thread::{self, LocalKey, ThreadId},
8};
9
10use stable_deref_trait::StableDeref;
11
12use super::{PoolError, PoolGuard, PoolItem, PoolProvider, core::Storage};
13
14mod remote_returns;
15
16use remote_returns::{DetachedReturns, RemoteReturns};
17
18struct Entry<T: PoolItem, const N: usize> {
19 value: T,
20 metadata: ThreadLocalMetadata<T, N>,
21}
22
23impl<T: PoolItem, const N: usize> Entry<T, N> {
24 const fn new(value: T, metadata: ThreadLocalMetadata<T, N>) -> Self {
25 Self { value, metadata }
26 }
27}
28
29struct ThreadLocalMetadata<T: PoolItem, const N: usize> {
31 pool_id: ThreadId,
32 return_queue: Weak<RemoteReturns<T, N>>,
33 remote_next: AtomicPtr<Entry<T, N>>,
34}
35
36#[doc(hidden)]
39pub struct FixedThreadLocalEntry<T: PoolItem, const N: usize> {
40 node: Box<Entry<T, N>>,
41}
42
43impl<T: PoolItem, const N: usize> Deref for FixedThreadLocalEntry<T, N> {
44 type Target = T;
45
46 #[inline(always)]
47 fn deref(&self) -> &T {
48 &self.node.value
49 }
50}
51
52impl<T: PoolItem, const N: usize> DerefMut for FixedThreadLocalEntry<T, N> {
53 #[inline(always)]
54 fn deref_mut(&mut self) -> &mut T {
55 &mut self.node.value
56 }
57}
58
59unsafe impl<T: PoolItem, const N: usize> StableDeref for FixedThreadLocalEntry<T, N> {}
61
62impl<T: PoolItem, const N: usize> ThreadLocalMetadata<T, N> {
63 #[inline(always)]
64 const fn new(pool_id: ThreadId, return_queue: Weak<RemoteReturns<T, N>>) -> Self {
65 Self {
66 pool_id,
67 return_queue,
68 remote_next: AtomicPtr::new(std::ptr::null_mut()),
69 }
70 }
71}
72
73struct ThreadLocalStorage<T: PoolItem, const N: usize> {
74 storage: Storage<Box<Entry<T, N>>, N>,
75 remote: Arc<RemoteReturns<T, N>>,
76}
77
78impl<T: PoolItem, const N: usize> ThreadLocalStorage<T, N> {
79 fn new() -> Self {
80 Self {
81 storage: Storage::new(),
82 remote: Arc::new(RemoteReturns::new()),
83 }
84 }
85}
86
87pub struct FixedThreadLocalPool<T: PoolItem, const N: usize>(UnsafeCell<ThreadLocalStorage<T, N>>);
106
107impl<T: PoolItem, const N: usize> FixedThreadLocalPool<T, N> {
108 #[must_use]
110 pub fn new() -> Self {
111 Self(UnsafeCell::new(ThreadLocalStorage::new()))
112 }
113
114 #[inline(always)]
115 fn try_take_stored(&self) -> Option<Box<Entry<T, N>>> {
116 unsafe { (&mut *self.0.get()).storage.pop() }
119 }
120
121 #[inline(always)]
122 fn try_store(&self, entry: Box<Entry<T, N>>) -> Result<(), Box<Entry<T, N>>> {
123 unsafe { (&mut *self.0.get()).storage.try_push(entry) }
125 }
126
127 #[inline(always)]
128 fn stored_len(&self) -> usize {
129 unsafe { (&*self.0.get()).storage.len() }
131 }
132
133 #[inline(always)]
134 fn take_remote_returns(&self) -> DetachedReturns<T, N> {
135 unsafe { (&*self.0.get()).remote.take_all() }
137 }
138
139 #[inline(always)]
140 fn metadata(&self) -> ThreadLocalMetadata<T, N> {
141 let return_queue = unsafe { Arc::downgrade(&(&*self.0.get()).remote) };
144 ThreadLocalMetadata::new(thread::current().id(), return_queue)
145 }
146
147 #[cold]
148 #[inline(never)]
149 fn refill_from_remote(&self) {
150 let mut returned = self.take_remote_returns();
151 {
152 let storage = unsafe { &mut *self.0.get() };
157 storage.storage.extend_newest_first(&mut returned);
158 }
159 drop(returned);
162 }
163
164 #[inline(always)]
165 fn try_take(&self) -> Option<Box<Entry<T, N>>> {
166 if let Some(entry) = self.try_take_stored() {
167 return Some(entry);
168 }
169 self.refill_from_remote();
170 self.try_take_stored()
171 }
172}
173
174impl<T: PoolItem, const N: usize> Default for FixedThreadLocalPool<T, N> {
175 fn default() -> Self {
176 Self::new()
177 }
178}
179
180pub type FixedThreadLocalPoolGuard<T, const N: usize> =
182 PoolGuard<T, &'static LocalKey<FixedThreadLocalPool<T, N>>>;
183
184impl<T: PoolItem, const N: usize> PoolProvider<T>
188 for &'static LocalKey<FixedThreadLocalPool<T, N>>
189{
190 type Entry = FixedThreadLocalEntry<T, N>;
191
192 #[inline(always)]
193 fn take<F, E>(&self, create: F) -> Result<Self::Entry, PoolError<E>>
194 where
195 F: FnOnce() -> Result<T, E>,
196 {
197 PoolError::catch(|| {
198 if let Ok(Some(entry)) = self.try_with(FixedThreadLocalPool::try_take) {
199 return Ok(FixedThreadLocalEntry { node: entry });
200 }
201
202 let value = create()?;
204 let metadata = self
205 .try_with(FixedThreadLocalPool::metadata)
206 .unwrap_or_else(|_| ThreadLocalMetadata::new(thread::current().id(), Weak::new()));
207 Ok(FixedThreadLocalEntry {
208 node: Box::new(Entry::new(value, metadata)),
209 })
210 })
211 }
212
213 #[inline(always)]
214 fn return_entry(&self, entry: Self::Entry) -> Result<(), Self::Entry> {
215 try_return_to_origin(self, entry.node).map_err(|node| FixedThreadLocalEntry { node })
216 }
217
218 fn warm<F, E>(&self, count: usize, mut create: F) -> Result<usize, PoolError<E>>
219 where
220 F: FnMut() -> Result<T, E>,
221 {
222 PoolError::catch(|| {
223 let target = count.min(N);
224 let Ok(missing) = self.try_with(|pool| {
225 pool.refill_from_remote();
226 target.saturating_sub(pool.stored_len())
227 }) else {
228 return Ok(0);
229 };
230 let mut inserted = 0;
231
232 for _ in 0..missing {
233 let value = create()?;
236 let Ok(metadata) = self.try_with(FixedThreadLocalPool::metadata) else {
237 drop(value);
238 break;
239 };
240 let entry = Box::new(Entry::new(value, metadata));
241 match try_return_to_origin(self, entry) {
242 Ok(()) => inserted += 1,
243 Err(entry) => {
244 drop(entry);
245 break;
246 }
247 }
248 }
249 Ok(inserted)
250 })
251 }
252}
253
254#[inline(always)]
255fn try_return_to_origin<T: PoolItem, const N: usize>(
256 pool: &'static LocalKey<FixedThreadLocalPool<T, N>>,
257 entry: Box<Entry<T, N>>,
258) -> Result<(), Box<Entry<T, N>>> {
259 if entry.metadata.pool_id == thread::current().id() {
260 let mut entry = Some(entry);
261 let _ = pool.try_with(|pool| {
262 if let Some(returned) = entry.take() {
263 entry = pool.try_store(returned).err();
264 }
265 });
266 return entry.map_or(Ok(()), Err);
267 }
268 return_remote(entry)
269}
270
271#[cold]
273#[inline(never)]
274fn return_remote<T: PoolItem, const N: usize>(
275 entry: Box<Entry<T, N>>,
276) -> Result<(), Box<Entry<T, N>>> {
277 let Some(queue) = entry.metadata.return_queue.upgrade() else {
279 return Err(entry);
280 };
281 queue.push(entry);
282 Ok(())
283}
284
285#[cfg(test)]
286mod tests;