use crate::ThreadBound;
use crate::foundation::Error;
use crate::metal::Device;
use crate::metal::generated_object_types::metal::LogState;
use crate::metal::generated_value_types::LogLevel;
use block2::{DynBlock, RcBlock};
use objc2::rc::Retained;
use objc2::runtime::{NSObjectProtocol, ProtocolObject};
use objc2::{msg_send, sel};
use objc2_foundation::NSString;
use objc2_metal::{
MTLCopyAllDevices, MTLCopyAllDevicesWithObserver, MTLDevice, MTLRemoveDeviceObserver,
};
use std::cell::RefCell;
use std::panic::{AssertUnwindSafe, catch_unwind};
use std::ptr::NonNull;
use std::sync::{Arc, Condvar, Mutex};
thread_local! {
static EXECUTING_DEVICE_CALLBACKS: RefCell<Vec<usize>> = const { RefCell::new(Vec::new()) };
}
struct DeliveryStatus {
active: bool,
in_flight: usize,
}
struct DeliveryState {
status: Mutex<DeliveryStatus>,
drained: Condvar,
}
impl DeliveryState {
fn new() -> Arc<Self> {
Arc::new(Self {
status: Mutex::new(DeliveryStatus {
active: true,
in_flight: 0,
}),
drained: Condvar::new(),
})
}
fn enter(state: &Arc<Self>) -> Option<DeliveryGuard> {
let mut status = state
.status
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
if !status.active {
return None;
}
status.in_flight = status.in_flight.checked_add(1)?;
drop(status);
let identity = Arc::as_ptr(state) as usize;
EXECUTING_DEVICE_CALLBACKS.with(|callbacks| callbacks.borrow_mut().push(identity));
Some(DeliveryGuard {
state: Arc::clone(state),
identity,
})
}
fn deactivate_and_wait(&self) {
let identity = self as *const Self as usize;
let executing_here = EXECUTING_DEVICE_CALLBACKS.with(|callbacks| {
callbacks
.borrow()
.iter()
.filter(|candidate| **candidate == identity)
.count()
});
let mut status = self
.status
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
status.active = false;
while status.in_flight > executing_here {
status = self
.drained
.wait(status)
.unwrap_or_else(std::sync::PoisonError::into_inner);
}
}
}
struct DeliveryGuard {
state: Arc<DeliveryState>,
identity: usize,
}
impl Drop for DeliveryGuard {
fn drop(&mut self) {
EXECUTING_DEVICE_CALLBACKS.with(|callbacks| {
let mut callbacks = callbacks.borrow_mut();
if let Some(position) = callbacks
.iter()
.rposition(|identity| *identity == self.identity)
{
callbacks.remove(position);
}
});
let mut status = self
.state
.status
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
status.in_flight = status.in_flight.saturating_sub(1);
self.state.drained.notify_all();
}
}
pub struct DeviceObserverRegistration {
observer: Retained<ProtocolObject<dyn NSObjectProtocol>>,
delivery: Arc<DeliveryState>,
handler: Arc<Mutex<Option<Arc<DeviceNotificationHandler>>>>,
_thread_bound: ThreadBound,
}
type DeviceNotificationHandler = dyn Fn(Device, String) + Send + Sync + 'static;
type DeviceNotificationBlock = dyn Fn(NonNull<ProtocolObject<dyn MTLDevice>>, NonNull<NSString>);
pub struct LogHandlerRegistration {
delivery: Arc<DeliveryState>,
handler: Arc<Mutex<Option<Arc<LogHandler>>>>,
_thread_bound: ThreadBound,
}
type LogHandler = dyn Fn(String, String, LogLevel, String) + Send + Sync + 'static;
impl Drop for LogHandlerRegistration {
fn drop(&mut self) {
self.delivery.deactivate_and_wait();
self.handler
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.take();
}
}
impl LogState {
pub fn add_log_handler(
&self,
handler: impl Fn(String, String, LogLevel, String) + Send + Sync + 'static,
) -> Result<LogHandlerRegistration, Error> {
let available: bool =
unsafe { msg_send![self.as_inner(), respondsToSelector: sel!(addLogHandler:)] };
if !available {
return Err(Error::unsupported("Metal log handlers are unavailable"));
}
let delivery = DeliveryState::new();
let callback_delivery = Arc::clone(&delivery);
let handler: Arc<LogHandler> = Arc::new(handler);
let handler = Arc::new(Mutex::new(Some(handler)));
let callback_handler = Arc::clone(&handler);
let block = RcBlock::new(
move |subsystem: NonNull<NSString>,
category: NonNull<NSString>,
level: isize,
message: NonNull<NSString>| {
let Some(_guard) = DeliveryState::enter(&callback_delivery) else {
return;
};
let subsystem = unsafe { subsystem.as_ref() }.to_string();
let category = unsafe { category.as_ref() }.to_string();
let message = unsafe { message.as_ref() }.to_string();
let handler = callback_handler
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.clone();
let Some(handler) = handler else {
return;
};
let _ = catch_unwind(AssertUnwindSafe(|| {
handler(
subsystem,
category,
LogLevel::from_system_raw(level),
message,
)
}));
},
);
unsafe {
let _: () = msg_send![self.as_inner(), addLogHandler: &*block];
}
Ok(LogHandlerRegistration {
delivery,
handler,
_thread_bound: ThreadBound::new(),
})
}
}
impl Drop for DeviceObserverRegistration {
fn drop(&mut self) {
{
let mut status = self
.delivery
.status
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
status.active = false;
}
unsafe { MTLRemoveDeviceObserver(&self.observer) };
self.delivery.deactivate_and_wait();
self.handler
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.take();
}
}
#[must_use]
pub fn all_devices() -> Vec<Device> {
MTLCopyAllDevices()
.to_vec()
.into_iter()
.map(Device::from_inner)
.collect()
}
pub fn all_devices_with_observer(
handler: impl Fn(Device, String) + Send + Sync + 'static,
) -> Result<(Vec<Device>, DeviceObserverRegistration), Error> {
let delivery = DeliveryState::new();
let callback_delivery = Arc::clone(&delivery);
let handler: Arc<DeviceNotificationHandler> = Arc::new(handler);
let handler = Arc::new(Mutex::new(Some(handler)));
let callback_handler = Arc::clone(&handler);
let block = RcBlock::new(
move |device: NonNull<ProtocolObject<dyn MTLDevice>>, name: NonNull<NSString>| {
let Some(_guard) = DeliveryState::enter(&callback_delivery) else {
return;
};
let device = unsafe { Retained::retain(device.as_ptr()) }.map(Device::from_inner);
let name = unsafe { name.as_ref() }.to_string();
if let Some(device) = device {
let handler = callback_handler
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.clone();
if let Some(handler) = handler {
let _ = catch_unwind(AssertUnwindSafe(|| handler(device, name)));
}
}
},
);
let block: &DynBlock<DeviceNotificationBlock> = █
let mut observer = std::ptr::null_mut::<ProtocolObject<dyn NSObjectProtocol>>();
let devices = unsafe {
MTLCopyAllDevicesWithObserver(
NonNull::from(&mut observer),
std::ptr::from_ref(block).cast_mut(),
)
};
let observer = unsafe { Retained::from_raw(observer) }
.ok_or_else(|| Error::unsupported("Metal did not return a device observer"))?;
let devices = devices
.to_vec()
.into_iter()
.map(Device::from_inner)
.collect();
Ok((
devices,
DeviceObserverRegistration {
observer,
delivery,
handler,
_thread_bound: ThreadBound::new(),
},
))
}
pub const METAL_CPP_VERSION_MAJOR: u32 = 381;
pub const METAL_CPP_VERSION_MINOR: u32 = 0;
pub const METAL_CPP_VERSION_PATCH: u32 = 0;
const fn version_key(major: u32, minor: u32, patch: u32) -> u128 {
((major as u128) << 64) | ((minor as u128) << 32) | patch as u128
}
#[must_use]
pub const fn metal_cpp_supports_version(major: u32, minor: u32, patch: u32) -> bool {
version_key(major, minor, patch)
<= version_key(
METAL_CPP_VERSION_MAJOR,
METAL_CPP_VERSION_MINOR,
METAL_CPP_VERSION_PATCH,
)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn deactivation_closes_callback_entry() {
let state = DeliveryState::new();
let guard = DeliveryState::enter(&state).expect("initial callback entry");
state.deactivate_and_wait();
drop(guard);
assert!(DeliveryState::enter(&state).is_none());
}
#[test]
fn version_query_matches_metal_cpp_macro_ordering() {
assert!(metal_cpp_supports_version(381, 0, 0));
assert!(metal_cpp_supports_version(380, u32::MAX, u32::MAX));
assert!(!metal_cpp_supports_version(381, 0, 1));
}
}