use std::{
cell::UnsafeCell,
ops::{Deref, DerefMut},
sync::{Arc, Weak, atomic::AtomicPtr},
thread::{self, LocalKey, ThreadId},
};
use stable_deref_trait::StableDeref;
use super::{PoolError, PoolGuard, PoolItem, PoolProvider, core::Storage};
mod remote_returns;
use remote_returns::{DetachedReturns, RemoteReturns};
struct Entry<T: PoolItem, const N: usize> {
value: T,
metadata: ThreadLocalMetadata<T, N>,
}
impl<T: PoolItem, const N: usize> Entry<T, N> {
const fn new(value: T, metadata: ThreadLocalMetadata<T, N>) -> Self {
Self { value, metadata }
}
}
struct ThreadLocalMetadata<T: PoolItem, const N: usize> {
pool_id: ThreadId,
return_queue: Weak<RemoteReturns<T, N>>,
remote_next: AtomicPtr<Entry<T, N>>,
}
#[doc(hidden)]
pub struct FixedThreadLocalEntry<T: PoolItem, const N: usize> {
node: Box<Entry<T, N>>,
}
impl<T: PoolItem, const N: usize> Deref for FixedThreadLocalEntry<T, N> {
type Target = T;
#[inline(always)]
fn deref(&self) -> &T {
&self.node.value
}
}
impl<T: PoolItem, const N: usize> DerefMut for FixedThreadLocalEntry<T, N> {
#[inline(always)]
fn deref_mut(&mut self) -> &mut T {
&mut self.node.value
}
}
unsafe impl<T: PoolItem, const N: usize> StableDeref for FixedThreadLocalEntry<T, N> {}
impl<T: PoolItem, const N: usize> ThreadLocalMetadata<T, N> {
#[inline(always)]
const fn new(pool_id: ThreadId, return_queue: Weak<RemoteReturns<T, N>>) -> Self {
Self {
pool_id,
return_queue,
remote_next: AtomicPtr::new(std::ptr::null_mut()),
}
}
}
struct ThreadLocalStorage<T: PoolItem, const N: usize> {
storage: Storage<Box<Entry<T, N>>, N>,
remote: Arc<RemoteReturns<T, N>>,
}
impl<T: PoolItem, const N: usize> ThreadLocalStorage<T, N> {
fn new() -> Self {
Self {
storage: Storage::new(),
remote: Arc::new(RemoteReturns::new()),
}
}
}
pub struct FixedThreadLocalPool<T: PoolItem, const N: usize>(UnsafeCell<ThreadLocalStorage<T, N>>);
impl<T: PoolItem, const N: usize> FixedThreadLocalPool<T, N> {
#[must_use]
pub fn new() -> Self {
Self(UnsafeCell::new(ThreadLocalStorage::new()))
}
#[inline(always)]
fn try_take_stored(&self) -> Option<Box<Entry<T, N>>> {
unsafe { (&mut *self.0.get()).storage.pop() }
}
#[inline(always)]
fn try_store(&self, entry: Box<Entry<T, N>>) -> Result<(), Box<Entry<T, N>>> {
unsafe { (&mut *self.0.get()).storage.try_push(entry) }
}
#[inline(always)]
fn stored_len(&self) -> usize {
unsafe { (&*self.0.get()).storage.len() }
}
#[inline(always)]
fn take_remote_returns(&self) -> DetachedReturns<T, N> {
unsafe { (&*self.0.get()).remote.take_all() }
}
#[inline(always)]
fn metadata(&self) -> ThreadLocalMetadata<T, N> {
let return_queue = unsafe { Arc::downgrade(&(&*self.0.get()).remote) };
ThreadLocalMetadata::new(thread::current().id(), return_queue)
}
#[cold]
#[inline(never)]
fn refill_from_remote(&self) {
let mut returned = self.take_remote_returns();
{
let storage = unsafe { &mut *self.0.get() };
storage.storage.extend_newest_first(&mut returned);
}
drop(returned);
}
#[inline(always)]
fn try_take(&self) -> Option<Box<Entry<T, N>>> {
if let Some(entry) = self.try_take_stored() {
return Some(entry);
}
self.refill_from_remote();
self.try_take_stored()
}
}
impl<T: PoolItem, const N: usize> Default for FixedThreadLocalPool<T, N> {
fn default() -> Self {
Self::new()
}
}
pub type FixedThreadLocalPoolGuard<T, const N: usize> =
PoolGuard<T, &'static LocalKey<FixedThreadLocalPool<T, N>>>;
impl<T: PoolItem, const N: usize> PoolProvider<T>
for &'static LocalKey<FixedThreadLocalPool<T, N>>
{
type Entry = FixedThreadLocalEntry<T, N>;
#[inline(always)]
fn take<F, E>(&self, create: F) -> Result<Self::Entry, PoolError<E>>
where
F: FnOnce() -> Result<T, E>,
{
PoolError::catch(|| {
if let Ok(Some(entry)) = self.try_with(FixedThreadLocalPool::try_take) {
return Ok(FixedThreadLocalEntry { node: entry });
}
let value = create()?;
let metadata = self
.try_with(FixedThreadLocalPool::metadata)
.unwrap_or_else(|_| ThreadLocalMetadata::new(thread::current().id(), Weak::new()));
Ok(FixedThreadLocalEntry {
node: Box::new(Entry::new(value, metadata)),
})
})
}
#[inline(always)]
fn return_entry(&self, entry: Self::Entry) -> Result<(), Self::Entry> {
try_return_to_origin(self, entry.node).map_err(|node| FixedThreadLocalEntry { node })
}
fn warm<F, E>(&self, count: usize, mut create: F) -> Result<usize, PoolError<E>>
where
F: FnMut() -> Result<T, E>,
{
PoolError::catch(|| {
let target = count.min(N);
let Ok(missing) = self.try_with(|pool| {
pool.refill_from_remote();
target.saturating_sub(pool.stored_len())
}) else {
return Ok(0);
};
let mut inserted = 0;
for _ in 0..missing {
let value = create()?;
let Ok(metadata) = self.try_with(FixedThreadLocalPool::metadata) else {
drop(value);
break;
};
let entry = Box::new(Entry::new(value, metadata));
match try_return_to_origin(self, entry) {
Ok(()) => inserted += 1,
Err(entry) => {
drop(entry);
break;
}
}
}
Ok(inserted)
})
}
}
#[inline(always)]
fn try_return_to_origin<T: PoolItem, const N: usize>(
pool: &'static LocalKey<FixedThreadLocalPool<T, N>>,
entry: Box<Entry<T, N>>,
) -> Result<(), Box<Entry<T, N>>> {
if entry.metadata.pool_id == thread::current().id() {
let mut entry = Some(entry);
let _ = pool.try_with(|pool| {
if let Some(returned) = entry.take() {
entry = pool.try_store(returned).err();
}
});
return entry.map_or(Ok(()), Err);
}
return_remote(entry)
}
#[cold]
#[inline(never)]
fn return_remote<T: PoolItem, const N: usize>(
entry: Box<Entry<T, N>>,
) -> Result<(), Box<Entry<T, N>>> {
let Some(queue) = entry.metadata.return_queue.upgrade() else {
return Err(entry);
};
queue.push(entry);
Ok(())
}
#[cfg(test)]
mod tests;