use std::sync::Arc;
use std::sync::atomic::{AtomicU64, Ordering};
use crate::metrics::{Counters, PoolStats, snapshot};
use crate::size_class::{ClassTable, SizeClass};
use crate::tls::TlsState;
static NEXT_POOL_ID: AtomicU64 = AtomicU64::new(1);
#[derive(Debug)]
pub(crate) struct PoolState {
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 counters: Counters,
}
#[derive(Clone, Debug)]
pub struct BufferPool {
pub(crate) state: Arc<PoolState>,
}
impl BufferPool {
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: Arc::new(PoolState {
id: NEXT_POOL_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),
counters: Counters::new(),
}),
}
}
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)
}
fn rebuild(mut self, f: impl FnOnce(&mut PoolState)) -> Self {
let state = Arc::get_mut(&mut self.state)
.expect("cannot reconfigure a shared pool — call config methods before cloning");
f(state);
self
}
#[inline]
#[must_use]
pub fn get(&self, size: usize) -> crate::PooledBuffer {
self.state.counters.gets.fetch_add(1, Ordering::Relaxed);
let Some((class_idx, class)) = self.state.table.route(size) else {
self.state.counters.oversize.fetch_add(1, Ordering::Relaxed);
self.state.counters.allocations.fetch_add(1, Ordering::Relaxed);
return crate::PooledBuffer::new(ClassTable::oversize(size), self.clone(), u8::MAX);
};
let tls_result = TlsState::with(|state| {
if !state.owns(self.state.id) {
state.bind(self.state.id, self.state.tls_cache_size);
}
if let Some(buf) = state.caches[class_idx].pop() {
return Some((buf, true)); }
state.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 {
self.state.counters.tls_hits.fetch_add(1, Ordering::Relaxed);
} else {
self.state.counters.shared_hits.fetch_add(1, Ordering::Relaxed);
}
SizeClass::resize(&mut buf, size);
return crate::PooledBuffer::new(buf, self.clone(), ci);
}
self.state.counters.allocations.fetch_add(1, Ordering::Relaxed);
crate::PooledBuffer::new(class.allocate(size), self.clone(), ci)
}
#[inline(always)]
pub(crate) fn put(&self, mut buffer: Vec<u8>, class_hint: u8) {
self.state.counters.puts.fetch_add(1, Ordering::Relaxed);
buffer.clear();
if class_hint == u8::MAX {
self.state.counters.discards.fetch_add(1, Ordering::Relaxed);
return;
}
let cap = buffer.capacity();
if cap < self.state.min_buffer_size {
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(|state| {
if !state.owns(self.state.id) {
state.bind(self.state.id, self.state.tls_cache_size);
}
let class = &self.state.table[class_idx];
if state.caches[class_idx].len() >= state.limit {
state.spill(class_idx, class, self.state.batch_size);
}
if state.caches[class_idx].len() < state.limit {
state.caches[class_idx].push(buffer);
return None;
}
Some(buffer)
});
if let Some(buf) = overflow {
let _ = self.state.table[class_idx].push(buf);
}
}
pub fn preallocate(&self, count: usize, size: usize) {
let Some((_, class)) = self.state.table.route(size) else {
return;
};
for _ in 0..count {
let mut buf = Vec::with_capacity(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 clear(&self) {
self.state.table.clear_all();
}
#[inline]
pub fn stats(&self) -> PoolStats {
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;
}
unsafe { buffer.set_len(buffer.capacity()) };
let _ = region::lock(buffer.as_ptr(), buffer.len());
buffer.clear();
}
}
impl Default for BufferPool {
fn default() -> Self {
Self::new()
}
}