use std::{
collections::{HashMap, VecDeque}, sync::{
Arc, Mutex, atomic::{AtomicU64, Ordering}
}, time::{SystemTime, UNIX_EPOCH}
};
use anyhow::{Result, bail};
use wgpu::{Buffer, BufferUsages, Device};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, PartialOrd, Ord)]
pub enum BufferSizeTier {
Tiny, Small, Medium, Large, XLarge, }
impl BufferSizeTier {
pub const fn size_bytes(self) -> u64 {
match self {
Self::Tiny => 4 * 1024,
Self::Small => 64 * 1024,
Self::Medium => 1024 * 1024,
Self::Large => 16 * 1024 * 1024,
Self::XLarge => 128 * 1024 * 1024,
}
}
pub const fn for_size(size: u64) -> Self {
match size {
0..=4096 => Self::Tiny,
4097..=65536 => Self::Small,
65_537..=1_048_576 => Self::Medium,
1_048_577..=16_777_216 => Self::Large,
_ => Self::XLarge,
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub enum PooledBufferType {
Storage, Staging, Uniform, }
#[derive(Debug, Clone)]
pub struct BufferPoolConfig {
pub max_pool_memory: u64,
pub eviction_timeout_secs: u64,
}
impl Default for BufferPoolConfig {
fn default() -> Self {
Self { max_pool_memory: 512 * 1024 * 1024, eviction_timeout_secs: 30 }
}
}
struct PooledBuffer {
buffer: Arc<Buffer>,
size: u64,
last_used_timestamp: u64,
}
impl PooledBuffer {
fn new(buffer: Arc<Buffer>, size: u64) -> Self {
Self { buffer, size, last_used_timestamp: current_timestamp() }
}
}
#[derive(Debug, Clone, Default)]
pub struct BufferPoolStats {
pub total_buffers: usize,
pub total_allocated_bytes: u64,
pub total_allocations: u64,
pub total_reuses: u64,
}
#[derive(Debug)]
pub struct AcquiredBuffer {
pub buffer: Arc<Buffer>,
#[allow(dead_code)] pub is_reused: bool,
pub tier_size: u64,
}
pub struct BufferPool {
device: Arc<Device>,
config: BufferPoolConfig,
pools: Mutex<HashMap<(PooledBufferType, BufferSizeTier), VecDeque<PooledBuffer>>>,
total_allocated: Arc<AtomicU64>,
total_allocations: Arc<AtomicU64>,
total_reuses: Arc<AtomicU64>,
}
impl std::fmt::Debug for BufferPool {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("BufferPool").field("config", &self.config).field("total_allocated", &self.total_allocated.load(std::sync::atomic::Ordering::Relaxed)).field("total_allocations", &self.total_allocations.load(std::sync::atomic::Ordering::Relaxed)).field("total_reuses", &self.total_reuses.load(std::sync::atomic::Ordering::Relaxed)).finish_non_exhaustive()
}
}
impl BufferPool {
pub fn new(device: Arc<Device>, config: BufferPoolConfig) -> Self {
Self { device, config, pools: Mutex::new(HashMap::new()), total_allocated: Arc::new(AtomicU64::new(0)), total_allocations: Arc::new(AtomicU64::new(0)), total_reuses: Arc::new(AtomicU64::new(0)) }
}
#[allow(clippy::significant_drop_tightening)] pub fn acquire(&self, size: u64, buffer_type: PooledBufferType, label: Option<&str>) -> Result<(Arc<Buffer>, bool)> {
if size == 0 {
bail!("Cannot acquire buffer of size 0");
}
let tier = BufferSizeTier::for_size(size);
let tier_size = tier.size_bytes();
self.evict_if_needed()?;
let mut pools = self.pools.lock().unwrap();
let key = (buffer_type, tier);
if let Some(deque) = pools.get_mut(&key)
&& let Some(mut pooled) = deque.pop_front()
{
pooled.last_used_timestamp = current_timestamp();
let buffer = pooled.buffer;
self.total_reuses.fetch_add(1, Ordering::Relaxed);
return Ok((buffer, false));
}
let buffer_usage = match buffer_type {
PooledBufferType::Storage => BufferUsages::STORAGE | BufferUsages::COPY_SRC | BufferUsages::COPY_DST,
PooledBufferType::Staging => BufferUsages::COPY_DST | BufferUsages::MAP_READ,
PooledBufferType::Uniform => BufferUsages::UNIFORM | BufferUsages::COPY_DST,
};
let buffer = Arc::new(self.device.create_buffer(&wgpu::BufferDescriptor { label, size: tier_size, usage: buffer_usage, mapped_at_creation: false }));
let current_total = self.total_allocated.fetch_add(tier_size, Ordering::Relaxed);
if current_total + tier_size > self.config.max_pool_memory {
bail!("Buffer pool memory limit exceeded: {} + {} > {}", current_total, tier_size, self.config.max_pool_memory);
}
self.total_allocations.fetch_add(1, Ordering::Relaxed);
Ok((buffer, true))
}
#[allow(clippy::significant_drop_tightening)] pub fn acquire_with_metadata(&self, size: u64, buffer_type: PooledBufferType, label: Option<&str>) -> Result<AcquiredBuffer> {
let tier_size = BufferSizeTier::for_size(size).size_bytes();
let (buffer, is_new) = self.acquire(size, buffer_type, label)?;
Ok(AcquiredBuffer { buffer, is_reused: !is_new, tier_size })
}
#[allow(clippy::significant_drop_tightening)] pub fn release(&self, buffer: Arc<Buffer>) {
if Arc::strong_count(&buffer) != 1 {
return; }
let buffer_size = buffer.size();
let tier = BufferSizeTier::for_size(buffer_size);
let buffer_type = if buffer.usage().contains(BufferUsages::UNIFORM) {
PooledBufferType::Uniform
} else if buffer.usage().contains(BufferUsages::MAP_READ) {
PooledBufferType::Staging
} else {
PooledBufferType::Storage
};
let mut pools = self.pools.lock().unwrap();
let key = (buffer_type, tier);
let deque = pools.entry(key).or_default();
deque.push_back(PooledBuffer::new(buffer, buffer_size));
}
#[allow(clippy::significant_drop_tightening)] pub fn stats(&self) -> BufferPoolStats {
let pools = self.pools.lock().unwrap();
let total_buffers: usize = pools.values().map(std::collections::VecDeque::len).sum();
let total_allocated = self.total_allocated.load(Ordering::Relaxed);
BufferPoolStats { total_buffers, total_allocated_bytes: total_allocated, total_allocations: self.total_allocations.load(Ordering::Relaxed), total_reuses: self.total_reuses.load(Ordering::Relaxed) }
}
#[allow(clippy::unnecessary_wraps, clippy::significant_drop_tightening)] fn evict_if_needed(&self) -> Result<()> {
let current_total = self.total_allocated.load(Ordering::Relaxed);
let target_memory = (self.config.max_pool_memory * 90) / 100;
if current_total <= target_memory {
return Ok(());
}
let mut pools = self.pools.lock().unwrap();
let current_time = current_timestamp();
let timeout_secs = self.config.eviction_timeout_secs;
let mut evicted = 0u64;
for deque in pools.values_mut() {
deque.retain(|pooled| {
let age_secs = (current_time - pooled.last_used_timestamp) / 1000;
if age_secs > timeout_secs {
evicted += pooled.size;
false
} else {
true
}
});
}
self.total_allocated.fetch_sub(evicted, Ordering::Relaxed);
Ok(())
}
#[allow(dead_code, clippy::significant_drop_tightening)] pub fn clear(&self) {
let mut pools = self.pools.lock().unwrap();
let total_size: u64 = pools.values().flat_map(|deque| deque.iter().map(|b| b.size)).sum();
pools.clear();
self.total_allocated.fetch_sub(total_size, Ordering::Relaxed);
}
#[allow(dead_code, clippy::significant_drop_tightening)] pub fn clear_all(&self) {
let mut pools = self.pools.lock().unwrap();
pools.clear();
self.total_allocated.store(0, Ordering::Relaxed);
self.total_allocations.store(0, Ordering::Relaxed);
self.total_reuses.store(0, Ordering::Relaxed);
}
}
#[allow(clippy::cast_possible_truncation)] fn current_timestamp() -> u64 {
SystemTime::now().duration_since(UNIX_EPOCH).unwrap_or_default().as_millis() as u64
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_buffer_size_tier_selection() {
assert_eq!(BufferSizeTier::for_size(1024), BufferSizeTier::Tiny);
assert_eq!(BufferSizeTier::for_size(4096), BufferSizeTier::Tiny);
assert_eq!(BufferSizeTier::for_size(8192), BufferSizeTier::Small);
assert_eq!(BufferSizeTier::for_size(65536), BufferSizeTier::Small);
assert_eq!(BufferSizeTier::for_size(100_000), BufferSizeTier::Medium);
assert_eq!(BufferSizeTier::for_size(1_048_576), BufferSizeTier::Medium);
assert_eq!(BufferSizeTier::for_size(10_000_000), BufferSizeTier::Large);
assert_eq!(BufferSizeTier::for_size(17_000_000), BufferSizeTier::XLarge);
}
#[test]
fn test_buffer_pool_config_default() {
let config = BufferPoolConfig::default();
assert_eq!(config.max_pool_memory, 512 * 1024 * 1024);
assert_eq!(config.eviction_timeout_secs, 30);
}
}