#![cfg(feature = "metal")]
use cera::backend::metal::params::*;
use cera::backend::metal::shaders;
fn msl_struct_bytes(src: &str, name: &str) -> Option<usize> {
let start = src.find(&format!("struct {name} "))?;
let open = start + src[start..].find('{')?;
let close = open + src[open..].find("};")?;
let body = &src[open + 1..close];
let code: String = body
.lines()
.map(|l| l.split("//").next().unwrap_or(""))
.collect::<Vec<_>>()
.join(" ");
let mut bytes = 0usize;
for decl in code.split(';') {
let ty = decl.split_whitespace().next().unwrap_or("");
match ty {
"" => {}
"uint" | "int" | "float" => bytes += 4,
"half" => bytes += 2,
other => panic!(
"struct {name}: unrecognized MSL field type `{other}` — teach \
msl_struct_bytes its width before trusting this test"
),
}
}
Some(bytes)
}
fn cases() -> Vec<(usize, &'static str, &'static str, &'static str)> {
vec![
(
size_of::<QkNormRopeParams>(),
shaders::QK_NORM_ROPE,
"Params",
"QkNormRopeParams",
),
(
size_of::<QkNormRopeBatchParams>(),
shaders::QK_NORM_ROPE_BATCH,
"BatchParams",
"QkNormRopeBatchParams",
),
(
size_of::<KvShiftKParams>(),
shaders::KV_SHIFT,
"KParams",
"KvShiftKParams",
),
(
size_of::<KvCopyParams>(),
shaders::KV_SHIFT,
"CopyParams",
"KvCopyParams",
),
(
size_of::<GemmF32Params>(),
shaders::GEMM_F32,
"GemmParams",
"GemmF32Params",
),
(
size_of::<QuantGemmParams>(),
shaders::GEMM_Q4_0,
"GemmParams",
"QuantGemmParams (q4_0)",
),
(
size_of::<QuantGemmParams>(),
shaders::GEMM_Q8_0,
"GemmParams",
"QuantGemmParams (q8_0)",
),
(
size_of::<QuantGemmParams>(),
shaders::GEMM_Q4_K,
"GemmParams",
"QuantGemmParams (q4_k)",
),
(
size_of::<QuantGemmParams>(),
shaders::GEMM_Q5_K,
"GemmParams",
"QuantGemmParams (q5_k)",
),
(
size_of::<QuantGemmParams>(),
shaders::GEMM_Q6_K,
"GemmParams",
"QuantGemmParams (q6_k)",
),
(
size_of::<GemvBatchParams>(),
shaders::GEMV_Q4_0_BATCH,
"BatchParams",
"GemvBatchParams (q4_0)",
),
(
size_of::<GemvBatchParams>(),
shaders::GEMV_Q8_0_BATCH,
"BatchParams",
"GemvBatchParams (q8_0)",
),
(
size_of::<GemvQkvParams>(),
shaders::GEMV_Q4_0_FAST,
"ParamsQKV",
"GemvQkvParams",
),
(
size_of::<GemvRmsParams>(),
shaders::GEMV_Q4_0_FAST,
"RMSParams",
"GemvRmsParams",
),
(
size_of::<GemvSplitKParams>(),
shaders::GEMV_Q4_0_FAST,
"SplitKParams",
"GemvSplitKParams",
),
(
size_of::<FlashAttnParams>(),
shaders::FLASH_ATTENTION,
"Params",
"FlashAttnParams",
),
(
size_of::<FlashAttnParams>(),
shaders::ATTENTION,
"Params",
"FlashAttnParams (classic)",
),
(
size_of::<FlashAttnParams>(),
shaders::ATTENTION_GQA,
"Params",
"FlashAttnParams (gqa)",
),
(
size_of::<SplitAttnParams>(),
shaders::ATTENTION_SPLITK,
"SplitParams",
"SplitAttnParams",
),
(
size_of::<PrefillAttnParams>(),
shaders::ATTENTION_PREFILL,
"PrefillAttnParams",
"PrefillAttnParams",
),
(
size_of::<ElementwiseParams>(),
shaders::ELEMENTWISE,
"Params",
"ElementwiseParams",
),
(
size_of::<ScaleParams>(),
shaders::ELEMENTWISE,
"ScaleParams",
"ScaleParams",
),
(
size_of::<VitLinearParams>(),
shaders::VIT_LINEAR,
"Params",
"VitLinearParams",
),
(
size_of::<VitAttnParams>(),
shaders::VIT_ATTENTION,
"Params",
"VitAttnParams (scalar)",
),
(
size_of::<VitAttnParams>(),
shaders::VIT_ATTENTION_MMA,
"VitAttnParams",
"VitAttnParams (mma)",
),
(
size_of::<TqParams>(),
shaders::TURBOQUANT,
"TqParams",
"TqParams",
),
(
size_of::<TqAttnParams>(),
shaders::FLASH_ATTENTION_TQ,
"TqAttnParams",
"TqAttnParams",
),
]
}
fn slang_params_bytes(src: &str) -> Option<usize> {
let candidates: Vec<&str> = src
.lines()
.filter(|l| {
let t = l.trim_start();
!t.starts_with("//")
&& t.contains("StructuredBuffer<uint")
&& !t.contains("RWStructuredBuffer<uint")
})
.collect();
let vector_typed: Vec<&str> = candidates
.iter()
.copied()
.filter(|l| !l.contains("StructuredBuffer<uint>"))
.collect();
let decl = match (candidates.as_slice(), vector_typed.as_slice()) {
(_, [only]) => only,
([only], _) => only,
(many, _) => *many
.iter()
.find(|l| l.contains("_metal") || l.contains("metal_"))?,
};
let after = decl.split("StructuredBuffer<uint").nth(1)?;
let (comp_txt, rest) = after.split_once('>')?;
let components: usize = if comp_txt.is_empty() {
1
} else {
comp_txt.parse().ok()?
};
let name = rest.split_whitespace().next()?;
let code: String = src
.lines()
.map(|l| l.split("//").next().unwrap_or(""))
.collect::<Vec<_>>()
.join("\n");
let code = code.as_str();
let switch_starts: Vec<usize> = code
.match_indices("__target_switch")
.filter(|(i, _)| {
code[i + "__target_switch".len()..]
.trim_start()
.starts_with('{')
})
.map(|(i, _)| i)
.collect();
let bytes = code.as_bytes();
let mut cuts: Vec<(usize, usize)> = Vec::new();
for sw in switch_starts {
let open = sw + code[sw..].find('{')?;
let (mut depth, mut end) = (0usize, 0usize);
for (i, &b) in bytes.iter().enumerate().skip(open) {
match b {
b'{' => depth += 1,
b'}' => {
depth -= 1;
if depth == 0 {
end = i + 1;
break;
}
}
_ => {}
}
}
if end == 0 {
return None;
}
let body = &code[open..end];
let Some(d) = body.find("default:") else {
continue; };
match body.find("case metal:") {
Some(m) if m < d => cuts.push((open + d, end - 1)),
_ => return None,
}
}
let mut scope = code.to_string();
for (from, to) in cuts.into_iter().rev() {
scope.replace_range(from..to, "");
}
let scope = scope.as_str();
let mut seen = false;
let max_idx = scope
.match_indices(name)
.filter_map(|(i, _)| {
scope[i + name.len()..]
.strip_prefix('[')
.and_then(|t| t.split_once(']'))
.map(|(inner, _)| inner.trim().parse::<usize>())
})
.try_fold(0usize, |acc, parsed| {
seen = true;
parsed.ok().map(|n| acc.max(n))
})?;
seen.then(|| components * (max_idx + 1) * 4)
}
fn slang_cases() -> Vec<(usize, &'static str, &'static str)> {
vec![
(
size_of::<RopeParams>(),
include_str!("../src/backend/shaders/slang/rope.slang"),
"RopeParams",
),
(
size_of::<BiasAddParams>(),
include_str!("../src/backend/shaders/slang/bias_add.slang"),
"BiasAddParams",
),
(
size_of::<ElementwiseParams>(),
include_str!("../src/backend/shaders/slang/gelu.slang"),
"ElementwiseParams (gelu)",
),
(
size_of::<LayerNormBatchParams>(),
include_str!("../src/backend/shaders/slang/layernorm_batch.slang"),
"LayerNormBatchParams",
),
(
size_of::<RmsNormBatchParams>(),
include_str!("../src/backend/shaders/slang/rmsnorm_batch.slang"),
"RmsNormBatchParams",
),
(
size_of::<Conv1dBatchParams>(),
include_str!("../src/backend/shaders/slang/conv1d_fused_batch.slang"),
"Conv1dBatchParams",
),
(
size_of::<ArgmaxParams>(),
include_str!("../src/backend/shaders/slang/argmax_f32.slang"),
"ArgmaxParams",
),
(
size_of::<ElementwiseParams>(),
include_str!("../src/backend/shaders/slang/activations.slang"),
"ElementwiseParams (activations)",
),
(
size_of::<Conv2dDirectParams>(),
include_str!("../src/backend/shaders/slang/conv2d_direct.slang"),
"Conv2dDirectParams",
),
(
size_of::<TransposeBlockedParams>(),
include_str!("../src/backend/shaders/slang/transpose_blocked.slang"),
"TransposeBlockedParams",
),
(
size_of::<Batch2dParams>(),
include_str!("../src/backend/shaders/slang/glu_split.slang"),
"Batch2dParams (glu_split)",
),
(
size_of::<Batch2dParams>(),
include_str!("../src/backend/shaders/slang/chan_affine_silu.slang"),
"Batch2dParams (chan_affine_silu)",
),
(
size_of::<AudioXlAttnParams>(),
include_str!("../src/backend/shaders/slang/audio_xl_attention.slang"),
"AudioXlAttnParams",
),
(
size_of::<StftFrameParams>(),
include_str!("../src/backend/shaders/slang/stft_frame.slang"),
"StftFrameParams",
),
(
size_of::<PowerSpecParams>(),
include_str!("../src/backend/shaders/slang/power_spec.slang"),
"PowerSpecParams",
),
(
size_of::<MelProjectParams>(),
include_str!("../src/backend/shaders/slang/mel_project.slang"),
"MelProjectParams",
),
(
size_of::<MelNormParams>(),
include_str!("../src/backend/shaders/slang/mel_norm.slang"),
"MelNormParams",
),
(
size_of::<MoeRouteParams>(),
include_str!("../src/backend/shaders/slang/moe_route.slang"),
"MoeRouteParams",
),
(
size_of::<MoeGemvParams>(),
include_str!("../src/backend/shaders/slang/moe_gemv_q4_0.slang"),
"MoeGemvParams",
),
(
size_of::<MoeCombineParams>(),
include_str!("../src/backend/shaders/slang/moe_combine.slang"),
"MoeCombineParams",
),
]
}
#[test]
fn metal_uploads_cover_what_the_slang_kernels_read() {
let cases: &[(usize, &str, &str)] = &[
(
size_of::<ArgmaxParams>(),
include_str!("../src/backend/shaders/slang/argmax_f32.slang"),
"ArgmaxParams -> argmax_f32.slang",
),
(
size_of::<NormParams>(),
include_str!("../src/backend/shaders/slang/rmsnorm.slang"),
"NormParams -> rmsnorm.slang",
),
(
size_of::<NormParams>(),
include_str!("../src/backend/shaders/slang/per_head_rmsnorm.slang"),
"NormParams -> per_head_rmsnorm.slang",
),
(
size_of::<Conv1dParams>(),
include_str!("../src/backend/shaders/slang/conv1d.slang"),
"Conv1dParams -> conv1d.slang",
),
(
size_of::<ElementwiseParams>(),
include_str!("../src/backend/shaders/slang/elementwise.slang"),
"ElementwiseParams -> elementwise.slang",
),
];
let failures: Vec<String> = cases
.iter()
.filter_map(|&(upload, src, label)| match slang_params_bytes(src) {
None => Some(format!("{label}: params binding not resolvable")),
Some(reads) if upload < reads => Some(format!(
"{label}: host uploads {upload} B but the kernel reads {reads} B, so it reads past the end of the buffer"
)),
Some(_) => None,
})
.collect();
assert!(
failures.is_empty(),
"Metal params upload too small:\n {}",
failures.join("\n ")
);
}
#[test]
fn every_slang_params_binding_stays_parseable() {
let sources: &[(&str, &str, usize)] = &[
(
include_str!("../src/backend/shaders/slang/argmax_f32.slang"),
"argmax_f32",
8,
),
(
include_str!("../src/backend/shaders/slang/bias_add.slang"),
"bias_add",
8,
),
(
include_str!("../src/backend/shaders/slang/conv1d.slang"),
"conv1d",
12,
),
(
include_str!("../src/backend/shaders/slang/conv1d_fused.slang"),
"conv1d_fused",
12,
),
(
include_str!("../src/backend/shaders/slang/conv1d_fused_batch.slang"),
"conv1d_fused_batch",
24,
),
(
include_str!("../src/backend/shaders/slang/elementwise.slang"),
"elementwise",
8,
),
(
include_str!("../src/backend/shaders/slang/gelu.slang"),
"gelu",
8,
),
(
include_str!("../src/backend/shaders/slang/layernorm_batch.slang"),
"layernorm_batch",
16,
),
(
include_str!("../src/backend/shaders/slang/per_head_rmsnorm.slang"),
"per_head_rmsnorm",
16,
),
(
include_str!("../src/backend/shaders/slang/moe_combine.slang"),
"moe_combine",
16,
),
(
include_str!("../src/backend/shaders/slang/moe_gemv_q4_0.slang"),
"moe_gemv_q4_0",
32,
),
(
include_str!("../src/backend/shaders/slang/moe_route.slang"),
"moe_route",
16,
),
(
include_str!("../src/backend/shaders/slang/rmsnorm.slang"),
"rmsnorm",
16,
),
(
include_str!("../src/backend/shaders/slang/rmsnorm_batch.slang"),
"rmsnorm_batch",
20,
),
(
include_str!("../src/backend/shaders/slang/rope.slang"),
"rope",
20,
),
(
include_str!("../src/backend/shaders/slang/softmax.slang"),
"softmax",
8,
),
];
let failures: Vec<String> = sources
.iter()
.filter_map(|&(src, name, want)| match slang_params_bytes(src) {
None => Some(format!(
"{name}.slang: params binding not found or not resolvable"
)),
Some(got) if got != want => Some(format!(
"{name}.slang: kernel reads {got} B of params, expected {want} B"
)),
Some(_) => None,
})
.collect();
assert!(
failures.is_empty(),
"Slang params drift:\n {}",
failures.join("\n ")
);
}
#[test]
fn rust_param_mirrors_match_msl_structs() {
let mut failures = Vec::new();
for (rust_bytes, src, msl_name, rust_name) in cases() {
match msl_struct_bytes(src, msl_name) {
None => failures.push(format!(
"{rust_name}: MSL `struct {msl_name}` not found — renamed or deleted?"
)),
Some(msl_bytes) if msl_bytes != rust_bytes => failures.push(format!(
"{rust_name}: Rust mirror is {rust_bytes} B but MSL `struct {msl_name}` is \
{msl_bytes} B. A kernel reading a struct wider than the upload reads past \
the end of it — undefined behaviour, not a crash. Update the Rust mirror \
so its width matches, keeping the fields in the same order as the shader."
)),
Some(_) => {}
}
}
for (rust_bytes, slang_src, rust_name) in slang_cases() {
match slang_params_bytes(slang_src) {
None => failures.push(format!(
"{rust_name}: no `StructuredBuffer<uint*>` params binding found in the Slang \
source: renamed, or the port changed how it takes parameters?"
)),
Some(slang_bytes) if slang_bytes != rust_bytes => failures.push(format!(
"{rust_name}: Rust mirror is {rust_bytes} B but the Slang kernel reads \
{slang_bytes} B of params. Same hazard as the MSL case: the kernel reads \
past the end of the upload."
)),
Some(_) => {}
}
}
assert!(
failures.is_empty(),
"MSL/Rust params layout drift:\n {}",
failures.join("\n ")
);
}
#[test]
fn parser_counts_fields_and_ignores_comments() {
let src = "
struct Foo {
uint a;
int b; // int c; <- a decoy inside a comment
float d;
half e;
};
struct OneLiner { uint n; uint _pad; };
";
assert_eq!(msl_struct_bytes(src, "Foo"), Some(4 + 4 + 4 + 2));
assert_eq!(msl_struct_bytes(src, "Missing"), None);
assert_eq!(msl_struct_bytes(src, "OneLiner"), Some(8));
assert_eq!(
msl_struct_bytes(shaders::QK_NORM_ROPE, "Params"),
Some(36),
"qk_norm_rope Params is 9 uints"
);
}
#[test]
fn slang_parser_refuses_to_guess() {
let one = |body: &str| {
format!("[[vk::binding(0)]] StructuredBuffer<uint> par_buf : register(t0);\n{body}")
};
assert_eq!(slang_params_bytes(&one("x = par_buf[2];")), Some(3 * 4));
assert_eq!(
slang_params_bytes(
"[[vk::binding(0)]] StructuredBuffer<uint2> par_buf : register(t0);\nx = par_buf[1].y;"
),
Some(2 * 2 * 4)
);
assert_eq!(
slang_params_bytes(&one("x = par_buf[1]; // par_buf[9] is not a read")),
Some(2 * 4)
);
assert_eq!(slang_params_bytes(&one("x = par_buf[i];")), None);
assert_eq!(slang_params_bytes(&one("x = par_buf[base + 1];")), None);
assert_eq!(slang_params_bytes(&one("x = 1;")), None);
assert_eq!(
slang_params_bytes(
"[[vk::binding(0)]] RWStructuredBuffer<uint> out_buf : register(u0);\n\
[[vk::binding(1)]] StructuredBuffer<uint> par_buf : register(t1);\n\
out_buf[0] = par_buf[3];"
),
Some(4 * 4)
);
assert_eq!(
slang_params_bytes(
"[[vk::binding(0)]] StructuredBuffer<uint4> p_wgsl : register(t0);\n\
[[vk::binding(1)]] StructuredBuffer<uint4> p_metal : register(t1);\n\
x = p_wgsl[7]; y = p_metal[0];"
),
Some(4 * 4)
);
assert_eq!(
slang_params_bytes(
"[[vk::binding(0)]] StructuredBuffer<uint> w : register(t0);\n\
[[vk::binding(1)]] StructuredBuffer<uint> sel_expert : register(t1);\n\
[[vk::binding(2)]] StructuredBuffer<uint4> params : register(t2);\n\
x = w[0] + sel_expert[0]; y = params[1].x;"
),
Some(4 * 2 * 4)
);
assert_eq!(
slang_params_bytes(
"[[vk::binding(0)]] StructuredBuffer<float> par_buf : register(t0);\n\
x = par_buf[1];"
),
None
);
assert_eq!(
slang_params_bytes(
"[[vk::binding(0)]] StructuredBuffer<uint> pa : register(t0);\n\
[[vk::binding(1)]] StructuredBuffer<uint> pb : register(t1);\n\
x = pa[1];"
),
None
);
let switched = "[[vk::binding(0)]] StructuredBuffer<uint> params : register(t0);\n\
void f() {\n __target_switch {\n case metal:\n x = params[1];\n break;\n\
\n default:\n x = params[7];\n break;\n }\n}";
assert_eq!(slang_params_bytes(switched), Some(2 * 4));
let inverted = "[[vk::binding(0)]] StructuredBuffer<uint> params : register(t0);\n\
void f() {\n __target_switch {\n default:\n x = params[7];\n break;\n\
\n case metal:\n x = params[1];\n break;\n }\n}";
assert_eq!(slang_params_bytes(inverted), None);
let two = format!(
"{switched}\nvoid g() {{\n __target_switch {{\n case metal:\n y = params[2];\n break;\n\n default:\n y = params[9];\n break;\n }}\n}}"
);
assert_eq!(slang_params_bytes(&two), Some(3 * 4));
assert_eq!(
slang_params_bytes(include_str!("../src/backend/shaders/slang/rope.slang")),
Some(20),
"rope.slang's metal arm reads params[0..4]"
);
}