use crate::foundation::Error;
use crate::metal::generated_object_types::metal::{
AccelerationStructure, AccelerationStructureBoundingBoxGeometryDescriptor,
AccelerationStructureCommandEncoder as GeneratedEncoder,
AccelerationStructureCurveGeometryDescriptor, AccelerationStructureDescriptor,
AccelerationStructureGeometryDescriptor,
AccelerationStructureMotionBoundingBoxGeometryDescriptor,
AccelerationStructureMotionCurveGeometryDescriptor,
AccelerationStructureMotionTriangleGeometryDescriptor,
AccelerationStructurePassSampleBufferAttachmentDescriptor,
AccelerationStructurePassSampleBufferAttachmentDescriptorArray,
AccelerationStructureTriangleGeometryDescriptor, CounterSampleBuffer, Fence, Heap,
IndirectInstanceAccelerationStructureDescriptor, InstanceAccelerationStructureDescriptor,
PrimitiveAccelerationStructureDescriptor,
};
use crate::metal::generated_struct_types::ResourceID;
use crate::metal::generated_value_types::{
AccelerationStructureRefitOptions, DataType, ResourceUsage,
};
use crate::metal::{Buffer, CommandBuffer, Texture};
use objc2::rc::Retained;
use objc2::runtime::{AnyClass, AnyObject, MessageReceiver, NSObjectProtocol, ProtocolObject, Sel};
use objc2::{msg_send, sel};
use objc2_metal::{MTLAccelerationStructureCommandEncoder, MTLCommandBuffer as _, MTLResourceID};
use std::marker::PhantomData;
const SCRATCH_OFFSET_ALIGNMENT: usize = 256;
pub const ACCELERATION_STRUCTURE_SAMPLE_ATTACHMENT_CAPACITY: usize = 4;
pub struct AccelerationStructureEncoder<'a> {
inner: GeneratedEncoder,
_command_buffer: PhantomData<&'a mut CommandBuffer>,
ended: bool,
}
impl<'a> AccelerationStructureEncoder<'a> {
pub(super) fn from_protocol(
inner: Retained<ProtocolObject<dyn MTLAccelerationStructureCommandEncoder>>,
_command_buffer: &'a mut CommandBuffer,
) -> Self {
let inner = GeneratedEncoder::from_inner(unsafe { Retained::cast_unchecked(inner) });
Self {
inner,
_command_buffer: PhantomData,
ended: false,
}
}
fn as_inner(&self) -> &AnyObject {
self.inner.as_inner()
}
}
impl CommandBuffer {
pub fn acceleration_structure_encoder(
&mut self,
) -> Result<AccelerationStructureEncoder<'_>, Error> {
if !self
.inner
.respondsToSelector(sel!(accelerationStructureCommandEncoder))
{
return Err(Error::unsupported(
"MTLCommandBuffer::accelerationStructureCommandEncoder is unavailable",
));
}
self.inner
.accelerationStructureCommandEncoder()
.map(|inner| AccelerationStructureEncoder::from_protocol(inner, self))
.ok_or_else(|| {
Error::unsupported("Metal could not create an acceleration-structure encoder")
})
}
pub fn acceleration_structure_encoder_with_descriptor(
&mut self,
descriptor: &crate::metal::generated_object_types::metal::AccelerationStructurePassDescriptor,
) -> Result<AccelerationStructureEncoder<'_>, Error> {
if !self
.inner
.respondsToSelector(sel!(accelerationStructureCommandEncoderWithDescriptor:))
{
return Err(Error::unsupported(
"MTLCommandBuffer::accelerationStructureCommandEncoderWithDescriptor is unavailable",
));
}
let inner: Option<Retained<ProtocolObject<dyn MTLAccelerationStructureCommandEncoder>>> = unsafe {
msg_send![
&*self.inner,
accelerationStructureCommandEncoderWithDescriptor: descriptor.as_inner()
]
};
inner
.map(|inner| AccelerationStructureEncoder::from_protocol(inner, self))
.ok_or_else(|| {
Error::unsupported("Metal could not create an acceleration-structure encoder")
})
}
}
#[derive(Clone, Copy)]
pub enum AccelerationStructureBuildDescriptor<'a> {
Base(&'a AccelerationStructureDescriptor),
Primitive(&'a PrimitiveAccelerationStructureDescriptor),
Instance(&'a InstanceAccelerationStructureDescriptor),
IndirectInstance(&'a IndirectInstanceAccelerationStructureDescriptor),
}
impl<'a> AccelerationStructureBuildDescriptor<'a> {
fn as_object(self) -> &'a AnyObject {
match self {
Self::Base(value) => value.as_inner(),
Self::Primitive(value) => value.as_inner(),
Self::Instance(value) => value.as_inner(),
Self::IndirectInstance(value) => value.as_inner(),
}
}
}
#[derive(Clone, Copy)]
pub enum AccelerationStructureResource<'a> {
Buffer(&'a Buffer),
Texture(&'a Texture),
AccelerationStructure(&'a AccelerationStructure),
}
#[derive(Clone, Copy)]
pub enum AccelerationStructureGeometry<'a> {
Triangle(&'a AccelerationStructureTriangleGeometryDescriptor),
BoundingBox(&'a AccelerationStructureBoundingBoxGeometryDescriptor),
Curve(&'a AccelerationStructureCurveGeometryDescriptor),
MotionTriangle(&'a AccelerationStructureMotionTriangleGeometryDescriptor),
MotionBoundingBox(&'a AccelerationStructureMotionBoundingBoxGeometryDescriptor),
MotionCurve(&'a AccelerationStructureMotionCurveGeometryDescriptor),
}
impl<'a> AccelerationStructureGeometry<'a> {
fn as_object(self) -> &'a AnyObject {
match self {
Self::Triangle(value) => value.as_inner(),
Self::BoundingBox(value) => value.as_inner(),
Self::Curve(value) => value.as_inner(),
Self::MotionTriangle(value) => value.as_inner(),
Self::MotionBoundingBox(value) => value.as_inner(),
Self::MotionCurve(value) => value.as_inner(),
}
}
}
impl<'a> AccelerationStructureResource<'a> {
fn as_object(self) -> &'a AnyObject {
match self {
Self::Buffer(value) => value.as_any_object(),
Self::Texture(value) => value.as_any_object(),
Self::AccelerationStructure(value) => value.as_inner(),
}
}
}
fn supports(object: &AnyObject, selector: Sel) -> bool {
unsafe { msg_send![object, respondsToSelector: selector] }
}
fn require_selector(object: &AnyObject, selector: Sel, operation: &str) -> Result<(), Error> {
if supports(object, selector) {
Ok(())
} else {
Err(Error::unsupported(format!(
"MTL::AccelerationStructureCommandEncoder::{operation} is unavailable"
)))
}
}
fn read_object_array(
owner: &AnyObject,
getter: Sel,
operation: &str,
) -> Result<Vec<Retained<AnyObject>>, Error> {
require_selector(owner, getter, operation)?;
let array = unsafe {
Retained::retain_autoreleased(owner.send_message::<_, *mut AnyObject>(getter, ()))
};
let Some(array) = array else {
return Ok(Vec::new());
};
let count: usize = unsafe { msg_send![&*array, count] };
let mut values = Vec::with_capacity(count);
for index in 0..count {
let value: Retained<AnyObject> = unsafe { msg_send![&*array, objectAtIndex: index] };
values.push(value);
}
Ok(values)
}
fn write_object_array(
owner: &AnyObject,
setter: Sel,
operation: &str,
values: &[&AnyObject],
) -> Result<(), Error> {
require_selector(owner, setter, operation)?;
let class = AnyClass::get(c"NSMutableArray")
.ok_or_else(|| Error::unsupported("NSMutableArray is unavailable on this system"))?;
let array: Retained<AnyObject> = unsafe { msg_send![class, new] };
for value in values {
unsafe {
let _: () = msg_send![&*array, addObject: *value];
}
}
unsafe {
let _: () = owner.send_message(setter, (&*array,));
}
Ok(())
}
fn checked_buffer_range(
buffer: &Buffer,
offset: usize,
length: usize,
operation: &str,
) -> Result<(), Error> {
if length == 0 {
return Err(Error::invalid_argument(format!(
"{operation} requires a non-zero buffer range"
)));
}
let end = offset.checked_add(length).ok_or_else(|| {
Error::invalid_argument(format!("{operation} buffer range overflows usize"))
})?;
if end > buffer.length() {
return Err(Error::invalid_argument(format!(
"{operation} buffer range is out of bounds"
)));
}
Ok(())
}
fn checked_attachment_index(index: usize) -> Result<(), Error> {
if index >= ACCELERATION_STRUCTURE_SAMPLE_ATTACHMENT_CAPACITY {
Err(Error::invalid_argument(
"acceleration-pass sample attachment index must be below 4",
))
} else {
Ok(())
}
}
fn validate_scratch(
scratch: &Buffer,
offset: usize,
required_size: usize,
operation: &str,
) -> Result<(), Error> {
if !offset.is_multiple_of(SCRATCH_OFFSET_ALIGNMENT) {
return Err(Error::invalid_argument(format!(
"{operation} scratch offset must be 256-byte aligned"
)));
}
checked_buffer_range(scratch, offset, required_size, operation)
}
impl AccelerationStructure {
pub fn gpu_resource_id(&self) -> Result<ResourceID, Error> {
require_selector(self.as_inner(), sel!(gpuResourceID), "gpuResourceID")?;
let raw: MTLResourceID = unsafe { msg_send![self.as_inner(), gpuResourceID] };
let value = unsafe { std::mem::transmute::<MTLResourceID, u64>(raw) };
Ok(ResourceID { _impl: value })
}
}
impl AccelerationStructurePassSampleBufferAttachmentDescriptorArray {
pub fn attachment(
&self,
index: usize,
) -> Result<AccelerationStructurePassSampleBufferAttachmentDescriptor, Error> {
checked_attachment_index(index)?;
require_selector(
self.as_inner(),
sel!(objectAtIndexedSubscript:),
"objectAtIndexedSubscript",
)?;
let value: Retained<AnyObject> =
unsafe { msg_send![self.as_inner(), objectAtIndexedSubscript: index] };
Ok(AccelerationStructurePassSampleBufferAttachmentDescriptor::from_inner(value))
}
pub fn set_attachment(
&self,
index: usize,
attachment: Option<&AccelerationStructurePassSampleBufferAttachmentDescriptor>,
) -> Result<(), Error> {
checked_attachment_index(index)?;
if let Some(attachment) = attachment
&& let Some(sample_buffer) = attachment.sample_buffer()?
{
let sample_count = sample_buffer.sample_count()?;
for sample_index in [
attachment.start_of_encoder_sample_index()?,
attachment.end_of_encoder_sample_index()?,
] {
if sample_index != usize::MAX && sample_index >= sample_count {
return Err(Error::invalid_argument(
"acceleration-pass counter sample index is out of bounds",
));
}
}
}
require_selector(
self.as_inner(),
sel!(setObject:atIndexedSubscript:),
"setObjectAtIndexedSubscript",
)?;
unsafe {
let _: () = msg_send![
self.as_inner(),
setObject: attachment.map(
AccelerationStructurePassSampleBufferAttachmentDescriptor::as_inner
),
atIndexedSubscript: index
];
}
Ok(())
}
}
impl PrimitiveAccelerationStructureDescriptor {
pub fn geometry_descriptors_vec(
&self,
) -> Result<Vec<AccelerationStructureGeometryDescriptor>, Error> {
read_object_array(
self.as_inner(),
sel!(geometryDescriptors),
"geometryDescriptors",
)
.map(|values| {
values
.into_iter()
.map(AccelerationStructureGeometryDescriptor::from_inner)
.collect()
})
}
pub fn set_geometry_descriptors_slice(
&self,
values: &[AccelerationStructureGeometry<'_>],
) -> Result<(), Error> {
let values = values
.iter()
.copied()
.map(AccelerationStructureGeometry::as_object)
.collect::<Vec<_>>();
write_object_array(
self.as_inner(),
sel!(setGeometryDescriptors:),
"setGeometryDescriptors",
&values,
)
}
}
macro_rules! buffer_array_property {
($owner:ty, $getter_fn:ident, $setter_fn:ident, $getter:ident, $setter:ident) => {
impl $owner {
#[doc = concat!("Returns `", stringify!($getter), "` as owned safe buffers.")]
pub fn $getter_fn(&self) -> Result<Vec<Buffer>, Error> {
read_object_array(self.as_inner(), sel!($getter), stringify!($getter))?
.into_iter()
.map(Buffer::from_any_object)
.collect()
}
#[doc = concat!("Writes `", stringify!($setter), "` from a borrowed safe slice.")]
pub fn $setter_fn(&self, values: &[&Buffer]) -> Result<(), Error> {
let values = values
.iter()
.map(|value| value.as_any_object())
.collect::<Vec<_>>();
write_object_array(
self.as_inner(),
sel!($setter:),
stringify!($setter),
&values,
)
}
}
};
}
buffer_array_property!(
AccelerationStructureMotionTriangleGeometryDescriptor,
vertex_buffers_vec,
set_vertex_buffers_slice,
vertexBuffers,
setVertexBuffers
);
buffer_array_property!(
AccelerationStructureMotionBoundingBoxGeometryDescriptor,
bounding_box_buffers_vec,
set_bounding_box_buffers_slice,
boundingBoxBuffers,
setBoundingBoxBuffers
);
buffer_array_property!(
AccelerationStructureMotionCurveGeometryDescriptor,
control_point_buffers_vec,
set_control_point_buffers_slice,
controlPointBuffers,
setControlPointBuffers
);
buffer_array_property!(
AccelerationStructureMotionCurveGeometryDescriptor,
radius_buffers_vec,
set_radius_buffers_slice,
radiusBuffers,
setRadiusBuffers
);
impl InstanceAccelerationStructureDescriptor {
pub fn instanced_acceleration_structures_vec(
&self,
) -> Result<Vec<AccelerationStructure>, Error> {
read_object_array(
self.as_inner(),
sel!(instancedAccelerationStructures),
"instancedAccelerationStructures",
)
.map(|values| {
values
.into_iter()
.map(AccelerationStructure::from_inner)
.collect()
})
}
pub fn set_instanced_acceleration_structures_slice(
&self,
values: &[&AccelerationStructure],
) -> Result<(), Error> {
let values = values
.iter()
.map(|value| value.as_inner())
.collect::<Vec<_>>();
write_object_array(
self.as_inner(),
sel!(setInstancedAccelerationStructures:),
"setInstancedAccelerationStructures",
&values,
)
}
}
impl AccelerationStructureEncoder<'_> {
pub fn build(
&self,
destination: &AccelerationStructure,
descriptor: AccelerationStructureBuildDescriptor<'_>,
required_acceleration_structure_size: usize,
scratch: &Buffer,
scratch_offset: usize,
required_scratch_size: usize,
) -> Result<(), Error> {
let encoder = self.as_inner();
require_selector(
encoder,
sel!(buildAccelerationStructure:descriptor:scratchBuffer:scratchBufferOffset:),
"buildAccelerationStructure",
)?;
if required_acceleration_structure_size == 0
|| destination.size()? < required_acceleration_structure_size
{
return Err(Error::invalid_argument(
"build destination is smaller than the required acceleration-structure size",
));
}
validate_scratch(
scratch,
scratch_offset,
required_scratch_size,
"acceleration-structure build",
)?;
unsafe {
let _: () = msg_send![
encoder,
buildAccelerationStructure: destination.as_inner(),
descriptor: descriptor.as_object(),
scratchBuffer: scratch.as_any_object(),
scratchBufferOffset: scratch_offset
];
}
Ok(())
}
#[allow(clippy::too_many_arguments)]
pub fn refit(
&self,
source: &AccelerationStructure,
descriptor: AccelerationStructureBuildDescriptor<'_>,
destination: Option<&AccelerationStructure>,
required_destination_size: usize,
scratch: &Buffer,
scratch_offset: usize,
required_scratch_size: usize,
options: AccelerationStructureRefitOptions,
) -> Result<(), Error> {
let encoder = self.as_inner();
require_selector(
encoder,
sel!(refitAccelerationStructure:descriptor:destination:scratchBuffer:scratchBufferOffset:options:),
"refitAccelerationStructure",
)?;
if !options.is_valid() {
return Err(Error::invalid_argument(
"refit options contain undeclared bits",
));
}
if required_destination_size == 0 {
return Err(Error::invalid_argument(
"refit requires a non-zero destination size",
));
}
let target = destination.unwrap_or(source);
if target.size()? < required_destination_size {
return Err(Error::invalid_argument(
"refit destination is smaller than the required size",
));
}
validate_scratch(
scratch,
scratch_offset,
required_scratch_size,
"acceleration-structure refit",
)?;
unsafe {
let _: () = msg_send![
encoder,
refitAccelerationStructure: source.as_inner(),
descriptor: descriptor.as_object(),
destination: destination.map(AccelerationStructure::as_inner),
scratchBuffer: scratch.as_any_object(),
scratchBufferOffset: scratch_offset,
options: options.as_raw()
];
}
Ok(())
}
pub fn copy(
&self,
source: &AccelerationStructure,
destination: &AccelerationStructure,
) -> Result<(), Error> {
let encoder = self.as_inner();
require_selector(
encoder,
sel!(copyAccelerationStructure:toAccelerationStructure:),
"copyAccelerationStructure",
)?;
if std::ptr::eq(source.as_inner(), destination.as_inner()) {
return Err(Error::invalid_argument(
"copy source and destination must be distinct",
));
}
if destination.size()? < source.size()? {
return Err(Error::invalid_argument(
"copy destination is smaller than the source",
));
}
unsafe {
let _: () = msg_send![
encoder,
copyAccelerationStructure: source.as_inner(),
toAccelerationStructure: destination.as_inner()
];
}
Ok(())
}
pub fn write_compacted_size_u32(
&self,
source: &AccelerationStructure,
destination: &Buffer,
offset: usize,
) -> Result<(), Error> {
self.write_compacted_size(source, destination, offset, 4, false)
}
pub fn write_compacted_size_u64(
&self,
source: &AccelerationStructure,
destination: &Buffer,
offset: usize,
) -> Result<(), Error> {
self.write_compacted_size(source, destination, offset, 8, true)
}
fn write_compacted_size(
&self,
source: &AccelerationStructure,
destination: &Buffer,
offset: usize,
width: usize,
wide: bool,
) -> Result<(), Error> {
if !offset.is_multiple_of(width) {
return Err(Error::invalid_argument(
"compacted-size destination offset is not naturally aligned",
));
}
checked_buffer_range(destination, offset, width, "compacted-size write")?;
let encoder = self.as_inner();
if wide {
require_selector(
encoder,
sel!(writeCompactedAccelerationStructureSize:toBuffer:offset:sizeDataType:),
"writeCompactedAccelerationStructureSize:sizeDataType",
)?;
let ulong_data_type = DataType::DataTypeULong.as_raw();
unsafe {
let _: () = msg_send![
encoder,
writeCompactedAccelerationStructureSize: source.as_inner(),
toBuffer: destination.as_any_object(),
offset: offset,
sizeDataType: ulong_data_type
];
}
} else {
require_selector(
encoder,
sel!(writeCompactedAccelerationStructureSize:toBuffer:offset:),
"writeCompactedAccelerationStructureSize",
)?;
unsafe {
let _: () = msg_send![
encoder,
writeCompactedAccelerationStructureSize: source.as_inner(),
toBuffer: destination.as_any_object(),
offset: offset
];
}
}
Ok(())
}
pub fn copy_and_compact(
&self,
source: &AccelerationStructure,
destination: &AccelerationStructure,
compacted_size: usize,
) -> Result<(), Error> {
let encoder = self.as_inner();
require_selector(
encoder,
sel!(copyAndCompactAccelerationStructure:toAccelerationStructure:),
"copyAndCompactAccelerationStructure",
)?;
if std::ptr::eq(source.as_inner(), destination.as_inner()) {
return Err(Error::invalid_argument(
"compaction source and destination must be distinct",
));
}
if compacted_size == 0 || destination.size()? < compacted_size {
return Err(Error::invalid_argument(
"compaction destination is smaller than the compacted size",
));
}
unsafe {
let _: () = msg_send![
encoder,
copyAndCompactAccelerationStructure: source.as_inner(),
toAccelerationStructure: destination.as_inner()
];
}
Ok(())
}
pub fn update_fence(&self, fence: &Fence) -> Result<(), Error> {
let encoder = self.as_inner();
require_selector(encoder, sel!(updateFence:), "updateFence")?;
unsafe {
let _: () = msg_send![encoder, updateFence: fence.as_inner()];
}
Ok(())
}
pub fn wait_for_fence(&self, fence: &Fence) -> Result<(), Error> {
let encoder = self.as_inner();
require_selector(encoder, sel!(waitForFence:), "waitForFence")?;
unsafe {
let _: () = msg_send![encoder, waitForFence: fence.as_inner()];
}
Ok(())
}
pub fn use_resource(
&self,
resource: AccelerationStructureResource<'_>,
usage: ResourceUsage,
) -> Result<(), Error> {
if !usage.is_valid() {
return Err(Error::invalid_argument(
"resource usage contains undeclared bits",
));
}
let encoder = self.as_inner();
require_selector(encoder, sel!(useResource:usage:), "useResource")?;
unsafe {
let _: () =
msg_send![encoder, useResource: resource.as_object(), usage: usage.as_raw()];
}
Ok(())
}
pub fn use_resources(
&self,
resources: &[AccelerationStructureResource<'_>],
usage: ResourceUsage,
) -> Result<(), Error> {
if resources.is_empty() {
return Err(Error::invalid_argument(
"resource usage requires at least one resource",
));
}
for resource in resources {
self.use_resource(*resource, usage)?;
}
Ok(())
}
pub fn use_heap(&self, heap: &Heap) -> Result<(), Error> {
let encoder = self.as_inner();
require_selector(encoder, sel!(useHeap:), "useHeap")?;
unsafe {
let _: () = msg_send![encoder, useHeap: heap.as_inner()];
}
Ok(())
}
pub fn use_heaps(&self, heaps: &[&Heap]) -> Result<(), Error> {
if heaps.is_empty() {
return Err(Error::invalid_argument(
"heap usage requires at least one heap",
));
}
for heap in heaps {
self.use_heap(heap)?;
}
Ok(())
}
pub fn sample_counters(
&self,
sample_buffer: &CounterSampleBuffer,
sample_index: usize,
barrier: bool,
) -> Result<(), Error> {
if sample_index >= sample_buffer.sample_count()? {
return Err(Error::invalid_argument(
"counter sample index is out of bounds",
));
}
let encoder = self.as_inner();
require_selector(
encoder,
sel!(sampleCountersInBuffer:atSampleIndex:withBarrier:),
"sampleCountersInBuffer",
)?;
unsafe {
let _: () = msg_send![
encoder,
sampleCountersInBuffer: sample_buffer.as_inner(),
atSampleIndex: sample_index,
withBarrier: barrier
];
}
Ok(())
}
pub fn end_encoding(mut self) {
unsafe {
let _: () = msg_send![self.inner.as_inner(), endEncoding];
}
self.ended = true;
}
}
impl Drop for AccelerationStructureEncoder<'_> {
fn drop(&mut self) {
if !self.ended {
unsafe {
let _: () = msg_send![self.inner.as_inner(), endEncoding];
}
self.ended = true;
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn attachment_indices_are_limited_to_proven_capacity() {
assert!(checked_attachment_index(0).is_ok());
assert!(checked_attachment_index(3).is_ok());
assert!(checked_attachment_index(4).is_err());
assert!(checked_attachment_index(usize::MAX).is_err());
}
}