lumen_engine/node/processing/
wgsl_shader.rs1use 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#[derive(Debug, Clone, Default, lumen_macros::Delegate)]
17pub struct WgslShaderParams {
18 #[meta(
20 name = "Shader source",
21 format = "wgsl",
22 multiline,
23 recommended_rows = 10
24 )]
25 pub shader: String,
26 #[meta(
28 name = "Shader bindings",
29 format = "json",
30 multiline,
31 recommended_rows = 6
32 )]
33 pub bindings: String,
34}
35
36#[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 ¶ms.shader,
115 ¶ms.bindings,
116 )?,
117 };
118 bound.write_buffer(self.buffer, 0, bytemuck::bytes_of(¶ms));
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}