#![deny(unsafe_op_in_unsafe_fn)]
use bytemuck::NoUninit;
use objc2::runtime::ProtocolObject;
use objc2_metal::{
MTLAccelerationStructure, MTLBuffer, MTLComputeCommandEncoder, MTLComputePipelineState,
MTLDepthStencilState, MTLRenderCommandEncoder, MTLRenderPipelineState, MTLSamplerState,
MTLTexture,
};
const BUFFER_TABLE_LEN: usize = 31;
const TEXTURE_TABLE_LEN: usize = 128;
const SAMPLER_TABLE_LEN: usize = 16;
const MAX_INLINE_BYTES: usize = 4096;
fn check_buffer_index(index: usize) {
assert!(
index < BUFFER_TABLE_LEN,
"buffer index {index} exceeds the {BUFFER_TABLE_LEN}-entry argument table"
);
}
fn check_texture_index(index: usize) {
assert!(
index < TEXTURE_TABLE_LEN,
"texture index {index} exceeds the {TEXTURE_TABLE_LEN}-entry argument table"
);
}
fn check_sampler_index(index: usize) {
assert!(
index < SAMPLER_TABLE_LEN,
"sampler index {index} exceeds the {SAMPLER_TABLE_LEN}-entry argument table"
);
}
fn check_inline_len(len: usize) {
assert!(
len <= MAX_INLINE_BYTES,
"inline constant of {len} bytes exceeds the {MAX_INLINE_BYTES}-byte limit"
);
}
pub(super) trait RenderEncode {
fn set_pipeline(&self, pso: &ProtocolObject<dyn MTLRenderPipelineState>);
fn set_depth_stencil(&self, state: &ProtocolObject<dyn MTLDepthStencilState>);
fn set_vertex_buffer(
&self,
buffer: &ProtocolObject<dyn MTLBuffer>,
offset: usize,
index: usize,
);
fn set_fragment_buffer(
&self,
buffer: &ProtocolObject<dyn MTLBuffer>,
offset: usize,
index: usize,
);
fn set_vertex_value<T: NoUninit>(&self, value: &T, index: usize);
fn set_fragment_value<T: NoUninit>(&self, value: &T, index: usize);
fn set_fragment_texture(&self, texture: &ProtocolObject<dyn MTLTexture>, index: usize);
fn set_fragment_sampler(&self, sampler: &ProtocolObject<dyn MTLSamplerState>, index: usize);
fn set_fragment_acceleration_structure(
&self,
structure: &ProtocolObject<dyn MTLAccelerationStructure>,
index: usize,
);
}
impl RenderEncode for ProtocolObject<dyn MTLRenderCommandEncoder> {
fn set_pipeline(&self, pso: &ProtocolObject<dyn MTLRenderPipelineState>) {
self.setRenderPipelineState(pso);
}
fn set_depth_stencil(&self, state: &ProtocolObject<dyn MTLDepthStencilState>) {
self.setDepthStencilState(Some(state));
}
fn set_vertex_buffer(
&self,
buffer: &ProtocolObject<dyn MTLBuffer>,
offset: usize,
index: usize,
) {
check_buffer_index(index);
unsafe { self.setVertexBuffer_offset_atIndex(Some(buffer), offset, index) };
}
fn set_fragment_buffer(
&self,
buffer: &ProtocolObject<dyn MTLBuffer>,
offset: usize,
index: usize,
) {
check_buffer_index(index);
unsafe { self.setFragmentBuffer_offset_atIndex(Some(buffer), offset, index) };
}
fn set_vertex_value<T: NoUninit>(&self, value: &T, index: usize) {
check_buffer_index(index);
let len = size_of::<T>();
check_inline_len(len);
unsafe {
self.setVertexBytes_length_atIndex(std::ptr::NonNull::from(value).cast(), len, index);
}
}
fn set_fragment_value<T: NoUninit>(&self, value: &T, index: usize) {
check_buffer_index(index);
let len = size_of::<T>();
check_inline_len(len);
unsafe {
self.setFragmentBytes_length_atIndex(std::ptr::NonNull::from(value).cast(), len, index);
}
}
fn set_fragment_texture(&self, texture: &ProtocolObject<dyn MTLTexture>, index: usize) {
check_texture_index(index);
unsafe { self.setFragmentTexture_atIndex(Some(texture), index) };
}
fn set_fragment_sampler(&self, sampler: &ProtocolObject<dyn MTLSamplerState>, index: usize) {
check_sampler_index(index);
unsafe { self.setFragmentSamplerState_atIndex(Some(sampler), index) };
}
fn set_fragment_acceleration_structure(
&self,
structure: &ProtocolObject<dyn MTLAccelerationStructure>,
index: usize,
) {
check_buffer_index(index);
unsafe { self.setFragmentAccelerationStructure_atBufferIndex(Some(structure), index) };
}
}
pub(super) trait ComputeEncode {
fn set_pipeline(&self, pso: &ProtocolObject<dyn MTLComputePipelineState>);
fn set_buffer(&self, buffer: &ProtocolObject<dyn MTLBuffer>, offset: usize, index: usize);
fn set_value<T: NoUninit>(&self, value: &T, index: usize);
fn set_texture(&self, texture: &ProtocolObject<dyn MTLTexture>, index: usize);
fn set_sampler(&self, sampler: &ProtocolObject<dyn MTLSamplerState>, index: usize);
}
impl ComputeEncode for ProtocolObject<dyn MTLComputeCommandEncoder> {
fn set_pipeline(&self, pso: &ProtocolObject<dyn MTLComputePipelineState>) {
self.setComputePipelineState(pso);
}
fn set_buffer(&self, buffer: &ProtocolObject<dyn MTLBuffer>, offset: usize, index: usize) {
check_buffer_index(index);
unsafe { self.setBuffer_offset_atIndex(Some(buffer), offset, index) };
}
fn set_value<T: NoUninit>(&self, value: &T, index: usize) {
check_buffer_index(index);
let len = size_of::<T>();
check_inline_len(len);
unsafe {
self.setBytes_length_atIndex(std::ptr::NonNull::from(value).cast(), len, index);
}
}
fn set_texture(&self, texture: &ProtocolObject<dyn MTLTexture>, index: usize) {
check_texture_index(index);
unsafe { self.setTexture_atIndex(Some(texture), index) };
}
fn set_sampler(&self, sampler: &ProtocolObject<dyn MTLSamplerState>, index: usize) {
check_sampler_index(index);
unsafe { self.setSamplerState_atIndex(Some(sampler), index) };
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn table_bounds_accept_the_last_slot_and_reject_the_next() {
check_buffer_index(BUFFER_TABLE_LEN - 1);
check_texture_index(TEXTURE_TABLE_LEN - 1);
check_sampler_index(SAMPLER_TABLE_LEN - 1);
check_inline_len(MAX_INLINE_BYTES);
for over in [
std::panic::catch_unwind(|| check_buffer_index(BUFFER_TABLE_LEN)),
std::panic::catch_unwind(|| check_texture_index(TEXTURE_TABLE_LEN)),
std::panic::catch_unwind(|| check_sampler_index(SAMPLER_TABLE_LEN)),
std::panic::catch_unwind(|| check_inline_len(MAX_INLINE_BYTES + 1)),
] {
assert!(over.is_err());
}
}
}