use naga::back::spv::{Options, PipelineOptions, WriterFlags};
use naga::valid::{Capabilities, ValidationFlags, Validator};
pub struct SpirvEmitter;
impl SpirvEmitter {
pub fn emit(module: &naga::Module, entry: &str) -> Result<Vec<u32>, String> {
let mut validator = Validator::new(ValidationFlags::all(), Capabilities::all());
let info = validator.validate(module).map_err(|e| {
format!(
"SPIR-V emit: naga validation failed: {e}. Fix: reject the naga::Module and look for a malformed lowering."
)
})?;
let options = Options::default();
let pipeline = PipelineOptions {
shader_stage: naga::ShaderStage::Compute,
entry_point: entry.to_owned(),
};
let mut out = Vec::new();
let mut writer = naga::back::spv::Writer::new(&options).map_err(|e| {
format!(
"SPIR-V emit: could not construct writer: {e}. Fix: upgrade naga or lower spv-out feature flags."
)
})?;
writer
.write(module, &info, Some(&pipeline), &None, &mut out)
.map_err(|e| {
format!(
"SPIR-V emit: writer.write failed: {e}. Fix: inspect the naga::Module for capabilities this adapter lacks."
)
})?;
Ok(out)
}
#[must_use]
pub fn default_flags() -> WriterFlags {
WriterFlags::empty()
}
}
pub const SPIRV_BACKEND_ID: &str = "spirv";
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn emit_returns_nonempty_words_for_empty_module() {
let mut module = naga::Module::default();
let entry = naga::EntryPoint {
name: "main".to_owned(),
stage: naga::ShaderStage::Compute,
early_depth_test: None,
workgroup_size: [1, 1, 1],
workgroup_size_overrides: None,
function: naga::Function::default(),
};
module.entry_points.push(entry);
match SpirvEmitter::emit(&module, "main") {
Ok(words) => {
assert!(!words.is_empty(), "SPIR-V output must not be empty");
assert_eq!(words[0], 0x0723_0203, "first word must be SPIR-V magic");
}
Err(msg) => {
assert!(
msg.contains("Fix:"),
"emit error must carry Fix: remediation: {msg}"
);
}
}
}
#[test]
fn backend_id_is_stable() {
assert_eq!(SPIRV_BACKEND_ID, "spirv");
}
#[test]
fn default_flags_are_empty() {
assert_eq!(SpirvEmitter::default_flags(), WriterFlags::empty());
}
}