use crate::foundation::{Error, metal_error};
use crate::metal::generated_object_types::{metal, metal4};
use crate::metal::{Buffer, ComputePipelineState, Device, RenderPipelineState, Texture};
use objc2::rc::Retained;
use objc2::runtime::AnyObject;
use objc2::{msg_send, sel};
use objc2_foundation::{NSError, NSString};
use objc2_metal::MTLResourceID;
const MAX_BUFFER_BIND_COUNT: usize = 31;
const MAX_TEXTURE_BIND_COUNT: usize = 128;
const MAX_SAMPLER_BIND_COUNT: usize = 16;
fn require_selector(
object: &AnyObject,
selector: objc2::runtime::Sel,
context: &str,
) -> Result<(), Error> {
let supported: bool = unsafe { msg_send![object, respondsToSelector: selector] };
if supported {
Ok(())
} else {
Err(Error::unsupported(format!("{context} is unavailable")))
}
}
fn resource_id(object: &AnyObject, context: &str) -> Result<MTLResourceID, Error> {
require_selector(object, sel!(gpuResourceID), context)?;
Ok(unsafe { msg_send![object, gpuResourceID] })
}
#[derive(Clone, Debug, Default, PartialEq, Eq)]
pub struct ArgumentTableDescriptor {
label: Option<String>,
max_buffer_bind_count: usize,
max_texture_bind_count: usize,
max_sampler_state_bind_count: usize,
initialize_bindings: bool,
support_attribute_strides: bool,
}
impl ArgumentTableDescriptor {
#[must_use]
pub fn new() -> Self {
Self::default()
}
#[must_use]
pub fn label(&self) -> Option<&str> {
self.label.as_deref()
}
pub fn set_label(&mut self, label: Option<&str>) {
self.label = label.map(str::to_owned);
}
#[must_use]
pub const fn max_buffer_bind_count(&self) -> usize {
self.max_buffer_bind_count
}
pub fn set_max_buffer_bind_count(&mut self, count: usize) -> Result<(), Error> {
if count > MAX_BUFFER_BIND_COUNT {
return Err(Error::invalid_argument(
"Metal 4 argument tables support at most 31 buffer bindings",
));
}
self.max_buffer_bind_count = count;
Ok(())
}
#[must_use]
pub const fn max_texture_bind_count(&self) -> usize {
self.max_texture_bind_count
}
pub fn set_max_texture_bind_count(&mut self, count: usize) -> Result<(), Error> {
if count > MAX_TEXTURE_BIND_COUNT {
return Err(Error::invalid_argument(
"Metal 4 argument tables support at most 128 texture bindings",
));
}
self.max_texture_bind_count = count;
Ok(())
}
#[must_use]
pub const fn max_sampler_state_bind_count(&self) -> usize {
self.max_sampler_state_bind_count
}
pub fn set_max_sampler_state_bind_count(&mut self, count: usize) -> Result<(), Error> {
if count > MAX_SAMPLER_BIND_COUNT {
return Err(Error::invalid_argument(
"Metal 4 argument tables support at most 16 sampler bindings",
));
}
self.max_sampler_state_bind_count = count;
Ok(())
}
#[must_use]
pub const fn initialize_bindings(&self) -> bool {
self.initialize_bindings
}
pub fn set_initialize_bindings(&mut self, initialize: bool) {
self.initialize_bindings = initialize;
}
#[must_use]
pub const fn support_attribute_strides(&self) -> bool {
self.support_attribute_strides
}
pub fn set_support_attribute_strides(&mut self, support: bool) {
self.support_attribute_strides = support;
}
fn make_native(&self) -> Result<metal4::ArgumentTableDescriptor, Error> {
let descriptor = metal4::ArgumentTableDescriptor::new()?;
descriptor.set_max_buffer_bind_count(self.max_buffer_bind_count)?;
descriptor.set_max_texture_bind_count(self.max_texture_bind_count)?;
descriptor.set_max_sampler_state_bind_count(self.max_sampler_state_bind_count)?;
descriptor.set_initialize_bindings(self.initialize_bindings)?;
descriptor.set_support_attribute_strides(self.support_attribute_strides)?;
if let Some(label) = &self.label {
descriptor.set_label(label)?;
}
Ok(descriptor)
}
}
#[derive(Clone)]
pub struct ArgumentTable {
inner: metal4::ArgumentTable,
device: Device,
label: Option<String>,
max_buffer_bind_count: usize,
max_texture_bind_count: usize,
max_sampler_state_bind_count: usize,
support_attribute_strides: bool,
}
#[derive(Clone, Copy)]
pub enum ArgumentTableResource<'a> {
Buffer(&'a Buffer),
Tensor(&'a metal::Tensor),
AccelerationStructure(&'a metal::AccelerationStructure),
}
impl<'a> ArgumentTableResource<'a> {
fn object(self) -> (&'a AnyObject, &'static str) {
match self {
Self::Buffer(value) => (value.as_any_object(), "MTLBuffer.gpuResourceID"),
Self::Tensor(value) => (value.as_inner(), "MTLTensor.gpuResourceID"),
Self::AccelerationStructure(value) => {
(value.as_inner(), "MTLAccelerationStructure.gpuResourceID")
}
}
}
}
impl Device {
pub fn new_mtl4_argument_table(
&self,
descriptor: &ArgumentTableDescriptor,
) -> Result<ArgumentTable, Error> {
require_selector(
self.as_any_object(),
sel!(newArgumentTableWithDescriptor:error:),
"MTLDevice.newArgumentTableWithDescriptor:error:",
)?;
let native = descriptor.make_native()?;
let result: Result<Retained<AnyObject>, Retained<NSError>> = unsafe {
msg_send![self.as_any_object(), newArgumentTableWithDescriptor: native.as_inner(), error: _]
};
let inner = result
.map(metal4::ArgumentTable::from_inner)
.map_err(|error| metal_error(&error))?;
Ok(ArgumentTable {
inner,
device: self.clone(),
label: descriptor.label.clone(),
max_buffer_bind_count: descriptor.max_buffer_bind_count,
max_texture_bind_count: descriptor.max_texture_bind_count,
max_sampler_state_bind_count: descriptor.max_sampler_state_bind_count,
support_attribute_strides: descriptor.support_attribute_strides,
})
}
}
impl ArgumentTable {
#[allow(dead_code)]
pub(crate) const fn as_generated(&self) -> &metal4::ArgumentTable {
&self.inner
}
#[must_use]
pub fn device(&self) -> Device {
self.device.clone()
}
#[must_use]
pub fn label(&self) -> Option<&str> {
self.label.as_deref()
}
fn ensure_same_device(&self, resource: &AnyObject, context: &str) -> Result<(), Error> {
require_selector(resource, sel!(device), context)?;
let resource_device: Retained<AnyObject> = unsafe { msg_send![resource, device] };
if std::ptr::eq(&*resource_device, self.device.as_any_object()) {
Ok(())
} else {
Err(Error::invalid_argument(
"argument-table resources must belong to the table's device",
))
}
}
fn check_buffer_index(&self, index: usize) -> Result<(), Error> {
if index >= self.max_buffer_bind_count {
Err(Error::invalid_argument(
"buffer binding index exceeds the table descriptor limit",
))
} else {
Ok(())
}
}
fn check_texture_index(&self, index: usize) -> Result<(), Error> {
if index >= self.max_texture_bind_count {
Err(Error::invalid_argument(
"texture binding index exceeds the table descriptor limit",
))
} else {
Ok(())
}
}
fn check_sampler_index(&self, index: usize) -> Result<(), Error> {
if index >= self.max_sampler_state_bind_count {
Err(Error::invalid_argument(
"sampler binding index exceeds the table descriptor limit",
))
} else {
Ok(())
}
}
pub fn set_buffer(&self, buffer: &Buffer, offset: usize, index: usize) -> Result<(), Error> {
self.check_buffer_index(index)?;
self.ensure_same_device(buffer.as_any_object(), "MTLBuffer.device")?;
if offset > buffer.length() {
return Err(Error::invalid_argument(
"buffer binding offset is out of bounds",
));
}
let address = buffer
.gpu_address()
.checked_add(offset as u64)
.ok_or_else(|| Error::invalid_argument("buffer GPU address overflow"))?;
require_selector(
self.inner.as_inner(),
sel!(setAddress:atIndex:),
"MTL4ArgumentTable.setAddress:atIndex:",
)?;
unsafe {
let _: () = msg_send![self.inner.as_inner(), setAddress: address, atIndex: index];
}
Ok(())
}
pub fn set_buffer_with_stride(
&self,
buffer: &Buffer,
offset: usize,
stride: usize,
index: usize,
) -> Result<(), Error> {
self.check_buffer_index(index)?;
self.ensure_same_device(buffer.as_any_object(), "MTLBuffer.device")?;
if !self.support_attribute_strides {
return Err(Error::invalid_argument(
"the table descriptor did not enable attribute strides",
));
}
let remaining = buffer
.length()
.checked_sub(offset)
.ok_or_else(|| Error::invalid_argument("buffer binding offset is out of bounds"))?;
if stride > remaining {
return Err(Error::invalid_argument(
"attribute stride exceeds the remaining buffer length",
));
}
let address = buffer
.gpu_address()
.checked_add(offset as u64)
.ok_or_else(|| Error::invalid_argument("buffer GPU address overflow"))?;
require_selector(
self.inner.as_inner(),
sel!(setAddress:attributeStride:atIndex:),
"MTL4ArgumentTable.setAddress:attributeStride:atIndex:",
)?;
unsafe {
let _: () = msg_send![self.inner.as_inner(), setAddress: address, attributeStride: stride, atIndex: index];
}
Ok(())
}
pub fn set_buffers(
&self,
bindings: &[(&Buffer, usize)],
start_index: usize,
) -> Result<(), Error> {
let end = start_index
.checked_add(bindings.len())
.ok_or_else(|| Error::invalid_argument("buffer binding range overflow"))?;
if end > self.max_buffer_bind_count {
return Err(Error::invalid_argument(
"buffer binding slice exceeds the table descriptor limit",
));
}
for (relative, &(buffer, offset)) in bindings.iter().enumerate() {
self.set_buffer(buffer, offset, start_index + relative)?;
}
Ok(())
}
fn set_buffer_resource_object(
&self,
object: &AnyObject,
index: usize,
context: &str,
) -> Result<(), Error> {
self.check_buffer_index(index)?;
self.ensure_same_device(object, context)?;
let id = resource_id(object, context)?;
require_selector(
self.inner.as_inner(),
sel!(setResource:atBufferIndex:),
"MTL4ArgumentTable.setResource:atBufferIndex:",
)?;
unsafe {
let _: () = msg_send![self.inner.as_inner(), setResource: id, atBufferIndex: index];
}
Ok(())
}
pub fn set_resource(
&self,
resource: ArgumentTableResource<'_>,
index: usize,
) -> Result<(), Error> {
let (object, context) = resource.object();
self.set_buffer_resource_object(object, index, context)
}
pub fn set_tensor(&self, tensor: &metal::Tensor, index: usize) -> Result<(), Error> {
self.set_resource(ArgumentTableResource::Tensor(tensor), index)
}
pub fn set_acceleration_structure(
&self,
structure: &metal::AccelerationStructure,
index: usize,
) -> Result<(), Error> {
self.set_resource(
ArgumentTableResource::AccelerationStructure(structure),
index,
)
}
pub fn set_texture(&self, texture: &Texture, index: usize) -> Result<(), Error> {
self.check_texture_index(index)?;
self.ensure_same_device(texture.as_any_object(), "MTLTexture.device")?;
require_selector(
self.inner.as_inner(),
sel!(setTexture:atIndex:),
"MTL4ArgumentTable.setTexture:atIndex:",
)?;
let raw = texture.gpu_resource_id()._impl;
let id = unsafe { std::mem::transmute::<u64, MTLResourceID>(raw) };
unsafe {
let _: () = msg_send![self.inner.as_inner(), setTexture: id, atIndex: index];
}
Ok(())
}
pub fn set_textures(&self, textures: &[&Texture], start_index: usize) -> Result<(), Error> {
let end = start_index
.checked_add(textures.len())
.ok_or_else(|| Error::invalid_argument("texture binding range overflow"))?;
if end > self.max_texture_bind_count {
return Err(Error::invalid_argument(
"texture binding slice exceeds the table descriptor limit",
));
}
for (relative, texture) in textures.iter().enumerate() {
self.set_texture(texture, start_index + relative)?;
}
Ok(())
}
pub fn set_sampler_state(
&self,
sampler: &metal::SamplerState,
index: usize,
) -> Result<(), Error> {
self.check_sampler_index(index)?;
self.ensure_same_device(sampler.as_inner(), "MTLSamplerState.device")?;
let id = resource_id(sampler.as_inner(), "MTLSamplerState.gpuResourceID")?;
require_selector(
self.inner.as_inner(),
sel!(setSamplerState:atIndex:),
"MTL4ArgumentTable.setSamplerState:atIndex:",
)?;
unsafe {
let _: () = msg_send![self.inner.as_inner(), setSamplerState: id, atIndex: index];
}
Ok(())
}
pub fn set_sampler_states(
&self,
samplers: &[&metal::SamplerState],
start_index: usize,
) -> Result<(), Error> {
let end = start_index
.checked_add(samplers.len())
.ok_or_else(|| Error::invalid_argument("sampler binding range overflow"))?;
if end > self.max_sampler_state_bind_count {
return Err(Error::invalid_argument(
"sampler binding slice exceeds the table descriptor limit",
));
}
for (relative, sampler) in samplers.iter().enumerate() {
self.set_sampler_state(sampler, start_index + relative)?;
}
Ok(())
}
pub fn set_tensors(&self, tensors: &[&metal::Tensor], start_index: usize) -> Result<(), Error> {
let end = start_index
.checked_add(tensors.len())
.ok_or_else(|| Error::invalid_argument("tensor binding range overflow"))?;
if end > self.max_buffer_bind_count {
return Err(Error::invalid_argument(
"tensor binding slice exceeds the table descriptor limit",
));
}
for (relative, tensor) in tensors.iter().enumerate() {
self.set_tensor(tensor, start_index + relative)?;
}
Ok(())
}
pub fn set_acceleration_structures(
&self,
structures: &[&metal::AccelerationStructure],
start_index: usize,
) -> Result<(), Error> {
let end = start_index
.checked_add(structures.len())
.ok_or_else(|| Error::invalid_argument("acceleration-structure range overflow"))?;
if end > self.max_buffer_bind_count {
return Err(Error::invalid_argument(
"acceleration-structure slice exceeds the table descriptor limit",
));
}
for (relative, structure) in structures.iter().enumerate() {
self.set_acceleration_structure(structure, start_index + relative)?;
}
Ok(())
}
}
#[derive(Clone)]
pub struct Archive {
inner: metal4::Archive,
}
impl Archive {
#[allow(dead_code)]
pub(crate) const fn from_generated(inner: metal4::Archive) -> Self {
Self { inner }
}
pub fn label(&self) -> Result<Option<String>, Error> {
self.inner.label()
}
pub fn set_label(&self, label: Option<&str>) -> Result<(), Error> {
require_selector(
self.inner.as_inner(),
sel!(setLabel:),
"MTL4Archive.setLabel:",
)?;
let label = label.map(NSString::from_str);
unsafe {
let _: () = msg_send![self.inner.as_inner(), setLabel: label.as_deref()];
}
Ok(())
}
pub fn new_binary_function(
&self,
descriptor: &metal4::BinaryFunctionDescriptor,
) -> Result<metal4::BinaryFunction, Error> {
require_selector(
self.inner.as_inner(),
sel!(newBinaryFunctionWithDescriptor:error:),
"MTL4Archive.newBinaryFunctionWithDescriptor:error:",
)?;
let result: Result<Retained<AnyObject>, Retained<NSError>> = unsafe {
msg_send![self.inner.as_inner(), newBinaryFunctionWithDescriptor: descriptor.as_inner(), error: _]
};
result
.map(metal4::BinaryFunction::from_inner)
.map_err(|error| metal_error(&error))
}
pub fn new_compute_pipeline_state(
&self,
descriptor: &metal4::ComputePipelineDescriptor,
linking: Option<&metal4::PipelineStageDynamicLinkingDescriptor>,
) -> Result<ComputePipelineState, Error> {
let result: Result<Retained<AnyObject>, Retained<NSError>> = if let Some(linking) = linking
{
require_selector(
self.inner.as_inner(),
sel!(newComputePipelineStateWithDescriptor:dynamicLinkingDescriptor:error:),
"MTL4Archive.newComputePipelineStateWithDescriptor:dynamicLinkingDescriptor:error:",
)?;
unsafe {
msg_send![self.inner.as_inner(), newComputePipelineStateWithDescriptor: descriptor.as_inner(), dynamicLinkingDescriptor: linking.as_inner(), error: _]
}
} else {
require_selector(
self.inner.as_inner(),
sel!(newComputePipelineStateWithDescriptor:error:),
"MTL4Archive.newComputePipelineStateWithDescriptor:error:",
)?;
unsafe {
msg_send![self.inner.as_inner(), newComputePipelineStateWithDescriptor: descriptor.as_inner(), error: _]
}
};
let inner = result.map_err(|error| metal_error(&error))?;
let inner = unsafe { Retained::cast_unchecked(inner) };
Ok(ComputePipelineState::new(inner))
}
pub fn new_render_pipeline_state(
&self,
descriptor: &metal4::PipelineDescriptor,
linking: Option<&metal4::RenderPipelineDynamicLinkingDescriptor>,
) -> Result<RenderPipelineState, Error> {
let result: Result<Retained<AnyObject>, Retained<NSError>> = if let Some(linking) = linking
{
require_selector(
self.inner.as_inner(),
sel!(newRenderPipelineStateWithDescriptor:dynamicLinkingDescriptor:error:),
"MTL4Archive.newRenderPipelineStateWithDescriptor:dynamicLinkingDescriptor:error:",
)?;
unsafe {
msg_send![self.inner.as_inner(), newRenderPipelineStateWithDescriptor: descriptor.as_inner(), dynamicLinkingDescriptor: linking.as_inner(), error: _]
}
} else {
require_selector(
self.inner.as_inner(),
sel!(newRenderPipelineStateWithDescriptor:error:),
"MTL4Archive.newRenderPipelineStateWithDescriptor:error:",
)?;
unsafe {
msg_send![self.inner.as_inner(), newRenderPipelineStateWithDescriptor: descriptor.as_inner(), error: _]
}
};
let inner = result.map_err(|error| metal_error(&error))?;
let inner = unsafe { Retained::cast_unchecked(inner) };
Ok(RenderPipelineState::new(inner))
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn descriptor_accepts_documented_binding_limits() {
let mut descriptor = ArgumentTableDescriptor::new();
descriptor
.set_max_buffer_bind_count(MAX_BUFFER_BIND_COUNT)
.unwrap();
descriptor
.set_max_texture_bind_count(MAX_TEXTURE_BIND_COUNT)
.unwrap();
descriptor
.set_max_sampler_state_bind_count(MAX_SAMPLER_BIND_COUNT)
.unwrap();
assert_eq!(descriptor.max_buffer_bind_count(), 31);
assert_eq!(descriptor.max_texture_bind_count(), 128);
assert_eq!(descriptor.max_sampler_state_bind_count(), 16);
}
#[test]
fn descriptor_rejects_counts_above_documented_limits() {
let mut descriptor = ArgumentTableDescriptor::new();
assert!(descriptor.set_max_buffer_bind_count(32).is_err());
assert!(descriptor.set_max_texture_bind_count(129).is_err());
assert!(descriptor.set_max_sampler_state_bind_count(17).is_err());
}
#[test]
fn descriptor_owns_label_and_boolean_state() {
let mut descriptor = ArgumentTableDescriptor::new();
descriptor.set_label(Some("bindings"));
descriptor.set_initialize_bindings(true);
descriptor.set_support_attribute_strides(true);
assert_eq!(descriptor.label(), Some("bindings"));
assert!(descriptor.initialize_bindings());
assert!(descriptor.support_attribute_strides());
descriptor.set_label(None);
assert_eq!(descriptor.label(), None);
}
}