const PAD_FIELD: &str = "_wisp_align_pad";
pub(crate) fn pad_uniform_structs(source: &str, module: &naga::Module) -> String {
let mut targets: Vec<(usize, u32)> = module
.global_variables
.iter()
.filter_map(|(_, var)| pad_target(source, module, var))
.collect();
targets.sort_by_key(|&(close, _)| std::cmp::Reverse(close));
let mut out = source.to_string();
for (close, floats) in targets {
let field: String = (0..floats)
.map(|i| format!("\n {PAD_FIELD}_{i}: f32,"))
.collect();
out.insert_str(close, &format!("{field}\n"));
}
out
}
fn pad_target(
source: &str,
module: &naga::Module,
var: &naga::GlobalVariable,
) -> Option<(usize, u32)> {
if var.space != naga::AddressSpace::Uniform {
return None;
}
let binding = var.binding.as_ref()?;
if binding.binding != 0 || binding.group > 1 {
return None;
}
let naga::TypeInner::Struct { span, .. } = module.types[var.ty].inner else {
return None;
};
let remainder = span % 16;
if remainder == 0 {
return None;
}
let range = module.types.get_span(var.ty).to_range()?;
let close = source.get(..range.end)?.rfind('}')?;
Some((close, (16 - remainder) / 4))
}
#[cfg(test)]
mod tests {
use super::*;
use crate::reflect::parse_and_validate;
fn uniform_size(source: &str, group: u32) -> u32 {
let module = parse_and_validate(source).unwrap().module;
module
.global_variables
.iter()
.find_map(|(_, var)| {
let binding = var.binding.as_ref()?;
(var.space == naga::AddressSpace::Uniform
&& binding.group == group
&& binding.binding == 0)
.then_some(())?;
match module.types[var.ty].inner {
naga::TypeInner::Struct { span, .. } => Some(span),
_ => None,
}
})
.unwrap()
}
const SHADER: &str = "\
struct Params {
wobble: f32,
}
@group(1) @binding(0) var<uniform> params: Params;
@fragment
fn frag() -> @location(0) vec4<f32> {
return vec4<f32>(params.wobble);
}
";
#[test]
fn pads_small_params_to_16() {
let module = parse_and_validate(SHADER).unwrap().module;
assert_eq!(uniform_size(SHADER, 1), 4, "single f32 starts at 4 bytes");
let padded = pad_uniform_structs(SHADER, &module);
assert_eq!(uniform_size(&padded, 1), 16);
}
#[test]
fn leaves_aligned_structs_untouched() {
let source = "\
struct Params {
tint: vec4<f32>,
}
@group(1) @binding(0) var<uniform> params: Params;
@fragment
fn frag() -> @location(0) vec4<f32> {
return params.tint;
}
";
let module = parse_and_validate(source).unwrap().module;
assert_eq!(uniform_size(source, 1), 16);
assert_eq!(pad_uniform_structs(source, &module), source);
}
}