use bytes::BytesMut;
use std::cell::RefCell;
use std::ops::{Deref, DerefMut};
use std::sync::atomic::{AtomicU64, Ordering};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub enum BufferSize {
Tiny,
Small,
Medium,
Large,
Huge,
Custom(usize),
}
impl BufferSize {
#[inline]
pub const fn capacity(self) -> usize {
match self {
Self::Tiny => 256,
Self::Small => 4096,
Self::Medium => 16384,
Self::Large => 65536,
Self::Huge => 262144,
Self::Custom(n) => n.next_power_of_two(),
}
}
#[inline]
pub const fn for_bytes(bytes: usize) -> Self {
if bytes <= 256 {
Self::Tiny
} else if bytes <= 4096 {
Self::Small
} else if bytes <= 16384 {
Self::Medium
} else if bytes <= 65536 {
Self::Large
} else if bytes <= 262144 {
Self::Huge
} else {
Self::Custom(bytes)
}
}
#[inline]
const fn pool_index(self) -> usize {
match self {
Self::Tiny => 0,
Self::Small => 1,
Self::Medium => 2,
Self::Large => 3,
Self::Huge => 4,
Self::Custom(_) => 5, }
}
}
#[derive(Debug, Clone)]
pub struct PoolConfig {
pub max_per_size: usize,
pub preallocate: usize,
pub max_age: u32,
pub collect_stats: bool,
}
impl Default for PoolConfig {
fn default() -> Self {
Self {
max_per_size: 64,
preallocate: 4,
max_age: 10000,
collect_stats: true,
}
}
}
impl PoolConfig {
pub fn high_performance() -> Self {
Self {
max_per_size: 128,
preallocate: 16,
max_age: 50000,
collect_stats: false,
}
}
pub fn memory_efficient() -> Self {
Self {
max_per_size: 16,
preallocate: 2,
max_age: 1000,
collect_stats: true,
}
}
}
#[derive(Debug, Default)]
pub struct PoolStats {
hits: AtomicU64,
misses: AtomicU64,
returns: AtomicU64,
discards: AtomicU64,
bytes_allocated: AtomicU64,
pooled_count: AtomicU64,
}
impl PoolStats {
pub fn new() -> Self {
Self::default()
}
#[inline]
pub fn record_hit(&self) {
self.hits.fetch_add(1, Ordering::Relaxed);
}
#[inline]
pub fn record_miss(&self, bytes: usize) {
self.misses.fetch_add(1, Ordering::Relaxed);
self.bytes_allocated
.fetch_add(bytes as u64, Ordering::Relaxed);
}
#[inline]
pub fn record_return(&self) {
self.returns.fetch_add(1, Ordering::Relaxed);
self.pooled_count.fetch_add(1, Ordering::Relaxed);
}
#[inline]
pub fn record_discard(&self) {
self.discards.fetch_add(1, Ordering::Relaxed);
}
#[inline]
pub fn record_taken(&self) {
self.pooled_count.fetch_sub(1, Ordering::Relaxed);
}
pub fn hits(&self) -> u64 {
self.hits.load(Ordering::Relaxed)
}
pub fn misses(&self) -> u64 {
self.misses.load(Ordering::Relaxed)
}
pub fn returns(&self) -> u64 {
self.returns.load(Ordering::Relaxed)
}
pub fn discards(&self) -> u64 {
self.discards.load(Ordering::Relaxed)
}
pub fn bytes_allocated(&self) -> u64 {
self.bytes_allocated.load(Ordering::Relaxed)
}
pub fn pooled_count(&self) -> u64 {
self.pooled_count.load(Ordering::Relaxed)
}
pub fn hit_rate(&self) -> f64 {
let hits = self.hits() as f64;
let total = hits + self.misses() as f64;
if total > 0.0 {
(hits / total) * 100.0
} else {
0.0
}
}
}
static POOL_STATS: PoolStats = PoolStats {
hits: AtomicU64::new(0),
misses: AtomicU64::new(0),
returns: AtomicU64::new(0),
discards: AtomicU64::new(0),
bytes_allocated: AtomicU64::new(0),
pooled_count: AtomicU64::new(0),
};
pub fn pool_stats() -> &'static PoolStats {
&POOL_STATS
}
struct SizePool {
buffers: Vec<(BytesMut, u32)>,
max_size: usize,
capacity: usize,
}
impl SizePool {
fn new(capacity: usize, max_size: usize, preallocate: usize) -> Self {
let preallocate = preallocate.min(max_size);
let mut buffers = Vec::with_capacity(max_size);
for _ in 0..preallocate {
buffers.push((BytesMut::with_capacity(capacity), 0));
}
Self {
buffers,
max_size,
capacity,
}
}
#[inline]
fn acquire(&mut self, collect_stats: bool) -> Option<(BytesMut, u32)> {
self.buffers.pop().map(|(mut buf, age)| {
buf.clear();
if collect_stats {
POOL_STATS.record_hit();
POOL_STATS.record_taken();
}
(buf, age)
})
}
#[inline]
fn release(&mut self, mut buf: BytesMut, age: u32, max_age: u32, collect_stats: bool) {
let age = age.saturating_add(1);
if age < max_age
&& self.buffers.len() < self.max_size
&& buf.capacity() <= self.capacity * 2
{
buf.clear();
self.buffers.push((buf, age));
if collect_stats {
POOL_STATS.record_return();
}
} else if collect_stats {
POOL_STATS.record_discard();
}
}
#[allow(dead_code)] fn allocate(&self) -> BytesMut {
POOL_STATS.record_miss(self.capacity);
BytesMut::with_capacity(self.capacity)
}
}
struct ThreadLocalPool {
pools: [SizePool; 6],
config: PoolConfig,
}
impl ThreadLocalPool {
fn new(config: PoolConfig) -> Self {
Self {
pools: [
SizePool::new(
BufferSize::Tiny.capacity(),
config.max_per_size,
config.preallocate,
),
SizePool::new(
BufferSize::Small.capacity(),
config.max_per_size,
config.preallocate,
),
SizePool::new(
BufferSize::Medium.capacity(),
config.max_per_size,
config.preallocate,
),
SizePool::new(
BufferSize::Large.capacity(),
config.max_per_size,
config.preallocate,
),
SizePool::new(
BufferSize::Huge.capacity(),
config.max_per_size,
config.preallocate,
),
SizePool::new(
BufferSize::Custom(1024 * 1024).capacity(),
config.max_per_size / 4,
0, ),
],
config,
}
}
#[inline]
fn acquire(&mut self, size: BufferSize) -> (BytesMut, u32) {
let idx = size.pool_index();
let pool = &mut self.pools[idx];
let collect_stats = self.config.collect_stats;
pool.acquire(collect_stats).unwrap_or_else(|| {
let capacity = if matches!(size, BufferSize::Custom(n) if n > 0) {
size.capacity()
} else {
pool.capacity
};
if collect_stats {
POOL_STATS.record_miss(capacity);
}
(BytesMut::with_capacity(capacity), 0)
})
}
#[inline]
fn release(&mut self, buf: BytesMut, age: u32, size: BufferSize) {
let idx = size.pool_index();
self.pools[idx].release(buf, age, self.config.max_age, self.config.collect_stats);
}
}
thread_local! {
static BUFFER_POOL: RefCell<ThreadLocalPool> = RefCell::new(
ThreadLocalPool::new(PoolConfig::default())
);
}
pub struct PooledBuffer {
inner: Option<BytesMut>,
size: BufferSize,
age: u32,
}
impl PooledBuffer {
fn new(buf: BytesMut, size: BufferSize, age: u32) -> Self {
Self {
inner: Some(buf),
size,
age,
}
}
#[inline]
pub fn take(mut self) -> BytesMut {
self.inner.take().expect("buffer already taken")
}
#[inline]
pub fn freeze(mut self) -> bytes::Bytes {
self.inner.take().expect("buffer already taken").freeze()
}
#[inline]
pub fn size_category(&self) -> BufferSize {
self.size
}
#[inline]
pub fn capacity(&self) -> usize {
self.inner.as_ref().map(|b| b.capacity()).unwrap_or(0)
}
}
impl Deref for PooledBuffer {
type Target = BytesMut;
#[inline]
fn deref(&self) -> &Self::Target {
self.inner.as_ref().expect("buffer already taken")
}
}
impl DerefMut for PooledBuffer {
#[inline]
fn deref_mut(&mut self) -> &mut Self::Target {
self.inner.as_mut().expect("buffer already taken")
}
}
impl Drop for PooledBuffer {
fn drop(&mut self) {
if let Some(buf) = self.inner.take() {
BUFFER_POOL.with(|pool| {
pool.borrow_mut().release(buf, self.age, self.size);
});
}
}
}
impl std::fmt::Debug for PooledBuffer {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("PooledBuffer")
.field("size", &self.size)
.field("len", &self.inner.as_ref().map(|b| b.len()))
.field("capacity", &self.inner.as_ref().map(|b| b.capacity()))
.finish()
}
}
#[inline]
pub fn acquire_buffer(size: BufferSize) -> PooledBuffer {
BUFFER_POOL.with(|pool| {
let (buf, age) = pool.borrow_mut().acquire(size);
PooledBuffer::new(buf, size, age)
})
}
#[inline]
pub fn acquire_buffer_for_bytes(bytes: usize) -> PooledBuffer {
acquire_buffer(BufferSize::for_bytes(bytes))
}
#[inline]
pub fn acquire_with_data(data: &[u8]) -> PooledBuffer {
let mut buf = acquire_buffer(BufferSize::for_bytes(data.len()));
buf.extend_from_slice(data);
buf
}
#[inline]
pub fn acquire_json_buffer() -> PooledBuffer {
acquire_buffer(BufferSize::Small)
}
#[inline]
pub fn acquire_body_buffer() -> PooledBuffer {
acquire_buffer(BufferSize::Medium)
}
#[inline]
pub fn acquire_response_buffer() -> PooledBuffer {
acquire_buffer(BufferSize::Small)
}
#[inline]
pub fn acquire_streaming_buffer() -> PooledBuffer {
acquire_buffer(BufferSize::Large)
}
#[inline]
pub fn with_buffer<F, R>(size: BufferSize, f: F) -> R
where
F: FnOnce(&mut BytesMut) -> R,
{
let mut buf = acquire_buffer(size);
f(&mut buf)
}
#[inline]
pub fn buffer_to_bytes<F, E>(size: BufferSize, f: F) -> Result<bytes::Bytes, E>
where
F: FnOnce(&mut BytesMut) -> Result<(), E>,
{
let mut buf = acquire_buffer(size);
f(&mut buf)?;
Ok(buf.freeze())
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_buffer_sizes() {
assert_eq!(BufferSize::Tiny.capacity(), 256);
assert_eq!(BufferSize::Small.capacity(), 4096);
assert_eq!(BufferSize::Medium.capacity(), 16384);
assert_eq!(BufferSize::Large.capacity(), 65536);
assert_eq!(BufferSize::Huge.capacity(), 262144);
assert_eq!(BufferSize::Custom(1000).capacity(), 1024); }
#[test]
fn test_size_for_bytes() {
assert_eq!(BufferSize::for_bytes(100), BufferSize::Tiny);
assert_eq!(BufferSize::for_bytes(1000), BufferSize::Small);
assert_eq!(BufferSize::for_bytes(10000), BufferSize::Medium);
assert_eq!(BufferSize::for_bytes(50000), BufferSize::Large);
assert_eq!(BufferSize::for_bytes(200000), BufferSize::Huge);
}
#[test]
fn test_acquire_and_release() {
let mut buf = acquire_buffer(BufferSize::Small);
assert!(buf.capacity() >= BufferSize::Small.capacity());
buf.extend_from_slice(b"Hello, World!");
assert_eq!(&buf[..], b"Hello, World!");
drop(buf);
let buf2 = acquire_buffer(BufferSize::Small);
assert!(buf2.is_empty()); }
#[test]
fn test_acquire_with_data() {
let buf = acquire_with_data(b"test data");
assert_eq!(&buf[..], b"test data");
}
#[test]
fn test_pooled_buffer_take() {
let buf = acquire_buffer(BufferSize::Tiny);
let inner = buf.take();
assert_eq!(inner.capacity(), BufferSize::Tiny.capacity());
}
#[test]
fn test_pooled_buffer_freeze() {
let mut buf = acquire_buffer(BufferSize::Tiny);
buf.extend_from_slice(b"frozen");
let bytes = buf.freeze();
assert_eq!(&bytes[..], b"frozen");
}
#[test]
fn test_with_buffer() {
let len = with_buffer(BufferSize::Small, |buf| {
buf.extend_from_slice(b"test");
buf.len()
});
assert_eq!(len, 4);
}
#[test]
fn test_buffer_to_bytes() {
let bytes: Result<bytes::Bytes, std::convert::Infallible> =
buffer_to_bytes(BufferSize::Tiny, |buf| {
buf.extend_from_slice(b"bytes");
Ok(())
});
assert_eq!(&bytes.unwrap()[..], b"bytes");
}
#[test]
fn test_pool_stats() {
let stats = pool_stats();
let initial_hits = stats.hits();
let _buf = acquire_buffer(BufferSize::Tiny);
drop(_buf);
let _buf = acquire_buffer(BufferSize::Tiny);
assert!(stats.hits() + stats.misses() > initial_hits);
}
#[test]
fn test_specialized_acquires() {
let json_buf = acquire_json_buffer();
assert!(json_buf.capacity() >= BufferSize::Small.capacity());
let body_buf = acquire_body_buffer();
assert!(body_buf.capacity() >= BufferSize::Medium.capacity());
let response_buf = acquire_response_buffer();
assert!(response_buf.capacity() >= BufferSize::Small.capacity());
let streaming_buf = acquire_streaming_buffer();
assert!(streaming_buf.capacity() >= BufferSize::Large.capacity());
}
#[test]
fn test_pool_config() {
let config = PoolConfig::high_performance();
assert_eq!(config.max_per_size, 128);
assert!(!config.collect_stats);
let config = PoolConfig::memory_efficient();
assert_eq!(config.max_per_size, 16);
assert!(config.collect_stats);
}
}