use naga_oil::compose::ShaderDefValue;
use std::collections::HashMap;
const COMMON_WGSL: &str = include_str!("../shaders/common.wgsl");
const PBR_EXT_WGSL: &str = include_str!("../shaders/pbr_ext.wgsl");
pub(crate) fn native_render_defs() -> HashMap<String, ShaderDefValue> {
HashMap::from([
("SHADOWS".to_string(), ShaderDefValue::Bool(true)),
("SKELETON_GROUP".to_string(), ShaderDefValue::UInt(3)),
("INSTANCE_GROUP".to_string(), ShaderDefValue::UInt(4)),
])
}
#[allow(dead_code)] pub(crate) fn web_render_defs() -> HashMap<String, ShaderDefValue> {
HashMap::from([
("SKELETON_GROUP".to_string(), ShaderDefValue::UInt(2)),
("INSTANCE_GROUP".to_string(), ShaderDefValue::UInt(3)),
])
}
pub(crate) fn compose_module(
source: &str,
label: &str,
shader_defs: HashMap<String, ShaderDefValue>,
) -> Result<(naga::Module, naga::valid::ModuleInfo), String> {
use naga_oil::compose::{
ComposableModuleDescriptor, Composer, NagaModuleDescriptor, ShaderLanguage,
};
let mut composer = Composer::default();
composer
.add_composable_module(ComposableModuleDescriptor {
source: COMMON_WGSL,
file_path: "gizmo/common.wgsl",
language: ShaderLanguage::Wgsl,
..Default::default()
})
.map_err(|e| format!("composing common.wgsl failed: {e}"))?;
composer
.add_composable_module(ComposableModuleDescriptor {
source: PBR_EXT_WGSL,
file_path: "gizmo/pbr_ext.wgsl",
language: ShaderLanguage::Wgsl,
..Default::default()
})
.map_err(|e| format!("composing pbr_ext.wgsl failed: {e}"))?;
let module = composer
.make_naga_module(NagaModuleDescriptor {
source,
file_path: label,
shader_defs,
..Default::default()
})
.map_err(|e| format!("naga_oil compose of '{label}' failed: {e}"))?;
let info = naga::valid::Validator::new(
naga::valid::ValidationFlags::all(),
naga::valid::Capabilities::all(),
)
.validate(&module)
.map_err(|e| format!("validating composed '{label}' failed: {e:?}"))?;
Ok((module, info))
}
pub(crate) fn compose_wgsl(
source: &str,
label: &str,
shader_defs: HashMap<String, ShaderDefValue>,
) -> String {
let (module, info) = compose_module(source, label, shader_defs)
.unwrap_or_else(|e| panic!("{e}"));
naga::back::wgsl::write_string(&module, &info, naga::back::wgsl::WriterFlags::empty())
.unwrap_or_else(|e| panic!("emitting WGSL for '{label}' failed: {e}"))
}
pub fn load_shader(
device: &wgpu::Device,
file_path: &str,
fallback_src: &str,
label: &str,
) -> wgpu::ShaderModule {
let source = std::fs::read_to_string(file_path).unwrap_or_else(|_| fallback_src.to_string());
device.create_shader_module(wgpu::ShaderModuleDescriptor {
label: Some(label),
source: wgpu::ShaderSource::Wgsl(source.into()),
})
}
pub fn load_shader_composed(
device: &wgpu::Device,
file_path: &str,
fallback_src: &str,
label: &str,
) -> wgpu::ShaderModule {
let source = std::fs::read_to_string(file_path).unwrap_or_else(|_| fallback_src.to_string());
let composed = compose_wgsl(&source, label, native_render_defs());
device.create_shader_module(wgpu::ShaderModuleDescriptor {
label: Some(label),
source: wgpu::ShaderSource::Wgsl(composed.into()),
})
}
#[cfg(target_arch = "wasm32")]
pub fn load_shader_composed_web(
device: &wgpu::Device,
fallback_src: &str,
label: &str,
) -> wgpu::ShaderModule {
let composed = compose_wgsl(fallback_src, label, web_render_defs());
device.create_shader_module(wgpu::ShaderModuleDescriptor {
label: Some(label),
source: wgpu::ShaderSource::Wgsl(composed.into()),
})
}
#[cfg(test)]
mod tests {
use super::*;
const SHADER_SRC: &str = include_str!("../shaders/shader.wgsl");
#[test]
fn baked_lit_path_shaders_type_check_under_both_schemas() {
let web_capable: &[(&str, &str)] = &[
("baked_lit.wgsl", include_str!("../shaders/baked_lit.wgsl")),
("shader.wgsl", include_str!("../shaders/shader.wgsl")),
("unlit.wgsl", include_str!("../shaders/unlit.wgsl")),
("water.wgsl", include_str!("../shaders/water.wgsl")),
("sky.wgsl", include_str!("../shaders/sky.wgsl")),
("backdrop.wgsl", include_str!("../shaders/backdrop.wgsl")),
("grid.wgsl", include_str!("../shaders/grid.wgsl")),
];
for (name, src) in web_capable {
compose_wgsl(src, name, native_render_defs());
compose_wgsl(src, name, web_render_defs());
}
let native_only: &[(&str, &str)] = &[
("deferred_lighting.wgsl", include_str!("../shaders/deferred_lighting.wgsl")),
("gbuffer.wgsl", include_str!("../shaders/gbuffer.wgsl")),
("shadow.wgsl", include_str!("../shaders/shadow.wgsl")),
("point_shadow.wgsl", include_str!("../shaders/point_shadow.wgsl")),
];
for (name, src) in native_only {
compose_wgsl(src, name, native_render_defs());
}
}
#[test]
fn web_compose_strips_shadows_and_shifts_groups() {
let web = compose_wgsl(SHADER_SRC, "shader.wgsl", web_render_defs());
assert!(
!web.contains("#import") && !web.contains("#ifdef") && !web.contains("#{"),
"web compose left unresolved naga_oil tokens"
);
assert!(!web.contains("var t_shadow"), "web variant must not declare t_shadow");
assert!(
!web.contains("textureSampleCompare"),
"web variant must not sample shadows"
);
assert!(!web.contains("@group(4)"), "web variant must not use @group(4)");
}
#[test]
fn native_compose_keeps_shadows_and_groups() {
let native = compose_wgsl(SHADER_SRC, "shader.wgsl", native_render_defs());
assert!(
!native.contains("#import") && !native.contains("#ifdef") && !native.contains("#{"),
"native compose left unresolved naga_oil tokens"
);
assert!(native.contains("var t_shadow"), "native variant must keep the shadow bindings");
assert!(
native.contains("textureSampleCompare"),
"native variant must keep shadow sampling"
);
assert!(native.contains("@group(4)"), "native variant keeps instance at group 4");
}
async fn setup_headless_gpu() -> Option<(wgpu::Device, wgpu::Queue)> {
crate::test_gpu::headless_device().await
}
async fn read_mat_cols(device: &wgpu::Device, queue: &wgpu::Queue, buffer: &wgpu::Buffer) -> [f32; 16] {
let staging = device.create_buffer(&wgpu::BufferDescriptor {
label: Some("inv_readback"),
size: 64,
usage: wgpu::BufferUsages::MAP_READ | wgpu::BufferUsages::COPY_DST,
mapped_at_creation: false,
});
let mut enc = device.create_command_encoder(&wgpu::CommandEncoderDescriptor { label: None });
enc.copy_buffer_to_buffer(buffer, 0, &staging, 0, 64);
queue.submit(Some(enc.finish()));
let slice = staging.slice(..);
let (tx, rx) = std::sync::mpsc::channel();
slice.map_async(wgpu::MapMode::Read, move |v| tx.send(v).unwrap());
let _ = device.poll(wgpu::PollType::Wait { submission_index: None, timeout: None });
rx.recv().unwrap().unwrap();
let data = slice
.get_mapped_range()
.expect("a just-mapped buffer's full range is always valid");
let mut out = [0.0f32; 16];
out.copy_from_slice(bytemuck::cast_slice(&data));
drop(data);
staging.unmap();
out
}
#[test]
fn inverse_mat4_matches_glam_on_gpu() {
let _gpu = crate::test_gpu::gpu_lock();
use gizmo_math::{Mat4, Vec3};
use wgpu::util::DeviceExt;
let composed = compose_wgsl(
r#"
#import gizmo::common::inverse_mat4
@group(0) @binding(0) var<storage, read> m_in: mat4x4<f32>;
@group(0) @binding(1) var<storage, read_write> m_out: mat4x4<f32>;
@compute @workgroup_size(1)
fn main() { m_out = inverse_mat4(m_in); }
"#,
"inverse_mat4_test",
HashMap::new(),
);
pollster::block_on(async {
let Some((device, queue)) = setup_headless_gpu().await else { return };
let shader = device.create_shader_module(wgpu::ShaderModuleDescriptor {
label: Some("inverse_mat4_test"),
source: wgpu::ShaderSource::Wgsl(composed.into()),
});
let view = Mat4::look_at_rh(Vec3::new(6.0, 3.0, 7.0), Vec3::new(0.0, 2.2, 0.0), Vec3::Y);
let proj = Mat4::perspective_rh(45f32.to_radians(), 16.0 / 9.0, 0.1, 500.0);
let vp = proj * view;
let in_buf = device.create_buffer_init(&wgpu::util::BufferInitDescriptor {
label: Some("m_in"),
contents: bytemuck::cast_slice(&vp.to_cols_array()),
usage: wgpu::BufferUsages::STORAGE,
});
let out_buf = device.create_buffer(&wgpu::BufferDescriptor {
label: Some("m_out"),
size: 64,
usage: wgpu::BufferUsages::STORAGE | wgpu::BufferUsages::COPY_SRC,
mapped_at_creation: false,
});
let pipeline = device.create_compute_pipeline(&wgpu::ComputePipelineDescriptor {
label: Some("inverse_mat4_test"),
layout: None,
module: &shader,
entry_point: Some("main"),
compilation_options: Default::default(),
cache: None,
});
let bg = device.create_bind_group(&wgpu::BindGroupDescriptor {
label: None,
layout: &pipeline.get_bind_group_layout(0),
entries: &[
wgpu::BindGroupEntry { binding: 0, resource: in_buf.as_entire_binding() },
wgpu::BindGroupEntry { binding: 1, resource: out_buf.as_entire_binding() },
],
});
let mut enc = device.create_command_encoder(&wgpu::CommandEncoderDescriptor { label: None });
{
let mut cpass = enc.begin_compute_pass(&wgpu::ComputePassDescriptor {
label: None,
timestamp_writes: None,
});
cpass.set_pipeline(&pipeline);
cpass.set_bind_group(0, &bg, &[]);
cpass.dispatch_workgroups(1, 1, 1);
}
queue.submit(Some(enc.finish()));
let gpu_cols = read_mat_cols(&device, &queue, &out_buf).await;
let gpu_inv = Mat4::from_cols_array(&gpu_cols);
let want = vp.inverse();
let g = gpu_inv.to_cols_array();
let w = want.to_cols_array();
let max_err = g.iter().zip(w.iter()).map(|(a, b)| (a - b).abs()).fold(0.0f32, f32::max);
assert!(max_err < 1e-3, "GPU inverse_mat4 differs from glam by {max_err} (buggy formula?)");
let clip = vp * gizmo_math::Vec4::new(0.0, 2.0, 0.0, 1.0);
let ndc = clip / clip.w;
let unproj_h = gpu_inv * gizmo_math::Vec4::new(ndc.x, ndc.y, ndc.z, 1.0);
let world = unproj_h.truncate() / unproj_h.w;
assert!(
(world - Vec3::new(0.0, 2.0, 0.0)).length() < 1e-2,
"NDC→world round-trip landed at {world:?}, expected ~(0,2,0) (buggy inverse_mat4)"
);
});
}
}