use std::ffi::CStr;
use std::ptr;
use std::sync::atomic::{AtomicBool, AtomicU8, AtomicUsize, Ordering};
use std::sync::{Arc, LazyLock, Mutex};
use metal::{BufferRef as MTLBufferRef, CommandQueueRef, DeviceRef};
use objc::runtime::{Class, Object, BOOL, YES};
use objc::{msg_send, sel, sel_impl};
use crate::error::MlxError;
static RESIDENCY_DISABLED_FLAG: AtomicU8 = AtomicU8::new(0);
static TEST_ALLOCATION_COUNT: AtomicUsize = AtomicUsize::new(0);
static TEST_COMMIT_CALL_COUNT: AtomicUsize = AtomicUsize::new(0);
#[repr(C)]
#[allow(non_snake_case)]
struct NSOperatingSystemVersion {
majorVersion: isize,
minorVersion: isize,
patchVersion: isize,
}
struct ObjcResidencySet {
ptr: *mut Object,
}
unsafe impl Send for ObjcResidencySet {}
unsafe impl Sync for ObjcResidencySet {}
enum ResidencySetInner {
Active {
object: ObjcResidencySet,
lock: Mutex<()>,
pending: AtomicBool,
heartbeat_ticks: AtomicUsize,
},
Noop,
}
impl Drop for ResidencySetInner {
fn drop(&mut self) {
let Self::Active { object, lock, .. } = self else {
return;
};
if let Ok(_guard) = lock.lock() {
unsafe {
let _: () = msg_send![object.ptr, removeAllAllocations];
let _: () = msg_send![object.ptr, commit];
let _: () = msg_send![object.ptr, release];
}
}
}
}
#[derive(Clone)]
pub(crate) struct ResidencySet {
inner: Arc<ResidencySetInner>,
}
impl ResidencySet {
pub(crate) fn new(device: &DeviceRef) -> Result<Self, MlxError> {
if !macos_15_or_newer() {
return Ok(Self {
inner: Arc::new(ResidencySetInner::Noop),
});
}
let Some(descriptor_class) = Class::get("MTLResidencySetDescriptor") else {
return Err(MlxError::ResidencySetError(
"MTLResidencySetDescriptor is unavailable".into(),
));
};
unsafe {
let descriptor: *mut Object = msg_send![descriptor_class, alloc];
let descriptor: *mut Object = msg_send![descriptor, init];
if descriptor.is_null() {
return Err(MlxError::ResidencySetError(
"failed to allocate MTLResidencySetDescriptor".into(),
));
}
let _: () = msg_send![descriptor, setInitialCapacity: 256usize];
let label_ptr = b"mlx_native_default\0".as_ptr() as *const i8;
if let Some(nsstring_class) = Class::get("NSString") {
let label_ns: *mut Object =
msg_send![nsstring_class, stringWithUTF8String: label_ptr];
if !label_ns.is_null() {
let _: () = msg_send![descriptor, setLabel: label_ns];
}
}
let mut error: *mut Object = ptr::null_mut();
let set: *mut Object =
msg_send![device, newResidencySetWithDescriptor: descriptor error: &mut error];
let _: () = msg_send![descriptor, release];
if !error.is_null() {
return Err(MlxError::ResidencySetError(ns_error_message(error)));
}
if set.is_null() {
return Err(MlxError::ResidencySetError(
"newResidencySetWithDescriptor:error: returned nil".into(),
));
}
let inner = Arc::new(ResidencySetInner::Active {
object: ObjcResidencySet { ptr: set },
lock: Mutex::new(()),
pending: AtomicBool::new(false),
heartbeat_ticks: AtomicUsize::new(residency_heartbeat_ticks()),
});
spawn_residency_heartbeat(&inner);
Ok(Self { inner })
}
}
pub(crate) fn keep_alive(&self) {
let ResidencySetInner::Active {
heartbeat_ticks, ..
} = &*self.inner
else {
return;
};
if *RESIDENCY_HEARTBEAT_KEEP_ALIVE_SECONDS > 0 {
heartbeat_ticks.store(residency_heartbeat_ticks(), Ordering::Release);
}
}
#[inline]
pub(crate) fn is_noop(&self) -> bool {
matches!(&*self.inner, ResidencySetInner::Noop)
}
#[inline]
pub(crate) fn same_owner(&self, other: &Self) -> bool {
Arc::ptr_eq(&self.inner, &other.inner)
}
pub(crate) fn add_allocation(&self, buffer: &MTLBufferRef) {
self.with_active_set(|set, pending| unsafe {
let _: () = msg_send![set, addAllocation: buffer];
TEST_ALLOCATION_COUNT.fetch_add(1, Ordering::AcqRel);
pending.store(true, Ordering::Release);
});
}
pub(crate) fn remove_allocation(&self, buffer: &MTLBufferRef) {
self.with_active_set(|set, pending| unsafe {
let _: () = msg_send![set, removeAllocation: buffer];
let _ =
TEST_ALLOCATION_COUNT.fetch_update(Ordering::AcqRel, Ordering::Acquire, |count| {
count.checked_sub(1)
});
pending.store(true, Ordering::Release);
});
}
#[allow(dead_code)]
pub(crate) fn remove_all_allocations(&self) {
self.with_active_set(|set, pending| unsafe {
let _: () = msg_send![set, removeAllAllocations];
TEST_ALLOCATION_COUNT.store(0, Ordering::Release);
pending.store(true, Ordering::Release);
});
}
pub(crate) fn commit(&self) {
self.with_active_set(|set, pending| unsafe {
let _: () = msg_send![set, commit];
TEST_COMMIT_CALL_COUNT.fetch_add(1, Ordering::AcqRel);
pending.store(false, Ordering::Release);
});
}
pub(crate) fn flush_pending(&self) -> bool {
let mut committed = false;
self.with_active_set(|set, pending| {
if pending.swap(false, Ordering::AcqRel) {
unsafe {
let _: () = msg_send![set, commit];
}
TEST_COMMIT_CALL_COUNT.fetch_add(1, Ordering::AcqRel);
committed = true;
}
});
committed
}
pub(crate) fn register_with_queue(&self, queue: &CommandQueueRef) {
self.with_active_set(|set, _pending| unsafe {
let _: () = msg_send![queue, addResidencySet: set];
});
}
fn with_active_set(&self, f: impl FnOnce(*mut Object, &AtomicBool)) {
let ResidencySetInner::Active {
object,
lock,
pending,
..
} = &*self.inner
else {
return;
};
if let Ok(_guard) = lock.lock() {
f(object.ptr, pending);
}
}
}
const RESIDENCY_HEARTBEAT_INTERVAL_MS: u64 = 5;
static RESIDENCY_HEARTBEAT_KEEP_ALIVE_SECONDS: LazyLock<usize> = LazyLock::new(|| {
std::env::var("MLX_NATIVE_RESIDENCY_KEEP_ALIVE_SECONDS")
.ok()
.and_then(|value| value.parse::<usize>().ok())
.unwrap_or(180)
});
fn residency_heartbeat_ticks() -> usize {
residency_heartbeat_ticks_for(*RESIDENCY_HEARTBEAT_KEEP_ALIVE_SECONDS)
}
fn residency_heartbeat_ticks_for(seconds: usize) -> usize {
seconds.saturating_mul(1000) / RESIDENCY_HEARTBEAT_INTERVAL_MS as usize
}
fn spawn_residency_heartbeat(inner: &Arc<ResidencySetInner>) {
if *RESIDENCY_HEARTBEAT_KEEP_ALIVE_SECONDS == 0 {
return;
}
let weak = Arc::downgrade(inner);
let _ = std::thread::Builder::new()
.name("mlx-residency-heartbeat".into())
.spawn(move || loop {
std::thread::sleep(std::time::Duration::from_millis(
RESIDENCY_HEARTBEAT_INTERVAL_MS,
));
let Some(inner) = weak.upgrade() else {
break;
};
let ResidencySetInner::Active {
object,
lock,
heartbeat_ticks,
..
} = &*inner
else {
break;
};
if heartbeat_ticks
.fetch_update(Ordering::AcqRel, Ordering::Acquire, |ticks| {
ticks.checked_sub(1)
})
.is_err()
{
continue;
}
if let Ok(_guard) = lock.lock() {
unsafe {
let _: () = msg_send![object.ptr, requestResidency];
}
};
});
}
pub(crate) fn residency_disabled_by_env() -> bool {
match RESIDENCY_DISABLED_FLAG.load(Ordering::Acquire) {
1 => false,
2 => true,
_ => {
let disabled = std::env::var("HF2Q_NO_RESIDENCY")
.map(|value| value == "1")
.unwrap_or(false);
RESIDENCY_DISABLED_FLAG.store(if disabled { 2 } else { 1 }, Ordering::Release);
disabled
}
}
}
#[doc(hidden)]
pub fn residency_allocation_count_for_test() -> usize {
TEST_ALLOCATION_COUNT.load(Ordering::Acquire)
}
#[doc(hidden)]
pub fn residency_commit_call_count_for_test() -> usize {
TEST_COMMIT_CALL_COUNT.load(Ordering::Acquire)
}
#[doc(hidden)]
pub fn reset_residency_test_counters() {
TEST_ALLOCATION_COUNT.store(0, Ordering::Release);
TEST_COMMIT_CALL_COUNT.store(0, Ordering::Release);
}
#[doc(hidden)]
pub fn reset_residency_env_cache_for_test() {
RESIDENCY_DISABLED_FLAG.store(0, Ordering::Release);
}
#[doc(hidden)]
pub fn macos_15_or_newer_for_test() -> bool {
macos_15_or_newer()
}
#[cfg(target_os = "macos")]
pub(crate) fn macos_15_or_newer() -> bool {
let Some(process_info_class) = Class::get("NSProcessInfo") else {
return false;
};
let version = NSOperatingSystemVersion {
majorVersion: 15,
minorVersion: 0,
patchVersion: 0,
};
unsafe {
let process_info: *mut Object = msg_send![process_info_class, processInfo];
if process_info.is_null() {
return false;
}
let ok: BOOL = msg_send![process_info, isOperatingSystemAtLeastVersion: version];
ok == YES
}
}
#[cfg(not(target_os = "macos"))]
pub(crate) fn macos_15_or_newer() -> bool {
false
}
unsafe fn ns_error_message(error: *mut Object) -> String {
unsafe {
let desc: *mut Object = msg_send![error, localizedDescription];
if desc.is_null() {
return "MTLResidencySet creation failed".into();
}
let text: *const std::os::raw::c_char = msg_send![desc, UTF8String];
if text.is_null() {
return "MTLResidencySet creation failed".into();
}
CStr::from_ptr(text).to_string_lossy().into_owned()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn residency_heartbeat_uses_five_millisecond_ticks() {
assert_eq!(residency_heartbeat_ticks_for(0), 0);
assert_eq!(residency_heartbeat_ticks_for(1), 200);
assert_eq!(residency_heartbeat_ticks_for(180), 36_000);
}
}