#![deny(unsafe_op_in_unsafe_fn)]
use dispatch2::DispatchData;
use objc2::rc::Retained;
use objc2::runtime::ProtocolObject;
use objc2_foundation::NSString;
use objc2_metal::{
MTLDevice as _, MTLPixelFormat, MTLRenderPipelineDescriptor, MTLRenderPipelineState,
};
use crate::metal::descriptors::{VertexAttr, VertexLayout, vertex_descriptor};
use crate::metal::post::fullscreen::{FullscreenBlend, build_slang_fullscreen_pipeline};
pub(super) fn ns_str(s: &str) -> Retained<NSString> {
NSString::from_str(s)
}
const OBJECT_COMMON_MSL: &str = include_str!("shaders/object_common.msl");
fn object_common(hot_reload: bool) -> std::borrow::Cow<'static, str> {
if hot_reload {
let path = concat!(
env!("CARGO_MANIFEST_DIR"),
"/src/metal/shaders/object_common.msl"
);
match std::fs::read_to_string(path) {
Ok(s) => return std::borrow::Cow::Owned(s),
Err(e) => {
tracing::debug!("hot-reload: falling back to embedded object_common.msl ({e})");
}
}
}
std::borrow::Cow::Borrowed(OBJECT_COMMON_MSL)
}
pub(super) fn shader_source(hot_reload: bool, name: &str) -> std::borrow::Cow<'static, str> {
let embedded: &'static str = match name {
"cull.metal" => include_str!("shaders/cull.metal"),
"main.metal" => include_str!("shaders/main.metal"),
_ => panic!(
"shader_source: '{name}' is not a registered Metal shader. Add an \
`include_str!(\"shaders/{name}\")` arm to shader_source in \
metal/pipeline.rs -- every shipped shader must be registered."
),
};
let src = if hot_reload {
let path = format!("{}/src/metal/shaders/{}", env!("CARGO_MANIFEST_DIR"), name);
match std::fs::read_to_string(&path) {
Ok(s) => std::borrow::Cow::Owned(s),
Err(e) => {
tracing::debug!(
"hot-reload: falling back to embedded source for {} ({})",
name,
e
);
std::borrow::Cow::Borrowed(embedded)
}
}
} else {
std::borrow::Cow::Borrowed(embedded)
};
if src.contains("{OBJECT_DATA}") {
return std::borrow::Cow::Owned(src.replace("{OBJECT_DATA}", &object_common(hot_reload)));
}
src
}
pub(super) fn shader_library(
device: &ProtocolObject<dyn objc2_metal::MTLDevice>,
hot_reload: bool,
name: &str,
) -> Result<Retained<ProtocolObject<dyn objc2_metal::MTLLibrary>>, String> {
if !hot_reload && let Some(bytes) = crate::metal::metallib::embedded_metallib(name) {
return load_library(device, bytes)
.map_err(|e| format!("{name}: failed to load precompiled metallib: {e}"));
}
let msl = shader_source(hot_reload, name);
let options = objc2_metal::MTLCompileOptions::new();
device
.newLibraryWithSource_options_error(&ns_str(msl.as_ref()), Some(&options))
.map_err(|e| format!("{name}: shader compile error: {e:?}"))
}
pub(super) fn stage_library(
device: &ProtocolObject<dyn objc2_metal::MTLDevice>,
hot_reload: bool,
bytes: &[u8],
) -> Result<Retained<ProtocolObject<dyn objc2_metal::MTLLibrary>>, String> {
if bytes.is_empty() {
return shader_library(device, hot_reload, "main.metal");
}
load_library(device, bytes)
}
pub(super) fn load_library(
device: &ProtocolObject<dyn objc2_metal::MTLDevice>,
bytes: &[u8],
) -> Result<Retained<ProtocolObject<dyn objc2_metal::MTLLibrary>>, String> {
let data = DispatchData::from_bytes(bytes);
device
.newLibraryWithData_error(&data)
.map_err(|e| format!("{:?}", e))
}
pub(super) fn build_text_pipeline(
device: &ProtocolObject<dyn objc2_metal::MTLDevice>,
swap_pixel_format: MTLPixelFormat,
hot_reload: bool,
) -> Result<Retained<ProtocolObject<dyn MTLRenderPipelineState>>, String> {
use objc2_metal::{MTLBlendFactor, MTLVertexFormat, MTLVertexStepFunction};
let vert_fn = crate::metal::slang_shaders::entry_function(
device,
&crate::metal::slang_shaders::TEXT_VERT,
hot_reload,
)?;
let frag_fn = crate::metal::slang_shaders::entry_function(
device,
&crate::metal::slang_shaders::TEXT_FRAG,
hot_reload,
)?;
let vert_desc = vertex_descriptor(
&[
VertexAttr {
index: 0,
format: MTLVertexFormat::Float2,
offset: 0,
buffer_index: 1,
},
VertexAttr {
index: 1,
format: MTLVertexFormat::Float2,
offset: 8,
buffer_index: 1,
},
VertexAttr {
index: 2,
format: MTLVertexFormat::Float3,
offset: 16,
buffer_index: 1,
},
VertexAttr {
index: 3,
format: MTLVertexFormat::Float,
offset: 28,
buffer_index: 1,
},
],
&[VertexLayout {
buffer_index: 1,
stride: 32,
step: MTLVertexStepFunction::PerVertex,
}],
);
let pipeline_desc = MTLRenderPipelineDescriptor::new();
pipeline_desc.setVertexDescriptor(Some(&vert_desc));
pipeline_desc.setVertexFunction(Some(&vert_fn));
pipeline_desc.setFragmentFunction(Some(&frag_fn));
pipeline_desc.setRasterSampleCount(1);
unsafe {
let ca = pipeline_desc.colorAttachments().objectAtIndexedSubscript(0);
ca.setPixelFormat(swap_pixel_format);
ca.setBlendingEnabled(true);
ca.setSourceRGBBlendFactor(MTLBlendFactor::SourceAlpha);
ca.setDestinationRGBBlendFactor(MTLBlendFactor::OneMinusSourceAlpha);
ca.setSourceAlphaBlendFactor(MTLBlendFactor::One);
ca.setDestinationAlphaBlendFactor(MTLBlendFactor::OneMinusSourceAlpha);
}
device
.newRenderPipelineStateWithDescriptor_error(&pipeline_desc)
.map_err(|e| format!("failed to create text pipeline state: {:?}", e))
}
pub(super) fn build_post_pipeline(
device: &ProtocolObject<dyn objc2_metal::MTLDevice>,
swap_pixel_format: MTLPixelFormat,
hot_reload: bool,
) -> Result<Retained<ProtocolObject<dyn MTLRenderPipelineState>>, String> {
build_slang_fullscreen_pipeline(
device,
&super::slang_shaders::COMPOSITE_FRAG,
swap_pixel_format,
FullscreenBlend::Replace,
hot_reload,
)
}
#[cfg(test)]
mod shader_source_tests {
use super::shader_source;
#[test]
fn embedded_path_splices_object_data() {
let s = shader_source(false, "main.metal");
assert!(s.contains("vertex VertexOut vertex_main("));
assert!(!s.contains("{OBJECT_DATA}"));
}
#[test]
#[should_panic(expected = "not a registered Metal shader")]
fn unknown_name_panics() {
let _ = shader_source(false, "nope.metal");
}
#[test]
#[should_panic(expected = "not a registered Metal shader")]
fn unknown_name_panics_even_with_hot_reload() {
let _ = shader_source(true, "nope.metal");
}
#[test]
fn hot_reload_prefers_disk_when_present() {
let s = shader_source(true, "main.metal");
assert!(s.contains("vertex VertexOut vertex_main("));
}
#[test]
fn main_metal_reflection_cut_matches_canonical() {
let src = shader_source(false, "main.metal");
let decl = src
.lines()
.find(|l| l.contains("constant float REFL_RESOLVE_CUT"))
.expect("REFL_RESOLVE_CUT declaration in main.metal");
let value: f32 = decl
.split(';')
.next()
.and_then(|head| head.split('=').nth(1))
.map(str::trim)
.and_then(|s| s.parse().ok())
.expect("parse REFL_RESOLVE_CUT value from main.metal");
assert_eq!(
value,
concinnity_core::gfx::ssr::REFLECTION_ROUGHNESS_CUT,
"main.metal REFL_RESOLVE_CUT must equal REFLECTION_ROUGHNESS_CUT"
);
}
#[test]
fn object_data_shaders_splice_the_shared_record() {
for name in ["main.metal", "cull.metal"] {
for hot_reload in [false, true] {
let src = shader_source(hot_reload, name);
assert!(
src.contains("struct GpuObjectData"),
"{name}: object record missing (hot_reload = {hot_reload})"
);
assert!(
!src.contains("{OBJECT_DATA}"),
"{name}: left the OBJECT_DATA marker (hot_reload = {hot_reload})"
);
}
}
}
#[test]
fn shipped_shaders_are_registered() {
const ASSEMBLED_ELSEWHERE: &[&str] = &[
"raymarch_helpers.metal",
"raymarch_shadow.metal",
"raymarch_template.metal",
"raymarch_volumetric_template.metal",
];
let dir = concat!(env!("CARGO_MANIFEST_DIR"), "/src/metal/shaders");
let mut checked = 0usize;
for entry in std::fs::read_dir(dir).expect("read shaders dir") {
let file_name = entry.expect("dir entry").file_name();
let name = file_name.to_str().expect("utf8 shader filename");
if !name.ends_with(".metal") || ASSEMBLED_ELSEWHERE.contains(&name) {
continue;
}
assert!(
!shader_source(false, name).trim().is_empty(),
"{name}: shader_source(false) returned empty source",
);
assert!(
!shader_source(true, name).trim().is_empty(),
"{name}: shader_source(true) returned empty source",
);
checked += 1;
}
assert!(checked > 0, "no .metal shaders found under {dir}");
}
}