Skip to main content

lumen_engine/node/processing/
wgsl_shader.rs

1use crate::{
2    expr::Expression,
3    node::{NodeId, NodeParamEvalContext, NodeParams, PortRef},
4};
5
6use crate::gpu::{
7    BoundFrame, CompiledOutput, FrameBindContext, GpuCompileNode, GpuCompiledNode, RasterHandle,
8    compiler,
9};
10
11pub(crate) const SHADER: &str = include_str!("wgsl_shader.wgsl");
12
13// TODO: replace the stringly shader/bindings surface with typed shader
14// parameters once custom-node authoring settles.
15/// Runs a custom WGSL compute shader over a raster.
16#[derive(Debug, Clone, Default, lumen_macros::Delegate)]
17pub struct WgslShaderParams {
18    /// Custom WGSL compute shader source.
19    #[meta(
20        name = "Shader source",
21        format = "wgsl",
22        multiline,
23        recommended_rows = 10
24    )]
25    pub shader: String,
26    /// JSON object describing shader binding values.
27    #[meta(
28        name = "Shader bindings",
29        format = "json",
30        multiline,
31        recommended_rows = 6
32    )]
33    pub bindings: String,
34}
35
36/// Runs a custom WGSL compute shader over a raster.
37#[derive(Debug, Clone, lumen_macros::Node)]
38#[node(kind = "wgsl_shader", name = "WGSL Shader", category = "processing")]
39pub struct WgslShader {
40    pub id: NodeId,
41    #[params]
42    pub params: WgslShaderParamsDelegate,
43    #[input()]
44    pub source: PortRef,
45}
46
47impl Default for WgslShader {
48    fn default() -> Self {
49        Self {
50            id: NodeId::new(0),
51            params: WgslShaderParamsDelegate::default(),
52            source: PortRef::empty(),
53        }
54    }
55}
56
57impl GpuCompileNode for WgslShader {
58    fn compile_gpu(
59        &self,
60        ctx: &mut crate::gpu::CompileContext<'_>,
61        port: &PortRef,
62    ) -> crate::Result<CompiledOutput> {
63        let params = self.params.eval(&NodeParamEvalContext {
64            node_id: self.id,
65            expr: &ctx.expr_context(self.id, "params"),
66        })?;
67        let shader = if params.shader.trim().is_empty() {
68            SHADER
69        } else {
70            params.shader.as_str()
71        };
72        let (source, texture, params) = ctx.compile_unary_filter(
73            self.id,
74            &self.source,
75            port,
76            "wgsl-shader",
77            shader,
78            std::mem::size_of::<compiler::WgslShaderParams>() as u64,
79        )?;
80        ctx.register_compiled_node(CompiledWgslShader {
81            node_id: self.id,
82            params: self.params.clone(),
83            buffer: params,
84        });
85        Ok(CompiledOutput::Raster(RasterHandle {
86            texture,
87            domain: source.domain,
88            metadata: source.metadata,
89        }))
90    }
91}
92
93#[derive(Debug, Clone)]
94struct CompiledWgslShader {
95    node_id: NodeId,
96    params: WgslShaderParamsDelegate,
97    buffer: lumen_gpu::BufferId,
98}
99
100impl GpuCompiledNode for CompiledWgslShader {
101    fn node_id(&self) -> NodeId {
102        self.node_id
103    }
104
105    fn bind(&self, ctx: &FrameBindContext<'_>, bound: &mut BoundFrame) -> crate::Result<()> {
106        let params = self.params.eval(&NodeParamEvalContext {
107            node_id: self.node_id,
108            expr: &ctx.expr_context(self.node_id, "params"),
109        })?;
110        let params = compiler::WgslShaderParams {
111            values: resolve_shader_values(
112                self.node_id,
113                &ctx.expr_context(self.node_id, "bindings"),
114                &params.shader,
115                &params.bindings,
116            )?,
117        };
118        bound.write_buffer(self.buffer, 0, bytemuck::bytes_of(&params));
119        Ok(())
120    }
121}
122
123fn resolve_shader_values(
124    node_id: NodeId,
125    ctx: &crate::expr::ExpressionContext<'_>,
126    shader: &str,
127    bindings: &str,
128) -> crate::Result<[f32; 4]> {
129    let fields = uniform_field_names(shader);
130    let mut values = [0.0; 4];
131
132    for raw_line in bindings.lines() {
133        let line = raw_line.split('#').next().unwrap_or("").trim();
134        if line.is_empty() {
135            continue;
136        }
137        let Some((raw_name, raw_value)) = line.split_once('=') else {
138            continue;
139        };
140        let name = raw_name.trim();
141        let raw_value = raw_value.trim();
142        if name.is_empty() || raw_value.starts_with("input") {
143            continue;
144        }
145
146        let Some(offset) = field_offset(&fields, name) else {
147            continue;
148        };
149        for (index, component) in parse_binding_values(node_id, name, raw_value, ctx)?
150            .into_iter()
151            .take(4 - offset)
152            .enumerate()
153        {
154            values[offset + index] = component;
155        }
156    }
157
158    Ok(values)
159}
160
161fn parse_binding_values(
162    node_id: NodeId,
163    name: &str,
164    raw_value: &str,
165    ctx: &crate::expr::ExpressionContext<'_>,
166) -> crate::Result<Vec<f32>> {
167    let raw_value = raw_value.trim();
168    if let Some(expression) = raw_value.strip_prefix('=') {
169        let value = Expression::parse(expression)
170            .map_err(crate::error::LumenError::from)?
171            .evaluate(ctx)?
172            .as_f64()
173            .ok_or_else(|| crate::error::PropertyError::InvalidType {
174                node_id,
175                property_path: format!("bindings.{name}"),
176                expected: "Float",
177                actual: "Expression",
178            })?;
179        return Ok(vec![value as f32]);
180    }
181
182    let delimiter: &[_] = if raw_value.contains(',') {
183        &[',']
184    } else {
185        &[' ', '\t']
186    };
187    Ok(raw_value
188        .split(delimiter)
189        .filter(|part| !part.trim().is_empty())
190        .filter_map(|part| part.trim().parse::<f32>().ok())
191        .collect())
192}
193
194fn field_offset(fields: &[UniformField], name: &str) -> Option<usize> {
195    let mut offset = 0;
196    for field in fields {
197        if field.name == name {
198            return Some(offset);
199        }
200        offset += field.lanes;
201        if offset >= 4 {
202            break;
203        }
204    }
205    None
206}
207
208#[derive(Debug, Clone)]
209struct UniformField {
210    name: String,
211    lanes: usize,
212}
213
214fn uniform_field_names(shader: &str) -> Vec<UniformField> {
215    let struct_name = shader
216        .split("var<uniform>")
217        .nth(1)
218        .and_then(|tail| tail.split(';').next())
219        .and_then(|decl| decl.split(':').nth(1))
220        .map(str::trim)
221        .filter(|name| !name.is_empty())
222        .unwrap_or("ShaderParams");
223    let Some(struct_body) = extract_struct_body(shader, struct_name) else {
224        return vec![UniformField {
225            name: "values".to_string(),
226            lanes: 4,
227        }];
228    };
229
230    struct_body
231        .split([',', '\n'])
232        .filter_map(|line| {
233            let line = line.trim();
234            let (name, ty) = line.split_once(':')?;
235            let name = name.trim();
236            if name.is_empty() {
237                return None;
238            }
239            Some(UniformField {
240                name: name.to_string(),
241                lanes: lanes_for_wgsl_type(ty.trim()),
242            })
243        })
244        .collect()
245}
246
247fn extract_struct_body<'a>(shader: &'a str, struct_name: &str) -> Option<&'a str> {
248    let start = shader.find(&format!("struct {struct_name}"))?;
249    let body_start = shader[start..].find('{')? + start + 1;
250    let body_end = shader[body_start..].find('}')? + body_start;
251    Some(&shader[body_start..body_end])
252}
253
254fn lanes_for_wgsl_type(ty: &str) -> usize {
255    if ty.contains("vec4") {
256        4
257    } else if ty.contains("vec3") {
258        3
259    } else if ty.contains("vec2") {
260        2
261    } else {
262        1
263    }
264}