use crate::codegen::wgsl::types::{map_type_to_wgsl, WgslType};
use crate::parser_impl::StructDecl;
use anyhow::Result;
#[derive(Debug, Clone)]
pub struct LayoutField {
pub name: String,
pub wgsl_type: WgslType,
pub offset: usize,
pub is_padding: bool,
}
#[derive(Debug, Clone)]
pub struct StructLayout {
pub name: String,
pub fields: Vec<LayoutField>,
pub total_size: usize,
pub alignment: usize,
}
impl StructLayout {
pub fn from_struct_decl(decl: &StructDecl) -> Result<Self> {
let mut fields = Vec::new();
let mut current_offset = 0;
let mut max_alignment = 1;
let mut padding_counter = 0;
for field in &decl.fields {
let wgsl_type = map_type_to_wgsl(&field.field_type)?;
let field_alignment = wgsl_type.alignment_bytes();
let field_size = wgsl_type.size_bytes();
max_alignment = max_alignment.max(field_alignment);
let misalignment = current_offset % field_alignment;
if misalignment != 0 {
let padding_needed = field_alignment - misalignment;
let num_f32_pads = padding_needed / 4;
for _ in 0..num_f32_pads {
fields.push(LayoutField {
name: format!("_pad{}", padding_counter),
wgsl_type: WgslType::F32,
offset: current_offset,
is_padding: true,
});
padding_counter += 1;
current_offset += 4;
}
}
fields.push(LayoutField {
name: field.name.clone(),
wgsl_type,
offset: current_offset,
is_padding: false,
});
current_offset += field_size;
}
let total_size = align_up(current_offset, max_alignment);
if current_offset < total_size {
let end_padding = total_size - current_offset;
let num_f32_pads = end_padding / 4;
for _ in 0..num_f32_pads {
fields.push(LayoutField {
name: format!("_pad{}", padding_counter),
wgsl_type: WgslType::F32,
offset: current_offset,
is_padding: true,
});
padding_counter += 1;
current_offset += 4;
}
}
Ok(StructLayout {
name: decl.name.clone(),
fields,
total_size,
alignment: max_alignment,
})
}
pub fn to_wgsl_string(&self) -> String {
let mut output = String::new();
output.push_str("struct ");
output.push_str(&self.name);
output.push_str(" {\n");
for field in &self.fields {
output.push_str(" ");
output.push_str(&field.name);
output.push_str(": ");
output.push_str(&field.wgsl_type.to_wgsl_string());
output.push_str(",\n");
}
output.push('}');
output
}
}
fn align_up(value: usize, alignment: usize) -> usize {
(value + alignment - 1) & !(alignment - 1)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_align_up() {
assert_eq!(align_up(0, 4), 0);
assert_eq!(align_up(1, 4), 4);
assert_eq!(align_up(4, 4), 4);
assert_eq!(align_up(5, 4), 8);
assert_eq!(align_up(12, 16), 16);
assert_eq!(align_up(13, 16), 16);
assert_eq!(align_up(16, 16), 16);
assert_eq!(align_up(17, 16), 32);
}
#[test]
fn test_wgsl_type_alignment() {
assert_eq!(WgslType::U32.alignment_bytes(), 4);
assert_eq!(WgslType::Vec2F32.alignment_bytes(), 8);
assert_eq!(WgslType::Vec3F32.alignment_bytes(), 16); assert_eq!(WgslType::Vec4F32.alignment_bytes(), 16);
}
#[test]
fn test_wgsl_type_size() {
assert_eq!(WgslType::U32.size_bytes(), 4);
assert_eq!(WgslType::Vec2F32.size_bytes(), 8);
assert_eq!(WgslType::Vec3F32.size_bytes(), 12); assert_eq!(WgslType::Vec4F32.size_bytes(), 16);
}
}