mod ast;
mod codegen;
mod expression_type_checking;
mod lexer;
pub mod parser;
mod shader_type_rules;
mod statement_type_checking;
mod type_checker;
use std::collections::HashSet;
use std::path::{Path, PathBuf};
pub use ast::*;
pub use parser::parse_wjsl;
pub use parser::parse_wjsl_with_filename;
pub use type_checker::type_check_wjsl;
pub fn transpile_wjsl(source: &str) -> Result<String, anyhow::Error> {
let ast = parse_wjsl(source)?;
type_checker::check(&ast, source)?;
let wgsl = codegen::WjslCodegen::new(ast).generate()?;
Ok(wgsl)
}
pub fn transpile_wjsl_with_includes(
source: &str,
base_dir: &Path,
) -> Result<String, anyhow::Error> {
let resolved = resolve_includes(source, base_dir, &mut Vec::new())?;
transpile_wjsl(&resolved)
}
pub fn resolve_includes(
source: &str,
base_dir: &Path,
include_stack: &mut Vec<PathBuf>,
) -> Result<String, anyhow::Error> {
let mut seen = HashSet::new();
resolve_includes_inner(source, base_dir, include_stack, &mut seen)
}
fn resolve_includes_inner(
source: &str,
base_dir: &Path,
include_stack: &mut Vec<PathBuf>,
seen: &mut HashSet<PathBuf>,
) -> Result<String, anyhow::Error> {
let mut output = String::with_capacity(source.len());
for line in source.lines() {
let trimmed = line.trim();
if let Some(path_str) = parse_include_directive(trimmed) {
let include_path = base_dir.join(path_str);
let canonical = include_path
.canonicalize()
.unwrap_or_else(|_| include_path.clone());
if include_stack.contains(&canonical) {
let chain: Vec<String> = include_stack
.iter()
.map(|p| p.display().to_string())
.collect();
return Err(anyhow::anyhow!(
"Circular #include detected: {} -> {} (chain: {})",
include_stack
.last()
.map(|p| p.display().to_string())
.unwrap_or_default(),
canonical.display(),
chain.join(" -> ")
));
}
if seen.contains(&canonical) {
continue;
}
let content = std::fs::read_to_string(&include_path).map_err(|e| {
anyhow::anyhow!(
"Failed to read #include \"{}\": {} (resolved to: {})",
path_str,
e,
include_path.display()
)
})?;
seen.insert(canonical.clone());
include_stack.push(canonical.clone());
let nested_base = include_path.parent().unwrap_or(base_dir);
let resolved = resolve_includes_inner(&content, nested_base, include_stack, seen)?;
include_stack.pop();
output.push_str(&resolved);
output.push('\n');
} else {
output.push_str(line);
output.push('\n');
}
}
Ok(output)
}
fn parse_include_directive(line: &str) -> Option<&str> {
let rest = if let Some(r) = line.strip_prefix("use ") {
r
} else {
line.strip_prefix("#include")?
};
let rest = rest.trim();
if rest.starts_with('"') && rest.ends_with('"') && rest.len() >= 2 {
Some(&rest[1..rest.len() - 1])
} else {
None
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_parse_array_with_size() {
let source = "struct Data { values: array<f32, 16> }";
let ast = parse_wjsl(source).unwrap();
let field = &ast.structs[0].fields[0];
match &field.ty {
Type::Array(elem, size) => {
assert_eq!(*size, Some(16));
assert!(matches!(**elem, Type::Scalar(ScalarType::F32)));
}
_ => panic!("Expected array<f32, 16>"),
}
}
#[test]
fn test_array_indexing_in_body() {
let source = r#"
@group(0) @binding(0) storage read clusters: array<vec4>;
@group(0) @binding(1) storage read_write instances: array<u32>;
@compute @workgroup_size(64, 1, 1)
fn main(@builtin(global_invocation_id) id: vec3<u32>) {
let cluster_id = id.x;
let cluster = clusters[cluster_id];
instances[cluster_id] = 1u;
}
"#;
let wgsl = transpile_wjsl(source).unwrap();
assert!(wgsl.contains("clusters[cluster_id]"));
assert!(wgsl.contains("instances[cluster_id]"));
}
}