Skip to main content

lumen_engine/node/processing/
wgsl_shader.rs

1use crate::{
2    expr::Expression,
3    node::{NodeId, NodeProperty, PortRef},
4};
5
6use crate::gpu::{
7    BoundFrame, CompiledOutput, FrameBindContext, FrameBinding, GpuCompileNode, GpuFrameBindNode,
8    RasterHandle, compiler,
9};
10
11pub(crate) const SHADER: &str = include_str!("wgsl_shader.wgsl");
12
13/// Runs a custom WGSL compute shader over a raster.
14#[derive(Debug, Clone, lumen_macros::Node)]
15#[node(kind = "wgsl_shader", name = "WGSL Shader", category = "processing")]
16pub struct WgslShader {
17    pub id: NodeId,
18    /// Custom WGSL compute shader source.
19    #[property(
20        kind = "string",
21        name = "Shader source",
22        format = "wgsl",
23        multiline,
24        recommended_rows = 10
25    )]
26    pub shader: NodeProperty,
27    /// JSON object describing shader binding values.
28    #[property(
29        kind = "string",
30        name = "Shader bindings",
31        format = "json",
32        multiline,
33        recommended_rows = 6
34    )]
35    pub bindings: NodeProperty,
36    #[input()]
37    pub source: PortRef,
38}
39
40impl Default for WgslShader {
41    fn default() -> Self {
42        Self {
43            id: NodeId::new(0),
44            shader: NodeProperty::String(String::new()),
45            bindings: NodeProperty::String(String::new()),
46            source: PortRef::empty(),
47        }
48    }
49}
50
51impl GpuCompileNode for WgslShader {
52    fn compile_gpu(
53        &self,
54        ctx: &mut crate::gpu::CompileContext<'_>,
55        port: &PortRef,
56    ) -> crate::Result<CompiledOutput> {
57        let shader =
58            self.shader
59                .resolve_string(self.id, "shader", &ctx.expr_context(self.id, "shader"))?;
60        let shader = if shader.trim().is_empty() {
61            SHADER
62        } else {
63            shader.as_str()
64        };
65        let (source, texture, params) = ctx.compile_unary_filter(
66            self.id,
67            &self.source,
68            port,
69            "wgsl-shader",
70            shader,
71            std::mem::size_of::<compiler::WgslShaderParams>() as u64,
72        )?;
73        ctx.push_frame_binding(FrameBinding::WgslShader {
74            node_id: self.id,
75            shader: self.shader.clone(),
76            bindings: self.bindings.clone(),
77            buffer: params,
78        });
79        Ok(CompiledOutput::Raster(RasterHandle {
80            texture,
81            domain: source.domain,
82            metadata: source.metadata,
83        }))
84    }
85}
86
87impl GpuFrameBindNode for WgslShader {
88    fn bind_gpu_frame(
89        &self,
90        ctx: &FrameBindContext<'_>,
91        binding: &FrameBinding,
92        bound: &mut BoundFrame,
93    ) -> crate::Result<()> {
94        let FrameBinding::WgslShader {
95            node_id,
96            shader,
97            bindings,
98            buffer,
99        } = binding
100        else {
101            return Ok(());
102        };
103        let shader =
104            shader.resolve_string(*node_id, "shader", &ctx.expr_context(*node_id, "shader"))?;
105        let bindings = bindings.resolve_string(
106            *node_id,
107            "bindings",
108            &ctx.expr_context(*node_id, "bindings"),
109        )?;
110        let params = compiler::WgslShaderParams {
111            values: resolve_shader_values(
112                *node_id,
113                &ctx.expr_context(*node_id, "bindings"),
114                &shader,
115                &bindings,
116            )?,
117        };
118        bound.write_buffer(*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}