use objc2::rc::Retained;
use objc2::runtime::ProtocolObject;
use objc2_foundation::{NSArray, NSRange, NSString};
use objc2_metal::{
MTLCommonCounterSetTimestamp, MTLComputePassDescriptor, MTLCounterResultTimestamp,
MTLCounterSampleBuffer, MTLCounterSampleBufferDescriptor, MTLCounterSet, MTLDevice,
MTLRenderPassDescriptor, MTLStorageMode,
};
pub(super) use crate::gfx::render_graph::{PASS_COUNT, PASS_NAMES, PassId};
pub(super) const FRAMES_IN_FLIGHT: usize = 3;
const NO_SAMPLE: usize = usize::MAX;
pub(super) struct PassTimingResources {
buffers: [Retained<ProtocolObject<dyn MTLCounterSampleBuffer>>; FRAMES_IN_FLIGHT],
frame_slot: usize,
attached: std::sync::atomic::AtomicU64,
}
impl PassTimingResources {
pub(super) fn new(device: &ProtocolObject<dyn MTLDevice>) -> Option<Self> {
let sets: Retained<NSArray<ProtocolObject<dyn MTLCounterSet>>> = device.counterSets()?;
let target_name = unsafe { MTLCommonCounterSetTimestamp };
let mut timestamp_set: Option<Retained<ProtocolObject<dyn MTLCounterSet>>> = None;
for set in sets.iter() {
let name: Retained<NSString> = set.name();
if &*name == target_name {
timestamp_set = Some(set);
break;
}
}
let timestamp_set = timestamp_set?;
let make_buf = || -> Option<Retained<ProtocolObject<dyn MTLCounterSampleBuffer>>> {
let desc = MTLCounterSampleBufferDescriptor::new();
desc.setCounterSet(Some(×tamp_set));
desc.setStorageMode(MTLStorageMode::Shared);
unsafe { desc.setSampleCount(PASS_COUNT * 2) };
device
.newCounterSampleBufferWithDescriptor_error(&desc)
.ok()
};
let b0 = make_buf()?;
let b1 = make_buf()?;
let b2 = make_buf()?;
Some(Self {
buffers: [b0, b1, b2],
frame_slot: 0,
attached: std::sync::atomic::AtomicU64::new(0),
})
}
fn active(&self) -> &ProtocolObject<dyn MTLCounterSampleBuffer> {
&self.buffers[self.frame_slot]
}
pub(super) fn begin_frame(&mut self) -> usize {
let slot = self.frame_slot;
self.frame_slot = (self.frame_slot + 1) % FRAMES_IN_FLIGHT;
self.attached.store(0, std::sync::atomic::Ordering::Relaxed);
slot
}
pub(super) fn attached_mask(&self) -> u64 {
self.attached.load(std::sync::atomic::Ordering::Relaxed)
}
fn mark_attached(&self, pass: PassId) {
self.attached.fetch_or(
1u64 << (pass as usize),
std::sync::atomic::Ordering::Relaxed,
);
}
pub(super) fn buffer_for(
&self,
slot: usize,
) -> Retained<ProtocolObject<dyn MTLCounterSampleBuffer>> {
self.buffers[slot].clone()
}
pub(super) fn attach_render(&self, desc: &MTLRenderPassDescriptor, pass: PassId) {
self.mark_attached(pass);
let (start, end) = slot_pair(pass);
unsafe {
let arr = desc.sampleBufferAttachments();
let entry = arr.objectAtIndexedSubscript(0);
entry.setSampleBuffer(Some(self.active()));
entry.setStartOfVertexSampleIndex(start);
entry.setEndOfVertexSampleIndex(NO_SAMPLE);
entry.setStartOfFragmentSampleIndex(NO_SAMPLE);
entry.setEndOfFragmentSampleIndex(end);
}
}
pub(super) fn attach_render_first(&self, desc: &MTLRenderPassDescriptor, pass: PassId) {
self.mark_attached(pass);
let (start, _) = slot_pair(pass);
unsafe {
let arr = desc.sampleBufferAttachments();
let entry = arr.objectAtIndexedSubscript(0);
entry.setSampleBuffer(Some(self.active()));
entry.setStartOfVertexSampleIndex(start);
entry.setEndOfVertexSampleIndex(NO_SAMPLE);
entry.setStartOfFragmentSampleIndex(NO_SAMPLE);
entry.setEndOfFragmentSampleIndex(NO_SAMPLE);
}
}
pub(super) fn attach_render_last(&self, desc: &MTLRenderPassDescriptor, pass: PassId) {
let (_, end) = slot_pair(pass);
unsafe {
let arr = desc.sampleBufferAttachments();
let entry = arr.objectAtIndexedSubscript(0);
entry.setSampleBuffer(Some(self.active()));
entry.setStartOfVertexSampleIndex(NO_SAMPLE);
entry.setEndOfVertexSampleIndex(NO_SAMPLE);
entry.setStartOfFragmentSampleIndex(NO_SAMPLE);
entry.setEndOfFragmentSampleIndex(end);
}
}
pub(super) fn attach_compute(&self, desc: &MTLComputePassDescriptor, pass: PassId) {
self.mark_attached(pass);
let (start, end) = slot_pair(pass);
unsafe {
let arr = desc.sampleBufferAttachments();
let entry = arr.objectAtIndexedSubscript(0);
entry.setSampleBuffer(Some(self.active()));
entry.setStartOfEncoderSampleIndex(start);
entry.setEndOfEncoderSampleIndex(end);
}
}
}
fn slot_pair(pass: PassId) -> (usize, usize) {
debug_assert!(
(pass as usize) < PASS_COUNT,
"PassId {pass:?} (index {}) >= PASS_COUNT {PASS_COUNT}: register it in PASS_NAMES \
and bump PASS_COUNT",
pass as usize,
);
let base = pass as usize * 2;
(base, base + 1)
}
pub(super) fn resolve(buffer: &ProtocolObject<dyn MTLCounterSampleBuffer>) -> [u32; PASS_COUNT] {
let range = NSRange::new(0, PASS_COUNT * 2);
let Some(data) = (unsafe { buffer.resolveCounterRange(range) }) else {
return [0; PASS_COUNT];
};
let bytes = data.len();
let needed = std::mem::size_of::<MTLCounterResultTimestamp>() * PASS_COUNT * 2;
if bytes < needed {
return [0; PASS_COUNT];
}
let timestamps: &[MTLCounterResultTimestamp] = unsafe {
std::slice::from_raw_parts(
data.as_bytes_unchecked().as_ptr() as *const MTLCounterResultTimestamp,
PASS_COUNT * 2,
)
};
let mut out = [0u32; PASS_COUNT];
for i in 0..PASS_COUNT {
let start = timestamps[i * 2].timestamp;
let end = timestamps[i * 2 + 1].timestamp;
if end == 0 || start == 0 || end <= start {
out[i] = 0;
continue;
}
let ns = end - start;
out[i] = (ns / 1000).min(u32::MAX as u64) as u32;
}
out
}
pub(super) fn frame_span_us(buffer: &ProtocolObject<dyn MTLCounterSampleBuffer>) -> Option<u32> {
let range = NSRange::new(0, PASS_COUNT * 2);
let data = unsafe { buffer.resolveCounterRange(range) }?;
let needed = std::mem::size_of::<MTLCounterResultTimestamp>() * PASS_COUNT * 2;
if data.len() < needed {
return None;
}
let timestamps: &[MTLCounterResultTimestamp] = unsafe {
std::slice::from_raw_parts(
data.as_bytes_unchecked().as_ptr() as *const MTLCounterResultTimestamp,
PASS_COUNT * 2,
)
};
let mut min_start = u64::MAX;
let mut max_end = 0u64;
for i in 0..PASS_COUNT {
let start = timestamps[i * 2].timestamp;
let end = timestamps[i * 2 + 1].timestamp;
if start == 0 || end == 0 || end <= start {
continue;
}
min_start = min_start.min(start);
max_end = max_end.max(end);
}
if max_end <= min_start {
return None;
}
Some(((max_end - min_start) / 1000).min(u32::MAX as u64) as u32)
}