use crate::{StorageError, StorageResult};
use std::{
collections::VecDeque,
sync::{
atomic::{AtomicUsize, Ordering},
Arc, Mutex,
},
};
use tokio::sync::Semaphore;
pub struct PooledBuffer {
pub buffer: Vec<u8>,
pool: Arc<MemoryPoolInner>,
size: usize,
}
impl PooledBuffer {
pub fn clear(&mut self) {
self.buffer.clear();
}
pub fn capacity(&self) -> usize {
self.buffer.capacity()
}
pub fn write(&mut self, data: &[u8]) -> StorageResult<()> {
if self.buffer.len() + data.len() > self.buffer.capacity() {
return Err(StorageError::InsufficientSpace {
required: data.len() as u64,
available: (self.buffer.capacity() - self.buffer.len()) as u64,
});
}
self.buffer.extend_from_slice(data);
Ok(())
}
pub fn data(&self) -> &[u8] {
&self.buffer
}
pub fn as_mut(&mut self) -> &mut Vec<u8> {
&mut self.buffer
}
}
impl Drop for PooledBuffer {
fn drop(&mut self) {
self.buffer.clear();
let buffer = std::mem::take(&mut self.buffer);
if buffer.capacity() >= self.size {
let mut pool = self.pool.buffers.lock().unwrap();
if pool.len() < self.pool.max_buffers {
pool.push_back(buffer);
self.pool.available_count.fetch_add(1, Ordering::Relaxed);
}
}
}
}
#[derive(Debug)]
struct MemoryPoolInner {
buffers: Mutex<VecDeque<Vec<u8>>>,
available_count: AtomicUsize,
max_buffers: usize,
buffer_size: usize,
semaphore: Semaphore,
}
#[derive(Debug)]
pub struct MemoryPool {
inner: Arc<MemoryPoolInner>,
stats: Arc<PoolStats>,
}
#[derive(Debug, Default)]
pub struct PoolStats {
pub total_gets: AtomicUsize,
pub total_returns: AtomicUsize,
pub cache_hits: AtomicUsize,
pub cache_misses: AtomicUsize,
pub active_buffers: AtomicUsize,
pub pool_size: AtomicUsize,
}
impl PoolStats {
pub fn hit_rate(&self) -> f64 {
let hits = self.cache_hits.load(Ordering::Relaxed) as f64;
let total = hits + self.cache_misses.load(Ordering::Relaxed) as f64;
if total > 0.0 {
hits / total
} else {
0.0
}
}
pub fn utilization(&self) -> f64 {
let active = self.active_buffers.load(Ordering::Relaxed) as f64;
let total = self.pool_size.load(Ordering::Relaxed) as f64;
if total > 0.0 {
active / total
} else {
0.0
}
}
}
impl MemoryPool {
pub fn new(pool_size_bytes: usize, buffer_size: usize) -> StorageResult<Self> {
let max_buffers = pool_size_bytes / buffer_size;
if max_buffers == 0 {
return Err(StorageError::configuration(
"Pool size too small for requested buffer size"
));
}
let mut buffers = VecDeque::with_capacity(max_buffers);
for _ in 0..max_buffers {
let mut buffer = Vec::with_capacity(buffer_size);
buffer.reserve_exact(buffer_size);
buffers.push_back(buffer);
}
let stats = Arc::new(PoolStats::default());
stats.pool_size.store(max_buffers, Ordering::Relaxed);
let inner = Arc::new(MemoryPoolInner {
buffers: Mutex::new(buffers),
available_count: AtomicUsize::new(max_buffers),
max_buffers,
buffer_size,
semaphore: Semaphore::new(max_buffers),
});
tracing::info!(
"MemoryPool initialized: {} buffers, {} bytes each, {} MB total",
max_buffers,
buffer_size,
pool_size_bytes / (1024 * 1024)
);
Ok(Self { inner, stats })
}
pub async fn get_buffer(&self) -> StorageResult<PooledBuffer> {
let _permit = self.inner.semaphore.acquire().await
.map_err(|_| StorageError::internal("Failed to acquire buffer permit"))?;
self.stats.total_gets.fetch_add(1, Ordering::Relaxed);
if let Some(buffer) = self.try_get_from_pool() {
self.stats.cache_hits.fetch_add(1, Ordering::Relaxed);
self.stats.active_buffers.fetch_add(1, Ordering::Relaxed);
return Ok(PooledBuffer {
buffer,
pool: Arc::clone(&self.inner),
size: self.inner.buffer_size,
});
}
self.stats.cache_misses.fetch_add(1, Ordering::Relaxed);
self.stats.active_buffers.fetch_add(1, Ordering::Relaxed);
let mut buffer = Vec::with_capacity(self.inner.buffer_size);
buffer.reserve_exact(self.inner.buffer_size);
Ok(PooledBuffer {
buffer,
pool: Arc::clone(&self.inner),
size: self.inner.buffer_size,
})
}
fn try_get_from_pool(&self) -> Option<Vec<u8>> {
let mut buffers = self.inner.buffers.lock().ok()?;
if let Some(buffer) = buffers.pop_front() {
self.inner.available_count.fetch_sub(1, Ordering::Relaxed);
Some(buffer)
} else {
None
}
}
pub fn stats(&self) -> PoolStats {
PoolStats {
total_gets: AtomicUsize::new(self.stats.total_gets.load(Ordering::Relaxed)),
total_returns: AtomicUsize::new(self.stats.total_returns.load(Ordering::Relaxed)),
cache_hits: AtomicUsize::new(self.stats.cache_hits.load(Ordering::Relaxed)),
cache_misses: AtomicUsize::new(self.stats.cache_misses.load(Ordering::Relaxed)),
active_buffers: AtomicUsize::new(self.stats.active_buffers.load(Ordering::Relaxed)),
pool_size: AtomicUsize::new(self.stats.pool_size.load(Ordering::Relaxed)),
}
}
pub fn health_check(&self) -> bool {
let available = self.inner.available_count.load(Ordering::Relaxed);
let active = self.stats.active_buffers.load(Ordering::Relaxed);
let total = available + active;
total <= self.inner.max_buffers
}
pub async fn warmup(&self) -> StorageResult<()> {
tracing::info!("Warming up memory pool...");
let warmup_count = self.inner.max_buffers / 2;
let mut buffers = Vec::new();
for _ in 0..warmup_count {
if let Ok(buffer) = self.get_buffer().await {
buffers.push(buffer);
}
}
drop(buffers);
tracing::info!("Memory pool warmup completed");
Ok(())
}
}
pub struct ParallelBufferProcessor {
pool: Arc<MemoryPool>,
worker_count: usize,
}
impl ParallelBufferProcessor {
pub fn new(pool: Arc<MemoryPool>, worker_count: Option<usize>) -> Self {
let worker_count = worker_count.unwrap_or_else(|| {
std::thread::available_parallelism()
.map(|n| n.get() * 2)
.unwrap_or(8)
});
Self { pool, worker_count }
}
pub async fn process_batch<T, F, Fut>(
&self,
items: Vec<T>,
processor: F,
) -> StorageResult<Vec<StorageResult<()>>>
where
T: Send + 'static + Clone,
F: Fn(T, PooledBuffer) -> Fut + Send + Sync + 'static,
Fut: std::future::Future<Output = StorageResult<()>> + Send,
{
let processor = Arc::new(processor);
let chunk_size = (items.len() + self.worker_count - 1) / self.worker_count;
let mut handles = Vec::new();
for chunk in items.chunks(chunk_size) {
let chunk = chunk.to_vec();
let pool = Arc::clone(&self.pool);
let processor = Arc::clone(&processor);
let handle = tokio::spawn(async move {
let mut results = Vec::new();
for item in chunk {
match pool.get_buffer().await {
Ok(buffer) => {
let result = processor(item, buffer).await;
results.push(result);
}
Err(e) => {
results.push(Err(e));
}
}
}
results
});
handles.push(handle);
}
let mut all_results = Vec::new();
for handle in handles {
match handle.await {
Ok(results) => all_results.extend(results),
Err(e) => return Err(StorageError::internal(format!("Worker task failed: {}", e))),
}
}
Ok(all_results)
}
}
#[cfg(test)]
mod tests {
use super::*;
use tokio;
#[tokio::test]
async fn test_memory_pool_basic_operations() {
let pool = MemoryPool::new(1024 * 1024, 1024).unwrap();
let mut buffer = pool.get_buffer().await.unwrap();
assert_eq!(buffer.capacity(), 1024);
let test_data = b"Hello, Memory Pool!";
buffer.write(test_data).unwrap();
assert_eq!(buffer.data(), test_data);
drop(buffer);
let stats = pool.stats();
assert_eq!(stats.total_gets.load(Ordering::Relaxed), 1);
}
#[tokio::test]
async fn test_concurrent_buffer_access() {
let pool = Arc::new(MemoryPool::new(10 * 1024, 1024).unwrap());
let mut handles = Vec::new();
for i in 0..5 {
let pool_clone = Arc::clone(&pool);
let handle = tokio::spawn(async move {
let mut buffer = pool_clone.get_buffer().await.unwrap();
buffer.write(format!("Data {}", i).as_bytes()).unwrap();
tokio::time::sleep(tokio::time::Duration::from_millis(10)).await;
buffer
});
handles.push(handle);
}
for handle in handles {
let _buffer = handle.await.unwrap();
}
assert!(pool.health_check());
}
}