Skip to main content

rylv_pool/
thread_local.rs

1//! Thread-local provider with lock-free cross-thread returns.
2
3use 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
29/// Origin routing carried by entries managed by [`FixedThreadLocalPool`].
30struct 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/// Opaque entry owned by the thread-local provider.
37/// Its value remains address-stable while the entry moves between threads.
38#[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
59// SAFETY: the value lives in a Box and the entry never replaces the allocation.
60unsafe 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
87/// Per-thread pool with capacity `N` fixed at compile time.
88///
89/// `N` limits locally retained entries, not active guards. Remote returns may
90/// temporarily exceed `N`; older overflow is discarded when the queue drains.
91///
92/// `UnsafeCell` removes dynamic borrow tracking from the hot path. Every
93/// private access is short and never invokes client code or drops an entry.
94/// The storage itself cannot be shared between threads; use its TLS key.
95/// When calling [`PoolProvider`] directly, return entries through the same key
96/// used to acquire them. [`FixedThreadLocalPoolGuard`] retains that key automatically.
97///
98/// ```compile_fail,E0277
99/// use rylv_pool::{FixedThreadLocalPool, PoolItem};
100/// # struct Item;
101/// # impl PoolItem for Item { fn reset(&mut self) {} }
102/// fn require_sync<T: Sync>() {}
103/// require_sync::<FixedThreadLocalPool<Item, 1>>();
104/// ```
105pub struct FixedThreadLocalPool<T: PoolItem, const N: usize>(UnsafeCell<ThreadLocalStorage<T, N>>);
106
107impl<T: PoolItem, const N: usize> FixedThreadLocalPool<T, N> {
108    /// Create empty storage for each thread that initializes this value.
109    #[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        // SAFETY: this TLS value is reachable only on its owning thread. The
117        // exclusive reference does not escape and no client code is invoked.
118        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        // SAFETY: insertion neither invokes client code nor drops the entry.
124        unsafe { (&mut *self.0.get()).storage.try_push(entry) }
125    }
126
127    #[inline(always)]
128    fn stored_len(&self) -> usize {
129        // SAFETY: this read does not invoke client code or expose a reference.
130        unsafe { (&*self.0.get()).storage.len() }
131    }
132
133    #[inline(always)]
134    fn take_remote_returns(&self) -> DetachedReturns<T, N> {
135        // SAFETY: the shared access reaches only the atomic remote queue.
136        unsafe { (&*self.0.get()).remote.take_all() }
137    }
138
139    #[inline(always)]
140    fn metadata(&self) -> ThreadLocalMetadata<T, N> {
141        // SAFETY: cloning a Weak neither mutates storage nor invokes client
142        // code, and the returned metadata owns the Weak handle.
143        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            // SAFETY: storage is confined to this thread and the reference
153            // does not escape this scope. The detached iterator only
154            // reconstructs owned Boxes; filling empty slots and reversing them
155            // neither invokes client code nor drops entries.
156            let storage = unsafe { &mut *self.0.get() };
157            storage.storage.extend_newest_first(&mut returned);
158        }
159        // The batch now contains only older overflow. Drop it after releasing
160        // storage access so destructors can safely reenter the pool.
161        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
180/// Guard using a consumer-defined thread-local pool key as its provider.
181pub type FixedThreadLocalPoolGuard<T, const N: usize> =
182    PoolGuard<T, &'static LocalKey<FixedThreadLocalPool<T, N>>>;
183
184/// If storage is unavailable during thread shutdown, acquisition creates an
185/// unpooled entry and warming does nothing. Returns reject entries so they
186/// are destroyed outside storage access.
187impl<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            // User code and allocation run after TLS storage access has ended.
203            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                // User code and allocation deliberately run without an active TLS
234                // storage access, preserving reentrant acquisition.
235                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/// Publish an entry back to its origin thread without touching current TLS.
272#[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    // TODO:perf Benchmark the Weak::upgrade on remote returns.
278    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;