use crate::ThreadBound;
use crate::foundation::Error;
use crate::metal::generated_value_types::{
CompileSymbolVisibility, DataType, FloatingPointConversionRoundingMode, FunctionOptions,
FunctionType, LanguageVersion, LibraryOptimizationLevel, LibraryType,
MathFloatingPointFunctions, MathMode, PatchType, PrimitiveTopologyClass, ShaderValidation,
TessellationControlPointIndexType, TessellationFactorFormat, TessellationFactorStepFunction,
TessellationPartitionMode, Winding,
};
use crate::metal::{Device, PixelFormat, Size};
use block2::RcBlock;
use objc2::rc::Retained;
use objc2::runtime::{AnyObject, NSObjectProtocol, ProtocolObject};
use objc2::{msg_send, sel};
use objc2_foundation::{NSArray, NSDictionary, NSError, NSRange, NSString};
use objc2_metal::{
MTL4BinaryFunction, MTL4RenderPipelineBinaryFunctionsDescriptor, MTLCompileOptions,
MTLCompileSymbolVisibility, MTLComputePipelineState, MTLDynamicLibrary, MTLFunction,
MTLFunctionConstantValues, MTLFunctionDescriptor, MTLIntersectionFunctionDescriptor,
MTLLanguageVersion, MTLLibrary, MTLLibraryOptimizationLevel, MTLLibraryType,
MTLMathFloatingPointFunctions, MTLMathMode, MTLPrimitiveTopologyClass,
MTLRenderPipelineColorAttachmentDescriptorArray, MTLRenderPipelineDescriptor,
MTLRenderPipelineState, MTLShaderValidation, MTLTessellationControlPointIndexType,
MTLTessellationFactorFormat, MTLTessellationFactorStepFunction, MTLTessellationPartitionMode,
MTLWinding,
};
use std::collections::HashMap;
use std::panic::{AssertUnwindSafe, catch_unwind};
use std::ptr::NonNull;
use std::sync::{Arc, Mutex};
fn function_constant_abi_size(data_type: DataType) -> Option<usize> {
match data_type.as_raw() {
3 | 29 | 33 => Some(4),
4 | 30 | 34 => Some(8),
5 | 31 | 35 => Some(16),
6 | 32 | 36 => Some(16),
7 => Some(16),
8 => Some(32),
9 => Some(32),
10 => Some(24),
11 => Some(48),
12 => Some(48),
13 => Some(32),
14 => Some(64),
15 => Some(64),
16 | 37 | 41 | 121 => Some(2),
17 | 38 | 42 | 122 => Some(4),
18 | 39 | 43 | 123 => Some(8),
19 | 40 | 44 | 124 => Some(8),
20 => Some(8),
21 => Some(16),
22 => Some(16),
23 => Some(12),
24 => Some(24),
25 => Some(24),
26 => Some(16),
27 => Some(32),
28 => Some(32),
45 | 49 | 53 => Some(1),
46 | 50 | 54 => Some(2),
47 | 51 | 55 => Some(4),
48 | 52 | 56 => Some(4),
81 | 85 => Some(8),
82 | 86 => Some(16),
83 | 87 => Some(32),
84 | 88 => Some(32),
_ => None,
}
}
fn aligned_constant_bytes(bytes: &[u8]) -> Result<Vec<u128>, Error> {
if bytes.is_empty() {
return Err(Error::invalid_argument("function constant bytes are empty"));
}
let words = bytes
.len()
.checked_add(15)
.ok_or_else(|| Error::invalid_argument("function constant byte length overflow"))?
/ 16;
let mut storage = vec![0_u128; words];
unsafe {
std::ptr::copy_nonoverlapping(
bytes.as_ptr(),
storage.as_mut_ptr().cast::<u8>(),
bytes.len(),
);
}
Ok(storage)
}
fn validate_function_constant_bytes(
data_type: DataType,
bytes: &[u8],
count: usize,
) -> Result<(), Error> {
let components = match data_type.as_raw() {
53 => 1,
54 => 2,
55 => 3,
56 => 4,
_ => return Ok(()),
};
let stride = function_constant_abi_size(data_type)
.ok_or_else(|| Error::invalid_argument("invalid function constant type"))?;
for item in 0..count {
let start = item
.checked_mul(stride)
.ok_or_else(|| Error::invalid_argument("function constant byte offset overflow"))?;
if bytes[start..start + components]
.iter()
.any(|value| *value > 1)
{
return Err(Error::invalid_argument(
"boolean function constants must contain only 0 or 1",
));
}
}
Ok(())
}
fn callback_function(value: *mut AnyObject, error: *mut NSError) -> Result<Function, Error> {
if let Some(error) = unsafe { error.as_ref() } {
return Err(crate::foundation::metal_error(error));
}
let value = unsafe { Retained::retain(value) }
.ok_or_else(|| Error::unsupported("Metal completed without a function or NSError"))?;
Function::from_any_object(value)
}
#[derive(Clone)]
pub struct Library {
pub(super) inner: Retained<ProtocolObject<dyn MTLLibrary>>,
_thread_bound: ThreadBound,
}
impl Library {
pub(super) const fn new(inner: Retained<ProtocolObject<dyn MTLLibrary>>) -> Self {
Self {
inner,
_thread_bound: ThreadBound::new(),
}
}
pub(crate) fn from_any_object(inner: Retained<AnyObject>) -> Result<Self, Error> {
Ok(Self {
inner: unsafe { Retained::cast_unchecked(inner) },
_thread_bound: ThreadBound::new(),
})
}
pub(crate) fn as_any_object(&self) -> &AnyObject {
unsafe { &*(std::ptr::from_ref(&*self.inner).cast::<AnyObject>()) }
}
pub fn function(&self, name: &str) -> Result<Function, Error> {
if name.is_empty() || name.as_bytes().contains(&0) {
return Err(Error::invalid_argument("shader function name is invalid"));
}
let name = NSString::from_str(name);
self.inner
.newFunctionWithName(&name)
.map(Function::new)
.ok_or_else(|| Error::unsupported("the shader function was not found"))
}
#[must_use]
pub fn function_names(&self) -> Vec<String> {
self.inner
.functionNames()
.iter()
.map(|value| value.to_string())
.collect()
}
pub fn specialized_function(
&self,
name: &str,
constants: &crate::metal::generated_object_types::metal::FunctionConstantValues,
) -> Result<Function, Error> {
if name.is_empty() || name.as_bytes().contains(&0) {
return Err(Error::invalid_argument("shader function name is invalid"));
}
if !self
.inner
.respondsToSelector(sel!(newFunctionWithName:constantValues:error:))
{
return Err(Error::unsupported("function specialization is unavailable"));
}
let name = NSString::from_str(name);
let constants = unsafe {
&*(std::ptr::from_ref(constants.as_inner()).cast::<MTLFunctionConstantValues>())
};
self.inner
.newFunctionWithName_constantValues_error(&name, constants)
.map(Function::new)
.map_err(|error| crate::foundation::metal_error(&error))
}
pub fn specialized_function_async(
&self,
name: &str,
constants: &crate::metal::generated_object_types::metal::FunctionConstantValues,
handler: impl FnOnce(Result<Function, Error>) + Send + 'static,
) -> Result<(), Error> {
if name.is_empty() || name.as_bytes().contains(&0) {
return Err(Error::invalid_argument("shader function name is invalid"));
}
if !self.inner.respondsToSelector(sel!(
newFunctionWithName:constantValues:completionHandler:
)) {
return Err(Error::unsupported(
"asynchronous function specialization is unavailable",
));
}
let name = NSString::from_str(name);
let state = Arc::new(Mutex::new(Some(handler)));
let callback_state = Arc::clone(&state);
let block = RcBlock::new(move |value: *mut AnyObject, error: *mut NSError| {
let callback = callback_state
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.take();
let Some(callback) = callback else { return };
let result = callback_function(value, error);
let _ = catch_unwind(AssertUnwindSafe(|| callback(result)));
});
unsafe {
let _: () = msg_send![&*self.inner,
newFunctionWithName: &*name,
constantValues: constants.as_inner(),
completionHandler: &*block
];
}
Ok(())
}
pub fn function_with_descriptor(
&self,
descriptor: &crate::metal::generated_object_types::metal::FunctionDescriptor,
) -> Result<Function, Error> {
if !self
.inner
.respondsToSelector(sel!(newFunctionWithDescriptor:error:))
{
return Err(Error::unsupported(
"descriptor-based function creation is unavailable",
));
}
let descriptor = unsafe {
&*(std::ptr::from_ref(descriptor.as_inner()).cast::<MTLFunctionDescriptor>())
};
self.inner
.newFunctionWithDescriptor_error(descriptor)
.map(Function::new)
.map_err(|error| crate::foundation::metal_error(&error))
}
pub fn function_with_descriptor_async(
&self,
descriptor: &crate::metal::generated_object_types::metal::FunctionDescriptor,
handler: impl FnOnce(Result<Function, Error>) + Send + 'static,
) -> Result<(), Error> {
if !self
.inner
.respondsToSelector(sel!(newFunctionWithDescriptor:completionHandler:))
{
return Err(Error::unsupported(
"asynchronous descriptor-based function creation is unavailable",
));
}
let state = Arc::new(Mutex::new(Some(handler)));
let callback_state = Arc::clone(&state);
let block = RcBlock::new(move |value: *mut AnyObject, error: *mut NSError| {
let callback = callback_state
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.take();
let Some(callback) = callback else { return };
let result = callback_function(value, error);
let _ = catch_unwind(AssertUnwindSafe(|| callback(result)));
});
unsafe {
let _: () = msg_send![&*self.inner,
newFunctionWithDescriptor: descriptor.as_inner(),
completionHandler: &*block
];
}
Ok(())
}
pub fn intersection_function_with_descriptor(
&self,
descriptor: &crate::metal::generated_object_types::metal::IntersectionFunctionDescriptor,
) -> Result<Function, Error> {
if !self
.inner
.respondsToSelector(sel!(newIntersectionFunctionWithDescriptor:error:))
{
return Err(Error::unsupported(
"intersection-function creation is unavailable",
));
}
let descriptor = unsafe {
&*(std::ptr::from_ref(descriptor.as_inner())
.cast::<MTLIntersectionFunctionDescriptor>())
};
self.inner
.newIntersectionFunctionWithDescriptor_error(descriptor)
.map(Function::new)
.map_err(|error| crate::foundation::metal_error(&error))
}
pub fn intersection_function_with_descriptor_async(
&self,
descriptor: &crate::metal::generated_object_types::metal::IntersectionFunctionDescriptor,
handler: impl FnOnce(Result<Function, Error>) + Send + 'static,
) -> Result<(), Error> {
if !self.inner.respondsToSelector(sel!(
newIntersectionFunctionWithDescriptor:completionHandler:
)) {
return Err(Error::unsupported(
"asynchronous intersection-function creation is unavailable",
));
}
let state = Arc::new(Mutex::new(Some(handler)));
let callback_state = Arc::clone(&state);
let block = RcBlock::new(move |value: *mut AnyObject, error: *mut NSError| {
let callback = callback_state
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.take();
let Some(callback) = callback else { return };
let result = callback_function(value, error);
let _ = catch_unwind(AssertUnwindSafe(|| callback(result)));
});
unsafe {
let _: () = msg_send![&*self.inner,
newIntersectionFunctionWithDescriptor: descriptor.as_inner(),
completionHandler: &*block
];
}
Ok(())
}
pub fn reflection_for_function(
&self,
name: &str,
) -> Result<Option<crate::metal::generated_object_types::metal::FunctionReflection>, Error>
{
if name.is_empty() || name.as_bytes().contains(&0) {
return Err(Error::invalid_argument("shader function name is invalid"));
}
if !self
.inner
.respondsToSelector(sel!(reflectionForFunctionWithName:))
{
return Err(Error::unsupported("function reflection is unavailable"));
}
let name = NSString::from_str(name);
Ok(self
.inner
.reflectionForFunctionWithName(&name)
.map(|value| {
let value = unsafe { Retained::cast_unchecked(value) };
crate::metal::generated_object_types::metal::FunctionReflection::from_inner(value)
}))
}
#[must_use]
pub fn device(&self) -> Device {
Device::from_inner(self.inner.device())
}
#[must_use]
pub fn install_name(&self) -> Option<String> {
self.inner.installName().map(|value| value.to_string())
}
#[must_use]
pub fn label(&self) -> Option<String> {
self.inner.label().map(|value| value.to_string())
}
pub fn set_label(&self, value: Option<&str>) {
let value = value.map(NSString::from_str);
self.inner.setLabel(value.as_deref());
}
#[must_use]
pub fn library_type(&self) -> LibraryType {
LibraryType::from_system_raw(self.inner.r#type().0)
}
}
#[derive(Clone)]
pub struct Function {
pub(super) inner: Retained<ProtocolObject<dyn MTLFunction>>,
_thread_bound: ThreadBound,
}
impl Function {
pub(super) const fn new(inner: Retained<ProtocolObject<dyn MTLFunction>>) -> Self {
Self {
inner,
_thread_bound: ThreadBound::new(),
}
}
pub(crate) fn from_any_object(inner: Retained<AnyObject>) -> Result<Self, Error> {
Ok(Self {
inner: unsafe { Retained::cast_unchecked(inner) },
_thread_bound: ThreadBound::new(),
})
}
pub(crate) fn as_any_object(&self) -> &AnyObject {
unsafe { &*(std::ptr::from_ref(&*self.inner).cast::<AnyObject>()) }
}
#[must_use]
pub fn name(&self) -> String {
self.inner.name().to_string()
}
#[must_use]
pub fn device(&self) -> Device {
Device::from_inner(self.inner.device())
}
#[must_use]
pub fn label(&self) -> Option<String> {
self.inner.label().map(|value| value.to_string())
}
pub fn set_label(&self, value: Option<&str>) {
let value = value.map(NSString::from_str);
self.inner.setLabel(value.as_deref());
}
#[must_use]
pub fn function_type(&self) -> FunctionType {
FunctionType::from_system_raw(self.inner.functionType().0)
}
#[must_use]
pub fn options(&self) -> FunctionOptions {
FunctionOptions::from_system_raw(self.inner.options().0)
}
#[must_use]
pub fn patch_control_point_count(&self) -> isize {
self.inner.patchControlPointCount()
}
#[must_use]
pub fn patch_type(&self) -> PatchType {
PatchType::from_system_raw(self.inner.patchType().0)
}
#[must_use]
pub fn function_constant_names(&self) -> Vec<String> {
self.inner
.functionConstantsDictionary()
.keys()
.map(|value| value.to_string())
.collect()
}
#[must_use]
pub fn vertex_attribute_names(&self) -> Vec<String> {
self.inner
.vertexAttributes()
.map(|values| {
values
.iter()
.map(|value| value.name().to_string())
.collect()
})
.unwrap_or_default()
}
#[must_use]
pub fn stage_input_attribute_names(&self) -> Vec<String> {
self.inner
.stageInputAttributes()
.map(|values| {
values
.iter()
.map(|value| value.name().to_string())
.collect()
})
.unwrap_or_default()
}
}
#[derive(Clone)]
pub struct RenderPipelineState {
pub(super) inner: Retained<ProtocolObject<dyn MTLRenderPipelineState>>,
_thread_bound: ThreadBound,
}
impl RenderPipelineState {
pub(super) const fn new(inner: Retained<ProtocolObject<dyn MTLRenderPipelineState>>) -> Self {
Self {
inner,
_thread_bound: ThreadBound::new(),
}
}
#[must_use]
pub fn max_total_threads_per_threadgroup(&self) -> usize {
self.inner.maxTotalThreadsPerThreadgroup()
}
#[must_use]
pub fn device(&self) -> Device {
Device::from_inner(self.inner.device())
}
#[must_use]
pub fn label(&self) -> Option<String> {
self.inner.label().map(|value| value.to_string())
}
#[must_use]
pub fn supports_indirect_command_buffers(&self) -> bool {
self.inner.supportIndirectCommandBuffers()
}
#[must_use]
pub fn threadgroup_size_matches_tile_size(&self) -> bool {
self.inner.threadgroupSizeMatchesTileSize()
}
#[must_use]
pub fn imageblock_sample_length(&self) -> usize {
self.inner.imageblockSampleLength()
}
pub fn shader_validation(&self) -> Result<ShaderValidation, Error> {
if !self.inner.respondsToSelector(sel!(shaderValidation)) {
return Err(Error::unsupported(
"pipeline shader validation is unavailable",
));
}
Ok(ShaderValidation::from_system_raw(
unsafe { self.inner.shaderValidation() }.0,
))
}
#[must_use]
pub fn required_threadgroups(&self) -> (Size, Size, Size) {
let tile = self.inner.requiredThreadsPerTileThreadgroup();
let mesh = self.inner.requiredThreadsPerMeshThreadgroup();
let object = self.inner.requiredThreadsPerObjectThreadgroup();
(
Size::new(tile.width, tile.height, tile.depth),
Size::new(mesh.width, mesh.height, mesh.depth),
Size::new(object.width, object.height, object.depth),
)
}
pub fn imageblock_memory_length(&self, dimensions: Size) -> Result<usize, Error> {
if dimensions.width == 0 || dimensions.height == 0 || dimensions.depth == 0 {
return Err(Error::invalid_argument(
"imageblock dimensions must be non-zero",
));
}
if !self
.inner
.respondsToSelector(sel!(imageblockMemoryLengthForDimensions:))
{
return Err(Error::unsupported(
"pipeline imageblock memory queries are unavailable",
));
}
Ok(self
.inner
.imageblockMemoryLengthForDimensions(dimensions.into()))
}
pub fn gpu_resource_id(&self) -> Result<u64, Error> {
if !self.inner.respondsToSelector(sel!(gpuResourceID)) {
return Err(Error::unsupported(
"pipeline GPU resource identifiers are unavailable",
));
}
let value = self.inner.gpuResourceID();
Ok(unsafe { std::ptr::read_unaligned(std::ptr::from_ref(&value).cast::<u64>()) })
}
pub fn max_total_threads_per_object_threadgroup(&self) -> Result<usize, Error> {
if !self
.inner
.respondsToSelector(sel!(maxTotalThreadsPerObjectThreadgroup))
{
return Err(Error::unsupported("object shaders are unavailable"));
}
Ok(self.inner.maxTotalThreadsPerObjectThreadgroup())
}
pub fn max_total_threads_per_mesh_threadgroup(&self) -> Result<usize, Error> {
if !self
.inner
.respondsToSelector(sel!(maxTotalThreadsPerMeshThreadgroup))
{
return Err(Error::unsupported("mesh shaders are unavailable"));
}
Ok(self.inner.maxTotalThreadsPerMeshThreadgroup())
}
pub fn object_thread_execution_width(&self) -> Result<usize, Error> {
if !self
.inner
.respondsToSelector(sel!(objectThreadExecutionWidth))
{
return Err(Error::unsupported("object shaders are unavailable"));
}
Ok(self.inner.objectThreadExecutionWidth())
}
pub fn mesh_thread_execution_width(&self) -> Result<usize, Error> {
if !self
.inner
.respondsToSelector(sel!(meshThreadExecutionWidth))
{
return Err(Error::unsupported("mesh shaders are unavailable"));
}
Ok(self.inner.meshThreadExecutionWidth())
}
pub fn max_total_threadgroups_per_mesh_grid(&self) -> Result<usize, Error> {
if !self
.inner
.respondsToSelector(sel!(maxTotalThreadgroupsPerMeshGrid))
{
return Err(Error::unsupported("mesh shaders are unavailable"));
}
Ok(self.inner.maxTotalThreadgroupsPerMeshGrid())
}
pub fn reflection(
&self,
) -> Result<Option<crate::metal::generated_object_types::metal::RenderPipelineReflection>, Error>
{
if !self.inner.respondsToSelector(sel!(reflection)) {
return Err(Error::unsupported("pipeline reflection is unavailable"));
}
Ok(self.inner.reflection().map(|value| {
let value = unsafe { Retained::cast_unchecked(value) };
crate::metal::generated_object_types::metal::RenderPipelineReflection::from_inner(value)
}))
}
pub fn function_handle_by_name(
&self,
name: &str,
stage: crate::metal::generated_value_types::RenderStages,
) -> Result<Option<crate::metal::generated_object_types::metal::FunctionHandle>, Error> {
if name.is_empty() || name.as_bytes().contains(&0) {
return Err(Error::invalid_argument("function handle name is invalid"));
}
if !stage.is_valid() {
return Err(Error::invalid_argument("render stage flags are invalid"));
}
if !self
.inner
.respondsToSelector(sel!(functionHandleWithName:stage:))
{
return Err(Error::unsupported("function handles are unavailable"));
}
let name = NSString::from_str(name);
let raw = stage.as_raw();
let value: Option<Retained<AnyObject>> =
unsafe { msg_send![&*self.inner, functionHandleWithName: &*name, stage: raw] };
Ok(value.map(crate::metal::generated_object_types::metal::FunctionHandle::from_inner))
}
pub fn function_handle_for_function(
&self,
function: &Function,
stage: crate::metal::generated_value_types::RenderStages,
) -> Result<Option<crate::metal::generated_object_types::metal::FunctionHandle>, Error> {
self.function_handle_by_name(&function.name(), stage)
}
pub fn function_handle_for_binary_function(
&self,
function: &crate::metal::generated_object_types::metal4::BinaryFunction,
stage: crate::metal::generated_value_types::RenderStages,
) -> Result<Option<crate::metal::generated_object_types::metal::FunctionHandle>, Error> {
let name = function
.name()?
.ok_or_else(|| Error::invalid_argument("binary function has no name"))?;
self.function_handle_by_name(&name, stage)
}
pub fn new_visible_function_table(
&self,
descriptor: &crate::metal::generated_object_types::metal::VisibleFunctionTableDescriptor,
stage: crate::metal::generated_value_types::RenderStages,
) -> Result<Option<crate::metal::generated_object_types::metal::VisibleFunctionTable>, Error>
{
if !stage.is_valid() {
return Err(Error::invalid_argument("render stage flags are invalid"));
}
if !self.inner.respondsToSelector(sel!(
newVisibleFunctionTableWithDescriptor:stage:
)) {
return Err(Error::unsupported(
"visible function tables are unavailable",
));
}
let value: Option<Retained<AnyObject>> = unsafe {
msg_send![&*self.inner,
newVisibleFunctionTableWithDescriptor: descriptor.as_inner(),
stage: stage.as_raw()
]
};
Ok(
value
.map(crate::metal::generated_object_types::metal::VisibleFunctionTable::from_inner),
)
}
pub fn new_intersection_function_table(
&self,
descriptor: &crate::metal::generated_object_types::metal::IntersectionFunctionTableDescriptor,
stage: crate::metal::generated_value_types::RenderStages,
) -> Result<Option<crate::metal::generated_object_types::metal::IntersectionFunctionTable>, Error>
{
if !stage.is_valid() {
return Err(Error::invalid_argument("render stage flags are invalid"));
}
if !self.inner.respondsToSelector(sel!(
newIntersectionFunctionTableWithDescriptor:stage:
)) {
return Err(Error::unsupported(
"intersection function tables are unavailable",
));
}
let value: Option<Retained<AnyObject>> = unsafe {
msg_send![&*self.inner,
newIntersectionFunctionTableWithDescriptor: descriptor.as_inner(),
stage: stage.as_raw()
]
};
Ok(value.map(
crate::metal::generated_object_types::metal::IntersectionFunctionTable::from_inner,
))
}
pub fn new_render_pipeline_descriptor(
&self,
) -> Result<crate::metal::generated_object_types::metal4::PipelineDescriptor, Error> {
if !self
.inner
.respondsToSelector(sel!(newRenderPipelineDescriptorForSpecialization))
{
return Err(Error::unsupported(
"render pipeline specialization descriptors are unavailable",
));
}
let value: Retained<AnyObject> =
unsafe { msg_send![&*self.inner, newRenderPipelineDescriptorForSpecialization] };
Ok(crate::metal::generated_object_types::metal4::PipelineDescriptor::from_inner(value))
}
pub fn with_metal4_binary_functions(
&self,
descriptor: &crate::metal::generated_object_types::metal4::RenderPipelineBinaryFunctionsDescriptor,
) -> Result<Self, Error> {
if !self
.inner
.respondsToSelector(sel!(newRenderPipelineStateWithBinaryFunctions:error:))
{
return Err(Error::unsupported(
"Metal 4 render pipeline binary linking is unavailable",
));
}
let descriptor = unsafe {
&*(std::ptr::from_ref(descriptor.as_inner())
.cast::<MTL4RenderPipelineBinaryFunctionsDescriptor>())
};
self.inner
.newRenderPipelineStateWithBinaryFunctions_error(descriptor)
.map(Self::new)
.map_err(|error| crate::foundation::metal_error(&error))
}
pub fn with_additional_binary_functions(
&self,
descriptor: &crate::metal::generated_object_types::metal::RenderPipelineFunctionsDescriptor,
) -> Result<Self, Error> {
if !self.inner.respondsToSelector(sel!(
newRenderPipelineStateWithAdditionalBinaryFunctions:error:
)) {
return Err(Error::unsupported(
"render pipeline binary-function linking is unavailable",
));
}
let descriptor = unsafe {
&*(std::ptr::from_ref(descriptor.as_inner())
.cast::<objc2_metal::MTLRenderPipelineFunctionsDescriptor>())
};
self.inner
.newRenderPipelineStateWithAdditionalBinaryFunctions_error(descriptor)
.map(Self::new)
.map_err(|error| crate::foundation::metal_error(&error))
}
}
#[derive(Clone)]
pub struct ComputePipelineState {
pub(super) inner: Retained<ProtocolObject<dyn MTLComputePipelineState>>,
_thread_bound: ThreadBound,
}
impl ComputePipelineState {
pub(super) const fn new(inner: Retained<ProtocolObject<dyn MTLComputePipelineState>>) -> Self {
Self {
inner,
_thread_bound: ThreadBound::new(),
}
}
#[must_use]
pub fn execution_width(&self) -> usize {
self.inner.threadExecutionWidth()
}
#[must_use]
pub fn max_total_threads_per_threadgroup(&self) -> usize {
self.inner.maxTotalThreadsPerThreadgroup()
}
#[must_use]
pub fn device(&self) -> Device {
Device::from_inner(self.inner.device())
}
#[must_use]
pub fn label(&self) -> Option<String> {
self.inner.label().map(|value| value.to_string())
}
#[must_use]
pub fn supports_indirect_command_buffers(&self) -> bool {
self.inner.supportIndirectCommandBuffers()
}
#[must_use]
pub fn static_threadgroup_memory_length(&self) -> usize {
self.inner.staticThreadgroupMemoryLength()
}
#[must_use]
pub fn required_threads_per_threadgroup(&self) -> Size {
let value = self.inner.requiredThreadsPerThreadgroup();
Size::new(value.width, value.height, value.depth)
}
pub fn shader_validation(&self) -> Result<ShaderValidation, Error> {
if !self.inner.respondsToSelector(sel!(shaderValidation)) {
return Err(Error::unsupported(
"pipeline shader validation is unavailable",
));
}
Ok(ShaderValidation::from_system_raw(
unsafe { self.inner.shaderValidation() }.0,
))
}
pub fn imageblock_memory_length(&self, dimensions: Size) -> Result<usize, Error> {
if dimensions.width == 0 || dimensions.height == 0 || dimensions.depth == 0 {
return Err(Error::invalid_argument(
"imageblock dimensions must be non-zero",
));
}
if !self
.inner
.respondsToSelector(sel!(imageblockMemoryLengthForDimensions:))
{
return Err(Error::unsupported(
"pipeline imageblock memory queries are unavailable",
));
}
Ok(unsafe {
self.inner
.imageblockMemoryLengthForDimensions(dimensions.into())
})
}
pub fn gpu_resource_id(&self) -> Result<u64, Error> {
if !self.inner.respondsToSelector(sel!(gpuResourceID)) {
return Err(Error::unsupported(
"pipeline GPU resource identifiers are unavailable",
));
}
let value = self.inner.gpuResourceID();
Ok(unsafe { std::ptr::read_unaligned(std::ptr::from_ref(&value).cast::<u64>()) })
}
pub fn reflection(
&self,
) -> Result<Option<crate::metal::generated_object_types::metal::ComputePipelineReflection>, Error>
{
if !self.inner.respondsToSelector(sel!(reflection)) {
return Err(Error::unsupported("pipeline reflection is unavailable"));
}
Ok(self.inner.reflection().map(|value| {
let value = unsafe { Retained::cast_unchecked(value) };
crate::metal::generated_object_types::metal::ComputePipelineReflection::from_inner(
value,
)
}))
}
pub fn function_handle_by_name(
&self,
name: &str,
) -> Result<Option<crate::metal::generated_object_types::metal::FunctionHandle>, Error> {
if name.is_empty() || name.as_bytes().contains(&0) {
return Err(Error::invalid_argument("function handle name is invalid"));
}
if !self.inner.respondsToSelector(sel!(functionHandleWithName:)) {
return Err(Error::unsupported("function handles are unavailable"));
}
let name = NSString::from_str(name);
let value: Option<Retained<AnyObject>> =
unsafe { msg_send![&*self.inner, functionHandleWithName: &*name] };
Ok(value.map(crate::metal::generated_object_types::metal::FunctionHandle::from_inner))
}
pub fn function_handle_for_function(
&self,
function: &Function,
) -> Result<Option<crate::metal::generated_object_types::metal::FunctionHandle>, Error> {
self.function_handle_by_name(&function.name())
}
pub fn function_handle_for_binary_function(
&self,
function: &crate::metal::generated_object_types::metal4::BinaryFunction,
) -> Result<Option<crate::metal::generated_object_types::metal::FunctionHandle>, Error> {
let name = function
.name()?
.ok_or_else(|| Error::invalid_argument("binary function has no name"))?;
self.function_handle_by_name(&name)
}
pub fn new_visible_function_table(
&self,
descriptor: &crate::metal::generated_object_types::metal::VisibleFunctionTableDescriptor,
) -> Result<Option<crate::metal::generated_object_types::metal::VisibleFunctionTable>, Error>
{
if !self
.inner
.respondsToSelector(sel!(newVisibleFunctionTableWithDescriptor:))
{
return Err(Error::unsupported(
"visible function tables are unavailable",
));
}
let value: Option<Retained<AnyObject>> = unsafe {
msg_send![&*self.inner, newVisibleFunctionTableWithDescriptor: descriptor.as_inner()]
};
Ok(
value
.map(crate::metal::generated_object_types::metal::VisibleFunctionTable::from_inner),
)
}
pub fn new_intersection_function_table(
&self,
descriptor: &crate::metal::generated_object_types::metal::IntersectionFunctionTableDescriptor,
) -> Result<Option<crate::metal::generated_object_types::metal::IntersectionFunctionTable>, Error>
{
if !self
.inner
.respondsToSelector(sel!(newIntersectionFunctionTableWithDescriptor:))
{
return Err(Error::unsupported(
"intersection function tables are unavailable",
));
}
let value: Option<Retained<AnyObject>> = unsafe {
msg_send![&*self.inner, newIntersectionFunctionTableWithDescriptor: descriptor.as_inner()]
};
Ok(value.map(
crate::metal::generated_object_types::metal::IntersectionFunctionTable::from_inner,
))
}
pub fn with_metal4_binary_functions(
&self,
functions: &[crate::metal::generated_object_types::metal4::BinaryFunction],
) -> Result<Self, Error> {
if functions.is_empty() {
return Err(Error::invalid_argument(
"at least one binary function is required",
));
}
if !self
.inner
.respondsToSelector(sel!(newComputePipelineStateWithBinaryFunctions:error:))
{
return Err(Error::unsupported(
"Metal 4 compute pipeline binary linking is unavailable",
));
}
let functions: Vec<Retained<ProtocolObject<dyn MTL4BinaryFunction>>> = functions
.iter()
.map(|function| {
let retained =
unsafe { Retained::retain(std::ptr::from_ref(function.as_inner()).cast_mut()) }
.expect("a borrowed Objective-C wrapper cannot be null");
unsafe { Retained::cast_unchecked(retained) }
})
.collect();
let functions = NSArray::from_retained_slice(&functions);
self.inner
.newComputePipelineStateWithBinaryFunctions_error(&functions)
.map(Self::new)
.map_err(|error| crate::foundation::metal_error(&error))
}
pub fn with_additional_functions(&self, functions: &[Function]) -> Result<Self, Error> {
if functions.is_empty() {
return Err(Error::invalid_argument(
"at least one additional function is required",
));
}
if !self.inner.respondsToSelector(sel!(
newComputePipelineStateWithAdditionalBinaryFunctions:error:
)) {
return Err(Error::unsupported(
"compute pipeline binary-function linking is unavailable",
));
}
let functions: Vec<Retained<ProtocolObject<dyn MTLFunction>>> =
functions.iter().map(|value| value.inner.clone()).collect();
let functions = NSArray::from_retained_slice(&functions);
self.inner
.newComputePipelineStateWithAdditionalBinaryFunctions_error(&functions)
.map(Self::new)
.map_err(|error| crate::foundation::metal_error(&error))
}
}
pub struct RenderPipelineDescriptor {
pub(super) inner: Retained<MTLRenderPipelineDescriptor>,
_thread_bound: ThreadBound,
}
#[derive(Clone, Copy, Debug)]
pub struct RenderPipelineOptions {
pub alpha_to_coverage: bool,
pub alpha_to_one: bool,
pub rasterization_enabled: bool,
pub tessellation_factor_scale_enabled: bool,
pub support_adding_fragment_binary_functions: bool,
pub support_adding_vertex_binary_functions: bool,
pub support_indirect_command_buffers: bool,
pub max_fragment_call_stack_depth: usize,
pub max_tessellation_factor: usize,
pub max_vertex_amplification_count: usize,
pub max_vertex_call_stack_depth: usize,
pub raster_sample_count: usize,
pub input_primitive_topology: PrimitiveTopologyClass,
pub depth_attachment_pixel_format: PixelFormat,
pub stencil_attachment_pixel_format: PixelFormat,
pub shader_validation: ShaderValidation,
pub tessellation_control_point_index_type: TessellationControlPointIndexType,
pub tessellation_factor_format: TessellationFactorFormat,
pub tessellation_factor_step_function: TessellationFactorStepFunction,
pub tessellation_output_winding_order: Winding,
pub tessellation_partition_mode: TessellationPartitionMode,
}
impl RenderPipelineDescriptor {
#[must_use]
pub fn empty() -> Self {
Self {
inner: MTLRenderPipelineDescriptor::new(),
_thread_bound: ThreadBound::new(),
}
}
#[must_use]
pub fn new(vertex: &Function, fragment: Option<&Function>, color_format: PixelFormat) -> Self {
let inner = MTLRenderPipelineDescriptor::new();
inner.setVertexFunction(Some(&vertex.inner));
inner.setFragmentFunction(fragment.map(|function| &*function.inner));
inner.setRasterSampleCount(1);
let attachments: Retained<MTLRenderPipelineColorAttachmentDescriptorArray> =
inner.colorAttachments();
let attachment = unsafe { attachments.objectAtIndexedSubscript(0) };
attachment.setPixelFormat(color_format.as_objc());
Self {
inner,
_thread_bound: ThreadBound::new(),
}
}
#[must_use]
pub fn label(&self) -> Option<String> {
self.inner.label().map(|value| value.to_string())
}
pub fn set_label(&self, value: Option<&str>) {
let value = value.map(NSString::from_str);
self.inner.setLabel(value.as_deref());
}
#[must_use]
pub fn vertex_function(&self) -> Option<Function> {
self.inner.vertexFunction().map(Function::new)
}
pub fn set_vertex_function(&self, value: Option<&Function>) {
self.inner
.setVertexFunction(value.map(|value| &*value.inner));
}
#[must_use]
pub fn fragment_function(&self) -> Option<Function> {
self.inner.fragmentFunction().map(Function::new)
}
pub fn set_fragment_function(&self, value: Option<&Function>) {
self.inner
.setFragmentFunction(value.map(|value| &*value.inner));
}
pub fn reset(&self) {
self.inner.reset();
}
pub fn options(&self) -> Result<RenderPipelineOptions, Error> {
if !self.inner.respondsToSelector(sel!(shaderValidation)) {
return Err(Error::unsupported(
"render pipeline shader validation is unavailable",
));
}
let shader_validation = self.inner.shaderValidation();
Ok(RenderPipelineOptions {
alpha_to_coverage: self.inner.isAlphaToCoverageEnabled(),
alpha_to_one: self.inner.isAlphaToOneEnabled(),
rasterization_enabled: self.inner.isRasterizationEnabled(),
tessellation_factor_scale_enabled: self.inner.isTessellationFactorScaleEnabled(),
support_adding_fragment_binary_functions: self
.inner
.supportAddingFragmentBinaryFunctions(),
support_adding_vertex_binary_functions: self.inner.supportAddingVertexBinaryFunctions(),
support_indirect_command_buffers: self.inner.supportIndirectCommandBuffers(),
max_fragment_call_stack_depth: self.inner.maxFragmentCallStackDepth(),
max_tessellation_factor: self.inner.maxTessellationFactor(),
max_vertex_amplification_count: self.inner.maxVertexAmplificationCount(),
max_vertex_call_stack_depth: self.inner.maxVertexCallStackDepth(),
raster_sample_count: self.inner.rasterSampleCount(),
input_primitive_topology: PrimitiveTopologyClass::from_system_raw(
self.inner.inputPrimitiveTopology().0,
),
depth_attachment_pixel_format: PixelFormat::from_system_raw(
self.inner.depthAttachmentPixelFormat().0,
),
stencil_attachment_pixel_format: PixelFormat::from_system_raw(
self.inner.stencilAttachmentPixelFormat().0,
),
shader_validation: ShaderValidation::from_system_raw(shader_validation.0),
tessellation_control_point_index_type:
TessellationControlPointIndexType::from_system_raw(
self.inner.tessellationControlPointIndexType().0,
),
tessellation_factor_format: TessellationFactorFormat::from_system_raw(
self.inner.tessellationFactorFormat().0,
),
tessellation_factor_step_function: TessellationFactorStepFunction::from_system_raw(
self.inner.tessellationFactorStepFunction().0,
),
tessellation_output_winding_order: Winding::from_system_raw(
self.inner.tessellationOutputWindingOrder().0,
),
tessellation_partition_mode: TessellationPartitionMode::from_system_raw(
self.inner.tessellationPartitionMode().0,
),
})
}
pub fn set_options(&self, value: &RenderPipelineOptions) -> Result<(), Error> {
if value.raster_sample_count == 0 || value.max_tessellation_factor == 0 {
return Err(Error::invalid_argument(
"sample and tessellation counts must be non-zero",
));
}
if value.max_tessellation_factor > 64 {
return Err(Error::invalid_argument(
"maximum tessellation factor cannot exceed 64",
));
}
for (selector, property) in [
(
sel!(setMaxVertexAmplificationCount:),
"maximum vertex amplification count",
),
(sel!(setInputPrimitiveTopology:), "input primitive topology"),
(
sel!(setTessellationPartitionMode:),
"tessellation partition mode",
),
(
sel!(setMaxTessellationFactor:),
"maximum tessellation factor",
),
(
sel!(setTessellationControlPointIndexType:),
"tessellation control-point index type",
),
(sel!(setShaderValidation:), "shader validation"),
] {
if !self.inner.respondsToSelector(selector) {
return Err(Error::unsupported(format!(
"render pipeline {property} is unavailable"
)));
}
}
self.inner
.setAlphaToCoverageEnabled(value.alpha_to_coverage);
self.inner.setAlphaToOneEnabled(value.alpha_to_one);
self.inner
.setRasterizationEnabled(value.rasterization_enabled);
self.inner
.setTessellationFactorScaleEnabled(value.tessellation_factor_scale_enabled);
self.inner.setSupportAddingFragmentBinaryFunctions(
value.support_adding_fragment_binary_functions,
);
self.inner
.setSupportAddingVertexBinaryFunctions(value.support_adding_vertex_binary_functions);
self.inner
.setSupportIndirectCommandBuffers(value.support_indirect_command_buffers);
self.inner
.setMaxFragmentCallStackDepth(value.max_fragment_call_stack_depth);
unsafe {
self.inner
.setMaxTessellationFactor(value.max_tessellation_factor)
};
unsafe {
self.inner
.setMaxVertexAmplificationCount(value.max_vertex_amplification_count)
};
self.inner
.setMaxVertexCallStackDepth(value.max_vertex_call_stack_depth);
self.inner.setRasterSampleCount(value.raster_sample_count);
unsafe {
self.inner
.setInputPrimitiveTopology(MTLPrimitiveTopologyClass(
value.input_primitive_topology.as_raw(),
))
};
self.inner
.setDepthAttachmentPixelFormat(value.depth_attachment_pixel_format.as_objc());
self.inner
.setStencilAttachmentPixelFormat(value.stencil_attachment_pixel_format.as_objc());
unsafe {
self.inner
.setTessellationControlPointIndexType(MTLTessellationControlPointIndexType(
value.tessellation_control_point_index_type.as_raw(),
))
};
self.inner
.setTessellationFactorFormat(MTLTessellationFactorFormat(
value.tessellation_factor_format.as_raw(),
));
self.inner
.setTessellationFactorStepFunction(MTLTessellationFactorStepFunction(
value.tessellation_factor_step_function.as_raw(),
));
self.inner.setTessellationOutputWindingOrder(MTLWinding(
value.tessellation_output_winding_order.as_raw(),
));
unsafe {
self.inner
.setTessellationPartitionMode(MTLTessellationPartitionMode(
value.tessellation_partition_mode.as_raw(),
))
};
self.inner
.setShaderValidation(MTLShaderValidation(value.shader_validation.as_raw()));
Ok(())
}
}
pub struct CompileOptions {
pub(super) inner: Retained<MTLCompileOptions>,
_thread_bound: ThreadBound,
}
impl CompileOptions {
pub(crate) fn from_any_object(inner: Retained<AnyObject>) -> Result<Self, Error> {
Ok(Self {
inner: unsafe { Retained::cast_unchecked(inner) },
_thread_bound: ThreadBound::new(),
})
}
pub(crate) fn as_any_object(&self) -> &AnyObject {
&self.inner
}
#[must_use]
pub fn new() -> Self {
Self {
inner: MTLCompileOptions::new(),
_thread_bound: ThreadBound::new(),
}
}
#[must_use]
pub fn strict() -> Self {
let inner = MTLCompileOptions::new();
inner.setMathMode(objc2_metal::MTLMathMode::Safe);
Self {
inner,
_thread_bound: ThreadBound::new(),
}
}
#[must_use]
#[allow(deprecated)]
pub fn fast_math_enabled(&self) -> bool {
self.inner.fastMathEnabled()
}
#[allow(deprecated)]
pub fn set_fast_math_enabled(&self, value: bool) {
self.inner.setFastMathEnabled(value);
}
#[must_use]
pub fn allows_referencing_undefined_symbols(&self) -> bool {
self.inner.allowReferencingUndefinedSymbols()
}
pub fn set_allows_referencing_undefined_symbols(&self, value: bool) {
self.inner.setAllowReferencingUndefinedSymbols(value);
}
#[must_use]
pub fn logging_enabled(&self) -> bool {
self.inner.enableLogging()
}
pub fn set_logging_enabled(&self, value: bool) {
self.inner.setEnableLogging(value);
}
#[must_use]
pub fn preserves_invariance(&self) -> bool {
self.inner.preserveInvariance()
}
pub fn set_preserves_invariance(&self, value: bool) {
self.inner.setPreserveInvariance(value);
}
#[must_use]
pub fn install_name(&self) -> Option<String> {
self.inner.installName().map(|value| value.to_string())
}
pub fn set_install_name(&self, value: Option<&str>) {
let value = value.map(NSString::from_str);
self.inner.setInstallName(value.as_deref());
}
#[must_use]
pub fn max_total_threads_per_threadgroup(&self) -> usize {
self.inner.maxTotalThreadsPerThreadgroup()
}
pub fn set_max_total_threads_per_threadgroup(&self, value: usize) -> Result<(), Error> {
if value == 0 {
return Err(Error::invalid_argument(
"maximum total threads must be non-zero",
));
}
self.inner.setMaxTotalThreadsPerThreadgroup(value);
Ok(())
}
#[must_use]
pub fn required_threads_per_threadgroup(&self) -> Size {
let value = self.inner.requiredThreadsPerThreadgroup();
Size::new(value.width, value.height, value.depth)
}
pub fn set_required_threads_per_threadgroup(&self, value: Size) -> Result<(), Error> {
if value.width == 0 || value.height == 0 || value.depth == 0 {
return Err(Error::invalid_argument(
"required threadgroup dimensions must be non-zero",
));
}
self.inner.setRequiredThreadsPerThreadgroup(value.into());
Ok(())
}
#[must_use]
pub fn math_mode(&self) -> MathMode {
MathMode::from_system_raw(self.inner.mathMode().0)
}
pub fn set_math_mode(&self, value: MathMode) {
self.inner.setMathMode(MTLMathMode(value.as_raw()));
}
#[must_use]
pub fn math_floating_point_functions(&self) -> MathFloatingPointFunctions {
MathFloatingPointFunctions::from_system_raw(self.inner.mathFloatingPointFunctions().0)
}
pub fn set_math_floating_point_functions(&self, value: MathFloatingPointFunctions) {
self.inner
.setMathFloatingPointFunctions(MTLMathFloatingPointFunctions(value.as_raw()));
}
#[must_use]
pub fn language_version(&self) -> LanguageVersion {
LanguageVersion::from_system_raw(self.inner.languageVersion().0)
}
pub fn set_language_version(&self, value: LanguageVersion) {
self.inner
.setLanguageVersion(MTLLanguageVersion(value.as_raw()));
}
#[must_use]
pub fn library_type(&self) -> LibraryType {
LibraryType::from_system_raw(self.inner.libraryType().0)
}
pub fn set_library_type(&self, value: LibraryType) {
self.inner.setLibraryType(MTLLibraryType(value.as_raw()));
}
#[must_use]
pub fn optimization_level(&self) -> LibraryOptimizationLevel {
LibraryOptimizationLevel::from_system_raw(self.inner.optimizationLevel().0)
}
pub fn set_optimization_level(&self, value: LibraryOptimizationLevel) {
self.inner
.setOptimizationLevel(MTLLibraryOptimizationLevel(value.as_raw()));
}
#[must_use]
pub fn compile_symbol_visibility(&self) -> CompileSymbolVisibility {
CompileSymbolVisibility::from_system_raw(self.inner.compileSymbolVisibility().0)
}
pub fn set_compile_symbol_visibility(&self, value: CompileSymbolVisibility) {
self.inner
.setCompileSymbolVisibility(MTLCompileSymbolVisibility(value.as_raw()));
}
pub fn libraries(
&self,
) -> Result<Vec<crate::metal::generated_object_types::metal::DynamicLibrary>, Error> {
if !self.inner.respondsToSelector(sel!(libraries)) {
return Err(Error::unsupported(
"compile-option libraries are unavailable",
));
}
Ok(self
.inner
.libraries()
.map(|values| {
values
.into_iter()
.map(|value| {
let value = unsafe { Retained::cast_unchecked(value) };
crate::metal::generated_object_types::metal::DynamicLibrary::from_inner(
value,
)
})
.collect()
})
.unwrap_or_default())
}
pub fn set_libraries(
&self,
libraries: Option<&[crate::metal::generated_object_types::metal::DynamicLibrary]>,
) -> Result<(), Error> {
if !self.inner.respondsToSelector(sel!(setLibraries:)) {
return Err(Error::unsupported(
"compile-option libraries are unavailable",
));
}
let libraries = libraries.map(|values| {
let retained: Vec<Retained<ProtocolObject<dyn MTLDynamicLibrary>>> = values
.iter()
.map(|value| {
let value = unsafe {
Retained::retain(std::ptr::from_ref(value.as_inner()).cast_mut())
}
.expect("a borrowed Objective-C wrapper cannot be null");
unsafe { Retained::cast_unchecked(value) }
})
.collect();
NSArray::from_retained_slice(&retained)
});
self.inner.setLibraries(libraries.as_deref());
Ok(())
}
pub fn preprocessor_macros(&self) -> Result<HashMap<String, String>, Error> {
if !self.inner.respondsToSelector(sel!(preprocessorMacros)) {
return Err(Error::unsupported(
"compile-option preprocessor macros are unavailable",
));
}
let mut result = HashMap::new();
if let Some(values) = self.inner.preprocessorMacros() {
for key in values.keys() {
if let Some(value) = values.objectForKey(&key) {
let description: Retained<NSString> =
unsafe { msg_send![&*value, description] };
result.insert(key.to_string(), description.to_string());
}
}
}
Ok(result)
}
pub fn set_preprocessor_macros(
&self,
macros: Option<&HashMap<String, String>>,
) -> Result<(), Error> {
if let Some(macros) = macros
&& macros
.iter()
.any(|(key, value)| key.is_empty() || key.contains('\0') || value.contains('\0'))
{
return Err(Error::invalid_argument("preprocessor macro is invalid"));
}
if !self.inner.respondsToSelector(sel!(setPreprocessorMacros:)) {
return Err(Error::unsupported(
"compile-option preprocessor macros are unavailable",
));
}
let dictionary = macros.map(|values| {
let keys: Vec<Retained<NSString>> = values
.keys()
.map(|value| NSString::from_str(value))
.collect();
let objects: Vec<Retained<NSString>> = values
.values()
.map(|value| NSString::from_str(value))
.collect();
let key_refs: Vec<&NSString> = keys.iter().map(|value| &**value).collect();
NSDictionary::<NSString, NSString>::from_retained_objects(&key_refs, &objects)
});
unsafe {
let _: () = msg_send![&*self.inner, setPreprocessorMacros: dictionary.as_deref()];
}
Ok(())
}
pub fn floating_point_conversion_rounding_mode(
&self,
) -> Result<FloatingPointConversionRoundingMode, Error> {
if !self
.inner
.respondsToSelector(sel!(floatingPointConversionRoundingMode))
{
return Err(Error::unsupported(
"floating-point conversion rounding mode is unavailable",
));
}
let raw: isize = unsafe { msg_send![&*self.inner, floatingPointConversionRoundingMode] };
Ok(FloatingPointConversionRoundingMode::from_system_raw(raw))
}
pub fn set_floating_point_conversion_rounding_mode(
&self,
value: FloatingPointConversionRoundingMode,
) -> Result<(), Error> {
if !value.is_valid() {
return Err(Error::invalid_argument(
"floating-point conversion rounding mode is invalid",
));
}
if !self
.inner
.respondsToSelector(sel!(setFloatingPointConversionRoundingMode:))
{
return Err(Error::unsupported(
"floating-point conversion rounding mode is unavailable",
));
}
let raw = value.as_raw();
unsafe {
let _: () = msg_send![&*self.inner, setFloatingPointConversionRoundingMode: raw];
}
Ok(())
}
}
impl Default for CompileOptions {
fn default() -> Self {
Self::new()
}
}
impl crate::metal::generated_object_types::metal::ComputePipelineDescriptor {
pub fn required_threads_per_threadgroup(&self) -> Result<Size, Error> {
let object = self.as_inner();
let available: bool =
unsafe { msg_send![object, respondsToSelector: sel!(requiredThreadsPerThreadgroup)] };
if !available {
return Err(Error::unsupported(
"required compute threadgroup dimensions are unavailable",
));
}
let value: objc2_metal::MTLSize =
unsafe { msg_send![object, requiredThreadsPerThreadgroup] };
Ok(Size::new(value.width, value.height, value.depth))
}
pub fn set_required_threads_per_threadgroup(&self, value: Size) -> Result<(), Error> {
let all_zero = value.width == 0 && value.height == 0 && value.depth == 0;
let all_non_zero = value.width != 0 && value.height != 0 && value.depth != 0;
if !all_zero && !all_non_zero {
return Err(Error::invalid_argument(
"required threadgroup dimensions must be all zero or all non-zero",
));
}
let object = self.as_inner();
let available: bool = unsafe {
msg_send![object, respondsToSelector: sel!(setRequiredThreadsPerThreadgroup:)]
};
if !available {
return Err(Error::unsupported(
"required compute threadgroup dimensions are unavailable",
));
}
let value: objc2_metal::MTLSize = value.into();
unsafe {
let _: () = msg_send![object, setRequiredThreadsPerThreadgroup: value];
}
Ok(())
}
pub fn reset(&self) -> Result<(), Error> {
let object = self.as_inner();
let available: bool = unsafe { msg_send![object, respondsToSelector: sel!(reset)] };
if !available {
return Err(Error::unsupported(
"compute pipeline descriptor reset is unavailable",
));
}
unsafe {
let _: () = msg_send![object, reset];
}
Ok(())
}
}
impl crate::metal::generated_object_types::metal::FunctionConstantValues {
pub fn reset(&self) -> Result<(), Error> {
let object = self.as_inner();
let available: bool = unsafe { msg_send![object, respondsToSelector: sel!(reset)] };
if !available {
return Err(Error::unsupported("function constants are unavailable"));
}
unsafe {
let _: () = msg_send![object, reset];
}
Ok(())
}
pub fn set_constant_at_index(
&self,
data_type: DataType,
index: usize,
bytes: &[u8],
) -> Result<(), Error> {
let size = function_constant_abi_size(data_type).ok_or_else(|| {
Error::invalid_argument("data type is not a value-compatible function constant")
})?;
if bytes.len() != size {
return Err(Error::invalid_argument(format!(
"function constant requires {size} ABI bytes"
)));
}
validate_function_constant_bytes(data_type, bytes, 1)?;
let mut storage = aligned_constant_bytes(bytes)?;
let pointer: NonNull<std::ffi::c_void> = NonNull::new(storage.as_mut_ptr().cast())
.ok_or_else(|| Error::invalid_argument("function constant storage is empty"))?;
let object = self.as_inner();
let available: bool =
unsafe { msg_send![object, respondsToSelector: sel!(setConstantValue:type:atIndex:)] };
if !available {
return Err(Error::unsupported("function constants are unavailable"));
}
unsafe {
let _: () = msg_send![object,
setConstantValue: pointer,
r#type: data_type.as_raw(),
atIndex: index
];
}
Ok(())
}
pub fn set_constant_named(
&self,
data_type: DataType,
name: &str,
bytes: &[u8],
) -> Result<(), Error> {
if name.is_empty() || name.as_bytes().contains(&0) {
return Err(Error::invalid_argument("function constant name is invalid"));
}
let size = function_constant_abi_size(data_type).ok_or_else(|| {
Error::invalid_argument("data type is not a value-compatible function constant")
})?;
if bytes.len() != size {
return Err(Error::invalid_argument(format!(
"function constant requires {size} ABI bytes"
)));
}
validate_function_constant_bytes(data_type, bytes, 1)?;
let mut storage = aligned_constant_bytes(bytes)?;
let pointer: NonNull<std::ffi::c_void> = NonNull::new(storage.as_mut_ptr().cast())
.ok_or_else(|| Error::invalid_argument("function constant storage is empty"))?;
let name = NSString::from_str(name);
let object = self.as_inner();
let available: bool =
unsafe { msg_send![object, respondsToSelector: sel!(setConstantValue:type:withName:)] };
if !available {
return Err(Error::unsupported("function constants are unavailable"));
}
unsafe {
let _: () = msg_send![object,
setConstantValue: pointer,
r#type: data_type.as_raw(),
withName: &*name
];
}
Ok(())
}
pub fn set_constants(
&self,
data_type: DataType,
range: std::ops::Range<usize>,
bytes: &[u8],
) -> Result<(), Error> {
if range.start > range.end {
return Err(Error::invalid_argument(
"function constant range is invalid",
));
}
let count = range.end - range.start;
if count == 0 {
return Err(Error::invalid_argument("function constant range is empty"));
}
let size = function_constant_abi_size(data_type).ok_or_else(|| {
Error::invalid_argument("data type is not a value-compatible function constant")
})?;
let expected = size
.checked_mul(count)
.ok_or_else(|| Error::invalid_argument("function constant range overflow"))?;
if bytes.len() != expected {
return Err(Error::invalid_argument(format!(
"function constant range requires {expected} packed ABI bytes"
)));
}
validate_function_constant_bytes(data_type, bytes, count)?;
let mut storage = aligned_constant_bytes(bytes)?;
let pointer: NonNull<std::ffi::c_void> = NonNull::new(storage.as_mut_ptr().cast())
.ok_or_else(|| Error::invalid_argument("function constant storage is empty"))?;
let range = NSRange::new(range.start, count);
let object = self.as_inner();
let available: bool = unsafe {
msg_send![object, respondsToSelector: sel!(setConstantValues:type:withRange:)]
};
if !available {
return Err(Error::unsupported("function constants are unavailable"));
}
unsafe {
let _: () = msg_send![object,
setConstantValues: pointer,
r#type: data_type.as_raw(),
withRange: range
];
}
Ok(())
}
}