use std::ops::Range;
use crate::ThreadBound;
use crate::foundation::Error;
use objc2::rc::Retained;
use objc2::runtime::{AnyObject, ProtocolObject};
use objc2::{msg_send, sel};
use objc2_foundation::{NSRange, NSString};
use objc2_metal::{MTL4CounterHeap, MTL4CounterHeapType};
pub const MAX_TIMESTAMP_COUNTERS: usize = 4096;
pub struct TimestampCounterHeap {
inner: Retained<ProtocolObject<dyn MTL4CounterHeap>>,
count: usize,
_thread_bound: ThreadBound,
}
pub struct CounterReadback {
pub(super) heap: TimestampCounterHeap,
pub(super) buffer: crate::Buffer,
pub(super) submission_id: u64,
pub(super) count: usize,
}
impl TimestampCounterHeap {
pub(super) const fn new(
inner: Retained<ProtocolObject<dyn MTL4CounterHeap>>,
count: usize,
) -> Self {
Self {
inner,
count,
_thread_bound: ThreadBound::new(),
}
}
#[must_use]
pub const fn count(&self) -> usize {
self.count
}
pub fn label(&self) -> Result<Option<String>, Error> {
let object = self.as_any_object();
require_selector(object, sel!(label), "MTL4::CounterHeap::label")?;
let value: Option<Retained<NSString>> = unsafe { msg_send![object, label] };
Ok(value.map(|value| value.to_string()))
}
pub fn set_label(&self, label: Option<&str>) -> Result<(), Error> {
let object = self.as_any_object();
require_selector(object, sel!(setLabel:), "MTL4::CounterHeap::setLabel")?;
let label = label.map(NSString::from_str);
unsafe { msg_send![object, setLabel: label.as_deref()] }
Ok(())
}
pub fn heap_type(&self) -> Result<crate::metal::generated_value_types::CounterHeapType, Error> {
let object = self.as_any_object();
require_selector(object, sel!(type), "MTL4::CounterHeap::type")?;
let value: MTL4CounterHeapType = unsafe { msg_send![object, type] };
Ok(crate::metal::generated_value_types::CounterHeapType::from_system_raw(value.0))
}
pub fn invalidate_range(&mut self, range: Range<usize>) -> Result<(), Error> {
let range = checked_range(range, self.count, "invalidateCounterRange")?;
let object = self.as_any_object();
require_selector(
object,
sel!(invalidateCounterRange:),
"MTL4::CounterHeap::invalidateCounterRange",
)?;
unsafe { msg_send![object, invalidateCounterRange: range] }
Ok(())
}
pub(super) fn as_any_object(&self) -> &AnyObject {
unsafe { &*(std::ptr::from_ref(&*self.inner).cast::<AnyObject>()) }
}
}
fn require_selector(
object: &AnyObject,
selector: objc2::runtime::Sel,
name: &str,
) -> Result<(), Error> {
let available: bool = unsafe { msg_send![object, respondsToSelector: selector] };
if available {
Ok(())
} else {
Err(Error::unsupported(format!("{name} is unavailable")))
}
}
pub(super) fn checked_range(
range: Range<usize>,
count: usize,
name: &str,
) -> Result<NSRange, Error> {
let length = range
.end
.checked_sub(range.start)
.ok_or_else(|| Error::invalid_argument(format!("{name} has an inverted range")))?;
if range.end > count {
return Err(Error::invalid_argument(format!(
"{name} range {:?} exceeds counter heap count {count}",
range
)));
}
Ok(NSRange {
location: range.start,
length,
})
}
#[cfg(test)]
mod tests {
use super::checked_range;
use std::ops::Range;
#[test]
fn checked_range_accepts_empty_and_full_ranges() {
let empty = checked_range(3..3, 4, "test").unwrap();
assert_eq!((empty.location, empty.length), (3, 0));
let full = checked_range(0..4, 4, "test").unwrap();
assert_eq!((full.location, full.length), (0, 4));
}
#[test]
fn checked_range_rejects_inversion_and_overrun() {
assert!(checked_range(Range { start: 3, end: 2 }, 4, "test").is_err());
assert!(checked_range(2..5, 4, "test").is_err());
}
}