use std::collections::HashMap;
use std::sync::Mutex;
use objc2::rc::Retained;
use objc2::runtime::ProtocolObject;
use objc2_metal::{MTLBuffer, MTLDevice, MTLResourceOptions};
const CAPACITY_BYTES: usize = 256 * 1024 * 1024;
pub(super) struct Pool {
state: Mutex<PoolState>,
}
struct PoolState {
free: HashMap<usize, Vec<Retained<ProtocolObject<dyn MTLBuffer>>>>,
held_bytes: usize,
}
impl Pool {
pub(super) fn new() -> Self {
Self {
state: Mutex::new(PoolState {
free: HashMap::new(),
held_bytes: 0,
}),
}
}
pub(super) fn take(
&self,
device: &ProtocolObject<dyn MTLDevice>,
bytes: usize,
) -> Result<Retained<ProtocolObject<dyn MTLBuffer>>, String> {
let class = bytes.max(16).next_power_of_two();
{
let mut state = self.state.lock().expect("the buffer pool is poisoned");
if let Some(buffer) = state.free.get_mut(&class).and_then(Vec::pop) {
state.held_bytes -= class;
return Ok(buffer);
}
}
device
.newBufferWithLength_options(class, MTLResourceOptions::StorageModeShared)
.ok_or_else(|| format!("allocating a {class}-byte buffer failed"))
}
pub(super) fn give(&self, buffer: Retained<ProtocolObject<dyn MTLBuffer>>) {
let class = buffer.length();
let mut state = self.state.lock().expect("the buffer pool is poisoned");
if state.held_bytes + class > CAPACITY_BYTES {
return;
}
state.held_bytes += class;
state.free.entry(class).or_default().push(buffer);
}
}