lumen_engine/node/processing/
wgsl_shader.rs1use 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#[derive(Debug, Clone, lumen_macros::Node)]
15#[node(kind = "wgsl_shader", name = "WGSL Shader", category = "processing")]
16pub struct WgslShader {
17 pub id: NodeId,
18 #[property(
20 kind = "string",
21 name = "Shader source",
22 format = "wgsl",
23 multiline,
24 recommended_rows = 10
25 )]
26 pub shader: NodeProperty,
27 #[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(¶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}