1use serde::{Deserialize, Serialize};
2
3use crate::wgsl_forge::ForgeError;
4use crate::wgsl_forge::KernelSpec;
5use crate::wgsl_forge::TargetBackend;
6
7#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
8pub struct ValidationReport {
9 pub source_hash: String,
10 pub entry_points: Vec<String>,
11 pub binding_count: usize,
12 pub naga_validated: bool,
13 pub native_tool_validated: Option<String>,
14}
15
16pub fn validate_wgsl(source: &str) -> Result<ValidationReport, ForgeError> {
17 let module = naga::front::wgsl::parse_str(source)
18 .map_err(|error| ForgeError::WgslParse(error.emit_to_string(source)))?;
19 let mut validator = naga::valid::Validator::new(
22 naga::valid::ValidationFlags::all(),
23 naga::valid::Capabilities::RAY_QUERY
24 | naga::valid::Capabilities::COOPERATIVE_MATRIX
25 | naga::valid::Capabilities::SHADER_FLOAT16,
26 );
27 validator
28 .validate(&module)
29 .map_err(|error| ForgeError::WgslValidation(format!("{error:?}")))?;
30
31 let mut entry_points = module
32 .entry_points
33 .iter()
34 .map(|entry| entry.name.clone())
35 .collect::<Vec<_>>();
36 entry_points.sort();
37 let binding_count = module
38 .global_variables
39 .iter()
40 .filter(|(_, variable)| variable.binding.is_some())
41 .count();
42 Ok(ValidationReport {
43 source_hash: blake3::hash(source.as_bytes()).to_hex().to_string(),
44 entry_points,
45 binding_count,
46 naga_validated: true,
47 native_tool_validated: None,
48 })
49}
50
51pub fn validate_native(
56 source: &str,
57 target: TargetBackend,
58 kernel: Option<&KernelSpec>,
59) -> Result<ValidationReport, ForgeError> {
60 use std::process::Command;
61
62 let tool_name = match target {
63 TargetBackend::Ptx => "ptxas",
64 TargetBackend::Hlsl => "dxc",
65 TargetBackend::Msl => "xcrun",
66 _ => {
67 return Err(ForgeError::WgslValidation(
68 "Not a native target supported by offline validation".to_string(),
69 ))
70 }
71 };
72
73 let temp_file = tempfile::Builder::new()
74 .suffix(match target {
75 TargetBackend::Ptx => ".ptx",
76 TargetBackend::Hlsl => ".hlsl",
77 TargetBackend::Msl => ".metal",
78 _ => "",
79 })
80 .tempfile()
81 .map_err(|e| ForgeError::Io(format!("Failed to create temp file: {}", e)))?;
82
83 let temp_path = temp_file.path().to_path_buf();
84
85 std::fs::write(&temp_path, source)
87 .map_err(|e| ForgeError::Io(format!("Failed to write to temp file: {}", e)))?;
88
89 let mut cmd = Command::new(tool_name);
90 match target {
91 TargetBackend::Ptx => {
92 if Command::new("ptxas").arg("--version").output().is_err() {
94 return Err(ForgeError::WgslValidation(
95 "ptxas not found in PATH".to_string(),
96 ));
97 }
98 cmd.arg(&temp_path).arg("-c"); }
100 TargetBackend::Hlsl => {
101 if Command::new("dxc").arg("--help").output().is_err() {
102 return Err(ForgeError::WgslValidation(
103 "dxc not found in PATH".to_string(),
104 ));
105 }
106 cmd.arg("-T").arg("cs_6_0").arg(&temp_path);
107 }
108 TargetBackend::Msl => {
109 if Command::new("xcrun").arg("--version").output().is_err() {
110 return Err(ForgeError::WgslValidation(
111 "xcrun not found in PATH".to_string(),
112 ));
113 }
114 cmd.arg("-sdk")
115 .arg("macosx")
116 .arg("metal")
117 .arg("-c")
118 .arg(&temp_path);
119 }
120 _ => {}
121 }
122
123 let output = cmd
124 .output()
125 .map_err(|e| ForgeError::Io(format!("Failed to execute validation tool: {}", e)))?;
126
127 if !output.status.success() {
128 let stderr = String::from_utf8_lossy(&output.stderr);
129 return Err(ForgeError::WgslValidation(format!(
130 "{} validation failed: {}",
131 tool_name, stderr
132 )));
133 }
134
135 let (entry_points, binding_count) = match kernel {
138 Some(k) => (vec![k.entry_point.clone()], k.buffers.len()),
139 None => (Vec::new(), 0),
140 };
141 Ok(ValidationReport {
142 source_hash: blake3::hash(source.as_bytes()).to_hex().to_string(),
143 entry_points,
144 binding_count,
145 naga_validated: false,
146 native_tool_validated: Some(tool_name.to_string()),
147 })
148}
149
150#[cfg(test)]
151mod tests {
152 use super::*;
153 use crate::wgsl_forge::{generate_builtin, BuiltinKernel, Schedule, TargetBackend};
154
155 #[test]
156 fn generated_schedules_pass_full_naga_validation() {
157 for vector_width in [1, 2, 4] {
158 let generated = generate_builtin(
159 BuiltinKernel::AffineF32,
160 Schedule {
161 vector_width,
162 ..Schedule::default()
163 },
164 TargetBackend::Wgsl,
165 )
166 .unwrap();
167 let report = validate_wgsl(&generated.source).expect("Naga validation");
168 assert_eq!(report.entry_points, vec!["affine_f32"]);
169 assert_eq!(report.binding_count, 3);
170 assert_eq!(report.source_hash, generated.source_hash);
171 }
172 }
173
174 #[test]
175 fn generated_topk_passes_full_naga_validation() {
176 for workgroup_size in [32u32, 64, 128, 256] {
178 let generated = generate_builtin(
179 BuiltinKernel::TopK,
180 Schedule {
181 workgroup_size,
182 items_per_invocation: 1,
183 vector_width: 1,
184 ..Schedule::default()
185 },
186 TargetBackend::Wgsl,
187 )
188 .unwrap();
189 assert!(
190 generated.source.contains("workgroupBarrier()"),
191 "top-k must emit barriers"
192 );
193 assert!(
194 generated
195 .source
196 .contains(&format!("array<f32, {workgroup_size}>")),
197 "shared arrays sized to workgroup size"
198 );
199 let report = validate_wgsl(&generated.source).expect("Naga validation of top-k");
200 assert_eq!(report.entry_points, vec!["topk"]);
201 assert_eq!(report.binding_count, 3);
202 }
203 }
204
205 #[test]
206 fn generated_p64_passes_naga_validation() {
207 let generated = generate_builtin(
208 BuiltinKernel::P64Project,
209 Schedule {
210 workgroup_size: 64,
211 items_per_invocation: 1,
212 vector_width: 1,
213 ..Schedule::default()
214 },
215 TargetBackend::Wgsl,
216 )
217 .unwrap();
218 assert!(generated.source.contains("struct P64Words64"));
219 assert!(generated.source.contains("arrayLength(&output)"));
220 let report = validate_wgsl(&generated.source).expect("Naga validation of p64-project");
221 assert_eq!(report.entry_points, vec!["p64_project"]);
222 assert_eq!(report.binding_count, 3);
223 }
224
225 #[test]
226 fn generated_ffn_passes_naga_validation() {
227 let generated = generate_builtin(
228 BuiltinKernel::FusedFfn,
229 Schedule {
230 workgroup_size: 64,
231 items_per_invocation: 1,
232 vector_width: 1,
233 ..Schedule::default()
234 },
235 TargetBackend::Wgsl,
236 )
237 .unwrap();
238 assert!(generated.source.contains("struct FfnParams"));
239 assert!(generated.source.contains("tanh("));
240 let report = validate_wgsl(&generated.source).expect("Naga validation of fused-ffn");
241 assert_eq!(report.entry_points, vec!["fused_ffn"]);
242 assert_eq!(report.binding_count, 5);
243 }
244
245 #[test]
246 fn generated_ray_probe_passes_naga_validation() {
247 let generated = generate_builtin(
249 BuiltinKernel::RayProbe,
250 Schedule {
251 workgroup_size: 64,
252 items_per_invocation: 1,
253 vector_width: 1,
254 ..Schedule::default()
255 },
256 TargetBackend::Wgsl,
257 )
258 .unwrap();
259 assert!(generated
260 .source
261 .contains("var scene: acceleration_structure;"));
262 assert!(generated.source.contains("rayQueryInitialize"));
263 assert!(generated
264 .source
265 .contains("rayQueryGetCommittedIntersection"));
266 let report = validate_wgsl(&generated.source).expect("Naga validation of ray-probe");
267 assert_eq!(report.entry_points, vec!["ray_probe"]);
268 assert_eq!(report.binding_count, 3);
269 }
270
271 #[test]
272 fn generated_ternary_gemv_passes_naga_validation() {
273 let generated = generate_builtin(
274 BuiltinKernel::TernaryGemv,
275 Schedule {
276 workgroup_size: 64,
277 items_per_invocation: 1,
278 vector_width: 1,
279 ..Schedule::default()
280 },
281 TargetBackend::Wgsl,
282 )
283 .unwrap();
284 assert!(generated.source.contains("struct TernaryGemvParams"));
285 assert!(generated.source.contains("& 3u"));
287 let report = validate_wgsl(&generated.source).expect("Naga validation of ternary-gemv");
288 assert_eq!(report.entry_points, vec!["ternary_gemv"]);
289 assert_eq!(report.binding_count, 5);
290 }
291
292 #[test]
293 fn generated_gemm_passes_naga_validation() {
294 let generated = generate_builtin(
295 BuiltinKernel::Gemm,
296 Schedule {
297 workgroup_size: 64,
298 items_per_invocation: 1,
299 vector_width: 1,
300 ..Schedule::default()
301 },
302 TargetBackend::Wgsl,
303 )
304 .unwrap();
305 assert!(generated.source.contains("struct GemmParams"));
306 assert!(generated.source.contains("kk * params.n + col"));
308 let report = validate_wgsl(&generated.source).expect("Naga validation of gemm");
309 assert_eq!(report.entry_points, vec!["gemm"]);
310 assert_eq!(report.binding_count, 4);
311 }
312
313 #[test]
314 fn generated_gemv_passes_naga_validation() {
315 let generated = generate_builtin(
316 BuiltinKernel::Gemv,
317 Schedule {
318 workgroup_size: 64,
319 items_per_invocation: 1,
320 vector_width: 1,
321 ..Schedule::default()
322 },
323 TargetBackend::Wgsl,
324 )
325 .unwrap();
326 assert!(generated.source.contains("struct GemvParams"));
327 assert!(generated.source.contains("a[a_row + j] * x[j]"));
329 let report = validate_wgsl(&generated.source).expect("Naga validation of gemv");
330 assert_eq!(report.entry_points, vec!["gemv"]);
331 assert_eq!(report.binding_count, 4);
332 }
333
334 #[test]
335 fn generated_fft_passes_naga_validation() {
336 let generated = generate_builtin(
340 BuiltinKernel::Fft,
341 Schedule {
342 workgroup_size: 256,
343 items_per_invocation: 1,
344 vector_width: 1,
345 ..Schedule::default()
346 },
347 TargetBackend::Wgsl,
348 )
349 .unwrap();
350 assert!(generated.source.contains("struct FftParams"));
351 assert!(generated.source.contains("reverseBits("));
353 assert!(generated.source.contains("workgroupBarrier()"));
354 assert!(
355 generated.source.contains("array<f32, 256>"),
356 "shared arrays sized to workgroup size"
357 );
358 let report = validate_wgsl(&generated.source).expect("Naga validation of fft");
359 assert_eq!(report.entry_points, vec!["fft"]);
360 assert_eq!(report.binding_count, 3);
361 }
362
363 #[test]
364 fn cooperative_matrix_tile_validates() {
365 let source = crate::wgsl_forge::matmul_tc_wgsl();
367 let report = validate_wgsl(&source).expect("coopmat tile should validate");
368 assert_eq!(report.entry_points, vec!["matmul_tc"]);
369 }
370
371 #[test]
372 fn cooperative_matrix_tiled_gemm_validates() {
373 let source = crate::wgsl_forge::emit::matmul_tc_wgsl_tiled();
378 let report = validate_wgsl(&source).expect("tiled coopmat GEMM should validate");
379 assert_eq!(
380 report.entry_points,
381 vec![crate::wgsl_forge::emit::MATMUL_TC_TILED_ENTRY]
382 );
383 assert_eq!(report.binding_count, 4);
384 }
385
386 #[test]
387 fn semantic_errors_are_rejected() {
388 let source = "@compute @workgroup_size(64) fn broken() { let x: u32 = 1.0; }";
389 assert!(validate_wgsl(source).is_err());
390 }
391
392 #[test]
397 fn native_emitters_produce_non_empty_bodies() {
398 use crate::wgsl_forge::BuiltinKernel as K;
399 let schedule = Schedule::default();
400 let cases: &[(K, TargetBackend, &[&str])] = &[
401 (
403 K::Gemm,
404 TargetBackend::Hlsl,
405 &["RWStructuredBuffer", "for (uint kk"],
406 ),
407 (
408 K::Gemv,
409 TargetBackend::Hlsl,
410 &["RWStructuredBuffer", "for (uint j"],
411 ),
412 (K::Fft, TargetBackend::Hlsl, &["groupshared", "reversebits"]),
413 (
414 K::TernaryGemv,
415 TargetBackend::Hlsl,
416 &["StructuredBuffer<uint>", "tern"],
417 ),
418 (
419 K::P64Project,
420 TargetBackend::Hlsl,
421 &["StructuredBuffer<P64Words64>"],
422 ),
423 (K::Gemm, TargetBackend::Msl, &["device", "for (uint kk"]),
425 (K::Gemv, TargetBackend::Msl, &["device", "for (uint j"]),
426 (K::Fft, TargetBackend::Msl, &["threadgroup", "reverse_bits"]),
427 (
428 K::TernaryGemv,
429 TargetBackend::Msl,
430 &["device const uint", "tern"],
431 ),
432 (
433 K::P64Project,
434 TargetBackend::Msl,
435 &["P64Words64", "record_count"],
436 ),
437 (
439 K::TernaryGemv,
440 TargetBackend::CudaC,
441 &["__global__", "tern"],
442 ),
443 (
444 K::P64Project,
445 TargetBackend::CudaC,
446 &["__global__", "P64Words64"],
447 ),
448 (K::Fft, TargetBackend::CudaC, &["__shared__", "__brev"]),
449 ];
450 for (kernel, target, expected) in cases {
451 let generated = generate_builtin(*kernel, schedule, *target)
452 .unwrap_or_else(|e| panic!("{kernel:?}/{target:?} emission failed: {e}"));
453 let src = &generated.source;
454 assert!(
455 src.lines().count() > 10,
456 "{kernel:?}/{target:?} body too short ({} lines)",
457 src.lines().count()
458 );
459 for kw in *expected {
460 assert!(
461 src.contains(kw),
462 "{kernel:?}/{target:?} missing expected keyword `{kw}`"
463 );
464 }
465 }
466 }
467}