use super::super::shared::{ShaderDesc, ShaderStageCompileDesc};
use super::super::{DeviceHandle, ShaderHandle};
use super::types::{MetalState, ShaderState};
use crate::slang::{ShaderTarget, SlangCompiler, SlangStage};
use ::metal as mtl;
use anyhow::{Context, Result};
use mtl::{Device as MTLDevice, Library};
use std::collections::HashMap;
pub(super) fn patch_compute_msl(msl: &str) -> String {
let s = patch_compute_msl_entry_point_params(msl);
patch_msl_threadgroup_copies(&s)
}
fn patch_compute_msl_entry_point_params(msl: &str) -> String {
if !msl.contains("EntryPointParams_0") || msl.contains("EntryPointParams_0 constant*") {
return msl.to_string();
}
const EP_NEEDLE: &str = "EntryPointParams_0 ";
let mut search_from = 0usize;
let var_name = loop {
let rel = match msl[search_from..].find(EP_NEEDLE) {
Some(p) => p,
None => return msl.to_string(),
};
let abs = search_from + rel;
let after = &msl[abs + EP_NEEDLE.len()..];
let var_end = after
.find(|c: char| !c.is_alphanumeric() && c != '_')
.unwrap_or(after.len());
let var = &after[..var_end];
if !var.is_empty() && after[var_end..].starts_with(')') {
break var.to_string();
}
search_from = abs + EP_NEEDLE.len() + 1;
};
let old_param = format!("{}{}{}", EP_NEEDLE, var_name, ")");
let new_param = format!("{}constant* {} [[buffer(1)]])", EP_NEEDLE, var_name);
let mut s = msl.replacen(&old_param, &new_param, 1);
let dot = format!("{}.", var_name);
let arrow = format!("{}->", var_name);
s = s.replace(&dot, &arrow);
let assign = format!("= {};", var_name);
let deref = format!("= *{};", var_name);
s = s.replace(&assign, &deref);
s
}
fn patch_msl_threadgroup_copies(msl: &str) -> String {
if !msl.contains("= *kernelContext") {
return msl.to_string();
}
let mut tg_ref_vars: std::collections::HashSet<String> = std::collections::HashSet::new();
let mut result = String::with_capacity(msl.len() + 128);
for line in msl.lines() {
let trimmed = line.trim_start();
let is_p1_explicit = trimmed.starts_with("thread array<") && line.contains("= *kernelContext");
let is_p1_implicit =
!trimmed.starts_with("thread") && trimmed.starts_with("array<") && line.contains("= *kernelContext");
let is_p2 = trimmed.starts_with("thread array<")
&& !line.contains("= *kernelContext")
&& line.find(" = ").is_some_and(|eq| {
let rhs = line[eq + 3..].trim_end_matches(';').trim();
tg_ref_vars.contains(rhs)
});
let patched = if is_p1_explicit || is_p1_implicit || is_p2 {
let s = if is_p1_implicit {
let arr_pos = line.find("array<").unwrap();
format!("{}thread {}", &line[..arr_pos], &line[arr_pos..])
} else {
line.to_string()
};
let s = s.replacen("thread array<", "threadgroup array<", 1);
if let (Some(gt), Some(eq)) = (s.rfind("> "), s.find(" = ")) {
if gt < eq {
let mut out = String::with_capacity(s.len() + 1);
out.push_str(&s[..gt + 1]);
out.push('&');
out.push_str(&s[gt + 1..]);
let after_gt = &s[gt + 2..]; if let Some(end) = after_gt.find(|c: char| !c.is_alphanumeric() && c != '_') {
let var = &after_gt[..end];
if !var.is_empty() {
tg_ref_vars.insert(var.to_string());
}
}
out
} else {
s
}
} else {
s
}
} else {
line.to_string()
};
result.push_str(&patched);
result.push('\n');
}
if !msl.ends_with('\n') && result.ends_with('\n') {
result.pop();
}
result
}
pub(super) fn patch_vertex_stage_in_pointer(msl: &str) -> String {
const STAGE_IN: &str = "[[stage_in]]";
let Some(si_pos) = msl.find(STAGE_IN) else {
return msl.to_string();
};
let line_start = msl[..si_pos].rfind('\n').map_or(0, |p| p + 1);
let line_end = msl[si_pos..].find('\n').map_or(msl.len(), |p| si_pos + p);
let line = &msl[line_start..line_end];
if !line.contains("thread*") {
return msl.to_string();
}
let before = &msl[..si_pos];
let prefix_end = before.trim_end().len();
let si_end = si_pos + STAGE_IN.len();
let mut result = String::with_capacity(msl.len());
result.push_str(&msl[..prefix_end]);
result.push_str(&msl[si_end..]);
result
}
pub(super) fn patch_vertex_msl_entry_point_params(msl: &str) -> String {
if !msl.contains("EntryPointParams_0") {
return msl.to_string();
}
if msl.contains("EntryPointParams_0 constant*") && msl.contains("[[buffer(1)]]") {
return msl.to_string();
}
const SIG_NEEDLE: &str = "[[buffer(0)]])";
const SIG_REPLACEMENT: &str = "[[buffer(0)]], EntryPointParams_0 constant* _goldy_ep [[buffer(1)]])";
let patched = if let Some(pos) = msl.find(SIG_NEEDLE) {
let mut s = String::with_capacity(msl.len() + 80);
s.push_str(&msl[..pos]);
s.push_str(SIG_REPLACEMENT);
s.push_str(&msl[pos + SIG_NEEDLE.len()..]);
s
} else {
return msl.to_string();
};
const ASSIGN_NEEDLE: &str = ")->gGoldy_0 = ";
if let Some(arrow_pos) = patched.find(ASSIGN_NEEDLE) {
let prefix = &patched[..arrow_pos];
if let Some(amp_pos) = prefix.rfind("(&") {
let kctx_name = &prefix[amp_pos + 2..]; let after_arrow = &patched[arrow_pos + ASSIGN_NEEDLE.len()..];
if let Some(semi_rel) = after_arrow.find(';') {
let semi_abs = arrow_pos + ASSIGN_NEEDLE.len() + semi_rel;
let injection = format!("\n (&{})->entryPointParams_0 = _goldy_ep;", kctx_name);
let mut result = String::with_capacity(patched.len() + injection.len());
result.push_str(&patched[..=semi_abs]);
result.push_str(&injection);
result.push_str(&patched[semi_abs + 1..]);
return result;
}
}
}
patched
}
fn compile_stage_with_reflection(
slang_compiler: &SlangCompiler,
device: &MTLDevice,
desc: &ShaderStageCompileDesc<'_>,
) -> Result<(Library, Option<crate::slang::ShaderReflection>)> {
let compile_outcome = slang_compiler.compile_bindless_with_reflection_and_defines(
desc.slang_source,
ShaderTarget::Metal,
&[(desc.entry_point, desc.stage)],
desc.search_paths,
desc.extra_defines,
desc.layout_checks,
desc.optimization_level,
);
let result = compile_outcome.with_context(|| format!("Failed to compile {} shader stage", desc.entry_point))?;
if !result.reflection.parameter_blocks.is_empty() {
tracing::debug!(
"Shader {} has {} ParameterBlock(s):",
desc.entry_point,
result.reflection.parameter_blocks.len()
);
for pb in &result.reflection.parameter_blocks {
tracing::debug!(
" - {} at slot {} (size={}, alignment={}, fields={})",
pb.name,
pb.binding_slot,
pb.size,
pb.alignment,
pb.fields.len()
);
for field in &pb.fields {
tracing::debug!(
" - {}: {:?} at offset {} (size={})",
field.name,
field.resource_kind,
field.offset,
field.size
);
}
}
}
let raw_msl = result.shader.as_str().context("Failed to get MSL source")?.to_string();
let msl_source = if desc.stage == SlangStage::Vertex {
let patched = patch_vertex_msl_entry_point_params(&raw_msl);
if patched != raw_msl {
tracing::debug!(
"Applied vertex EntryPointParams [[buffer(1)]] patch for {}",
desc.entry_point
);
}
let patched2 = patch_vertex_stage_in_pointer(&patched);
if patched2 != patched {
tracing::debug!(
"Applied vertex [[stage_in]] pointer-to-value patch for {}",
desc.entry_point
);
}
patched2
} else if desc.stage == SlangStage::Compute {
let patched = patch_compute_msl(&raw_msl);
if patched != raw_msl {
tracing::debug!(
"Applied compute MSL patches (EntryPointParams / threadgroup copy) for {}",
desc.entry_point
);
}
patched
} else {
raw_msl
};
tracing::debug!("Compiled MSL {} shader ({} bytes)", desc.entry_point, msl_source.len());
if let Ok(dump_dir) = std::env::var("GOLDY_DUMP_SHADERS") {
use std::io::Write;
use std::sync::atomic::{AtomicU32, Ordering};
static DUMP_IDX: AtomicU32 = AtomicU32::new(0);
let idx = DUMP_IDX.fetch_add(1, Ordering::Relaxed);
let dir = std::path::Path::new(&dump_dir);
let _ = std::fs::create_dir_all(dir);
let filename = format!("{:03}_{}.metal", idx, desc.entry_point);
if let Ok(mut f) = std::fs::File::create(dir.join(&filename)) {
let _ = f.write_all(msl_source.as_bytes());
tracing::info!("Dumped MSL to {}/{}", dump_dir, filename);
}
}
let library = device
.new_library_with_source(&msl_source, &mtl::CompileOptions::new())
.map_err(|e| anyhow::anyhow!("Failed to create Metal library for {}: {}", desc.entry_point, e))?;
Ok((library, Some(result.reflection)))
}
pub(super) fn ensure_stage_compiled(
state: &mut MetalState,
shader_handle: ShaderHandle,
stage: SlangStage,
) -> Result<()> {
struct CompileScratch {
device_handle: DeviceHandle,
slang_source: String,
search_paths: Vec<String>,
optimization_level: crate::types::OptimizationLevel,
defines: Vec<(String, String)>,
layout_checks: Vec<crate::slang::OwnedLayoutCheck>,
entry_point: &'static str,
}
let maybe_scratch: Option<CompileScratch> = {
let shaders = &state.shaders;
let shader = shaders.get(&shader_handle).context("Invalid shader handle")?;
let (entry_point, need_compile) = match stage {
SlangStage::Vertex => ("vs_main", shader.vertex_library.is_none()),
SlangStage::Fragment => ("fs_main", shader.fragment_library.is_none()),
SlangStage::Compute => ("cs_main", shader.compute_library.is_none()),
_ => anyhow::bail!("Metal backend only supports Vertex, Fragment, and Compute stages"),
};
if !need_compile {
None
} else {
Some(CompileScratch {
device_handle: shader.device_handle,
slang_source: shader.slang_source.clone(),
search_paths: shader.search_paths.clone(),
optimization_level: shader.optimization_level,
defines: shader.defines.clone(),
layout_checks: shader.layout_checks.clone(),
entry_point,
})
}
};
let Some(scratch) = maybe_scratch else {
return Ok(());
};
let mtl_dev = state
.devices
.get(&scratch.device_handle)
.context("Shader's device no longer valid")?
.device
.clone();
let search_path_refs: Vec<&str> = scratch.search_paths.iter().map(|s| s.as_str()).collect();
let extra_defines: Vec<(&str, &str)> = scratch.defines.iter().map(|(k, v)| (k.as_str(), v.as_str())).collect();
let compile_desc = ShaderStageCompileDesc {
slang_source: &scratch.slang_source,
search_paths: &search_path_refs,
entry_point: scratch.entry_point,
stage,
extra_defines: &extra_defines,
layout_checks: &scratch.layout_checks,
optimization_level: scratch.optimization_level,
};
let compiler = state.slang_compiler_mut_or_init()?;
let (library, reflection) = compile_stage_with_reflection(compiler, &mtl_dev, &compile_desc)?;
let shader = state
.shaders
.get_mut(&shader_handle)
.expect("shader handle must be valid after ensure_stage_compiled");
match stage {
SlangStage::Vertex => shader.vertex_library = Some(library),
SlangStage::Fragment => shader.fragment_library = Some(library),
SlangStage::Compute => shader.compute_library = Some(library),
_ => unreachable!("stage already validated"),
}
if shader.reflection.is_none() {
let reflection = reflection.map(|mut r| {
if r.push_constant_categories.is_empty() {
r.push_constant_categories =
crate::slang::virtual_main::extract_push_constant_categories(&scratch.slang_source);
}
r
});
shader.reflection = reflection;
}
if !scratch.layout_checks.is_empty() {
shader.layout_checks.clear();
}
Ok(())
}
pub(super) fn create(
devices: &HashMap<DeviceHandle, super::types::SharedLogicalDevice>,
shaders: &mut HashMap<ShaderHandle, ShaderState>,
next_shader_handle: &mut ShaderHandle,
desc: ShaderDesc<'_>,
) -> Result<ShaderHandle> {
devices.get(&desc.device).context("Invalid device handle")?;
let handle = *next_shader_handle;
*next_shader_handle += 1;
shaders.insert(
handle,
ShaderState {
device_handle: desc.device,
slang_source: desc.slang_source.to_string(),
search_paths: desc.search_paths.iter().map(|s| s.to_string()).collect(),
defines: desc
.defines
.iter()
.map(|(k, v)| (k.to_string(), v.to_string()))
.collect(),
optimization_level: desc.optimization_level,
vertex_library: None,
fragment_library: None,
compute_library: None,
reflection: None,
layout_checks: desc.layout_checks,
},
);
tracing::debug!("Created shader handle {} (compilation deferred)", handle);
Ok(handle)
}
pub(super) fn destroy(
_devices: &HashMap<DeviceHandle, super::types::SharedLogicalDevice>,
shaders: &mut HashMap<ShaderHandle, ShaderState>,
shader_handle: ShaderHandle,
) {
shaders.remove(&shader_handle);
}
#[cfg(test)]
mod tests {
use super::{
patch_compute_msl, patch_compute_msl_entry_point_params, patch_msl_threadgroup_copies,
patch_vertex_msl_entry_point_params, patch_vertex_stage_in_pointer,
};
#[test]
fn stage_in_stripped_from_internal_thread_pointer_fn() {
let msl = concat!(
"StaticVarying_0 vs_main_0(",
"const StaticVertexIn_0 thread* _pt0_0 [[stage_in]], ",
"KernelContext_0 thread* kernelContext_4)\n",
"{\n return _pt0_0;\n}\n",
"[[vertex]] vs_main_Result_0 vs_main(",
"vertexInput_0 _S19 [[stage_in]], ",
"GoldyBindlessResources_default_0 constant* gGoldy_1 [[buffer(0)]], ",
"EntryPointParams_0 constant* _goldy_ep [[buffer(1)]])\n",
"{\n return _S19;\n}\n",
);
let out = patch_vertex_stage_in_pointer(msl);
assert!(
out.contains("thread* _pt0_0,"),
"thread pointer parameter should remain without [[stage_in]]"
);
assert!(
!out.contains("_pt0_0 [[stage_in]]"),
"[[stage_in]] must be stripped from the thread* pointer parameter"
);
assert!(
out.contains("_S19 [[stage_in]]"),
"[[stage_in]] on the real vertex entry point must not be touched"
);
}
#[test]
fn stage_in_preserved_when_thread_star_only_in_helper_fn() {
let msl = concat!(
"VertexOutput_0 _goldy_user_vs_main_0(const VertexInput_0 thread* input_0)\n",
"{\n return *input_0;\n}\n",
"[[vertex]] vs_main_Result_0 vs_main(vertexInput_0 _S1 [[stage_in]])\n",
"{\n return _S1;\n}\n",
);
let out = patch_vertex_stage_in_pointer(msl);
assert_eq!(
out, msl,
"MSL must be unchanged when [[stage_in]] is only on a value-type entry point"
);
}
#[test]
fn stage_in_noop_when_absent() {
let msl = "[[kernel]] void cs_main(device uint* buf [[buffer(0)]])\n{\n}\n";
let out = patch_vertex_stage_in_pointer(msl);
assert_eq!(out, msl, "MSL without [[stage_in]] must pass through unchanged");
}
#[test]
fn vertex_ep_params_injected_when_buffer1_missing() {
let msl = concat!(
"struct EntryPointParams_0 { uint _bw0_0; };\n",
"struct KernelContext_0 {\n",
" GoldyBindlessResources_default_0 constant* gGoldy_0;\n",
" EntryPointParams_0 constant* entryPointParams_0;\n",
"};\n",
"[[vertex]] vs_main_Result_0 vs_main(",
"vertexInput_0 _S19 [[stage_in]], ",
"GoldyBindlessResources_default_0 constant* gGoldy_1 [[buffer(0)]])\n",
"{\n",
" KernelContext_0 kernelContext_5;\n",
" (&kernelContext_5)->gGoldy_0 = gGoldy_1;\n",
" return _S19;\n",
"}\n",
);
let out = patch_vertex_msl_entry_point_params(msl);
assert!(
out.contains("EntryPointParams_0 constant* _goldy_ep [[buffer(1)]])"),
"[[buffer(1)]] parameter must be injected into the entry-point signature"
);
assert!(
out.contains("(&kernelContext_5)->entryPointParams_0 = _goldy_ep;"),
"entryPointParams_0 field must be wired to _goldy_ep in the entry-point body"
);
assert!(
out.contains("[[buffer(0)]]"),
"[[buffer(0)]] binding must remain after the patch"
);
}
#[test]
fn vertex_ep_params_noop_when_already_correct() {
let msl = concat!(
"struct EntryPointParams_0 { uint _bw0_0; };\n",
"[[vertex]] vs_main_Result_0 vs_main(",
"vertexInput_0 _S19 [[stage_in]], ",
"GoldyBindlessResources_default_0 constant* gGoldy_1 [[buffer(0)]], ",
"EntryPointParams_0 constant* _goldy_ep [[buffer(1)]])\n",
"{\n return _S19;\n}\n",
);
let out = patch_vertex_msl_entry_point_params(msl);
assert_eq!(out, msl, "MSL with correct [[buffer(1)]] must pass through unchanged");
}
#[test]
fn vertex_ep_params_noop_when_no_entry_point_params() {
let msl = concat!(
"[[vertex]] vs_main_Result_0 vs_main(vertexInput_0 _S1 [[stage_in]])\n",
"{\n return _S1;\n}\n",
);
let out = patch_vertex_msl_entry_point_params(msl);
assert_eq!(out, msl, "MSL without EntryPointParams_0 must pass through unchanged");
}
#[test]
fn vertex_ep_params_noop_for_correctly_bound_fragment_shader() {
let msl = concat!(
"struct EntryPointParams_0 { uint _bw0_0; };\n",
"[[fragment]] pixelOutput_0 fs_main(",
"pixelInput_0 _S11 [[stage_in]], ",
"GoldyBindlessResources_default_0 constant* gGoldy_1 [[buffer(0)]], ",
"EntryPointParams_0 constant* entryPointParams_1 [[buffer(1)]])\n",
"{\n return _S11;\n}\n",
);
let out = patch_vertex_msl_entry_point_params(msl);
assert_eq!(
out, msl,
"Fragment shader with correct [[buffer(1)]] must pass through unchanged"
);
}
#[test]
fn compute_ep_params_patched_when_missing_constant_ptr() {
let msl = concat!(
"struct EntryPointParams_0 { uint _bw0_0; };\n",
"[[kernel]] void cs_main(\n",
" device uint* buf [[buffer(0)]],\n",
" EntryPointParams_0 epVar0)\n",
"{\n",
" uint v = epVar0._bw0_0;\n",
" uint w = epVar0._bw0_0;\n",
" KernelContext_0 kc2 = epVar0;\n",
"}\n",
);
let out = patch_compute_msl_entry_point_params(msl);
assert!(
out.contains("EntryPointParams_0 constant* epVar0 [[buffer(1)]])"),
"entry-point parameter must be fixed to `constant* epVar0 [[buffer(1)]]`"
);
assert!(
out.contains("epVar0->_bw0_0"),
"member access must be rewritten from . to ->"
);
assert!(
out.contains("= *epVar0;"),
"struct-copy assignment must be rewritten to dereference the pointer"
);
}
#[test]
fn compute_ep_params_noop_when_already_correct() {
let msl = concat!(
"struct EntryPointParams_0 { uint _bw0_0; };\n",
"[[kernel]] void cs_main(\n",
" device uint* buf [[buffer(0)]],\n",
" EntryPointParams_0 constant* epVar0 [[buffer(1)]])\n",
"{\n}\n",
);
let out = patch_compute_msl_entry_point_params(msl);
assert_eq!(out, msl, "MSL with correct constant* must pass through unchanged");
}
#[test]
fn compute_ep_params_noop_when_absent() {
let msl = concat!(
"[[kernel]] void cs_main(device uint* buf [[buffer(0)]])\n",
"{\n buf[0] = 42;\n}\n",
);
let out = patch_compute_msl_entry_point_params(msl);
assert_eq!(out, msl, "MSL without EntryPointParams_0 must pass through unchanged");
}
#[test]
fn threadgroup_copies_explicit_thread_patched() {
let msl = concat!(
"[[kernel]] void cs_main()\n",
"{\n",
" thread array<uint, int(64)> scratch = *kernelContext_0->shared_data;\n",
" scratch[0] = 1;\n",
"}\n",
);
let out = patch_msl_threadgroup_copies(msl);
assert!(
out.contains("threadgroup array<uint, int(64)>& scratch"),
"explicit thread copy must become a threadgroup reference"
);
assert!(
!out.contains("thread array<uint"),
"thread address space must be replaced by threadgroup"
);
}
#[test]
fn threadgroup_copies_implicit_thread_patched() {
let msl = concat!(
"[[kernel]] void cs_main()\n",
"{\n",
" array<uint, int(64)> scratch = *kernelContext_0->shared_data;\n",
" scratch[0] = 1;\n",
"}\n",
);
let out = patch_msl_threadgroup_copies(msl);
assert!(
out.contains("threadgroup array<uint, int(64)>& scratch"),
"implicit thread copy must become a threadgroup reference"
);
}
#[test]
fn threadgroup_copies_secondary_copy_also_patched() {
let msl = concat!(
"[[kernel]] void cs_main()\n",
"{\n",
" thread array<uint, int(64)> scratch = *kernelContext_0->shared_data;\n",
" thread array<uint, int(64)> local = scratch;\n",
" local[0] = 1;\n",
"}\n",
);
let out = patch_msl_threadgroup_copies(msl);
assert!(
out.contains("threadgroup array<uint, int(64)>& scratch"),
"pattern-1 variable must become a threadgroup reference"
);
assert!(
out.contains("threadgroup array<uint, int(64)>& local"),
"pattern-2 copy of a tg-ref variable must also become a threadgroup reference"
);
assert!(
!out.contains("thread array<uint"),
"no thread array copies must remain after patching"
);
}
#[test]
fn threadgroup_copies_noop_when_no_kernelcontext_copy() {
let msl = concat!(
"[[kernel]] void cs_main()\n",
"{\n",
" thread uint x = 0;\n",
" x = 1;\n",
"}\n",
);
let out = patch_msl_threadgroup_copies(msl);
assert_eq!(out, msl, "MSL without kernelContext copies must pass through unchanged");
}
#[test]
fn patch_compute_msl_fixes_both_bugs() {
let msl = concat!(
"struct EntryPointParams_0 { uint _bw0_0; };\n",
"[[kernel]] void cs_main(\n",
" device uint* buf [[buffer(0)]],\n",
" EntryPointParams_0 epVar0)\n",
"{\n",
" thread array<uint, int(32)> scratch = *kernelContext_0->shared_data;\n",
" buf[0] = epVar0._bw0_0 + scratch[0];\n",
"}\n",
);
let out = patch_compute_msl(msl);
assert!(
out.contains("EntryPointParams_0 constant* epVar0 [[buffer(1)]])"),
"Bug A: entry-point parameter must be fixed to constant*"
);
assert!(
out.contains("epVar0->_bw0_0"),
"Bug A: member access must be rewritten to ->"
);
assert!(
out.contains("threadgroup array<uint, int(32)>& scratch"),
"Bug B: thread array copy must become a threadgroup reference"
);
}
#[test]
fn patch_compute_msl_noop_when_clean() {
let msl = concat!(
"[[kernel]] void cs_main(device uint* buf [[buffer(0)]])\n",
"{\n",
" buf[0] = 42;\n",
"}\n",
);
let out = patch_compute_msl(msl);
assert_eq!(out, msl, "clean MSL must pass through the orchestrator unchanged");
}
}