use bytes::{Bytes, BytesMut};
use std::collections::VecDeque;
use std::sync::{Arc, Mutex};
#[derive(Debug)]
pub struct BufferPool {
pool: Arc<Mutex<VecDeque<BytesMut>>>,
buffer_size: usize,
max_pool_size: usize,
}
impl BufferPool {
pub fn new(buffer_size: usize, max_pool_size: usize) -> Self {
Self {
pool: Arc::new(Mutex::new(VecDeque::with_capacity(max_pool_size))),
buffer_size,
max_pool_size,
}
}
pub fn get(&self) -> PooledBuffer {
let buffer = {
let mut pool = self.pool.lock().unwrap();
pool.pop_front()
.unwrap_or_else(|| BytesMut::with_capacity(self.buffer_size))
};
PooledBuffer {
buffer: Some(buffer),
pool: Arc::clone(&self.pool),
max_pool_size: self.max_pool_size,
}
}
pub fn stats(&self) -> PoolStats {
let pool = self.pool.lock().unwrap();
PoolStats {
available_buffers: pool.len(),
buffer_size: self.buffer_size,
max_pool_size: self.max_pool_size,
}
}
}
pub struct PooledBuffer {
buffer: Option<BytesMut>,
pool: Arc<Mutex<VecDeque<BytesMut>>>,
max_pool_size: usize,
}
impl PooledBuffer {
pub fn get_mut(&mut self) -> &mut BytesMut {
self.buffer.as_mut().unwrap()
}
pub fn get(&self) -> &BytesMut {
self.buffer.as_ref().unwrap()
}
pub fn freeze(mut self) -> Bytes {
let buffer = self.buffer.take().unwrap();
buffer.freeze()
}
pub fn clear(&mut self) {
if let Some(ref mut buffer) = self.buffer {
buffer.clear();
}
}
pub fn len(&self) -> usize {
self.buffer.as_ref().map_or(0, |b| b.len())
}
pub fn is_empty(&self) -> bool {
self.len() == 0
}
pub fn capacity(&self) -> usize {
self.buffer.as_ref().map_or(0, |b| b.capacity())
}
}
impl Drop for PooledBuffer {
fn drop(&mut self) {
if let Some(mut buffer) = self.buffer.take() {
buffer.clear();
let mut pool = self.pool.lock().unwrap();
if pool.len() < self.max_pool_size {
pool.push_back(buffer);
}
}
}
}
impl std::ops::Deref for PooledBuffer {
type Target = BytesMut;
fn deref(&self) -> &Self::Target {
self.buffer.as_ref().unwrap()
}
}
impl std::ops::DerefMut for PooledBuffer {
fn deref_mut(&mut self) -> &mut Self::Target {
self.buffer.as_mut().unwrap()
}
}
#[derive(Debug, Clone)]
pub struct PoolStats {
pub available_buffers: usize,
pub buffer_size: usize,
pub max_pool_size: usize,
}
static GLOBAL_POOL: std::sync::OnceLock<BufferPool> = std::sync::OnceLock::new();
pub fn global_pool() -> &'static BufferPool {
GLOBAL_POOL.get_or_init(|| {
BufferPool::new(8192, 100) })
}
pub fn init_global_pool(buffer_size: usize, max_pool_size: usize) -> Result<(), &'static str> {
GLOBAL_POOL
.set(BufferPool::new(buffer_size, max_pool_size))
.map_err(|_| "Global pool already initialized")
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_buffer_pool_basic() {
let pool = BufferPool::new(1024, 5);
let mut buffer = pool.get();
assert_eq!(buffer.capacity(), 1024);
assert!(buffer.is_empty());
buffer.extend_from_slice(b"hello");
assert_eq!(buffer.len(), 5);
let stats = pool.stats();
assert_eq!(stats.available_buffers, 0);
drop(buffer);
let stats = pool.stats();
assert_eq!(stats.available_buffers, 1);
}
#[test]
fn test_buffer_pool_reuse() {
let pool = BufferPool::new(1024, 5);
{
let mut buffer = pool.get();
buffer.extend_from_slice(b"test data");
}
let buffer = pool.get();
assert!(buffer.is_empty());
assert_eq!(buffer.capacity(), 1024);
}
#[test]
fn test_buffer_pool_max_size() {
let pool = BufferPool::new(1024, 2);
let _buffer1 = pool.get();
let _buffer2 = pool.get();
let _buffer3 = pool.get();
assert_eq!(pool.stats().available_buffers, 0);
drop(_buffer1);
drop(_buffer2);
drop(_buffer3);
assert_eq!(pool.stats().available_buffers, 2);
}
#[test]
fn test_global_pool() {
let pool = global_pool();
let buffer = pool.get();
assert!(buffer.capacity() > 0);
}
}