use std::collections::VecDeque;
use objc2::rc::Retained;
use objc2::runtime::ProtocolObject;
use objc2_metal::{MTLBuffer, MTLDevice, MTLResourceOptions};
use super::context::{bytes_of_slice, write_buffer_region};
pub(super) fn grow_to(have: usize, needed: usize) -> Option<usize> {
let needed = needed.max(1);
if have >= needed {
return None;
}
Some(needed.next_power_of_two().max(256))
}
pub(super) struct RetirePool<T> {
pending: VecDeque<(u64, T)>,
}
impl<T> RetirePool<T> {
pub(super) fn new() -> Self {
Self {
pending: VecDeque::new(),
}
}
pub(super) fn push(&mut self, retired_at: u64, payload: T) {
self.pending.push_back((retired_at, payload));
}
pub(super) fn collect(&mut self, frame_id: u64, depth: u64) {
while let Some(&(retired_at, _)) = self.pending.front() {
if retired_at.saturating_add(depth) <= frame_id {
self.pending.pop_front();
} else {
break;
}
}
}
#[cfg(test)]
pub(super) fn len(&self) -> usize {
self.pending.len()
}
}
pub(super) struct TransientRing {
slots: Vec<Option<Retained<ProtocolObject<dyn MTLBuffer>>>>,
}
impl TransientRing {
pub(super) fn new(depth: usize) -> Self {
Self {
slots: (0..depth.max(1)).map(|_| None).collect(),
}
}
pub(super) fn slot(
&mut self,
device: &ProtocolObject<dyn MTLDevice>,
slot: usize,
min_len: usize,
) -> Result<Retained<ProtocolObject<dyn MTLBuffer>>, String> {
let idx = slot % self.slots.len();
let have = self.slots[idx].as_ref().map_or(0, |buf| buf.length());
if let Some(cap) = grow_to(have, min_len) {
let buf = device
.newBufferWithLength_options(cap, MTLResourceOptions::StorageModeShared)
.ok_or("failed to allocate transient ring buffer")?;
self.slots[idx] = Some(buf);
}
Ok(self.slots[idx]
.as_ref()
.expect("ring slot was just ensured")
.clone())
}
pub(super) fn write(
&mut self,
device: &ProtocolObject<dyn MTLDevice>,
slot: usize,
bytes: &[u8],
) -> Result<Retained<ProtocolObject<dyn MTLBuffer>>, String> {
let buf = self.slot(device, slot, bytes.len().max(1))?;
write_buffer_region(&buf, 0, bytes)?;
Ok(buf)
}
}
type ObjectBuffer = Option<Retained<ProtocolObject<dyn MTLBuffer>>>;
#[derive(Default)]
struct SkinSlot {
palette: ObjectBuffer,
weights: ObjectBuffer,
}
pub(super) struct JointRing {
slots: Vec<Vec<SkinSlot>>,
}
impl JointRing {
pub(super) fn new(depth: usize) -> Self {
Self {
slots: (0..depth.max(1)).map(|_| Vec::new()).collect(),
}
}
pub(super) fn write_all(
&mut self,
device: &ProtocolObject<dyn MTLDevice>,
slot: usize,
palettes: &[Vec<[[f32; 4]; 4]>],
) -> Result<Vec<Retained<ProtocolObject<dyn MTLBuffer>>>, String> {
self.objects(slot, palettes.len())
.iter_mut()
.zip(palettes)
.map(|(o, mats)| fill(device, &mut o.palette, bytes_of_slice(mats), "joint"))
.collect()
}
pub(super) fn write_weights(
&mut self,
device: &ProtocolObject<dyn MTLDevice>,
slot: usize,
weights: &[Vec<f32>],
) -> Result<Vec<Retained<ProtocolObject<dyn MTLBuffer>>>, String> {
self.objects(slot, weights.len())
.iter_mut()
.zip(weights)
.map(|(o, w)| fill(device, &mut o.weights, bytes_of_slice(w), "morph weight"))
.collect()
}
fn objects(&mut self, slot: usize, count: usize) -> &mut [SkinSlot] {
let idx = slot % self.slots.len();
let slots = &mut self.slots[idx];
if slots.len() < count {
slots.resize_with(count, SkinSlot::default);
}
&mut slots[..count]
}
}
fn fill(
device: &ProtocolObject<dyn MTLDevice>,
cell: &mut ObjectBuffer,
bytes: &[u8],
what: &str,
) -> Result<Retained<ProtocolObject<dyn MTLBuffer>>, String> {
let have = cell.as_ref().map_or(0, |buf| buf.length());
if let Some(cap) = grow_to(have, bytes.len()) {
let buf = device
.newBufferWithLength_options(cap, MTLResourceOptions::StorageModeShared)
.ok_or_else(|| format!("failed to allocate {what} ring buffer"))?;
*cell = Some(buf);
}
let buf = cell.as_ref().expect("ring slot was just ensured");
write_buffer_region(buf, 0, bytes)?;
Ok(buf.clone())
}
#[cfg(test)]
mod tests {
use super::{RetirePool, grow_to};
#[test]
fn grow_to_keeps_a_slot_that_already_fits() {
assert_eq!(grow_to(256, 200), None);
assert_eq!(grow_to(256, 256), None);
assert_eq!(grow_to(256, 0), None);
}
#[test]
fn grow_to_rounds_up_past_the_floor() {
assert_eq!(grow_to(0, 1), Some(256));
assert_eq!(grow_to(0, 0), Some(256));
assert_eq!(grow_to(0, 300), Some(512));
assert_eq!(grow_to(512, 513), Some(1024));
assert_eq!(grow_to(1024, 4096), Some(4096));
}
#[test]
fn retire_pool_holds_payloads_for_depth_frames() {
let mut pool: RetirePool<u32> = RetirePool::new();
pool.push(0, 100);
pool.push(1, 101);
pool.push(2, 102);
pool.collect(1, 2);
assert_eq!(pool.len(), 3);
pool.collect(2, 2);
assert_eq!(pool.len(), 2);
pool.collect(3, 2);
assert_eq!(pool.len(), 1);
pool.collect(100, 2);
assert_eq!(pool.len(), 0);
}
#[test]
fn retire_pool_depth_one_frees_one_frame_later() {
let mut pool: RetirePool<u32> = RetirePool::new();
pool.push(5, 7);
pool.collect(5, 1);
assert_eq!(pool.len(), 1);
pool.collect(6, 1);
assert_eq!(pool.len(), 0);
}
#[test]
fn retire_pool_drains_multiple_same_frame_pushes() {
let mut pool: RetirePool<u32> = RetirePool::new();
pool.push(4, 1);
pool.push(4, 2);
pool.push(4, 3);
pool.collect(5, 2); assert_eq!(pool.len(), 3);
pool.collect(6, 2); assert_eq!(pool.len(), 0);
}
#[test]
fn retire_pool_collect_is_idempotent_and_empty_safe() {
let mut pool: RetirePool<u32> = RetirePool::new();
pool.collect(10, 2); assert_eq!(pool.len(), 0);
pool.push(0, 1);
pool.collect(10, 2);
pool.collect(10, 2); assert_eq!(pool.len(), 0);
}
}