1use std::hash::Hash;
4use std::{
5 cell::RefCell,
6 fmt::Write as _,
7 num::{NonZeroU32, NonZeroUsize},
8 ops::Deref,
9 rc::{Rc, Weak},
10};
11
12use super::parser;
13
14pub type ParentNodeReference = Weak<RefCell<Node>>;
15
16#[derive(Clone)]
17pub struct NodeReference(Rc<RefCell<Node>>);
18
19impl std::fmt::Debug for NodeReference {
20 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
21 self.0.borrow().fmt(f)
22 }
23}
24
25impl NodeReference {
26 pub fn new<F, E>(f: F) -> Result<NodeReference, E>
27 where
28 F: FnOnce(ParentNodeReference) -> Result<Node, E>,
29 {
30 let mut error = None;
31
32 let node = Rc::new_cyclic(|r| match f(r.clone()) {
33 Ok(node) => RefCell::new(node),
34 Err(e) => {
35 error = Some(e);
36 RefCell::new(Node::root())
37 }
38 });
39
40 if let Some(e) = error {
41 Err(e)
42 } else {
43 Ok(NodeReference(node))
44 }
45 }
46
47 pub fn get_descendant(&self, child_name: &str) -> Option<NodeReference> {
49 find_descendant(self, child_name, DescendantSearch::Any)
50 }
51
52 pub fn get_children(&self) -> Option<Vec<NodeReference>> {
53 self.borrow().get_children()
54 }
55
56 pub(crate) fn identity(&self) -> usize {
58 Rc::as_ptr(&self.0) as usize
59 }
60
61 pub fn get_main(&self) -> Option<NodeReference> {
63 if let Some(m) = self.get_descendant("main") {
64 return Some(m);
65 } else {
66 for child in self.get_children()? {
67 if let Some(m) = child.get_main() {
68 return Some(m);
69 }
70 }
71 }
72
73 None
74 }
75}
76
77impl From<Node> for NodeReference {
78 fn from(node: Node) -> Self {
79 NodeReference(Rc::new(RefCell::new(node)))
80 }
81}
82
83impl PartialEq for NodeReference {
84 fn eq(&self, other: &Self) -> bool {
85 Rc::ptr_eq(&self.0, &other.0)
86 }
87}
88
89impl Eq for NodeReference {}
90
91impl Hash for NodeReference {
92 fn hash<H>(&self, state: &mut H)
93 where
94 H: std::hash::Hasher,
95 {
96 Rc::as_ptr(&self.0).hash(state);
97 }
98}
99
100impl Deref for NodeReference {
101 type Target = RefCell<Node>;
102
103 fn deref(&self) -> &Self::Target {
104 &self.0
105 }
106}
107
108pub(super) fn lex(mut node: parser::Node) -> Result<NodeReference, LexError> {
109 node.sort();
110 lex_with_root(Node::root(), node)
111}
112
113pub(super) fn lex_with_root(root: Node, mut node: parser::Node) -> Result<NodeReference, LexError> {
114 node.sort();
115
116 let root: NodeReference = root.into();
117
118 match &node.node {
119 parser::Nodes::Scope { name, children } => {
120 assert_eq!(*name, "root");
121
122 for child in children {
123 let c = lex_parsed_node(vec![root.clone()], child)?;
124 root.borrow_mut().add_child(c);
125 }
126
127 Ok(root)
128 }
129 _ => Err(LexError::Undefined { message: None }),
130 }
131}
132
133#[derive(Clone)]
134pub struct Node {
135 node: Nodes,
137}
138
139impl Node {
140 fn internal_new(node: Node) -> NodeReference {
141 NodeReference(Rc::new(RefCell::new(node)))
142 }
143
144 pub fn root() -> Node {
146 let void = primitive_type("void");
147 let bool_t = primitive_type("bool");
148 let u8_t = primitive_type("u8");
149 let u16_t = primitive_type("u16");
150 let u32_t = primitive_type("u32");
151 let i32_t = primitive_type("i32");
152 let f32_t = primitive_type("f32");
153
154 let vec2u16 = record_type("vec2u16", [("x", u16_t.clone()), ("y", u16_t.clone())]);
155 let vec4u16 = record_type(
156 "vec4u16",
157 [
158 ("x", u16_t.clone()),
159 ("y", u16_t.clone()),
160 ("z", u16_t.clone()),
161 ("w", u16_t.clone()),
162 ],
163 );
164 let vec2u32 = record_type("vec2u", [("x", u32_t.clone()), ("y", u32_t.clone())]);
165 let vec2i32 = record_type("vec2i", [("x", i32_t.clone()), ("y", i32_t.clone())]);
166 let vec2f32 = record_type("vec2f", [("x", f32_t.clone()), ("y", f32_t.clone())]);
167 let vec3f32 = record_type("vec3f", [("x", f32_t.clone()), ("y", f32_t.clone()), ("z", f32_t.clone())]);
168 let vec3u32 = record_type("vec3u", [("x", u32_t.clone()), ("y", u32_t.clone()), ("z", u32_t.clone())]);
169 let vec4u32 = record_type(
170 "vec4u",
171 [
172 ("x", u32_t.clone()),
173 ("y", u32_t.clone()),
174 ("z", u32_t.clone()),
175 ("w", u32_t.clone()),
176 ],
177 );
178 let vec4f32 = record_type(
179 "vec4f",
180 [
181 ("x", f32_t.clone()),
182 ("y", f32_t.clone()),
183 ("z", f32_t.clone()),
184 ("w", f32_t.clone()),
185 ],
186 );
187 let mat4f32 = record_type(
188 "mat4f",
189 [
190 ("x", vec4f32.clone()),
191 ("y", vec4f32.clone()),
192 ("z", vec4f32.clone()),
193 ("w", vec4f32.clone()),
194 ],
195 );
196 let mat4x3f32 = record_type(
197 "mat4x3f",
198 [
199 ("x", vec3f32.clone()),
200 ("y", vec3f32.clone()),
201 ("z", vec3f32.clone()),
202 ("w", vec3f32.clone()),
203 ],
204 );
205
206 let texture_2d = primitive_type("Texture2D");
207 let texture_3d = primitive_type("Texture3D");
208 let array_texture_2d = primitive_type("ArrayTexture2D");
209 let atomic_u32 = primitive_type("atomicu32");
210
211 let builtins = vec![
212 void.clone(),
213 bool_t,
214 u8_t.clone(),
215 u16_t.clone(),
216 u32_t.clone(),
217 i32_t.clone(),
218 f32_t.clone(),
219 vec2u16,
220 vec4u16,
221 vec2u32.clone(),
222 vec2i32,
223 vec2f32.clone(),
224 vec3u32.clone(),
225 vec3f32.clone(),
226 vec4u32,
227 vec4f32.clone(),
228 mat4f32,
229 mat4x3f32,
230 texture_2d.clone(),
231 texture_3d.clone(),
232 array_texture_2d,
233 atomic_u32.clone(),
234 builtin_intrinsic(
235 "sample",
236 vec![("texture_sampler", texture_2d.clone()), ("uv", vec2f32.clone())],
237 vec4f32.clone(),
238 ),
239 builtin_intrinsic(
240 "texture_lod",
241 vec![("texture", texture_2d.clone()), ("uv", vec2f32.clone())],
242 vec4f32.clone(),
243 ),
244 builtin_intrinsic(
245 "texture_lod",
246 vec![("texture", texture_3d.clone()), ("uv", vec3f32.clone())],
247 vec4f32.clone(),
248 ),
249 builtin_intrinsic(
250 "fetch",
251 vec![("texture", texture_2d.clone()), ("coord", vec2u32.clone())],
252 vec4f32.clone(),
253 ),
254 builtin_intrinsic(
255 "fetch_u32",
256 vec![("texture", texture_2d.clone()), ("coord", vec2u32.clone())],
257 u32_t.clone(),
258 ),
259 builtin_intrinsic(
260 "dot",
261 vec![("left", vec2f32.clone()), ("right", vec2f32.clone())],
262 f32_t.clone(),
263 ),
264 builtin_intrinsic(
265 "dot",
266 vec![("left", vec4f32.clone()), ("right", vec4f32.clone())],
267 f32_t.clone(),
268 ),
269 builtin_intrinsic(
270 "dot",
271 vec![("left", vec3f32.clone()), ("right", vec3f32.clone())],
272 f32_t.clone(),
273 ),
274 builtin_intrinsic(
275 "cross",
276 vec![("left", vec3f32.clone()), ("right", vec3f32.clone())],
277 vec3f32.clone(),
278 ),
279 builtin_intrinsic("length", vec![("value", vec4f32.clone())], f32_t.clone()),
280 builtin_intrinsic("length", vec![("value", vec3f32.clone())], f32_t.clone()),
281 builtin_intrinsic("normalize", vec![("value", vec4f32.clone())], vec4f32.clone()),
282 builtin_intrinsic("normalize", vec![("value", vec3f32.clone())], vec3f32.clone()),
283 builtin_intrinsic("max", vec![("left", f32_t.clone()), ("right", f32_t.clone())], f32_t.clone()),
284 builtin_intrinsic("min", vec![("left", f32_t.clone()), ("right", f32_t.clone())], f32_t.clone()),
285 builtin_intrinsic(
286 "max",
287 vec![("left", vec2f32.clone()), ("right", vec2f32.clone())],
288 vec2f32.clone(),
289 ),
290 builtin_intrinsic(
291 "max",
292 vec![("left", vec3f32.clone()), ("right", vec3f32.clone())],
293 vec3f32.clone(),
294 ),
295 builtin_intrinsic(
296 "clamp",
297 vec![
298 ("value", f32_t.clone()),
299 ("minimum", f32_t.clone()),
300 ("maximum", f32_t.clone()),
301 ],
302 f32_t.clone(),
303 ),
304 builtin_intrinsic(
305 "clamp",
306 vec![
307 ("value", vec3f32.clone()),
308 ("minimum", vec3f32.clone()),
309 ("maximum", vec3f32.clone()),
310 ],
311 vec3f32.clone(),
312 ),
313 builtin_intrinsic("log2", vec![("value", vec3f32.clone())], vec3f32.clone()),
314 builtin_intrinsic(
315 "pow",
316 vec![("value", vec3f32.clone()), ("exponent", vec3f32.clone())],
317 vec3f32.clone(),
318 ),
319 builtin_intrinsic(
320 "pow",
321 vec![("value", f32_t.clone()), ("exponent", f32_t.clone())],
322 f32_t.clone(),
323 ),
324 builtin_intrinsic(
325 "reflect",
326 vec![("incident", vec4f32.clone()), ("normal", vec4f32.clone())],
327 vec4f32.clone(),
328 ),
329 builtin_intrinsic("abs", vec![("value", f32_t.clone())], f32_t.clone()),
330 builtin_intrinsic("abs", vec![("value", vec2f32.clone())], vec2f32.clone()),
331 builtin_intrinsic("sqrt", vec![("value", f32_t.clone())], f32_t.clone()),
332 builtin_intrinsic("exp", vec![("value", f32_t.clone())], f32_t.clone()),
333 builtin_intrinsic("exp", vec![("value", vec3f32.clone())], vec3f32.clone()),
334 builtin_intrinsic("sin", vec![("value", f32_t.clone())], f32_t.clone()),
335 builtin_intrinsic("cos", vec![("value", f32_t.clone())], f32_t.clone()),
336 builtin_intrinsic("tan", vec![("value", f32_t.clone())], f32_t.clone()),
337 builtin_intrinsic("round", vec![("value", vec2f32.clone())], vec2f32.clone()),
338 builtin_intrinsic("fract", vec![("value", f32_t.clone())], f32_t.clone()),
339 builtin_intrinsic("fwidth", vec![("value", f32_t.clone())], f32_t.clone()),
340 builtin_intrinsic("radians", vec![("value", f32_t.clone())], f32_t.clone()),
341 builtin_intrinsic("inversesqrt", vec![("value", f32_t.clone())], f32_t.clone()),
342 builtin_intrinsic("f32", vec![("value", u32_t.clone())], f32_t.clone()),
343 builtin_intrinsic("f32", vec![("value", i32_t.clone())], f32_t.clone()),
344 builtin_intrinsic("u32", vec![("value", u32_t.clone())], u32_t.clone()),
345 builtin_intrinsic("u32", vec![("value", u8_t.clone())], u32_t.clone()),
346 builtin_intrinsic("u32", vec![("value", u16_t.clone())], u32_t.clone()),
347 builtin_intrinsic("u32", vec![("value", i32_t)], u32_t.clone()),
348 builtin_intrinsic("u32", vec![("value", f32_t.clone())], u32_t.clone()),
349 builtin_intrinsic(
350 "smoothstep",
351 vec![("edge0", f32_t.clone()), ("edge1", f32_t.clone()), ("value", f32_t.clone())],
352 f32_t.clone(),
353 ),
354 builtin_intrinsic("step", vec![("edge", f32_t.clone()), ("value", f32_t.clone())], f32_t.clone()),
355 builtin_intrinsic(
356 "mix",
357 vec![("left", f32_t.clone()), ("right", f32_t.clone()), ("factor", f32_t.clone())],
358 f32_t.clone(),
359 ),
360 builtin_intrinsic("thread_idx", vec![], u32_t.clone()),
361 builtin_intrinsic("threadgroup_position", vec![], u32_t.clone()),
362 builtin_intrinsic("thread_position", vec![], u32_t.clone()),
363 builtin_intrinsic("workgroup_barrier", vec![], void.clone()),
364 builtin_intrinsic("set_task_mesh_output_count", vec![("count", u32_t.clone())], void.clone()),
365 builtin_intrinsic("thread_id", vec![], vec2u32.clone()),
366 builtin_intrinsic(
367 "set_mesh_output_counts",
368 vec![("vertex_count", u32_t.clone()), ("primitive_count", u32_t.clone())],
369 void.clone(),
370 ),
371 builtin_intrinsic(
372 "set_mesh_vertex_position",
373 vec![("vertex_index", u32_t.clone()), ("position", vec4f32.clone())],
374 void.clone(),
375 ),
376 builtin_intrinsic(
377 "set_mesh_triangle",
378 vec![("primitive_index", u32_t.clone()), ("triangle", vec3u32.clone())],
379 void.clone(),
380 ),
381 builtin_intrinsic(
382 "image_load",
383 vec![("image", texture_2d.clone()), ("coord", vec2u32.clone())],
384 vec4f32.clone(),
385 ),
386 builtin_intrinsic(
387 "image_load_u32",
388 vec![("image", texture_2d.clone()), ("coord", vec2u32.clone())],
389 u32_t.clone(),
390 ),
391 builtin_intrinsic(
392 "atomic_add",
393 vec![("value", atomic_u32.clone()), ("increment", u32_t.clone())],
394 u32_t.clone(),
395 ),
396 builtin_intrinsic("atomic_load", vec![("value", atomic_u32.clone())], u32_t.clone()),
397 builtin_intrinsic(
398 "atomic_store",
399 vec![("value", atomic_u32), ("stored", u32_t.clone())],
400 void.clone(),
401 ),
402 builtin_intrinsic("texture_size", vec![("texture", texture_2d.clone())], vec2u32.clone()),
403 builtin_intrinsic("image_size", vec![("image", texture_2d.clone())], vec2u32.clone()),
404 builtin_intrinsic(
405 "guard_image_bounds",
406 vec![("image", texture_2d.clone()), ("coord", vec2u32.clone())],
407 void.clone(),
408 ),
409 builtin_intrinsic(
410 "write",
411 vec![
412 ("image", texture_2d.clone()),
413 ("coord", vec2u32.clone()),
414 ("value", vec4f32.clone()),
415 ],
416 void.clone(),
417 ),
418 builtin_intrinsic(
419 "image_atomic_or",
420 vec![
421 ("image", texture_2d.clone()),
422 ("coord", vec2u32.clone()),
423 ("value", u32_t.clone()),
424 ],
425 u32_t.clone(),
426 ),
427 ];
428
429 let mut root = Node::scope("root".to_string());
430 root.add_children(builtins);
431
432 root
433 }
434
435 pub fn scope(name: String) -> Node {
437 Node {
438 node: Nodes::Scope {
440 name,
441 children: Vec::with_capacity(16),
442 },
443 }
444 }
445
446 pub fn r#struct(name: &str, fields: Vec<NodeReference>) -> Node {
448 Node {
449 node: Nodes::Struct {
450 name: name.to_string(),
451 template: None,
452 fields,
453 types: Vec::new(),
454 },
455 }
456 }
457
458 pub fn member(name: &str, r#type: NodeReference) -> Node {
459 Node {
460 node: Nodes::Member {
461 name: name.to_string(),
462 r#type,
463 count: None,
464 },
465 }
466 }
467
468 pub fn array(name: &str, r#type: NodeReference, size: usize) -> NodeReference {
469 Self::internal_new(Node {
470 node: Nodes::Member {
471 name: name.to_string(),
472 r#type,
473 count: Some(NonZeroUsize::new(size).expect("Invalid size")),
474 },
475 })
476 }
477
478 pub fn function(
479 name: &str,
480 params: Vec<NodeReference>,
481 return_type: NodeReference,
482 statements: Vec<NodeReference>,
483 ) -> Node {
484 Node {
485 node: Nodes::Function {
486 name: name.to_string(),
487 params,
488 return_type,
489 statements,
490 },
491 }
492 }
493
494 pub fn conditional(condition: NodeReference, statements: Vec<NodeReference>) -> Node {
495 Node {
496 node: Nodes::Conditional { condition, statements },
497 }
498 }
499
500 pub fn for_loop(
501 initializer: NodeReference,
502 condition: NodeReference,
503 update: NodeReference,
504 statements: Vec<NodeReference>,
505 ) -> Node {
506 Node {
507 node: Nodes::ForLoop {
508 initializer,
509 condition,
510 update,
511 statements,
512 },
513 }
514 }
515
516 pub fn expression(expression: Expressions) -> Node {
517 Node {
518 node: Nodes::Expression(expression),
519 }
520 }
521
522 pub fn glsl(code: String, inputs: Vec<NodeReference>, outputs: Vec<NodeReference>) -> Node {
523 Self::raw(Some(code), None, None, inputs, outputs)
524 }
525
526 pub fn hlsl(code: String, inputs: Vec<NodeReference>, outputs: Vec<NodeReference>) -> Node {
527 Self::raw(None, Some(code), None, inputs, outputs)
528 }
529
530 pub fn msl(code: String, inputs: Vec<NodeReference>, outputs: Vec<NodeReference>) -> Node {
531 Self::raw(None, None, Some(code), inputs, outputs)
532 }
533
534 pub fn raw(
536 glsl: Option<String>,
537 hlsl: Option<String>,
538 msl: Option<String>,
539 inputs: Vec<NodeReference>,
540 outputs: Vec<NodeReference>,
541 ) -> Node {
542 Node {
543 node: Nodes::Raw {
544 glsl,
545 hlsl,
546 msl,
547 input: inputs,
548 output: outputs,
549 },
550 }
551 }
552
553 pub fn r#macro(name: &str, body: NodeReference) -> Node {
554 Node {
555 node: Nodes::Expression(Expressions::Macro {
556 name: name.to_string(),
557 body,
558 }),
559 }
560 }
561
562 pub fn binding(name: &str, r#type: BindingTypes, slot: u32, read: bool, write: bool) -> Node {
563 Self::binding_with_count(name, r#type, slot, read, write, None)
564 }
565
566 fn binding_with_count(
567 name: &str,
568 r#type: BindingTypes,
569 slot: u32,
570 read: bool,
571 write: bool,
572 count: Option<NonZeroU32>,
573 ) -> Node {
574 Node {
575 node: Nodes::Binding {
576 name: name.to_string(),
577 r#type,
578 slot,
579 read,
580 write,
581 count,
582 },
583 }
584 }
585
586 pub fn binding_array(name: &str, r#type: BindingTypes, slot: u32, read: bool, write: bool, count: usize) -> Node {
587 let count = u32::try_from(count)
588 .expect("Invalid binding array count. The most likely cause is that a resource array exceeds u32::MAX elements.");
589 let count = NonZeroU32::new(count).expect(
590 "Invalid binding array count. The most likely cause is that a resource array was declared with zero elements.",
591 );
592 Self::binding_with_count(name, r#type, slot, read, write, Some(count))
593 }
594
595 pub fn push_constant(members: Vec<NodeReference>) -> Node {
596 Node {
597 node: Nodes::PushConstant { members },
598 }
599 }
600
601 pub fn intrinsic(name: &str, elements: Vec<NodeReference>, r#return: NodeReference) -> Node {
602 Node {
603 node: Nodes::Intrinsic {
604 name: name.to_string(),
605 elements,
606 r#return,
607 },
608 }
609 }
610
611 pub fn specialization(name: &str, r#type: NodeReference) -> Node {
612 Node {
613 node: Nodes::Specialization {
614 name: name.to_string(),
615 r#type,
616 },
617 }
618 }
619
620 pub fn constant(name: &str, r#type: NodeReference, value: NodeReference) -> Node {
621 Node {
622 node: Nodes::Const {
623 name: name.to_string(),
624 r#type,
625 value,
626 },
627 }
628 }
629
630 pub fn input(name: &str, format: NodeReference, location: u8) -> Node {
631 Node {
632 node: Nodes::Input {
633 name: name.to_string(),
634 format,
635 location,
636 },
637 }
638 }
639
640 pub fn output(name: &str, format: NodeReference, location: u8) -> Node {
641 Self::output_with_count(name, format, location, None)
642 }
643
644 pub fn output_array(name: &str, format: NodeReference, location: u8, count: u32) -> Node {
645 Self::output_with_count(name, format, location, NonZeroUsize::new(count as usize))
646 }
647
648 fn output_with_count(name: &str, format: NodeReference, location: u8, count: Option<NonZeroUsize>) -> Node {
649 Node {
650 node: Nodes::Output {
651 name: name.to_string(),
652 format,
653 location,
654 count,
655 },
656 }
657 }
658
659 pub fn task_payload(name: &str, format: NodeReference, count: u32) -> Node {
660 let count = NonZeroUsize::new(count as usize).expect(
661 "Invalid task-payload count. The most likely cause is that a task-payload array was declared with zero elements.",
662 );
663 Node {
664 node: Nodes::TaskPayload {
665 name: name.to_string(),
666 format,
667 count,
668 },
669 }
670 }
671
672 pub fn workgroup(name: &str, format: NodeReference) -> Node {
673 Node {
674 node: Nodes::Workgroup {
675 name: name.to_string(),
676 format,
677 },
678 }
679 }
680
681 pub fn new(node: Nodes) -> Node {
682 Node { node }
683 }
684
685 pub fn add_child(&mut self, child: NodeReference) -> NodeReference {
686 match &mut self.node {
687 Nodes::Scope { children, .. } => {
688 children.push(child.clone());
689 }
690 Nodes::Struct { fields, .. } => {
691 fields.push(child.clone());
692 }
693 Nodes::Function { statements, .. } => {
694 statements.push(child.clone());
695 }
696 Nodes::PushConstant { members } => {
697 members.push(child.clone());
698 }
699 Nodes::Intrinsic { elements, .. } => {
700 elements.push(child.clone());
701 }
702 _ => {}
703 }
704
705 child
706 }
707
708 pub fn add_children(&mut self, children: Vec<NodeReference>) -> Vec<NodeReference> {
709 let mut ch = Vec::with_capacity(children.len());
710
711 for child in children {
712 ch.push(self.add_child(child));
713 }
714
715 ch
716 }
717
718 pub fn node(&self) -> &Nodes {
719 &self.node
720 }
721
722 pub fn get_name(&self) -> Option<&str> {
723 match &self.node {
724 Nodes::Scope { name, .. }
725 | Nodes::Function { name, .. }
726 | Nodes::Member { name, .. }
727 | Nodes::Struct { name, .. }
728 | Nodes::Intrinsic { name, .. }
729 | Nodes::Binding { name, .. }
730 | Nodes::Parameter { name, .. }
731 | Nodes::Specialization { name, .. }
732 | Nodes::Literal { name, .. }
733 | Nodes::Const { name, .. } => Some(name),
734 Nodes::Input { name, .. }
735 | Nodes::Output { name, .. }
736 | Nodes::TaskPayload { name, .. }
737 | Nodes::Workgroup { name, .. } => Some(name),
738 Nodes::PushConstant { .. } => Some("push_constant"),
739 Nodes::Expression(Expressions::VariableDeclaration { name, .. } | Expressions::Member { name, .. }) => Some(name),
740 _ => None,
741 }
742 }
743
744 pub fn get_children(&self) -> Option<Vec<NodeReference>> {
745 match &self.node {
746 Nodes::Scope { children, .. }
747 | Nodes::Struct { fields: children, .. }
748 | Nodes::Intrinsic { elements: children, .. } => Some(children.clone()),
749 Nodes::Function { statements, .. } => Some(statements.clone()),
750 Nodes::Conditional { condition, statements } => {
751 let mut children = Vec::with_capacity(statements.len() + 1);
752 children.push(condition.clone());
753 children.extend(statements.iter().cloned());
754 Some(children)
755 }
756 Nodes::ForLoop {
757 initializer,
758 condition,
759 update,
760 statements,
761 } => {
762 let mut children = Vec::with_capacity(statements.len() + 3);
763 children.push(initializer.clone());
764 children.push(condition.clone());
765 children.push(update.clone());
766 children.extend(statements.iter().cloned());
767 Some(children)
768 }
769 Nodes::Expression(Expressions::IntrinsicCall { elements: children, .. }) => Some(children.clone()),
770 _ => None,
771 }
772 }
773
774 pub fn get_child(&self, child_name: &str) -> Option<NodeReference> {
775 self.get_children()?
776 .iter()
777 .find(|child| child.borrow().get_name() == Some(child_name))
778 .cloned()
779 }
780
781 pub fn node_mut(&mut self) -> &mut Nodes {
782 &mut self.node
783 }
784
785 pub fn null() -> Node {
786 Self { node: Nodes::Null }
787 }
788
789 fn sentence(elements: Vec<NodeReference>) -> Node {
790 Self {
791 node: Nodes::Expression(Expressions::Expression { elements }),
792 }
793 }
794}
795
796#[derive(Clone, Debug, PartialEq, Eq)]
797pub enum BindingTypes {
798 Buffer { members: Vec<NodeReference> },
799 CombinedImageSampler { format: String },
800 Image { format: String },
801}
802
803#[derive(Clone)]
804pub enum Nodes {
805 Null,
806 Scope {
807 name: String,
808 children: Vec<NodeReference>,
809 },
810 Struct {
811 name: String,
812 template: Option<NodeReference>,
813 fields: Vec<NodeReference>,
814 types: Vec<NodeReference>,
815 },
816 Member {
817 name: String,
818 r#type: NodeReference,
819 count: Option<NonZeroUsize>,
820 },
821 Function {
822 name: String,
823 params: Vec<NodeReference>,
824 return_type: NodeReference,
825 statements: Vec<NodeReference>,
826 },
827 Conditional {
828 condition: NodeReference,
829 statements: Vec<NodeReference>,
830 },
831 ForLoop {
832 initializer: NodeReference,
833 condition: NodeReference,
834 update: NodeReference,
835 statements: Vec<NodeReference>,
836 },
837 Specialization {
838 name: String,
839 r#type: NodeReference,
840 },
841 Expression(Expressions),
842 Raw {
843 glsl: Option<String>,
844 hlsl: Option<String>,
845 msl: Option<String>,
846 input: Vec<NodeReference>,
847 output: Vec<NodeReference>,
848 },
849 Binding {
850 name: String,
851 slot: u32,
852 read: bool,
853 write: bool,
854 r#type: BindingTypes,
855 count: Option<NonZeroU32>,
856 },
857 PushConstant {
858 members: Vec<NodeReference>,
859 },
860 Intrinsic {
861 name: String,
862 elements: Vec<NodeReference>,
863 r#return: NodeReference,
864 },
865 Input {
866 name: String,
867 format: NodeReference,
868 location: u8,
869 },
870 Output {
871 name: String,
872 format: NodeReference,
873 location: u8,
874 count: Option<NonZeroUsize>,
875 },
876 TaskPayload {
877 name: String,
878 format: NodeReference,
879 count: NonZeroUsize,
880 },
881 Workgroup {
882 name: String,
883 format: NodeReference,
884 },
885 Parameter {
886 name: String,
887 r#type: NodeReference,
888 },
889 Literal {
890 name: String,
891 value: NodeReference,
892 },
893 Const {
895 name: String,
896 r#type: NodeReference,
897 value: NodeReference,
898 },
899}
900
901impl Nodes {
902 pub fn is_leaf(&self) -> bool {
903 match self {
904 Nodes::Function { .. } => false,
905 Nodes::Conditional { .. } | Nodes::ForLoop { .. } => false,
906 Nodes::Struct { .. } => false,
907 Nodes::Binding { .. } => false,
908 Nodes::PushConstant { .. } => false,
909 Nodes::Input { .. } | Nodes::Output { .. } | Nodes::TaskPayload { .. } | Nodes::Workgroup { .. } => false,
910 Nodes::Specialization { .. } => false,
911 Nodes::Const { .. } => false,
912 Nodes::Literal { .. } => true,
913 Nodes::Parameter { .. } => true,
914 Nodes::Null => true,
915 Nodes::Scope { .. } => true,
916 Nodes::Intrinsic { .. } => true,
917 Nodes::Member { .. } => true,
918 Nodes::Expression { .. } => true,
919 Nodes::Raw { .. } => true,
920 }
921 }
922
923 pub fn is_indexable(&self) -> bool {
924 fn type_is_indexable(r#type: &NodeReference) -> bool {
925 let r#type = r#type.borrow();
926 matches!(r#type.node(), Nodes::Struct { template: Some(_), .. })
927 || r#type
928 .get_name()
929 .is_some_and(|name| name.starts_with("vec") || name.starts_with("mat"))
930 }
931
932 match self {
933 Nodes::Member { r#type, count, .. } => count.is_some() || type_is_indexable(r#type),
934 Nodes::Input { format, .. } => type_is_indexable(format),
935 Nodes::Output { format, count, .. } => count.is_some() || type_is_indexable(format),
936 Nodes::TaskPayload { .. } => true,
937 Nodes::Parameter { r#type, .. }
938 | Nodes::Specialization { r#type, .. }
939 | Nodes::Const { r#type, .. }
940 | Nodes::Expression(Expressions::VariableDeclaration { r#type, .. }) => type_is_indexable(r#type),
941 Nodes::Expression(Expressions::Member { source, .. }) => source.borrow().node().is_indexable(),
942 Nodes::Expression(Expressions::Accessor { right, .. }) => right.borrow().node().is_indexable(),
943 _ => false,
944 }
945 }
946
947 pub fn is_buffer_binding(&self) -> bool {
948 match self {
949 Nodes::Binding {
950 r#type: BindingTypes::Buffer { .. },
951 ..
952 } => true,
953 Nodes::Expression(Expressions::Member { source, .. }) => source.borrow().node().is_buffer_binding(),
954 _ => false,
955 }
956 }
957}
958
959impl std::fmt::Debug for Node {
960 fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
961 match &self.node {
962 Nodes::Null => {
963 write!(f, "Null")
964 }
965 Nodes::Scope { name, children } => {
966 write!(
967 f,
968 "Scope {{ name: {}, children: {:#?} }}",
969 name,
970 children.iter().map(|c| c.0.borrow().get_name().map(|e| e.to_string()))
971 )
972 }
973 Nodes::Struct { name, fields, .. } => {
974 write!(
975 f,
976 "Struct {{ name: {}, fields: {:?} }}",
977 name,
978 fields.iter().map(|c| c.0.borrow().get_name().map(|e| e.to_string()))
979 )
980 }
981 Nodes::Member { name, r#type, .. } => {
982 write!(
983 f,
984 "Member {{ name: {}, type: {:?} }}",
985 name,
986 r#type.0.borrow().get_name().map(|e| e.to_string())
987 )
988 }
989 Nodes::Function {
990 name,
991 params,
992 statements,
993 ..
994 } => {
995 write!(
996 f,
997 "Function {{ name: {}, parameters: {:?}, statements: {:?} }}",
998 name,
999 params.iter().map(|c| c.0.borrow().get_name().map(|e| e.to_string())),
1000 statements.iter().map(|c| c.0.borrow().get_name().map(|e| e.to_string()))
1001 )
1002 }
1003 Nodes::Conditional { condition, statements } => {
1004 write!(
1005 f,
1006 "Conditional {{ condition: {:?}, statements: {:?} }}",
1007 condition, statements
1008 )
1009 }
1010 Nodes::ForLoop {
1011 initializer,
1012 condition,
1013 update,
1014 statements,
1015 } => {
1016 write!(
1017 f,
1018 "ForLoop {{ initializer: {:?}, condition: {:?}, update: {:?}, statements: {:?} }}",
1019 initializer, condition, update, statements
1020 )
1021 }
1022 Nodes::Specialization { name, r#type } => {
1023 write!(
1024 f,
1025 "Specialization {{ name: {}, type: {:?} }}",
1026 name,
1027 r#type.0.borrow().get_name().map(|e| e.to_string())
1028 )
1029 }
1030 Nodes::Expression(expression) => {
1031 write!(f, "Expression {{ {:?} }}", expression)
1032 }
1033 Nodes::Raw {
1034 glsl,
1035 hlsl,
1036 msl,
1037 input,
1038 output,
1039 } => {
1040 write!(
1041 f,
1042 "RawCode {{ glsl: {:?}, hlsl: {:?}, msl: {:?}, input: {:?}, output: {:?} }}",
1043 glsl,
1044 hlsl,
1045 msl,
1046 input.iter().map(|c| c.0.borrow().get_name().map(|e| e.to_string())),
1047 output.iter().map(|c| c.0.borrow().get_name().map(|e| e.to_string()))
1048 )
1049 }
1050 Nodes::Binding {
1051 name,
1052 slot,
1053 read,
1054 write,
1055 r#type,
1056 count,
1057 } => {
1058 write!(
1059 f,
1060 "Binding {{ name: {}, slot: {}, read: {}, write: {}, type: {:?}, count: {:?} }}",
1061 name, slot, read, write, r#type, count
1062 )
1063 }
1064 Nodes::PushConstant { members } => {
1065 write!(
1066 f,
1067 "PushConstant {{ members: {:?} }}",
1068 members.iter().map(|c| c.0.borrow().get_name().map(|e| e.to_string()))
1069 )
1070 }
1071 Nodes::Intrinsic {
1072 name,
1073 elements,
1074 r#return,
1075 } => {
1076 write!(
1077 f,
1078 "Intrinsic {{ name: {}, elements: {:?}, return: {:?} }}",
1079 name,
1080 elements.iter().map(|c| c.0.borrow().get_name().map(|e| e.to_string())),
1081 r#return.0.borrow().get_name().map(|e| e.to_string())
1082 )
1083 }
1084 Nodes::Parameter { name, r#type } => {
1085 write!(
1086 f,
1087 "Parameter {{ name: {}, type: {:?} }}",
1088 name,
1089 r#type.0.borrow().get_name().map(|e| e.to_string())
1090 )
1091 }
1092 Nodes::Input { name, format, location } => {
1093 write!(
1094 f,
1095 "Input {{ name: {}, format: {:?}, location: {} }}",
1096 name,
1097 format.0.borrow().get_name().map(|e| e.to_string()),
1098 location
1099 )
1100 }
1101 Nodes::Output {
1102 name,
1103 format,
1104 location,
1105 count,
1106 } => {
1107 write!(
1108 f,
1109 "Output {{ name: {}, format: {:?}, location: {}, count: {:?} }}",
1110 name,
1111 format.0.borrow().get_name().map(|e| e.to_string()),
1112 location,
1113 count
1114 )
1115 }
1116 Nodes::TaskPayload { name, format, count } => {
1117 write!(
1118 f,
1119 "TaskPayload {{ name: {}, format: {:?}, count: {} }}",
1120 name,
1121 format.0.borrow().get_name().map(|e| e.to_string()),
1122 count
1123 )
1124 }
1125 Nodes::Workgroup { name, format } => {
1126 write!(
1127 f,
1128 "Workgroup {{ name: {}, format: {:?} }}",
1129 name,
1130 format.0.borrow().get_name().map(|e| e.to_string())
1131 )
1132 }
1133 Nodes::Literal { name, value } => {
1134 write!(
1135 f,
1136 "Literal {{ name: {}, value: {:?} }}",
1137 name,
1138 value.0.borrow().get_name().map(|e| e.to_string())
1139 )
1140 }
1141 Nodes::Const { name, r#type, value } => {
1142 write!(
1143 f,
1144 "Const {{ name: {}, type: {:?}, value: {:?} }}",
1145 name,
1146 r#type.0.borrow().get_name().map(|e| e.to_string()),
1147 value
1148 )
1149 }
1150 }
1151 }
1152}
1153
1154#[derive(Clone, Debug, PartialEq, Eq)]
1155pub enum Operators {
1156 Plus,
1157 Minus,
1158 Multiply,
1159 Divide,
1160 Modulo,
1161 ShiftLeft,
1162 ShiftRight,
1163 BitwiseAnd,
1164 BitwiseOr,
1165 Assignment,
1166 Equality,
1167 LessThan,
1168 Inequality,
1169 GreaterThan,
1170 LessThanOrEqual,
1171 GreaterThanOrEqual,
1172 LogicalAnd,
1173 LogicalOr,
1174}
1175
1176#[derive(Clone, Debug)]
1177pub enum Expressions {
1178 Return {
1179 value: Option<NodeReference>,
1180 },
1181 Continue,
1182 Member {
1183 name: String,
1184 source: NodeReference,
1185 },
1186 Expression {
1187 elements: Vec<NodeReference>,
1188 },
1189 Literal {
1190 value: String,
1191 },
1192 FunctionCall {
1193 function: NodeReference,
1194 parameters: Vec<NodeReference>,
1195 },
1196 IntrinsicCall {
1197 intrinsic: NodeReference,
1198 arguments: Vec<NodeReference>,
1199 elements: Vec<NodeReference>,
1200 },
1201 Operator {
1202 operator: Operators,
1203 left: NodeReference,
1204 right: NodeReference,
1205 },
1206 VariableDeclaration {
1207 name: String,
1208 r#type: NodeReference,
1209 },
1210 Accessor {
1211 left: NodeReference,
1212 right: NodeReference,
1213 },
1214 Macro {
1215 name: String,
1216 body: NodeReference,
1217 },
1218}
1219
1220#[derive(Debug, PartialEq, Eq)]
1221pub enum LexError {
1222 Undefined { message: Option<String> },
1223 FunctionCallParametersDoNotMatchFunctionParameters,
1224 AccessingUndeclaredMember { name: String },
1225 ReferenceToUndefinedType { type_name: String },
1226}
1227
1228#[derive(Clone, Copy, PartialEq, Eq)]
1229enum DescendantSearch {
1230 Any,
1231 NonIntrinsic,
1232}
1233
1234fn get_reference(chain: &[NodeReference], name: &str) -> Option<NodeReference> {
1236 for node in chain.iter().rev() {
1237 let reference = match node.borrow().node() {
1238 Nodes::Intrinsic { .. } => find_descendant(node, name, DescendantSearch::Any),
1239 _ => find_descendant(node, name, DescendantSearch::NonIntrinsic),
1240 };
1241
1242 if let Some(c) = reference {
1243 return Some(c);
1244 }
1245 }
1246
1247 None
1248}
1249
1250fn resolve_type(chain: &[NodeReference], type_name: &str) -> Result<NodeReference, LexError> {
1251 if let Some(existing) = get_reference(chain, type_name) {
1252 return Ok(existing);
1253 }
1254
1255 if type_name.contains('[') {
1256 let mut parts = type_name.split(['[', ']']);
1257 let element_type_name = parts.next().ok_or(LexError::Undefined {
1258 message: Some("No type name".to_string()),
1259 })?;
1260 let count = parts
1261 .next()
1262 .ok_or(LexError::Undefined {
1263 message: Some("No count".to_string()),
1264 })?
1265 .parse::<usize>()
1266 .map_err(|_| LexError::Undefined {
1267 message: Some("Invalid count".to_string()),
1268 })?;
1269
1270 let element_type = parser::TypeName::Named(element_type_name);
1271 return resolve_array_type(chain, &element_type, count);
1272 }
1273
1274 get_reference(chain, type_name).ok_or(LexError::ReferenceToUndefinedType {
1275 type_name: type_name.to_string(),
1276 })
1277}
1278
1279fn resolve_descriptor_type(
1281 chain: &[NodeReference],
1282 resource_type: &str,
1283 format: Option<&str>,
1284) -> Result<BindingTypes, LexError> {
1285 if format.is_some() && resource_type != "StorageImage" {
1286 return Err(LexError::Undefined {
1287 message: Some(format!(
1288 "Resource type {resource_type} cannot declare a storage image format. The most likely cause is that a format was attached to a non-StorageImage descriptor."
1289 )),
1290 });
1291 }
1292
1293 match resource_type {
1294 "Texture2D" => Ok(BindingTypes::CombinedImageSampler { format: String::new() }),
1295 "Texture2DArray" => Ok(BindingTypes::CombinedImageSampler {
1296 format: "ArrayTexture2D".to_string(),
1297 }),
1298 "Texture3D" => Ok(BindingTypes::CombinedImageSampler {
1299 format: "Texture3D".to_string(),
1300 }),
1301 "StorageImage" => Ok(BindingTypes::Image {
1302 format: format.unwrap_or("unknown").to_string(),
1303 }),
1304 struct_name => {
1305 let r#struct = resolve_type(chain, struct_name)?;
1306 let members = match r#struct.borrow().node() {
1307 Nodes::Struct { fields, .. } => fields.clone(),
1308 _ => {
1309 return Err(LexError::ReferenceToUndefinedType {
1310 type_name: struct_name.to_string(),
1311 });
1312 }
1313 };
1314 Ok(BindingTypes::Buffer { members })
1315 }
1316 }
1317}
1318
1319fn resolve_array_type(
1321 chain: &[NodeReference],
1322 element_type_name: &parser::TypeName,
1323 count: usize,
1324) -> Result<NodeReference, LexError> {
1325 let mut array_name = String::new();
1326 append_type_name(&mut array_name, element_type_name);
1327 let _ = write!(array_name, "[{count}]");
1328 if let Some(existing) = get_reference(chain, &array_name) {
1329 return Ok(existing);
1330 }
1331
1332 let element_type = resolve_type_name(chain, element_type_name)?;
1333 let array_type = Node::internal_new(Node {
1334 node: Nodes::Struct {
1335 name: array_name,
1336 template: Some(element_type.clone()),
1337 fields: (0..count)
1338 .map(|index| Node::member(&format!("value_{index}"), element_type.clone()).into())
1339 .collect(),
1340 types: Vec::new(),
1341 },
1342 });
1343
1344 Ok(array_type)
1345}
1346
1347fn append_type_name(name: &mut String, type_name: &parser::TypeName) {
1349 match type_name {
1350 parser::TypeName::Named(type_name) => name.push_str(type_name),
1351 parser::TypeName::Array { element, count } => {
1352 append_type_name(name, element);
1353 let _ = write!(name, "[{count}]");
1354 }
1355 }
1356}
1357
1358fn resolve_type_name(chain: &[NodeReference], type_name: &parser::TypeName) -> Result<NodeReference, LexError> {
1360 match type_name {
1361 parser::TypeName::Named(type_name) => resolve_type(chain, type_name),
1362 parser::TypeName::Array { element, count } => {
1363 let count = usize::try_from(*count).map_err(|_| LexError::Undefined {
1364 message: Some("Invalid count".to_string()),
1365 })?;
1366 resolve_array_type(chain, element, count)
1367 }
1368 }
1369}
1370fn resolve_member(chain: &[NodeReference], name: &str) -> Result<NodeReference, LexError> {
1371 if let Some(left) = chain.last() {
1374 let source = match left.borrow().node() {
1375 Nodes::Expression(Expressions::Member { source, .. }) => Some(source.clone()),
1376 _ => None,
1377 };
1378 if let Some(source) = source {
1379 if let Nodes::Binding {
1380 r#type: BindingTypes::Buffer { members },
1381 ..
1382 } = source.borrow().node()
1383 {
1384 if let Some(member) = find_named_child(members, name) {
1385 return Ok(member);
1386 }
1387 }
1388 }
1389 }
1390 get_reference(chain, name).ok_or(LexError::AccessingUndeclaredMember { name: name.to_string() })
1391}
1392
1393fn extend_chain(chain: &[NodeReference], parent: &NodeReference) -> Vec<NodeReference> {
1395 let mut extended = chain.to_vec();
1396 extended.push(parent.clone());
1397 extended
1398}
1399
1400fn lex_child_with_parent(
1402 chain: &[NodeReference],
1403 parent: &NodeReference,
1404 parser_node: &parser::Node,
1405) -> Result<NodeReference, LexError> {
1406 lex_parsed_node(extend_chain(chain, parent), parser_node)
1407}
1408
1409fn lex_raw_code(
1411 chain: &[NodeReference],
1412 glsl: Option<&str>,
1413 hlsl: Option<&str>,
1414 msl: Option<&str>,
1415 input: &[&str],
1416 output: &[&str],
1417) -> Result<Node, LexError> {
1418 let inputs = input
1419 .iter()
1420 .map(|name| resolve_member(chain, name))
1421 .collect::<Result<Vec<_>, _>>()?;
1422
1423 let vec3f = resolve_member(chain, "vec3f")?;
1424 let outputs = output
1425 .iter()
1426 .map(|name| {
1427 Node::expression(Expressions::VariableDeclaration {
1428 name: (*name).to_string(),
1429 r#type: vec3f.clone(),
1430 })
1431 .into()
1432 })
1433 .collect();
1434
1435 Ok(Node::raw(
1436 glsl.map(str::to_string),
1437 hlsl.map(str::to_string),
1438 msl.map(str::to_string),
1439 inputs,
1440 outputs,
1441 ))
1442}
1443
1444fn find_descendant(node: &NodeReference, child_name: &str, mode: DescendantSearch) -> Option<NodeReference> {
1445 let prefer_descendants_before_self = mode == DescendantSearch::NonIntrinsic
1446 && matches!(
1447 node.borrow().node(),
1448 Nodes::PushConstant { .. }
1449 | Nodes::Member { .. }
1450 | Nodes::Parameter { .. }
1451 | Nodes::Input { .. }
1452 | Nodes::Output { .. }
1453 | Nodes::TaskPayload { .. }
1454 | Nodes::Workgroup { .. }
1455 | Nodes::Expression(Expressions::Member { .. })
1456 );
1457
1458 if !prefer_descendants_before_self && node.borrow().get_name() == Some(child_name) {
1459 return Some(node.clone());
1460 }
1461
1462 let result = match node.borrow().node() {
1463 Nodes::Scope { children, .. } | Nodes::Struct { fields: children, .. } | Nodes::PushConstant { members: children } => {
1464 find_in_children(children, child_name, mode == DescendantSearch::NonIntrinsic, mode)
1465 }
1466 Nodes::Intrinsic { elements, .. } => {
1467 if mode == DescendantSearch::Any {
1468 find_in_children(elements, child_name, false, mode)
1469 } else {
1470 None
1471 }
1472 }
1473 Nodes::Member { r#type, .. } | Nodes::Parameter { r#type, .. } => find_descendant(r#type, child_name, mode),
1474 Nodes::Function { params, statements, .. } => find_in_function(params, statements, child_name, mode),
1475 Nodes::Conditional { condition, statements } if mode == DescendantSearch::NonIntrinsic => {
1476 find_descendant(condition, child_name, mode).or_else(|| find_in_descendants(statements, child_name, mode))
1477 }
1478 Nodes::ForLoop {
1479 initializer,
1480 condition,
1481 update,
1482 statements,
1483 } if mode == DescendantSearch::NonIntrinsic => find_descendant(initializer, child_name, mode)
1484 .or_else(|| find_descendant(condition, child_name, mode))
1485 .or_else(|| find_descendant(update, child_name, mode))
1486 .or_else(|| find_in_descendants(statements, child_name, mode)),
1487 Nodes::Expression(expression) => find_in_expression(expression, child_name, mode),
1488 Nodes::Raw { output, .. } => find_in_descendants(output, child_name, mode),
1489 Nodes::Binding {
1490 r#type: BindingTypes::Buffer { members },
1491 ..
1492 } => find_in_descendants(members, child_name, mode),
1493 Nodes::Input { format, .. }
1494 | Nodes::Output { format, .. }
1495 | Nodes::TaskPayload { format, .. }
1496 | Nodes::Workgroup { format, .. } => find_descendant(format, child_name, mode),
1497 _ => None,
1498 };
1499
1500 result.or_else(|| {
1501 if prefer_descendants_before_self && node.borrow().get_name() == Some(child_name) {
1502 Some(node.clone())
1503 } else {
1504 None
1505 }
1506 })
1507}
1508
1509fn find_in_children(
1510 children: &[NodeReference],
1511 child_name: &str,
1512 prefer_direct_children: bool,
1513 mode: DescendantSearch,
1514) -> Option<NodeReference> {
1515 if prefer_direct_children {
1516 find_named_child(children, child_name).or_else(|| find_in_descendants(children, child_name, mode))
1517 } else {
1518 find_in_descendants(children, child_name, mode)
1519 }
1520}
1521
1522fn find_named_child(children: &[NodeReference], child_name: &str) -> Option<NodeReference> {
1523 children
1524 .iter()
1525 .find(|child| child.borrow().get_name() == Some(child_name))
1526 .cloned()
1527}
1528
1529fn find_in_descendants(children: &[NodeReference], child_name: &str, mode: DescendantSearch) -> Option<NodeReference> {
1530 children.iter().find_map(|child| find_descendant(child, child_name, mode))
1531}
1532
1533fn find_in_function(
1534 params: &[NodeReference],
1535 statements: &[NodeReference],
1536 child_name: &str,
1537 mode: DescendantSearch,
1538) -> Option<NodeReference> {
1539 find_named_child(params, child_name).or_else(|| {
1540 statements
1541 .iter()
1542 .find_map(|statement| find_in_function_statement(statement, child_name, mode))
1543 })
1544}
1545
1546fn find_in_function_statement(statement: &NodeReference, child_name: &str, mode: DescendantSearch) -> Option<NodeReference> {
1547 match statement.borrow().node() {
1548 Nodes::Expression(expression) => find_in_function_expression(statement, expression, child_name, mode),
1549 Nodes::Raw { output, .. } if mode == DescendantSearch::Any => find_in_descendants(output, child_name, mode),
1550 _ => None,
1551 }
1552}
1553
1554fn find_in_function_expression(
1555 statement: &NodeReference,
1556 expression: &Expressions,
1557 child_name: &str,
1558 mode: DescendantSearch,
1559) -> Option<NodeReference> {
1560 match mode {
1561 DescendantSearch::Any => match expression {
1562 Expressions::Operator { left, right, .. } => {
1563 find_descendant(left, child_name, mode).or_else(|| find_descendant(right, child_name, mode))
1564 }
1565 Expressions::VariableDeclaration { name, .. } if child_name == name => Some(statement.clone()),
1566 Expressions::Accessor { left, right } => {
1567 find_descendant(left, child_name, mode).or_else(|| find_descendant(right, child_name, mode))
1568 }
1569 Expressions::Return { value } => value.as_ref().and_then(|value| find_descendant(value, child_name, mode)),
1570 _ => None,
1571 },
1572 DescendantSearch::NonIntrinsic => match expression {
1573 Expressions::VariableDeclaration { name, .. } if child_name == name => Some(statement.clone()),
1574 Expressions::Operator { left, .. } => find_descendant(left, child_name, mode),
1575 _ => None,
1576 },
1577 }
1578}
1579
1580fn find_in_expression(expression: &Expressions, child_name: &str, mode: DescendantSearch) -> Option<NodeReference> {
1581 match expression {
1582 Expressions::Operator { left, .. } if mode == DescendantSearch::NonIntrinsic => find_descendant(left, child_name, mode),
1584 Expressions::Operator { left, right, .. } => {
1585 find_descendant(left, child_name, mode).or_else(|| find_descendant(right, child_name, mode))
1586 }
1587 Expressions::Member { source, .. } => find_descendant(source, child_name, mode),
1588 Expressions::Expression { elements } => find_in_descendants(elements, child_name, mode),
1589 Expressions::VariableDeclaration { r#type, .. } => find_descendant(r#type, child_name, mode),
1590 Expressions::Accessor { left, right } => {
1591 find_descendant(right, child_name, mode).or_else(|| find_descendant(left, child_name, mode))
1592 }
1593 Expressions::IntrinsicCall { intrinsic, .. } => {
1594 let intrinsic = intrinsic.borrow();
1595 if let Nodes::Intrinsic { r#return, .. } = intrinsic.node() {
1596 find_descendant(r#return, child_name, mode)
1597 } else {
1598 None
1599 }
1600 }
1601 Expressions::Return { value } => value.as_ref().and_then(|value| find_descendant(value, child_name, mode)),
1602 _ => None,
1603 }
1604}
1605
1606fn lex_parsed_node(chain: Vec<NodeReference>, parser_node: &parser::Node) -> Result<NodeReference, LexError> {
1607 let node = match &parser_node.node {
1608 parser::Nodes::Null => Node::new(Nodes::Null).into(),
1609 parser::Nodes::Scope { name, children } => {
1610 assert_ne!(*name, "root"); let this: NodeReference = Node::scope(name.to_string()).into();
1613 for child in children {
1614 let child = lex_child_with_parent(&chain, &this, child)?;
1615 this.borrow_mut().add_child(child);
1616 }
1617
1618 this
1619 }
1620 parser::Nodes::Struct { name, fields } => {
1621 if let Some(n) = get_reference(&chain, name) {
1622 return Ok(n.clone());
1624 }
1625
1626 let this: NodeReference = Node::r#struct(name, Vec::new()).into();
1627 for field in fields {
1628 let field = lex_child_with_parent(&chain, &this, field)?;
1629 this.borrow_mut().add_child(field);
1630 }
1631
1632 this
1633 }
1634 parser::Nodes::Specialization { name, r#type } => {
1635 let t = resolve_type(&chain, r#type)?;
1636
1637 let this = Node::new(Nodes::Specialization {
1638 name: name.to_string(),
1639 r#type: t,
1640 });
1641
1642 this.into()
1643 }
1644 parser::Nodes::Member { name, r#type } => {
1645 let t = if r#type.contains('<') {
1646 let mut s = r#type.split(['<', '>']);
1647
1648 let outer_type_name = s.next().ok_or(LexError::Undefined {
1649 message: Some("No outer name".to_string()),
1650 })?;
1651
1652 let outer_type = resolve_type(&chain, outer_type_name)?;
1653
1654 let inner_type_name = s.next().ok_or(LexError::Undefined {
1655 message: Some("No inner name".to_string()),
1656 })?;
1657
1658 let inner_type = if let Some(stripped) = inner_type_name.strip_suffix('*') {
1659 let x = Node::internal_new(Node {
1660 node: Nodes::Struct {
1661 name: format!("{}*", stripped),
1662 template: Some(outer_type.clone()),
1663 fields: Vec::new(),
1664 types: Vec::new(),
1665 },
1666 });
1667
1668 x
1669 } else {
1670 resolve_type(&chain, inner_type_name)?
1671 };
1672
1673 if let Some(n) = get_reference(&chain, r#type) {
1674 return Ok(n.clone());
1676 }
1677
1678 let children = Vec::new();
1679
1680 let this = Node {
1681 node: Nodes::Struct {
1682 name: r#type.to_string(),
1683 template: Some(outer_type.clone()),
1684 fields: children,
1685 types: vec![inner_type],
1686 },
1687 };
1688
1689 let this: NodeReference = this.into();
1690
1691 return Ok(this);
1692 } else if r#type.contains('[') {
1693 let mut s = r#type.split(['[', ']']);
1694
1695 let type_name = s.next().ok_or(LexError::Undefined {
1696 message: Some("No type name".to_string()),
1697 })?;
1698
1699 let member_type = resolve_type(&chain, type_name)?;
1700
1701 let count = s
1702 .next()
1703 .ok_or(LexError::Undefined {
1704 message: Some("No count".to_string()),
1705 })?
1706 .parse()
1707 .map_err(|_| LexError::Undefined {
1708 message: Some("Invalid count".to_string()),
1709 })?;
1710
1711 return Ok(Node::array(name, member_type, count));
1712 } else {
1713 resolve_type(&chain, r#type)?
1714 };
1715
1716 let this: NodeReference = Node::member(name, t).into();
1717
1718 this
1719 }
1720 parser::Nodes::Parameter { name, r#type } => {
1721 let t = resolve_type(&chain, r#type)?;
1722
1723 let this = Node::new(Nodes::Parameter {
1724 name: name.to_string(),
1725 r#type: t,
1726 });
1727
1728 this.into()
1729 }
1730 parser::Nodes::Input { name, format, location } => {
1731 let t = resolve_type(&chain, format)?;
1732
1733 let this = Node::new(Nodes::Input {
1734 name: name.to_string(),
1735 format: t,
1736 location: *location,
1737 });
1738
1739 this.into()
1740 }
1741 parser::Nodes::Output {
1742 name,
1743 format,
1744 location,
1745 count,
1746 } => {
1747 let t = resolve_type(&chain, format)?;
1748
1749 let this = Node::new(Nodes::Output {
1750 name: name.to_string(),
1751 format: t,
1752 location: *location,
1753 count: *count,
1754 });
1755
1756 this.into()
1757 }
1758 parser::Nodes::TaskPayload { name, format, count } => {
1759 let format = resolve_type(&chain, format)?;
1760 Node::new(Nodes::TaskPayload {
1761 name: name.to_string(),
1762 format,
1763 count: *count,
1764 })
1765 .into()
1766 }
1767 parser::Nodes::Workgroup { name, format } => {
1768 let format = resolve_type(&chain, format)?;
1769 Node::new(Nodes::Workgroup {
1770 name: name.to_string(),
1771 format,
1772 })
1773 .into()
1774 }
1775 parser::Nodes::Function {
1776 name,
1777 return_type,
1778 statements,
1779 params,
1780 ..
1781 } => {
1782 let t = resolve_type(&chain, return_type)?;
1783
1784 let this: NodeReference = Node::function(name, Vec::new(), t, Vec::new()).into();
1785
1786 for param in params {
1787 let param = lex_child_with_parent(&chain, &this, param)?;
1788 match this.borrow_mut().node_mut() {
1789 Nodes::Function { params, .. } => {
1790 params.push(param);
1791 }
1792 _ => {
1793 panic!("Expected function");
1794 }
1795 }
1796 }
1797
1798 let mut scoped_chain = extend_chain(&chain, &this);
1799
1800 for statement in statements {
1801 let statement = lex_parsed_node(scoped_chain.clone(), statement)?;
1802 this.borrow_mut().add_child(statement);
1803 scoped_chain.push(
1804 this.borrow()
1805 .get_children()
1806 .and_then(|children| children.last().cloned())
1807 .unwrap(),
1808 );
1809 }
1810
1811 this
1812 }
1813 parser::Nodes::Conditional { condition, statements } => {
1814 let condition = lex_parsed_node(chain.clone(), condition)?;
1815 let mut lexed_statements = Vec::with_capacity(statements.len());
1816 let mut scoped_chain = chain.clone();
1817
1818 for statement in statements {
1819 let statement = lex_parsed_node(scoped_chain.clone(), statement)?;
1820 scoped_chain.push(statement.clone());
1821 lexed_statements.push(statement);
1822 }
1823
1824 Node::conditional(condition, lexed_statements).into()
1825 }
1826 parser::Nodes::ForLoop {
1827 initializer,
1828 condition,
1829 update,
1830 statements,
1831 } => {
1832 let initializer = lex_parsed_node(chain.clone(), initializer)?;
1833 let mut scoped_chain = chain.clone();
1834 scoped_chain.push(initializer.clone());
1835 let condition = lex_parsed_node(scoped_chain.clone(), condition)?;
1836 let update = lex_parsed_node(scoped_chain.clone(), update)?;
1837 let mut lexed_statements = Vec::with_capacity(statements.len());
1838
1839 for statement in statements {
1840 let statement = lex_parsed_node(scoped_chain.clone(), statement)?;
1841 scoped_chain.push(statement.clone());
1842 lexed_statements.push(statement);
1843 }
1844
1845 Node::for_loop(initializer, condition, update, lexed_statements).into()
1846 }
1847 parser::Nodes::PushConstant { members } => {
1848 let this: NodeReference = Node::push_constant(vec![]).into();
1849
1850 for member in members
1851 .iter()
1852 .filter(|member| matches!(member.node, parser::Nodes::Member { .. }))
1853 {
1854 let c = lex_child_with_parent(&chain, &this, member)?;
1855 this.borrow_mut().add_child(c);
1856 }
1857
1858 this
1859 }
1860 parser::Nodes::Binding {
1861 name,
1862 r#type,
1863 slot,
1864 read,
1865 write,
1866 count,
1867 } => {
1868 let r#type = match &r#type.node {
1869 parser::Nodes::Type { members, .. } => BindingTypes::Buffer {
1870 members: members
1871 .iter()
1872 .map(|m| lex_parsed_node(chain.clone(), m))
1873 .collect::<Result<Vec<NodeReference>, LexError>>()?,
1874 },
1875 parser::Nodes::Image { format } => BindingTypes::Image {
1876 format: format.to_string(),
1877 },
1878 parser::Nodes::CombinedImageSampler { format } => BindingTypes::CombinedImageSampler {
1879 format: format.to_string(),
1880 },
1881 _ => {
1882 return Err(LexError::Undefined {
1883 message: Some("Invalid binding type".to_string()),
1884 });
1885 }
1886 };
1887
1888 let this = if let Some(count) = count {
1889 Node::binding_array(name, r#type, *slot, *read, *write, count.get())
1890 } else {
1891 Node::binding(name, r#type, *slot, *read, *write)
1892 };
1893
1894 this.into()
1895 }
1896 parser::Nodes::Descriptor {
1897 name,
1898 resource_type,
1899 format,
1900 slot,
1901 read,
1902 write,
1903 count,
1904 } => Node::binding_with_count(
1905 name,
1906 resolve_descriptor_type(&chain, resource_type, *format)?,
1907 *slot,
1908 *read,
1909 *write,
1910 *count,
1911 )
1912 .into(),
1913 parser::Nodes::Type { name, members } => {
1914 let mut this = Node::r#struct(name, Vec::new());
1915
1916 for member in members {
1917 let c = lex_parsed_node(chain.clone(), member)?;
1918 this.add_child(c);
1919 }
1920
1921 this.into()
1922 }
1923 parser::Nodes::Image { format } => {
1924 let this = Node::binding(
1925 "image",
1926 BindingTypes::Image {
1927 format: format.to_string(),
1928 },
1929 0,
1930 false,
1931 false,
1932 );
1933
1934 this.into()
1935 }
1936 parser::Nodes::CombinedImageSampler { format } => {
1937 let this = Node::binding(
1938 "combined_image_sampler",
1939 BindingTypes::CombinedImageSampler {
1940 format: format.to_string(),
1941 },
1942 0,
1943 false,
1944 false,
1945 );
1946
1947 this.into()
1948 }
1949 parser::Nodes::RawCode {
1950 glsl,
1951 hlsl,
1952 msl,
1953 input,
1954 output,
1955 ..
1956 } => lex_raw_code(&chain, glsl.as_deref(), hlsl.as_deref(), msl.as_deref(), input, output)?.into(),
1957 parser::Nodes::Literal { name, body } => Node::new(Nodes::Literal {
1958 name: name.to_string(),
1959 value: lex_parsed_node(chain, body)?,
1960 })
1961 .into(),
1962 parser::Nodes::Expression(expression) => {
1963 let this = match expression {
1964 parser::Expressions::Return { value } => Node::expression(Expressions::Return {
1965 value: match value {
1966 Some(value) => Some(lex_parsed_node(chain.clone(), value)?),
1967 None => None,
1968 },
1969 }),
1970 parser::Expressions::Continue => Node::expression(Expressions::Continue),
1971 parser::Expressions::Accessor { left, right } => {
1972 let left = lex_parsed_node(chain.clone(), left)?;
1973
1974 let right = {
1975 let left = left.clone();
1976
1977 let mut chain = chain.clone();
1978 chain.push(left); lex_parsed_node(chain.clone(), right)?
1981 };
1982
1983 Node::expression(Expressions::Accessor { left, right })
1984 }
1985 parser::Expressions::Member { name } => Node::expression(Expressions::Member {
1986 source: resolve_member(&chain, name)?,
1987 name: name.to_string(),
1988 }),
1989 parser::Expressions::Literal { value } => Node::expression(Expressions::Literal {
1990 value: value.to_string(),
1991 }),
1992 parser::Expressions::Expression(elements) => Node::sentence(
1993 elements
1994 .iter()
1995 .map(|e| lex_parsed_node(chain.clone(), e))
1996 .collect::<Result<Vec<NodeReference>, LexError>>()?,
1997 ),
1998 parser::Expressions::Call { name, parameters } => {
1999 let parameters = parameters
2000 .iter()
2001 .map(|e| lex_parsed_node(chain.clone(), e))
2002 .collect::<Result<Vec<NodeReference>, LexError>>()?;
2003 let function = resolve_call_target(&chain, name, ¶meters)?;
2004 let r = function.clone(); {
2007 let b = RefCell::borrow(&function.0);
2009 match b.node() {
2010 Nodes::Function { params, .. } | Nodes::Struct { fields: params, .. } => {
2011 if params.len() != parameters.len() {
2012 return Err(LexError::FunctionCallParametersDoNotMatchFunctionParameters);
2013 }
2014 Node::expression(Expressions::FunctionCall { function: r, parameters })
2015 }
2016 Nodes::Intrinsic { elements, .. } => Node::expression(Expressions::IntrinsicCall {
2017 intrinsic: r,
2018 arguments: parameters.clone(),
2019 elements: build_intrinsic(elements, ¶meters)?,
2020 }),
2021 _ => {
2022 return Err(LexError::Undefined {
2023 message: Some("Encountered parsing error while evaluating function call. Expected Function | Struct | Intrinsic, but found other.".to_string()),
2024 });
2025 }
2026 }
2027 }
2028 }
2029 parser::Expressions::Operator { name, left, right } => Node::expression(Expressions::Operator {
2030 operator: match *name {
2031 "+" => Operators::Plus,
2032 "-" => Operators::Minus,
2033 "*" => Operators::Multiply,
2034 "/" => Operators::Divide,
2035 "%" => Operators::Modulo,
2036 "<<" => Operators::ShiftLeft,
2037 ">>" => Operators::ShiftRight,
2038 "&" => Operators::BitwiseAnd,
2039 "|" => Operators::BitwiseOr,
2040 "=" => Operators::Assignment,
2041 "==" => Operators::Equality,
2042 "<" => Operators::LessThan,
2043 "!=" => Operators::Inequality,
2044 ">" => Operators::GreaterThan,
2045 "<=" => Operators::LessThanOrEqual,
2046 ">=" => Operators::GreaterThanOrEqual,
2047 "&&" => Operators::LogicalAnd,
2048 "||" => Operators::LogicalOr,
2049 _ => {
2050 panic!("Invalid operator")
2051 }
2052 },
2053 left: lex_parsed_node(chain.clone(), left)?,
2054 right: lex_parsed_node(chain.clone(), right)?,
2055 }),
2056 parser::Expressions::VariableDeclaration { name, r#type } => {
2057 Node::expression(Expressions::VariableDeclaration {
2058 name: name.to_string(),
2059 r#type: resolve_type_name(&chain, r#type)?,
2060 })
2061 }
2062 parser::Expressions::RawCode {
2063 glsl,
2064 hlsl,
2065 msl,
2066 input,
2067 output,
2068 } => lex_raw_code(&chain, *glsl, *hlsl, *msl, input, output)?,
2069 parser::Expressions::Macro { name, body } => Node::r#macro(name, lex_parsed_node(chain, body)?),
2070 };
2071
2072 this.into()
2073 }
2074 parser::Nodes::Intrinsic {
2075 name,
2076 elements,
2077 r#return,
2078 ..
2079 } => {
2080 let this: NodeReference = Node::intrinsic(name, Vec::new(), resolve_type(&chain, r#return)?).into();
2081
2082 for element in elements {
2083 let element = lex_child_with_parent(&chain, &this, element)?;
2084 this.borrow_mut().add_child(element);
2085 }
2086
2087 this
2088 }
2089 parser::Nodes::Const { name, r#type, value } => {
2090 let t = resolve_type_name(&chain, r#type)?;
2091
2092 let v = lex_parsed_node(chain.clone(), value)?;
2093
2094 Node::constant(name, t, v).into()
2095 }
2096 };
2097
2098 Ok(node)
2099}
2100
2101fn build_intrinsic(elements: &[NodeReference], parameters: &[NodeReference]) -> Result<Vec<NodeReference>, LexError> {
2102 let expected_parameter_count = elements
2103 .iter()
2104 .filter(|element| matches!(element.borrow().node(), Nodes::Parameter { .. }))
2105 .count();
2106
2107 if expected_parameter_count != parameters.len() {
2108 return Err(LexError::FunctionCallParametersDoNotMatchFunctionParameters);
2109 }
2110
2111 let has_body = elements
2112 .iter()
2113 .any(|element| !matches!(element.borrow().node(), Nodes::Parameter { .. }));
2114
2115 if !has_body {
2116 return Ok(parameters.to_vec());
2117 }
2118
2119 build_intrinsic_elements(elements, &mut parameters.iter())
2120}
2121
2122fn intrinsic_matches_parameters(intrinsic: &NodeReference, parameters: &[NodeReference]) -> bool {
2123 let intrinsic = intrinsic.borrow();
2124 let Nodes::Intrinsic { elements, .. } = intrinsic.node() else {
2125 return false;
2126 };
2127
2128 let expected_parameters = elements
2129 .iter()
2130 .filter_map(|element| match element.borrow().node() {
2131 Nodes::Parameter { r#type, .. } => Some(r#type.clone()),
2132 _ => None,
2133 })
2134 .collect::<Vec<_>>();
2135
2136 if expected_parameters.len() != parameters.len() {
2137 return false;
2138 }
2139
2140 expected_parameters
2141 .iter()
2142 .zip(parameters.iter())
2143 .all(|(expected, parameter)| expression_matches_type(parameter, expected))
2144}
2145
2146fn expression_matches_type(expression: &NodeReference, expected_type: &NodeReference) -> bool {
2147 infer_expression_type(expression)
2148 .map(|actual_type| actual_type.borrow().get_name() == expected_type.borrow().get_name())
2149 .unwrap_or(false)
2150}
2151
2152fn infer_expression_type(expression: &NodeReference) -> Option<NodeReference> {
2153 match expression.borrow().node() {
2154 Nodes::Expression(Expressions::Expression { elements }) if elements.len() == 1 => infer_expression_type(&elements[0]),
2155 Nodes::Expression(Expressions::Literal { value }) => infer_literal_type(value),
2156 Nodes::Expression(Expressions::VariableDeclaration { r#type, .. }) => Some(r#type.clone()),
2157 Nodes::Expression(Expressions::Member { source, .. }) => infer_member_type(source),
2158 Nodes::Expression(Expressions::Accessor { right, .. }) => infer_expression_type(right),
2159 Nodes::Expression(Expressions::FunctionCall { function, .. }) => infer_callable_return_type(function),
2160 Nodes::Expression(Expressions::IntrinsicCall { intrinsic, .. }) => infer_callable_return_type(intrinsic),
2161 Nodes::Expression(Expressions::Operator { operator, left, right }) => match operator {
2162 Operators::Assignment => infer_expression_type(left),
2163 Operators::Equality
2164 | Operators::LessThan
2165 | Operators::Inequality
2166 | Operators::GreaterThan
2167 | Operators::LessThanOrEqual
2168 | Operators::GreaterThanOrEqual
2169 | Operators::LogicalAnd
2170 | Operators::LogicalOr => None,
2171 _ => infer_expression_type(left).or_else(|| infer_expression_type(right)),
2172 },
2173 _ => None,
2174 }
2175}
2176
2177fn infer_literal_type(value: &str) -> Option<NodeReference> {
2178 let root = Node::root();
2179 if matches!(value, "true" | "false") {
2180 root.get_child("bool")
2181 } else if value.contains(['.', 'e', 'E']) {
2182 root.get_child("f32")
2183 } else {
2184 root.get_child("u32")
2185 }
2186}
2187
2188fn infer_member_type(source: &NodeReference) -> Option<NodeReference> {
2189 match source.borrow().node() {
2190 Nodes::Member { r#type, .. }
2191 | Nodes::Parameter { r#type, .. }
2192 | Nodes::Input { format: r#type, .. }
2193 | Nodes::Output { format: r#type, .. }
2194 | Nodes::TaskPayload { format: r#type, .. }
2195 | Nodes::Workgroup { format: r#type, .. }
2196 | Nodes::Specialization { r#type, .. }
2197 | Nodes::Const { r#type, .. } => Some(r#type.clone()),
2198 Nodes::Expression(Expressions::VariableDeclaration { r#type, .. }) => Some(r#type.clone()),
2199 Nodes::Expression(Expressions::Member { source, name }) => {
2200 let parent_type = infer_member_type(source)?;
2201 find_named_member_type(&parent_type, name)
2202 }
2203 Nodes::Expression(Expressions::Accessor { right, .. }) => infer_expression_type(right),
2204 _ => None,
2205 }
2206}
2207
2208fn find_named_member_type(parent_type: &NodeReference, member_name: &str) -> Option<NodeReference> {
2209 match parent_type.borrow().node() {
2210 Nodes::Struct { fields, .. } => fields.iter().find_map(|field| match field.borrow().node() {
2211 Nodes::Member { name, r#type, .. } if name == member_name => Some(r#type.clone()),
2212 _ => None,
2213 }),
2214 _ => None,
2215 }
2216}
2217
2218fn infer_callable_return_type(callable: &NodeReference) -> Option<NodeReference> {
2219 match callable.borrow().node() {
2220 Nodes::Function { return_type, .. } => Some(return_type.clone()),
2221 Nodes::Struct { .. } => Some(callable.clone()),
2222 Nodes::Intrinsic { r#return, .. } => Some(r#return.clone()),
2223 _ => None,
2224 }
2225}
2226
2227fn resolve_call_target(
2228 chain: &[NodeReference],
2229 name: &parser::TypeName,
2230 parameters: &[NodeReference],
2231) -> Result<NodeReference, LexError> {
2232 let parser::TypeName::Named(name) = name else {
2233 return resolve_type_name(chain, name);
2234 };
2235
2236 for node in chain.iter().rev() {
2237 if let Some(candidate) = resolve_call_target_in_node(node, name, parameters) {
2238 return Ok(candidate);
2239 }
2240 }
2241
2242 if let Ok(r#type) = resolve_type(chain, name) {
2243 return Ok(r#type);
2244 }
2245 Err(LexError::FunctionCallParametersDoNotMatchFunctionParameters)
2246}
2247
2248fn resolve_call_target_in_node(node: &NodeReference, name: &str, parameters: &[NodeReference]) -> Option<NodeReference> {
2249 match node.borrow().node() {
2250 Nodes::Scope { children, .. } | Nodes::Struct { fields: children, .. } | Nodes::PushConstant { members: children } => {
2251 children.iter().find_map(|child| match child.borrow().node() {
2252 Nodes::Intrinsic {
2253 name: candidate_name, ..
2254 } if candidate_name == name && intrinsic_matches_parameters(child, parameters) => Some(child.clone()),
2255 Nodes::Function {
2256 name: candidate_name,
2257 params,
2258 ..
2259 } if candidate_name == name && params.len() == parameters.len() => Some(child.clone()),
2260 Nodes::Struct {
2261 name: candidate_name,
2262 fields,
2263 ..
2264 } if candidate_name == name && fields.len() == parameters.len() => Some(child.clone()),
2265 _ => resolve_call_target_in_node(child, name, parameters),
2266 })
2267 }
2268 _ => None,
2269 }
2270}
2271
2272fn build_intrinsic_elements<'a>(
2273 elements: &[NodeReference],
2274 parameters: &mut impl Iterator<Item = &'a NodeReference>,
2275) -> Result<Vec<NodeReference>, LexError> {
2276 let mut ret = Vec::new();
2277
2278 for e in elements
2279 .iter()
2280 .filter(|e| !matches!(e.borrow().node(), Nodes::Parameter { .. }))
2281 {
2282 let f = e.borrow();
2283 let e = match f.node() {
2284 Nodes::Expression(expression) => match expression {
2285 Expressions::Member { source, .. } => match source.deref().borrow().node() {
2286 Nodes::Parameter { .. } => parameters
2287 .next()
2288 .ok_or(LexError::Undefined {
2289 message: Some("Expected parameter".to_string()),
2290 })?
2291 .clone(),
2292 _ => e.clone(),
2293 },
2294 Expressions::Expression { elements } => NodeReference::from(Node::expression(Expressions::Expression {
2295 elements: build_intrinsic_elements(elements, parameters)?,
2296 })),
2297 _ => e.clone(),
2298 },
2299 _ => e.clone(),
2300 };
2301
2302 ret.push(e);
2303 }
2304
2305 Ok(ret)
2306}
2307
2308fn builtin_intrinsic(name: &str, parameters: Vec<(&str, NodeReference)>, r#return: NodeReference) -> NodeReference {
2309 let intrinsic: NodeReference = Node::intrinsic(name, Vec::new(), r#return).into();
2310
2311 for (parameter_name, parameter_type) in parameters {
2312 intrinsic.borrow_mut().add_child(
2313 Node::new(Nodes::Parameter {
2314 name: parameter_name.to_string(),
2315 r#type: parameter_type,
2316 })
2317 .into(),
2318 );
2319 }
2320
2321 intrinsic
2322}
2323
2324fn primitive_type(name: &str) -> NodeReference {
2325 Node::r#struct(name, Vec::new()).into()
2326}
2327
2328fn record_type<const N: usize>(name: &str, fields: [(&str, NodeReference); N]) -> NodeReference {
2329 Node::r#struct(
2330 name,
2331 fields
2332 .into_iter()
2333 .map(|(field_name, field_type)| Node::member(field_name, field_type).into())
2334 .collect(),
2335 )
2336 .into()
2337}
2338
2339#[cfg(test)]
2340mod tests {
2341 use super::*;
2342 use crate::tokenizer;
2343
2344 #[cfg(target_pointer_width = "64")]
2345 #[test]
2346 #[should_panic(expected = "resource array exceeds u32::MAX elements")]
2347 fn binding_array_rejects_count_larger_than_flat_metadata() {
2348 Node::binding_array(
2349 "textures",
2350 BindingTypes::CombinedImageSampler { format: String::new() },
2351 0,
2352 true,
2353 false,
2354 (u32::MAX as usize) + 1,
2355 );
2356 }
2357
2358 #[test]
2359 fn source_descriptors_lower_to_existing_flat_binding_types() {
2360 let source = r#"
2361 Data: struct {
2362 value: u32,
2363 weight: f32,
2364 }
2365 data: descriptor<Data, 2, read_write>;
2366 texture: descriptor<Texture2D, 5, read>;
2367 texture_array: descriptor<Texture2DArray, 7, read, 16>;
2368 volume: descriptor<Texture3D, 30, read>;
2369 result: descriptor<StorageImage<rgba16f>, 31, write>;
2370 unformatted_result: descriptor<StorageImage, 32, write>;
2371 main: fn () -> void {
2372 data.value = data.value;
2373 }
2374 "#;
2375
2376 let root = crate::compile_to_besl(source, None).expect("resource descriptors should lex");
2377 let data = root.borrow().get_child("data").expect("data descriptor should exist");
2378 assert!(matches!(
2379 data.borrow().node(),
2380 Nodes::Binding {
2381 slot: 2,
2382 read: true,
2383 write: true,
2384 r#type: BindingTypes::Buffer { members },
2385 count: None,
2386 ..
2387 } if members.iter().map(|member| member.borrow().get_name().map(str::to_owned)).collect::<Vec<_>>()
2388 == vec![Some("value".to_string()), Some("weight".to_string())]
2389 ));
2390
2391 let texture = root.borrow().get_child("texture").expect("texture descriptor should exist");
2392 assert!(matches!(
2393 texture.borrow().node(),
2394 Nodes::Binding {
2395 slot: 5,
2396 read: true,
2397 write: false,
2398 r#type: BindingTypes::CombinedImageSampler { format },
2399 ..
2400 } if format.is_empty()
2401 ));
2402
2403 let texture_array = root
2404 .borrow()
2405 .get_child("texture_array")
2406 .expect("texture array descriptor should exist");
2407 assert!(matches!(
2408 texture_array.borrow().node(),
2409 Nodes::Binding {
2410 slot: 7,
2411 r#type: BindingTypes::CombinedImageSampler { format },
2412 count: Some(count),
2413 ..
2414 } if format == "ArrayTexture2D" && count.get() == 16
2415 ));
2416
2417 let volume = root.borrow().get_child("volume").expect("volume descriptor should exist");
2418 assert!(matches!(
2419 volume.borrow().node(),
2420 Nodes::Binding {
2421 r#type: BindingTypes::CombinedImageSampler { format },
2422 ..
2423 } if format == "Texture3D"
2424 ));
2425
2426 let result = root
2427 .borrow()
2428 .get_child("result")
2429 .expect("storage image descriptor should exist");
2430 assert!(matches!(
2431 result.borrow().node(),
2432 Nodes::Binding {
2433 slot: 31,
2434 read: false,
2435 write: true,
2436 r#type: BindingTypes::Image { format },
2437 ..
2438 } if format == "rgba16f"
2439 ));
2440
2441 let unformatted_result = root
2442 .borrow()
2443 .get_child("unformatted_result")
2444 .expect("unformatted storage image descriptor should exist");
2445 assert!(matches!(
2446 unformatted_result.borrow().node(),
2447 Nodes::Binding {
2448 slot: 32,
2449 read: false,
2450 write: true,
2451 r#type: BindingTypes::Image { format },
2452 ..
2453 } if format == "unknown"
2454 ));
2455 }
2456
2457 #[test]
2458 fn source_atomic_buffers_and_push_constants_link_without_injected_rust_nodes() {
2459 let source = r#"
2460 Counters: struct {
2461 values: atomicu32[8],
2462 }
2463 counters: descriptor<Counters, 3, read_write>;
2464 push_constant: push_constant {
2465 index: u32,
2466 }
2467 main: fn () -> void {
2468 let old: u32 = atomic_add(counters.values[push_constant.index], 1);
2469 atomic_store(counters.values[push_constant.index], atomic_load(counters.values[old]));
2470 }
2471 "#;
2472
2473 let root = crate::compile_to_besl(source, None).expect("standalone atomic shader should link");
2474 root.get_main().expect("standalone atomic shader should have main");
2475 assert!(root.borrow().get_child("push_constant").is_some());
2476 }
2477
2478 #[test]
2479 fn source_task_storage_and_stage_interfaces_link_without_injected_rust_nodes() {
2480 let source = r#"
2481 instance_index: input<u32, 0>;
2482 primitive_index: output<u32, 1>;
2483 visible_meshlets: task_payload<u32, 32>;
2484 visible_count: workgroup<atomicu32>;
2485 main: fn () -> void {
2486 let position: u32 = thread_position();
2487 visible_meshlets[thread_idx()] = position;
2488 atomic_store(visible_count, position);
2489 workgroup_barrier();
2490 set_task_mesh_output_count(atomic_load(visible_count));
2491 primitive_index = instance_index;
2492 }
2493 "#;
2494
2495 let root = crate::compile_to_besl(source, None).expect("standalone task shader should link");
2496 let payload = root
2497 .borrow()
2498 .get_child("visible_meshlets")
2499 .expect("task payload declaration should be linked");
2500 assert!(matches!(
2501 payload.borrow().node(),
2502 Nodes::TaskPayload { count, format, .. }
2503 if count.get() == 32 && format.borrow().get_name() == Some("u32")
2504 ));
2505 assert!(payload.borrow().node().is_indexable());
2506
2507 let workgroup = root
2508 .borrow()
2509 .get_child("visible_count")
2510 .expect("workgroup declaration should be linked");
2511 assert!(matches!(
2512 workgroup.borrow().node(),
2513 Nodes::Workgroup { format, .. } if format.borrow().get_name() == Some("atomicu32")
2514 ));
2515 assert!(root.get_main().is_some());
2516 }
2517
2518 #[test]
2519 fn source_boolean_literals_link_as_bool_values() {
2520 let source = r#"
2521 main: fn () -> void {
2522 let enabled: bool = true;
2523 let disabled: bool = false;
2524 }
2525 "#;
2526
2527 let root = crate::compile_to_besl(source, None).expect("boolean literals should link");
2528 root.get_main().expect("boolean literal shader should have main");
2529 assert_eq!(infer_literal_type("true").unwrap().borrow().get_name(), Some("bool"));
2530 assert_eq!(infer_literal_type("false").unwrap().borrow().get_name(), Some("bool"));
2531 }
2532
2533 #[test]
2534 fn source_buffer_descriptor_requires_a_declared_type() {
2535 let tokens = tokenizer::tokenize("data: descriptor<Missing, 0, read>;").expect("descriptor should tokenize");
2536 let parsed = parser::parse(&tokens).expect("descriptor should parse");
2537 assert_eq!(
2538 lex(parsed),
2539 Err(LexError::ReferenceToUndefinedType {
2540 type_name: "Missing".to_string(),
2541 })
2542 );
2543 }
2544
2545 fn assert_type(node: &Node, type_name: &str) {
2546 match &node.node {
2547 Nodes::Struct { name, .. } => {
2548 assert_eq!(name, type_name);
2549 }
2550 _ => {
2551 panic!("Expected type");
2552 }
2553 }
2554 }
2555
2556 #[test]
2557 fn raw_code_constructors_select_only_the_requested_backend() {
2558 const EXPECTED: [(Option<&str>, Option<&str>, Option<&str>); 3] =
2559 [(Some("g"), None, None), (None, Some("h"), None), (None, None, Some("m"))];
2560
2561 let parser_nodes = [
2562 parser::Node::glsl("g", &[], &[]),
2563 parser::Node::hlsl("h", &[], &[]),
2564 parser::Node::msl("m", &[], &[]),
2565 ];
2566 let linked_nodes = [
2567 Node::glsl("g".into(), Vec::new(), Vec::new()),
2568 Node::hlsl("h".into(), Vec::new(), Vec::new()),
2569 Node::msl("m".into(), Vec::new(), Vec::new()),
2570 ];
2571
2572 for ((parser_node, linked_node), expected) in parser_nodes.into_iter().zip(linked_nodes).zip(EXPECTED) {
2573 let parser::Nodes::RawCode { glsl, hlsl, msl, .. } = parser_node.node() else {
2574 panic!("Expected parser raw-code node. The constructor returned a different node variant.");
2575 };
2576 assert_eq!((glsl.as_deref(), hlsl.as_deref(), msl.as_deref()), expected);
2577
2578 let Nodes::Raw { glsl, hlsl, msl, .. } = linked_node.node() else {
2579 panic!("Expected linked raw-code node. The constructor returned a different node variant.");
2580 };
2581 assert_eq!((glsl.as_deref(), hlsl.as_deref(), msl.as_deref()), expected);
2582 }
2583 }
2584
2585 #[test]
2586 fn lex_non_existant_function_struct_member_type() {
2587 let source = "
2588Foo: struct {
2589 bar: NonExistantType
2590}";
2591
2592 let tokens = tokenizer::tokenize(source).expect("Failed to tokenize");
2593 let node = parser::parse(&tokens).expect("Failed to parse");
2594 lex(node)
2595 .err()
2596 .filter(|e| {
2597 e == &LexError::ReferenceToUndefinedType {
2598 type_name: "NonExistantType".to_string(),
2599 }
2600 })
2601 .expect("Expected error");
2602 }
2603
2604 #[test]
2605 fn lex_non_existant_function_return_type() {
2606 let source = "
2607main: fn () -> NonExistantType {}";
2608
2609 let tokens = tokenizer::tokenize(source).expect("Failed to tokenize");
2610 let node = parser::parse(&tokens).expect("Failed to parse");
2611 lex(node)
2612 .err()
2613 .filter(|e| {
2614 e == &LexError::ReferenceToUndefinedType {
2615 type_name: "NonExistantType".to_string(),
2616 }
2617 })
2618 .expect("Expected error");
2619 }
2620
2621 #[test]
2622 fn lex_wrong_parameter_count() {
2623 let source = "
2624function: fn () -> void {}
2625main: fn () -> void {
2626 function(vec3f(1.0, 1.0, 1.0), vec3f(0.0, 0.0, 0.0));
2627}";
2628
2629 let tokens = tokenizer::tokenize(source).expect("Failed to tokenize");
2630 let node = parser::parse(&tokens).expect("Failed to parse");
2631 lex(node)
2632 .err()
2633 .filter(|e| e == &LexError::FunctionCallParametersDoNotMatchFunctionParameters)
2634 .expect("Expected error");
2635 }
2636
2637 #[test]
2638 fn lex_function() {
2639 let source = "
2640main: fn () -> void {
2641 let position: vec4f = vec4f(0.0, 0.0, 0.0, 1.0);
2642 position = position;
2643}";
2644
2645 let tokens = tokenizer::tokenize(source).expect("Failed to tokenize");
2646 let node = parser::parse(&tokens).expect("Failed to parse");
2647 let node = lex(node).expect("Failed to lex");
2648
2649 let vec4f = node.get_descendant("vec4f").expect("Expected vec4f");
2650
2651 let nb = node.borrow();
2652
2653 match &nb.node {
2654 Nodes::Scope { .. } => {
2655 let main = node.get_descendant("main").expect("Expected main");
2656 let main = RefCell::borrow(&main.0);
2657
2658 match main.node() {
2659 Nodes::Function {
2660 name,
2661 return_type,
2662 statements,
2663 ..
2664 } => {
2665 assert_eq!(name, "main");
2666 assert_type(&return_type.borrow(), "void");
2667
2668 let position = statements[0].borrow();
2669
2670 match position.node() {
2671 Nodes::Expression(Expressions::Operator { operator, left, right }) => {
2672 let position = left.borrow();
2673
2674 assert_eq!(operator, &Operators::Assignment);
2675
2676 match position.node() {
2677 Nodes::Expression(Expressions::VariableDeclaration { name, r#type }) => {
2678 assert_eq!(name, "position");
2679
2680 assert_eq!(r#type, &vec4f);
2681 }
2682 _ => {
2683 panic!("Expected expression");
2684 }
2685 }
2686
2687 let constructor = right.borrow();
2688
2689 match constructor.node() {
2690 Nodes::Expression(Expressions::FunctionCall {
2691 function, parameters, ..
2692 }) => {
2693 let function = RefCell::borrow(&function.0);
2694 let name = function.get_name().expect("Expected name");
2695
2696 assert_eq!(name, "vec4f");
2697 assert_eq!(parameters.len(), 4);
2698 }
2699 _ => {
2700 panic!("Expected expression");
2701 }
2702 }
2703 }
2704 _ => {
2705 panic!("Expected variable declaration");
2706 }
2707 }
2708 }
2709 _ => {
2710 panic!("Expected function.");
2711 }
2712 }
2713 }
2714 _ => {
2715 panic!("Expected scope");
2716 }
2717 }
2718 }
2719
2720 #[test]
2721 fn parse_script() {
2722 let script = r#"
2723 used: fn () -> void {
2724 return;
2725 }
2726
2727 not_used: fn () -> void {
2728 return;
2729 }
2730
2731 main: fn () -> void {
2732 used();
2733 }
2734 "#;
2735
2736 let tokens = tokenizer::tokenize(script).expect("Failed to tokenize");
2737 let node = parser::parse(&tokens).expect("Failed to parse");
2738 lex(node).expect("Failed to lex");
2739 }
2740
2741 #[test]
2742 fn lex_struct() {
2743 let script = r#"
2744 Vertex: struct {
2745 array: u32[3],
2746 position: vec3f,
2747 normal: vec3f,
2748 }
2749 "#;
2750
2751 let tokens = tokenizer::tokenize(script).expect("Failed to tokenize");
2752 let node = parser::parse(&tokens).expect("Failed to parse");
2753 let node = lex(node).expect("Failed to lex");
2754
2755 let nb = node.borrow();
2756
2757 match nb.node() {
2758 Nodes::Scope { name, .. } => {
2759 assert_eq!(name, "root");
2760
2761 let vertex = node.get_descendant("Vertex").expect("Expected Vertex");
2762 let vertex = RefCell::borrow(&vertex.0);
2763
2764 match vertex.node() {
2765 Nodes::Struct { name, fields, .. } => {
2766 assert_eq!(name, "Vertex");
2767 assert_eq!(fields.len(), 3);
2768
2769 let array = fields[0].borrow();
2770
2771 match array.node() {
2772 Nodes::Member { name, r#type, count } => {
2773 assert_eq!(name, "array");
2774 assert_type(&r#type.borrow(), "u32");
2775 assert_eq!(count, &Some(NonZeroUsize::new(3).expect("Invalid count")));
2776 }
2777 _ => {
2778 panic!("Expected member");
2779 }
2780 }
2781 }
2782 _ => {
2783 panic!("Expected struct");
2784 }
2785 }
2786 }
2787 _ => {
2788 panic!("Expected scope");
2789 }
2790 }
2791 }
2792
2793 #[test]
2794 fn lex_array_index_accessor() {
2795 let script = r#"
2796 main: fn () -> void {
2797 let value: f32 = buff.values[1];
2798 }
2799 "#;
2800
2801 let mut root = Node::root();
2802 let float_type = root.get_child("f32").expect("Expected f32");
2803 root.add_child(
2804 Node::binding(
2805 "buff",
2806 BindingTypes::Buffer {
2807 members: vec![Node::array("values", float_type, 3)],
2808 },
2809 0,
2810 true,
2811 false,
2812 )
2813 .into(),
2814 );
2815
2816 let node = crate::compile_to_besl(script, Some(root)).expect("Failed to lex");
2817 let main = node.get_descendant("main").expect("Expected main");
2818 let main = main.borrow();
2819
2820 let Nodes::Function { statements, .. } = main.node() else {
2821 panic!("Expected function");
2822 };
2823
2824 let statement = statements[0].borrow();
2825 let Nodes::Expression(Expressions::Operator { right, .. }) = statement.node() else {
2826 panic!("Expected assignment");
2827 };
2828 let right = right.borrow();
2829 let Nodes::Expression(Expressions::Accessor { left, right }) = right.node() else {
2830 panic!("Expected outer accessor");
2831 };
2832 assert!(matches!(
2833 right.borrow().node(),
2834 Nodes::Expression(Expressions::Expression { elements })
2835 if elements.len() == 1
2836 && matches!(elements[0].borrow().node(), Nodes::Expression(Expressions::Literal { value }) if value == "1")
2837 ));
2838 assert!(matches!(
2839 left.borrow().node(),
2840 Nodes::Expression(Expressions::Accessor { .. })
2841 ));
2842 }
2843
2844 #[test]
2845 fn lex_same_named_buffer_members_resolve_to_member_declarations() {
2846 let script = r#"
2847 main: fn () -> void {
2848 let material_index: u32 = meshes.meshes[0].material_index;
2849 let mapped: u32 = pixel_mapping.pixel_mapping[1];
2850 }
2851 "#;
2852
2853 let mut root = Node::root();
2854 let u32_type = root.get_child("u32").expect("Expected u32");
2855 let mesh = root.add_child(Node::r#struct("Mesh", vec![Node::member("material_index", u32_type.clone()).into()]).into());
2856
2857 root.add_children(vec![
2858 Node::binding(
2859 "meshes",
2860 BindingTypes::Buffer {
2861 members: vec![Node::array("meshes", mesh, 4)],
2862 },
2863 0,
2864 true,
2865 false,
2866 )
2867 .into(),
2868 Node::binding(
2869 "pixel_mapping",
2870 BindingTypes::Buffer {
2871 members: vec![Node::array("pixel_mapping", u32_type, 4)],
2872 },
2873 1,
2874 true,
2875 true,
2876 )
2877 .into(),
2878 ]);
2879
2880 let node = crate::compile_to_besl(script, Some(root)).expect("Failed to lex");
2881 let main = node.get_descendant("main").expect("Expected main");
2882 let main = main.borrow();
2883
2884 let Nodes::Function { statements, .. } = main.node() else {
2885 panic!("Expected function");
2886 };
2887
2888 let material_index_access = match statements[0].borrow().node() {
2889 Nodes::Expression(Expressions::Operator { right, .. }) => right.clone(),
2890 _ => panic!("Expected assignment"),
2891 };
2892 let (indexed_meshes, material_index_member) = match material_index_access.borrow().node() {
2893 Nodes::Expression(Expressions::Accessor { left, right }) => (left.clone(), right.clone()),
2894 _ => panic!("Expected struct member accessor"),
2895 };
2896 match material_index_member.borrow().node() {
2897 Nodes::Expression(Expressions::Member { name, source }) => {
2898 assert_eq!(name, "material_index");
2899 assert!(matches!(
2900 source.borrow().node(),
2901 Nodes::Member { name, count, .. } if name == "material_index" && count.is_none()
2902 ));
2903 }
2904 _ => panic!("Expected material_index member expression"),
2905 }
2906
2907 let meshes_member = match indexed_meshes.borrow().node() {
2908 Nodes::Expression(Expressions::Accessor { left, .. }) => match left.borrow().node() {
2909 Nodes::Expression(Expressions::Accessor { left, right }) => {
2910 assert_eq!(left.borrow().get_name(), Some("meshes"));
2911 assert!(
2912 right.borrow().node().is_indexable(),
2913 "Expected meshes.meshes to stay indexable"
2914 );
2915 right.clone()
2916 }
2917 _ => panic!("Expected meshes accessor"),
2918 },
2919 _ => panic!("Expected indexed meshes accessor"),
2920 };
2921 match meshes_member.borrow().node() {
2922 Nodes::Expression(Expressions::Member { name, source }) => {
2923 assert_eq!(name, "meshes");
2924 assert!(matches!(
2925 source.borrow().node(),
2926 Nodes::Member { name, count, .. } if name == "meshes" && count == &Some(NonZeroUsize::new(4).expect("Expected valid count"))
2927 ));
2928 }
2929 _ => panic!("Expected meshes member expression"),
2930 }
2931
2932 let pixel_mapping_access = match statements[1].borrow().node() {
2933 Nodes::Expression(Expressions::Operator { right, .. }) => right.clone(),
2934 _ => panic!("Expected assignment"),
2935 };
2936 let pixel_mapping_member = match pixel_mapping_access.borrow().node() {
2937 Nodes::Expression(Expressions::Accessor { left, .. }) => {
2938 assert!(left.borrow().node().is_indexable());
2939 match left.borrow().node() {
2940 Nodes::Expression(Expressions::Accessor { right, .. }) => right.clone(),
2941 _ => panic!("Expected pixel_mapping accessor"),
2942 }
2943 }
2944 _ => panic!("Expected indexed pixel_mapping accessor"),
2945 };
2946 match pixel_mapping_member.borrow().node() {
2947 Nodes::Expression(Expressions::Member { name, source }) => {
2948 assert_eq!(name, "pixel_mapping");
2949 assert!(matches!(
2950 source.borrow().node(),
2951 Nodes::Member { name, count, .. } if name == "pixel_mapping" && count == &Some(NonZeroUsize::new(4).expect("Expected valid count"))
2952 ));
2953 }
2954 _ => panic!("Expected pixel_mapping member expression"),
2955 };
2956 }
2957
2958 #[test]
2963 fn fragment_shader() {
2964 let source = r#"
2965 main: fn () -> void {
2966 let albedo: vec3f = vec3f(1.0, 0.0, 0.0);
2967 }
2968 "#;
2969
2970 let tokens = tokenizer::tokenize(source).expect("Failed to tokenize");
2971 let node = parser::parse(&tokens).expect("Failed to parse");
2972 let node = lex(node).expect("Failed to lex");
2973
2974 let nb = node.borrow();
2975
2976 let vec3f = node.get_descendant("vec3f").expect("Expected vec3f");
2977
2978 match nb.node() {
2979 Nodes::Scope { name, .. } => {
2980 assert_eq!(name, "root");
2981
2982 let main = node.get_descendant("main").expect("Expected main");
2983 let main = RefCell::borrow(&main.0);
2984
2985 match main.node() {
2986 Nodes::Function {
2987 name,
2988 return_type,
2989 statements,
2990 ..
2991 } => {
2992 assert_eq!(name, "main");
2993 assert_type(&return_type.borrow(), "void");
2994
2995 let albedo = statements[0].borrow();
2996
2997 match albedo.node() {
2998 Nodes::Expression(Expressions::Operator { operator, left, right }) => {
2999 let albedo = left.borrow();
3000
3001 assert_eq!(operator, &Operators::Assignment);
3002
3003 match albedo.node() {
3004 Nodes::Expression(Expressions::VariableDeclaration { name, r#type }) => {
3005 assert_eq!(name, "albedo");
3006 assert_eq!(r#type, &vec3f);
3007 }
3008 _ => {
3009 panic!("Expected expression");
3010 }
3011 }
3012
3013 let constructor = right.borrow();
3014
3015 match constructor.node() {
3016 Nodes::Expression(Expressions::FunctionCall {
3017 function, parameters, ..
3018 }) => {
3019 let function = RefCell::borrow(&function.0);
3020 let name = function.get_name().expect("Expected name");
3021
3022 assert_eq!(name, "vec3f");
3023 assert_eq!(parameters.len(), 3);
3024 }
3025 _ => {
3026 panic!("Expected expression");
3027 }
3028 }
3029 }
3030 _ => {
3031 panic!("Expected variable declaration");
3032 }
3033 }
3034 }
3035 _ => {
3036 panic!("Expected function.");
3037 }
3038 }
3039 }
3040 _ => {
3041 panic!("Expected scope");
3042 }
3043 }
3044 }
3045
3046 #[test]
3049 fn lex_intrinsic() {
3050 let source = "
3051main: fn () -> void {
3052 let n: f32 = intrinsic(0).y;
3053}";
3054
3055 let tokens = tokenizer::tokenize(source).expect("Failed to tokenize");
3056 let mut node = parser::parse(&tokens).expect("Failed to parse");
3057
3058 let intrinsic = parser::Node::intrinsic(
3059 "intrinsic",
3060 parser::Node::parameter("num", "u32"),
3061 parser::Node::sentence(vec![
3062 parser::Node::glsl("vec3(", &[], &[]),
3063 parser::Node::member_expression("num"),
3064 parser::Node::glsl(")", &[], &[]),
3065 ]),
3066 "vec3f",
3067 );
3068
3069 node.add(vec![intrinsic]);
3070
3071 let node = lex(node).expect("Failed to lex");
3072
3073 let nb = node.borrow();
3074
3075 match nb.node() {
3076 Nodes::Scope { name, .. } => {
3077 assert_eq!(name, "root");
3078
3079 let main = node.get_descendant("main").unwrap();
3080 let main = main.borrow();
3081
3082 match main.node() {
3083 Nodes::Function { name, statements, .. } => {
3084 assert_eq!(name, "main");
3085
3086 let n = statements[0].borrow();
3087
3088 match n.node() {
3089 Nodes::Expression(Expressions::Operator { operator, left, right }) => {
3090 assert_eq!(operator, &Operators::Assignment);
3091
3092 let n = left.borrow();
3093
3094 match n.node() {
3095 Nodes::Expression(Expressions::VariableDeclaration { name, r#type }) => {
3096 assert_eq!(name, "n");
3097 assert_type(&r#type.borrow(), "f32");
3098 }
3099 _ => {
3100 panic!("Expected variable declaration");
3101 }
3102 }
3103
3104 let intrinsic = right.borrow();
3105
3106 match intrinsic.node() {
3107 Nodes::Expression(Expressions::Accessor { left, right }) => {
3108 let left = left.borrow();
3109
3110 match left.node() {
3111 Nodes::Expression(Expressions::IntrinsicCall { intrinsic, .. }) => {
3112 let intrinsic = intrinsic.borrow();
3113
3114 match intrinsic.node() {
3115 Nodes::Intrinsic { name, elements, .. } => {
3116 assert_eq!(name, "intrinsic");
3117 assert_eq!(elements.len(), 2);
3118 }
3119 _ => {
3120 panic!("Expected intrinsic");
3121 }
3122 }
3123 }
3124 _ => {
3125 panic!("Expected intrinsic call");
3126 }
3127 }
3128
3129 let right = right.borrow();
3130
3131 match right.node() {
3132 Nodes::Expression(Expressions::Member { name, .. }) => {
3133 assert_eq!(name, "y");
3134 }
3135 _ => {
3136 panic!("Expected member");
3137 }
3138 }
3139 }
3140 _ => {
3141 panic!("Expected accessor");
3142 }
3143 }
3144 }
3145 _ => {
3146 panic!("Expected assignment");
3147 }
3148 }
3149 }
3150 _ => {
3151 panic!("Expected feature");
3152 }
3153 }
3154 }
3155 _ => {
3156 panic!("Expected scope");
3157 }
3158 }
3159 }
3160
3161 #[test]
3162 fn lex_builtin_texture_intrinsics() {
3163 let script = r#"
3164 main: fn () -> void {
3165 let uv: vec2f = vec2f(0.5, 0.5);
3166 let coord: vec2u = vec2u(1, 2);
3167 let color: vec4f = sample(texture_sampler, uv);
3168 let texel: vec4f = fetch(texture, coord);
3169 }
3170 "#;
3171
3172 let mut root = Node::root();
3173 root.add_child(
3174 Node::binding(
3175 "texture_sampler",
3176 BindingTypes::CombinedImageSampler { format: String::new() },
3177 0,
3178 true,
3179 false,
3180 )
3181 .into(),
3182 );
3183 root.add_child(
3184 Node::binding(
3185 "texture",
3186 BindingTypes::CombinedImageSampler { format: String::new() },
3187 1,
3188 true,
3189 false,
3190 )
3191 .into(),
3192 );
3193
3194 let node = crate::compile_to_besl(script, Some(root)).expect("Failed to lex");
3195 let main = node.get_descendant("main").expect("Expected main");
3196 let main = main.borrow();
3197
3198 let Nodes::Function { statements, .. } = main.node() else {
3199 panic!("Expected function");
3200 };
3201
3202 let sample_statement = statements[2].borrow();
3203 let fetch_statement = statements[3].borrow();
3204
3205 let assert_intrinsic_call = |statement: &Node, expected_name: &str| match statement.node() {
3206 Nodes::Expression(Expressions::Operator { right, .. }) => {
3207 let right = right.borrow();
3208 match right.node() {
3209 Nodes::Expression(Expressions::IntrinsicCall {
3210 intrinsic,
3211 arguments,
3212 elements,
3213 }) => {
3214 assert_eq!(arguments.len(), 2);
3215 assert_eq!(elements.len(), 2);
3216
3217 let intrinsic = intrinsic.borrow();
3218 match intrinsic.node() {
3219 Nodes::Intrinsic {
3220 name,
3221 r#return,
3222 elements,
3223 } => {
3224 assert_eq!(name, expected_name);
3225 assert_type(&r#return.borrow(), "vec4f");
3226 assert_eq!(elements.len(), 2);
3227 }
3228 _ => panic!("Expected intrinsic"),
3229 }
3230 }
3231 _ => panic!("Expected intrinsic call"),
3232 }
3233 }
3234 _ => panic!("Expected assignment"),
3235 };
3236
3237 assert_intrinsic_call(&sample_statement, "sample");
3238 assert_intrinsic_call(&fetch_statement, "fetch");
3239 }
3240
3241 #[test]
3242 fn lex_builtin_texture_intrinsics_validate_parameter_count() {
3243 let source = r#"
3244 main: fn () -> void {
3245 let color: vec4f = sample(texture_sampler);
3246 }
3247 "#;
3248
3249 let tokens = tokenizer::tokenize(source).expect("Failed to tokenize");
3250 let parsed = parser::parse(&tokens).expect("Failed to parse");
3251
3252 let mut root = Node::root();
3253 root.add_child(
3254 Node::binding(
3255 "texture_sampler",
3256 BindingTypes::CombinedImageSampler { format: String::new() },
3257 0,
3258 true,
3259 false,
3260 )
3261 .into(),
3262 );
3263
3264 lex_with_root(root, parsed)
3265 .err()
3266 .filter(|error| error == &LexError::FunctionCallParametersDoNotMatchFunctionParameters)
3267 .expect("Expected parameter count validation error");
3268 }
3269
3270 #[test]
3271 fn lex_builtin_image_write_intrinsic() {
3272 let script = r#"
3273 main: fn () -> void {
3274 write(image, vec2u(1, 2), vec4f(1.0, 0.0, 0.0, 1.0));
3275 }
3276 "#;
3277
3278 let mut root = Node::root();
3279 root.add_child(
3280 Node::binding(
3281 "image",
3282 BindingTypes::Image {
3283 format: "rgba8".to_string(),
3284 },
3285 0,
3286 false,
3287 true,
3288 )
3289 .into(),
3290 );
3291
3292 let node = crate::compile_to_besl(script, Some(root)).expect("Failed to lex");
3293 let main = node.get_descendant("main").expect("Expected main");
3294 let main = main.borrow();
3295
3296 let Nodes::Function { statements, .. } = main.node() else {
3297 panic!("Expected function");
3298 };
3299
3300 let write_statement = statements[0].borrow();
3301 match write_statement.node() {
3302 Nodes::Expression(Expressions::IntrinsicCall {
3303 intrinsic,
3304 arguments,
3305 elements,
3306 }) => {
3307 assert_eq!(arguments.len(), 3);
3308 assert_eq!(elements.len(), 3);
3309
3310 let intrinsic = intrinsic.borrow();
3311 match intrinsic.node() {
3312 Nodes::Intrinsic { name, r#return, .. } => {
3313 assert_eq!(name, "write");
3314 assert_type(&r#return.borrow(), "void");
3315 }
3316 _ => panic!("Expected intrinsic"),
3317 }
3318 }
3319 _ => panic!("Expected intrinsic call"),
3320 }
3321 }
3322
3323 #[test]
3324 fn lex_builtin_dot_intrinsic() {
3325 let script = r#"
3326 main: fn () -> void {
3327 let strength: f32 = dot(vec3f(1.0, 0.0, 0.0), vec3f(0.5, 0.5, 0.0));
3328 }
3329 "#;
3330
3331 let node = crate::compile_to_besl(script, None).expect("Failed to lex");
3332 let main = node.get_descendant("main").expect("Expected main");
3333 let main = main.borrow();
3334
3335 let Nodes::Function { statements, .. } = main.node() else {
3336 panic!("Expected function");
3337 };
3338
3339 let statement = statements[0].borrow();
3340 match statement.node() {
3341 Nodes::Expression(Expressions::Operator { right, .. }) => match right.borrow().node() {
3342 Nodes::Expression(Expressions::IntrinsicCall {
3343 intrinsic, arguments, ..
3344 }) => {
3345 assert_eq!(arguments.len(), 2);
3346 match intrinsic.borrow().node() {
3347 Nodes::Intrinsic { name, r#return, .. } => {
3348 assert_eq!(name, "dot");
3349 assert_type(&r#return.borrow(), "f32");
3350 }
3351 _ => panic!("Expected intrinsic"),
3352 }
3353 }
3354 _ => panic!("Expected intrinsic call"),
3355 },
3356 _ => panic!("Expected assignment"),
3357 }
3358 }
3359
3360 #[test]
3361 fn lex_builtin_cross_intrinsic() {
3362 let script = r#"
3363 main: fn () -> void {
3364 let normal: vec3f = cross(vec3f(1.0, 0.0, 0.0), vec3f(0.0, 1.0, 0.0));
3365 }
3366 "#;
3367
3368 let node = crate::compile_to_besl(script, None).expect("Failed to lex");
3369 let main = node.get_descendant("main").expect("Expected main");
3370 let main = main.borrow();
3371
3372 let Nodes::Function { statements, .. } = main.node() else {
3373 panic!("Expected function");
3374 };
3375
3376 let statement = statements[0].borrow();
3377 match statement.node() {
3378 Nodes::Expression(Expressions::Operator { right, .. }) => match right.borrow().node() {
3379 Nodes::Expression(Expressions::IntrinsicCall {
3380 intrinsic, arguments, ..
3381 }) => {
3382 assert_eq!(arguments.len(), 2);
3383 match intrinsic.borrow().node() {
3384 Nodes::Intrinsic { name, r#return, .. } => {
3385 assert_eq!(name, "cross");
3386 assert_type(&r#return.borrow(), "vec3f");
3387 }
3388 _ => panic!("Expected intrinsic"),
3389 }
3390 }
3391 _ => panic!("Expected intrinsic call"),
3392 },
3393 _ => panic!("Expected assignment"),
3394 }
3395 }
3396
3397 #[test]
3398 fn lex_builtin_length_and_normalize_intrinsics() {
3399 let script = r#"
3400 main: fn () -> void {
3401 let magnitude: f32 = length(vec3f(3.0, 4.0, 0.0));
3402 let direction: vec3f = normalize(vec3f(3.0, 4.0, 0.0));
3403 }
3404 "#;
3405
3406 let node = crate::compile_to_besl(script, None).expect("Failed to lex");
3407 let main = node.get_descendant("main").expect("Expected main");
3408 let main = main.borrow();
3409
3410 let Nodes::Function { statements, .. } = main.node() else {
3411 panic!("Expected function");
3412 };
3413
3414 let magnitude = statements[0].borrow();
3415 let direction = statements[1].borrow();
3416
3417 match magnitude.node() {
3418 Nodes::Expression(Expressions::Operator { right, .. }) => match right.borrow().node() {
3419 Nodes::Expression(Expressions::IntrinsicCall { intrinsic, .. }) => match intrinsic.borrow().node() {
3420 Nodes::Intrinsic { name, r#return, .. } => {
3421 assert_eq!(name, "length");
3422 assert_type(&r#return.borrow(), "f32");
3423 }
3424 _ => panic!("Expected intrinsic"),
3425 },
3426 _ => panic!("Expected intrinsic call"),
3427 },
3428 _ => panic!("Expected assignment"),
3429 }
3430
3431 match direction.node() {
3432 Nodes::Expression(Expressions::Operator { right, .. }) => match right.borrow().node() {
3433 Nodes::Expression(Expressions::IntrinsicCall { intrinsic, .. }) => match intrinsic.borrow().node() {
3434 Nodes::Intrinsic { name, r#return, .. } => {
3435 assert_eq!(name, "normalize");
3436 assert_type(&r#return.borrow(), "vec3f");
3437 }
3438 _ => panic!("Expected intrinsic"),
3439 },
3440 _ => panic!("Expected intrinsic call"),
3441 },
3442 _ => panic!("Expected assignment"),
3443 }
3444 }
3445
3446 #[test]
3447 fn lex_builtin_reflect_intrinsic() {
3448 let root = Node::root();
3449 let reflect = root.get_child("reflect").expect("Expected reflect builtin");
3450 match reflect.borrow().node() {
3451 Nodes::Intrinsic {
3452 name,
3453 elements,
3454 r#return,
3455 } => {
3456 assert_eq!(name, "reflect");
3457 assert_eq!(elements.len(), 2);
3458 assert_type(&r#return.borrow(), "vec4f");
3459 }
3460 _ => panic!("Expected intrinsic"),
3461 };
3462 }
3463
3464 #[test]
3465 fn lex_builtin_thread_idx_intrinsic() {
3466 let script = r#"
3467 main: fn () -> void {
3468 let index: u32 = thread_idx();
3469 }
3470 "#;
3471
3472 let node = crate::compile_to_besl(script, None).expect("Failed to lex");
3473 let main = node.get_descendant("main").expect("Expected main");
3474 let main = main.borrow();
3475
3476 let Nodes::Function { statements, .. } = main.node() else {
3477 panic!("Expected function");
3478 };
3479
3480 let statement = statements[0].borrow();
3481 match statement.node() {
3482 Nodes::Expression(Expressions::Operator { right, .. }) => match right.borrow().node() {
3483 Nodes::Expression(Expressions::IntrinsicCall {
3484 intrinsic, arguments, ..
3485 }) => {
3486 assert!(arguments.is_empty());
3487 match intrinsic.borrow().node() {
3488 Nodes::Intrinsic { name, r#return, .. } => {
3489 assert_eq!(name, "thread_idx");
3490 assert_type(&r#return.borrow(), "u32");
3491 }
3492 _ => panic!("Expected intrinsic"),
3493 }
3494 }
3495 _ => panic!("Expected intrinsic call"),
3496 },
3497 _ => panic!("Expected assignment"),
3498 }
3499 }
3500
3501 #[test]
3502 fn lex_const_variable() {
3503 let script = r#"
3504 PI: const f32 = 3.14;
3505
3506 main: fn () -> void {
3507 PI;
3508 }
3509 "#;
3510
3511 let node = crate::compile_to_besl(script, None).expect("Failed to lex");
3512
3513 let pi = node.get_descendant("PI").expect("Expected PI const");
3514 let pi = pi.borrow();
3515
3516 match pi.node() {
3517 Nodes::Const { name, r#type, value } => {
3518 assert_eq!(name, "PI");
3519 assert_eq!(r#type.borrow().get_name().unwrap(), "f32");
3520 match value.borrow().node() {
3521 Nodes::Expression(Expressions::Literal { value }) => {
3522 assert_eq!(value, "3.14");
3523 }
3524 _ => panic!("Expected a literal expression value"),
3525 }
3526 }
3527 _ => panic!("Expected Const node"),
3528 }
3529 }
3530
3531 #[test]
3532 fn lex_const_array_variable() {
3533 let script = r#"
3534 WEIGHTS: const f32[3] = f32[3](0.5, 0.25, 0.125);
3535
3536 main: fn () -> void {
3537 let value: f32 = WEIGHTS[1];
3538 }
3539 "#;
3540
3541 let node = crate::compile_to_besl(script, None).expect("Failed to lex");
3542
3543 let weights = node.get_descendant("WEIGHTS").expect("Expected WEIGHTS const");
3544 let weights = weights.borrow();
3545
3546 match weights.node() {
3547 Nodes::Const { name, r#type, value } => {
3548 assert_eq!(name, "WEIGHTS");
3549 assert_eq!(r#type.borrow().get_name().unwrap(), "f32[3]");
3550 assert!(weights.node().is_indexable());
3551 {
3552 let value = value.borrow();
3553 assert!(matches!(value.node(), Nodes::Expression(Expressions::FunctionCall { .. })));
3554 }
3555 }
3556 _ => panic!("Expected Const node"),
3557 }
3558
3559 let main = node.get_descendant("main").expect("Expected main");
3560 let statements = {
3561 let main = main.borrow();
3562 let Nodes::Function { statements, .. } = main.node() else {
3563 panic!("Expected function");
3564 };
3565 statements.clone()
3566 };
3567
3568 let statement = statements[0].clone();
3569 {
3570 let statement = statement.borrow();
3571 match statement.node() {
3572 Nodes::Expression(Expressions::Operator { right, .. }) => {
3573 let right = right.borrow();
3574 assert!(matches!(right.node(), Nodes::Expression(Expressions::Accessor { .. })));
3575 }
3576 _ => panic!("Expected assignment"),
3577 }
3578 };
3579 }
3580
3581 #[test]
3582 fn lex_array_constructor_call() {
3583 let script = r#"
3584 main: fn () -> void {
3585 let weights: f32[3] = f32[3](0.5, 0.25, 0.125);
3586 }
3587 "#;
3588
3589 let node = crate::compile_to_besl(script, None).expect("Failed to lex");
3590 let main = node.get_descendant("main").expect("Expected main");
3591 let statements = {
3592 let main = main.borrow();
3593 let Nodes::Function { statements, .. } = main.node() else {
3594 panic!("Expected function");
3595 };
3596 statements.clone()
3597 };
3598
3599 let statement = statements[0].clone();
3600 {
3601 let statement = statement.borrow();
3602 match statement.node() {
3603 Nodes::Expression(Expressions::Operator { left, right, .. }) => {
3604 match left.borrow().node() {
3605 Nodes::Expression(Expressions::VariableDeclaration { r#type, .. }) => {
3606 assert_eq!(r#type.borrow().get_name().unwrap(), "f32[3]");
3607 }
3608 _ => panic!("Expected variable declaration"),
3609 }
3610
3611 match right.borrow().node() {
3612 Nodes::Expression(Expressions::FunctionCall { function, parameters }) => {
3613 assert_eq!(parameters.len(), 3);
3614 assert_eq!(function.borrow().get_name().unwrap(), "f32[3]");
3615 }
3616 _ => panic!("Expected function call"),
3617 }
3618 }
3619 _ => panic!("Expected assignment"),
3620 }
3621 };
3622 }
3623
3624 #[test]
3625 fn lex_conditional_block() {
3626 let script = r#"
3627 main: fn () -> void {
3628 let n: u32 = 0;
3629 if (n < 1) {
3630 n = 2;
3631 }
3632 }
3633 "#;
3634
3635 let node = crate::compile_to_besl(script, None).expect("Failed to lex");
3636 let main = node.get_descendant("main").expect("Expected main");
3637 let main = main.borrow();
3638
3639 let Nodes::Function { statements, .. } = main.node() else {
3640 panic!("Expected function");
3641 };
3642
3643 let conditional = statements[1].borrow();
3644 match conditional.node() {
3645 Nodes::Conditional { condition, statements } => {
3646 assert_eq!(statements.len(), 1);
3647
3648 match condition.borrow().node() {
3649 Nodes::Expression(Expressions::Operator { operator, .. }) => {
3650 assert_eq!(operator, &Operators::LessThan);
3651 }
3652 _ => panic!("Expected less-than condition"),
3653 }
3654 }
3655 _ => panic!("Expected conditional node"),
3656 }
3657 }
3658
3659 #[test]
3660 fn lex_for_loop_block() {
3661 let script = r#"
3662 main: fn () -> void {
3663 let sum: u32 = 0;
3664 for (let i: u32 = 0; i < 4; i = i + 1) {
3665 sum = sum + i;
3666 }
3667 }
3668 "#;
3669
3670 let node = crate::compile_to_besl(script, None).expect("Failed to lex");
3671 let main = node.get_descendant("main").expect("Expected main");
3672 let main = main.borrow();
3673
3674 let Nodes::Function { statements, .. } = main.node() else {
3675 panic!("Expected function");
3676 };
3677
3678 let for_loop = statements[1].borrow();
3679 match for_loop.node() {
3680 Nodes::ForLoop {
3681 initializer,
3682 condition,
3683 update,
3684 statements,
3685 } => {
3686 assert_eq!(statements.len(), 1);
3687 assert!(matches!(
3688 initializer.borrow().node(),
3689 Nodes::Expression(Expressions::Operator { operator, .. }) if operator == &Operators::Assignment
3690 ));
3691 assert!(matches!(
3692 condition.borrow().node(),
3693 Nodes::Expression(Expressions::Operator { operator, .. }) if operator == &Operators::LessThan
3694 ));
3695 assert!(matches!(
3696 update.borrow().node(),
3697 Nodes::Expression(Expressions::Operator { operator, .. }) if operator == &Operators::Assignment
3698 ));
3699 }
3700 _ => panic!("Expected for loop node"),
3701 }
3702 }
3703
3704 #[test]
3705 fn lex_bitwise_expression() {
3706 let script = r#"
3707 main: fn () -> void {
3708 let packed: u32 = 1 << 8 | 2 & 255;
3709 }
3710 "#;
3711
3712 let node = crate::compile_to_besl(script, None).expect("Failed to lex");
3713 let main = node.get_descendant("main").expect("Expected main");
3714 let main = main.borrow();
3715
3716 let Nodes::Function { statements, .. } = main.node() else {
3717 panic!("Expected function");
3718 };
3719
3720 let statement = statements[0].borrow();
3721 match statement.node() {
3722 Nodes::Expression(Expressions::Operator { right, .. }) => match right.borrow().node() {
3723 Nodes::Expression(Expressions::Operator { operator, left, right }) => {
3724 assert_eq!(operator, &Operators::BitwiseOr);
3725 assert!(matches!(
3726 left.borrow().node(),
3727 Nodes::Expression(Expressions::Operator { operator, .. }) if operator == &Operators::ShiftLeft
3728 ));
3729 assert!(matches!(
3730 right.borrow().node(),
3731 Nodes::Expression(Expressions::Operator { operator, .. }) if operator == &Operators::BitwiseAnd
3732 ));
3733 }
3734 _ => panic!("Expected bitwise or expression"),
3735 },
3736 _ => panic!("Expected assignment"),
3737 }
3738 }
3739
3740 #[test]
3741 fn lex_comparison_and_continue() {
3742 let script = r#"
3743 main: fn () -> void {
3744 for (let i: u32 = 0; i <= 4; i = i + 1) {
3745 if (i >= 2) {
3746 continue;
3747 }
3748 }
3749 }
3750 "#;
3751
3752 let node = crate::compile_to_besl(script, None).expect("Failed to lex");
3753 let main = node.get_descendant("main").expect("Expected main");
3754 let main = main.borrow();
3755
3756 let Nodes::Function { statements, .. } = main.node() else {
3757 panic!("Expected function");
3758 };
3759
3760 let for_loop = statements[0].borrow();
3761 let Nodes::ForLoop {
3762 condition, statements, ..
3763 } = for_loop.node()
3764 else {
3765 panic!("Expected for loop");
3766 };
3767
3768 assert!(matches!(
3769 condition.borrow().node(),
3770 Nodes::Expression(Expressions::Operator { operator, .. }) if operator == &Operators::LessThanOrEqual
3771 ));
3772
3773 let conditional = statements[0].borrow();
3774 let Nodes::Conditional { condition, statements } = conditional.node() else {
3775 panic!("Expected conditional");
3776 };
3777
3778 assert!(matches!(
3779 condition.borrow().node(),
3780 Nodes::Expression(Expressions::Operator { operator, .. }) if operator == &Operators::GreaterThanOrEqual
3781 ));
3782 assert!(matches!(
3783 statements[0].borrow().node(),
3784 Nodes::Expression(Expressions::Continue)
3785 ));
3786 }
3787
3788 #[test]
3789 fn lex_scalar_intrinsic_overloads() {
3790 let script = r#"
3791 main: fn () -> void {
3792 let maximum: f32 = max(1.0, 2.0);
3793 let clamped: f32 = clamp(1.5, 0.0, 1.0);
3794 }
3795 "#;
3796
3797 let node = crate::compile_to_besl(script, None).expect("Failed to lex");
3798 let main = node.get_descendant("main").expect("Expected main");
3799 let main = main.borrow();
3800
3801 let Nodes::Function { statements, .. } = main.node() else {
3802 panic!("Expected function");
3803 };
3804
3805 for (statement, expected_name, expected_type) in [(&statements[0], "max", "f32"), (&statements[1], "clamp", "f32")] {
3806 match statement.borrow().node() {
3807 Nodes::Expression(Expressions::Operator { right, .. }) => match right.borrow().node() {
3808 Nodes::Expression(Expressions::IntrinsicCall { intrinsic, .. }) => match intrinsic.borrow().node() {
3809 Nodes::Intrinsic { name, r#return, .. } => {
3810 assert_eq!(name, expected_name);
3811 assert_type(&r#return.borrow(), expected_type);
3812 }
3813 _ => panic!("Expected intrinsic"),
3814 },
3815 _ => panic!("Expected intrinsic call"),
3816 },
3817 _ => panic!("Expected assignment"),
3818 }
3819 }
3820 }
3821
3822 #[test]
3824 fn lex_u32_widening_intrinsic_overloads() {
3825 let script = r#"
3826 main: fn () -> void {
3827 let byte: u8 = 7;
3828 let word: u16 = 513;
3829 let byte_wide: u32 = u32(byte);
3830 let word_wide: u32 = u32(word);
3831 }
3832 "#;
3833
3834 crate::compile_to_besl(script, None)
3835 .expect("Failed to resolve u32 widening calls. The most likely cause is a missing narrow-integer overload.");
3836 }
3837
3838 #[test]
3839 fn lex_vector_intrinsic_overloads_still_resolve() {
3840 let script = r#"
3841 main: fn () -> void {
3842 let maximum: vec3f = max(vec3f(1.0, 2.0, 3.0), vec3f(4.0, 5.0, 6.0));
3843 let clamped: vec3f = clamp(vec3f(1.5, 0.5, 0.0), vec3f(0.0, 0.0, 0.0), vec3f(1.0, 1.0, 1.0));
3844 }
3845 "#;
3846
3847 let node = crate::compile_to_besl(script, None).expect("Failed to lex");
3848 let main = node.get_descendant("main").expect("Expected main");
3849 let main = main.borrow();
3850
3851 let Nodes::Function { statements, .. } = main.node() else {
3852 panic!("Expected function");
3853 };
3854
3855 for (statement, expected_name, expected_type) in [(&statements[0], "max", "vec3f"), (&statements[1], "clamp", "vec3f")]
3856 {
3857 match statement.borrow().node() {
3858 Nodes::Expression(Expressions::Operator { right, .. }) => match right.borrow().node() {
3859 Nodes::Expression(Expressions::IntrinsicCall { intrinsic, .. }) => match intrinsic.borrow().node() {
3860 Nodes::Intrinsic { name, r#return, .. } => {
3861 assert_eq!(name, expected_name);
3862 assert_type(&r#return.borrow(), expected_type);
3863 }
3864 _ => panic!("Expected intrinsic"),
3865 },
3866 _ => panic!("Expected intrinsic call"),
3867 },
3868 _ => panic!("Expected assignment"),
3869 }
3870 }
3871 }
3872}