use std::sync::atomic::{AtomicU64, Ordering};
use crate::allocator::{Allocator, HeapAllocator};
use crate::size_class::{ClassTable, SizeClass};
use crate::stats::{Counters, Stats, snapshot};
use crate::tls::TlsState;
static NEXT_ID: AtomicU64 = AtomicU64::new(1);
#[derive(Debug)]
pub(crate) struct State {
pub id: u64,
pub table: ClassTable,
pub tls_cache_size: usize,
pub min_buffer_size: usize,
pub pinned_memory: bool,
pub batch_size: usize,
pub track_stats: bool,
pub counters: Counters,
pub allocator: Box<dyn Allocator>,
}
#[derive(Debug)]
pub struct ZeroPool {
pub(crate) state: State,
}
impl ZeroPool {
pub fn new() -> Self {
use crate::config::{
DEFAULT_MIN_BUFFER_SIZE, cpu_count, default_batch_size, default_max_buffers_per_class,
default_tls_cache_size,
};
let cpus = cpu_count();
let tls = default_tls_cache_size(cpus);
Self {
state: State {
id: NEXT_ID.fetch_add(1, Ordering::Relaxed),
table: ClassTable::new(default_max_buffers_per_class(cpus)),
tls_cache_size: tls,
min_buffer_size: DEFAULT_MIN_BUFFER_SIZE,
pinned_memory: false,
batch_size: default_batch_size(tls),
track_stats: false,
counters: Counters::new(),
allocator: Box::new(HeapAllocator),
},
}
}
pub fn allocator(self, alloc: impl Allocator) -> Self {
self.rebuild(|s| s.allocator = Box::new(alloc))
}
pub fn min_buffer_size(self, size: usize) -> Self {
self.rebuild(|s| s.min_buffer_size = size)
}
pub fn tls_cache_size(self, size: usize) -> Self {
assert!(size > 0, "tls_cache_size must be > 0");
self.rebuild(|s| {
s.tls_cache_size = size;
s.batch_size = crate::config::default_batch_size(size);
})
}
pub fn max_buffers_per_class(self, count: usize) -> Self {
assert!(count > 0, "max_buffers_per_class must be > 0");
self.rebuild(|s| s.table = ClassTable::new(count))
}
pub fn pinned_memory(self, enabled: bool) -> Self {
self.rebuild(|s| s.pinned_memory = enabled)
}
pub fn batch_size(self, size: usize) -> Self {
self.rebuild(|s| s.batch_size = size)
}
pub fn track_stats(self, enabled: bool) -> Self {
self.rebuild(|s| s.track_stats = enabled)
}
fn rebuild(mut self, f: impl FnOnce(&mut State)) -> Self {
f(&mut self.state);
self
}
#[inline]
#[must_use]
pub fn alloc(&self, size: usize) -> crate::Buf<'_> {
if self.state.track_stats {
self.state.counters.gets.fetch_add(1, Ordering::Relaxed);
}
let Some((class_idx, class)) = self.state.table.route(size) else {
if self.state.track_stats {
self.state.counters.oversize.fetch_add(1, Ordering::Relaxed);
self.state.counters.allocations.fetch_add(1, Ordering::Relaxed);
}
let mut buf = self.state.allocator.allocate(size);
buf.truncate(size);
self.pin(&mut buf);
return crate::Buf::new(buf, self, u8::MAX);
};
let tls_result = TlsState::with(|tls| {
if !tls.owns(self.state.id) {
tls.bind(self.state.id, self.state.tls_cache_size);
}
if let Some(buf) = tls.caches[class_idx].pop() {
return Some((buf, true));
}
tls.refill(class_idx, class, self.state.batch_size).map(|buf| (buf, false))
});
let ci = class_idx as u8;
if let Some((mut buf, from_tls)) = tls_result {
if from_tls {
if self.state.track_stats {
self.state.counters.tls_hits.fetch_add(1, Ordering::Relaxed);
}
} else if self.state.track_stats {
self.state.counters.shared_hits.fetch_add(1, Ordering::Relaxed);
}
SizeClass::resize_zeroed(&mut buf, size);
return crate::Buf::new(buf, self, ci);
}
if self.state.track_stats {
self.state.counters.allocations.fetch_add(1, Ordering::Relaxed);
}
let mut buf = self.state.allocator.allocate(class.class_size);
buf.truncate(size);
self.pin(&mut buf);
crate::Buf::new(buf, self, ci)
}
#[inline]
#[must_use]
pub fn alloc_uninit(&self, size: usize) -> crate::BufUninit<'_> {
if self.state.track_stats {
self.state.counters.gets.fetch_add(1, Ordering::Relaxed);
}
let Some((class_idx, class)) = self.state.table.route(size) else {
if self.state.track_stats {
self.state.counters.oversize.fetch_add(1, Ordering::Relaxed);
self.state.counters.allocations.fetch_add(1, Ordering::Relaxed);
}
let mut buf = Vec::with_capacity(size);
SizeClass::resize_uninit(&mut buf, size);
self.pin(&mut buf);
return crate::BufUninit::new(buf, self, u8::MAX);
};
let tls_result = TlsState::with(|tls| {
if !tls.owns(self.state.id) {
tls.bind(self.state.id, self.state.tls_cache_size);
}
if let Some(buf) = tls.caches[class_idx].pop() {
return Some((buf, true));
}
tls.refill(class_idx, class, self.state.batch_size).map(|buf| (buf, false))
});
let ci = class_idx as u8;
if let Some((mut buf, from_tls)) = tls_result {
if from_tls {
if self.state.track_stats {
self.state.counters.tls_hits.fetch_add(1, Ordering::Relaxed);
}
} else if self.state.track_stats {
self.state.counters.shared_hits.fetch_add(1, Ordering::Relaxed);
}
SizeClass::resize_uninit(&mut buf, size);
return crate::BufUninit::new(buf, self, ci);
}
if self.state.track_stats {
self.state.counters.allocations.fetch_add(1, Ordering::Relaxed);
}
let mut buf = Vec::with_capacity(class.class_size);
SizeClass::resize_uninit(&mut buf, size);
self.pin(&mut buf);
crate::BufUninit::new(buf, self, ci)
}
#[inline(always)]
pub(crate) fn dealloc(&self, mut buffer: Vec<u8>, class_hint: u8) {
if self.state.track_stats {
self.state.counters.puts.fetch_add(1, Ordering::Relaxed);
}
buffer.clear();
if class_hint == u8::MAX {
if self.state.track_stats {
self.state.counters.discards.fetch_add(1, Ordering::Relaxed);
}
return;
}
let cap = buffer.capacity();
if cap < self.state.min_buffer_size {
if self.state.track_stats {
self.state.counters.discards.fetch_add(1, Ordering::Relaxed);
}
return;
}
let class_idx = if cap >= ClassTable::boundary(class_hint as usize) {
class_hint as usize
} else {
let Some((idx, _)) = self.state.table.route_capacity(cap) else {
return;
};
idx
};
self.pin(&mut buffer);
let overflow = TlsState::with(|tls| {
if !tls.owns(self.state.id) {
tls.bind(self.state.id, self.state.tls_cache_size);
}
let class = &self.state.table[class_idx];
if tls.caches[class_idx].len() >= tls.limit {
tls.spill(class_idx, class, self.state.batch_size);
}
if tls.caches[class_idx].len() < tls.limit {
tls.caches[class_idx].push(buffer);
return None;
}
Some(buffer)
});
if let Some(buf) = overflow {
let _ = self.state.table[class_idx].push(buf);
}
}
pub fn warm(&self, count: usize, size: usize) {
let Some((_, class)) = self.state.table.route(size) else {
return;
};
for _ in 0..count {
let mut buf = self.state.allocator.allocate(class.class_size);
self.pin(&mut buf);
if class.push(buf).is_err() {
break;
}
}
}
#[inline]
#[must_use]
pub fn len(&self) -> usize {
self.state.table.total_buffered()
}
#[inline]
#[must_use]
pub fn is_empty(&self) -> bool {
self.state.table.all_empty()
}
pub fn drain(&self) {
self.state.table.clear_all();
}
#[inline]
pub fn stats(&self) -> Stats {
snapshot(&self.state.counters, self.state.table.classes())
}
pub fn reset_stats(&self) {
self.state.counters.reset();
}
#[inline(always)]
fn pin(&self, buffer: &mut Vec<u8>) {
if !self.state.pinned_memory {
return;
}
if buffer.capacity() == 0 {
return;
}
let _ = region::lock(buffer.as_ptr(), buffer.capacity());
}
}
impl Default for ZeroPool {
fn default() -> Self {
Self::new()
}
}