1use std::sync::Arc;
4
5use vyre_foundation::ir::model::expr::Ident;
6use vyre_foundation::ir::{BufferAccess, BufferDecl, DataType, Expr, Node, Program};
7use vyre_foundation::vast::{NODE_STRIDE_U32, SENTINEL};
8
9pub const PREORDER_OP_ID: &str = "vyre-primitives::graph::vast_walk_preorder";
11pub const POSTORDER_OP_ID: &str = "vyre-primitives::graph::vast_walk_postorder";
13
14#[derive(Clone, Copy, Debug, Eq, PartialEq)]
16pub enum VastWalkOrder {
17 Preorder,
19 Postorder,
21}
22
23impl VastWalkOrder {
24 fn op_id(self) -> &'static str {
25 match self {
26 Self::Preorder => PREORDER_OP_ID,
27 Self::Postorder => POSTORDER_OP_ID,
28 }
29 }
30}
31
32#[derive(Debug, Clone)]
34pub struct VastTreeWalkProgramPlan {
35 pub preorder: Program,
37 pub postorder: Program,
39}
40
41pub fn try_ast_walk_plan(
47 nodes: &str,
48 preorder_out: &str,
49 postorder_out: &str,
50 node_count: u32,
51 out_cap: u32,
52) -> Result<VastTreeWalkProgramPlan, String> {
53 Ok(VastTreeWalkProgramPlan {
54 preorder: try_ast_walk_preorder(nodes, preorder_out, node_count, out_cap)?,
55 postorder: try_ast_walk_postorder(nodes, postorder_out, node_count, out_cap)?,
56 })
57}
58
59#[must_use]
65pub fn ast_walk_preorder(nodes: &str, out: &str, node_count: u32, out_cap: u32) -> Program {
66 try_ast_walk_preorder(nodes, out, node_count, out_cap).unwrap_or_else(|error| panic!("{error}"))
70}
71
72pub fn try_ast_walk_preorder(
75 nodes: &str,
76 out: &str,
77 node_count: u32,
78 out_cap: u32,
79) -> Result<Program, String> {
80 try_ast_walk_order(VastWalkOrder::Preorder, nodes, out, node_count, out_cap)
81}
82
83#[must_use]
89pub fn ast_walk_postorder(nodes: &str, out: &str, node_count: u32, out_cap: u32) -> Program {
90 try_ast_walk_postorder(nodes, out, node_count, out_cap)
94 .unwrap_or_else(|error| panic!("{error}"))
95}
96
97pub fn try_ast_walk_postorder(
100 nodes: &str,
101 out: &str,
102 node_count: u32,
103 out_cap: u32,
104) -> Result<Program, String> {
105 try_ast_walk_order(VastWalkOrder::Postorder, nodes, out, node_count, out_cap)
106}
107
108pub fn try_ast_walk_order(
111 order: VastWalkOrder,
112 nodes: &str,
113 out: &str,
114 node_count: u32,
115 out_cap: u32,
116) -> Result<Program, String> {
117 let op_id = order.op_id();
118 let (stride, node_words, out_words) = checked_tree_walk_shape(node_count, out_cap, op_id)?;
119 let body = match order {
120 VastWalkOrder::Preorder => preorder_body(nodes, out, node_count, out_cap, stride),
121 VastWalkOrder::Postorder => postorder_body(nodes, out, node_count, out_cap, stride),
122 };
123
124 Ok(tree_walk_program(
125 op_id, nodes, out, node_words, out_words, body,
126 ))
127}
128
129fn preorder_body(nodes: &str, out: &str, node_count: u32, out_cap: u32, stride: u32) -> Vec<Node> {
130 let valid_node = |expr: Expr| valid_node_expr(expr, node_count);
131
132 vec![
133 Node::let_bind("oi", Expr::u32(0)),
134 Node::let_bind("n", Expr::u32(0)),
135 Node::let_bind("active", Expr::bool(true)),
136 Node::loop_for(
137 "step",
138 Expr::u32(0),
139 Expr::u32(node_count),
140 vec![Node::if_then(
141 Expr::and(
142 Expr::var("active"),
143 Expr::and(
144 Expr::lt(Expr::var("oi"), Expr::u32(out_cap)),
145 valid_node(Expr::var("n")),
146 ),
147 ),
148 vec![
149 Node::let_bind("base", Expr::mul(Expr::var("n"), Expr::u32(stride))),
150 Node::let_bind(
151 "fc",
152 Expr::load(nodes, Expr::add(Expr::var("base"), Expr::u32(2))),
153 ),
154 Node::store(out, Expr::var("oi"), Expr::var("n")),
155 Node::assign("oi", Expr::add(Expr::var("oi"), Expr::u32(1))),
156 Node::if_then(
157 valid_node(Expr::var("fc")),
158 vec![Node::assign("n", Expr::var("fc"))],
159 ),
160 Node::if_then(
161 Expr::not(valid_node(Expr::var("fc"))),
162 vec![
163 Node::let_bind("next", Expr::u32(SENTINEL)),
164 Node::let_bind("walk", Expr::var("n")),
165 Node::loop_for(
166 "climb",
167 Expr::u32(0),
168 Expr::u32(node_count),
169 vec![Node::if_then(
170 Expr::and(
171 Expr::eq(Expr::var("next"), Expr::u32(SENTINEL)),
172 valid_node(Expr::var("walk")),
173 ),
174 vec![
175 Node::let_bind(
176 "walk_base",
177 Expr::mul(Expr::var("walk"), Expr::u32(stride)),
178 ),
179 Node::let_bind(
180 "sib",
181 Expr::load(
182 nodes,
183 Expr::add(Expr::var("walk_base"), Expr::u32(3)),
184 ),
185 ),
186 Node::if_then(
187 valid_node(Expr::var("sib")),
188 vec![Node::assign("next", Expr::var("sib"))],
189 ),
190 Node::if_then(
191 Expr::not(valid_node(Expr::var("sib"))),
192 vec![
193 Node::let_bind(
194 "parent",
195 Expr::load(
196 nodes,
197 Expr::add(
198 Expr::var("walk_base"),
199 Expr::u32(1),
200 ),
201 ),
202 ),
203 Node::assign("walk", Expr::var("parent")),
204 ],
205 ),
206 ],
207 )],
208 ),
209 Node::if_then(
210 Expr::eq(Expr::var("next"), Expr::u32(SENTINEL)),
211 vec![Node::assign("active", Expr::bool(false))],
212 ),
213 Node::if_then(
214 valid_node(Expr::var("next")),
215 vec![Node::assign("n", Expr::var("next"))],
216 ),
217 ],
218 ),
219 ],
220 )],
221 ),
222 ]
223}
224
225fn postorder_body(nodes: &str, out: &str, node_count: u32, out_cap: u32, stride: u32) -> Vec<Node> {
226 let valid_node = |expr: Expr| valid_node_expr(expr, node_count);
227
228 vec![
229 Node::let_bind("oi", Expr::u32(0)),
230 Node::let_bind("n", Expr::u32(0)),
231 Node::let_bind("active", Expr::bool(true)),
232 descend_to_leftmost_leaf_node(nodes, node_count, stride),
233 Node::loop_for(
234 "emit",
235 Expr::u32(0),
236 Expr::u32(node_count),
237 vec![Node::if_then(
238 Expr::and(
239 Expr::var("active"),
240 Expr::and(
241 Expr::lt(Expr::var("oi"), Expr::u32(out_cap)),
242 valid_node(Expr::var("n")),
243 ),
244 ),
245 vec![
246 Node::store(out, Expr::var("oi"), Expr::var("n")),
247 Node::assign("oi", Expr::add(Expr::var("oi"), Expr::u32(1))),
248 Node::if_then(
249 Expr::eq(Expr::var("n"), Expr::u32(0)),
250 vec![Node::assign("active", Expr::bool(false))],
251 ),
252 Node::if_then(
253 Expr::ne(Expr::var("n"), Expr::u32(0)),
254 vec![
255 Node::let_bind("base", Expr::mul(Expr::var("n"), Expr::u32(stride))),
256 Node::let_bind(
257 "sib",
258 Expr::load(nodes, Expr::add(Expr::var("base"), Expr::u32(3))),
259 ),
260 Node::if_then(
261 valid_node(Expr::var("sib")),
262 vec![
263 Node::assign("n", Expr::var("sib")),
264 descend_to_leftmost_leaf_node(nodes, node_count, stride),
265 ],
266 ),
267 Node::if_then(
268 Expr::not(valid_node(Expr::var("sib"))),
269 vec![
270 Node::let_bind(
271 "parent",
272 Expr::load(
273 nodes,
274 Expr::add(Expr::var("base"), Expr::u32(1)),
275 ),
276 ),
277 Node::if_then(
278 valid_node(Expr::var("parent")),
279 vec![Node::assign("n", Expr::var("parent"))],
280 ),
281 Node::if_then(
282 Expr::not(valid_node(Expr::var("parent"))),
283 vec![Node::assign("active", Expr::bool(false))],
284 ),
285 ],
286 ),
287 ],
288 ),
289 ],
290 )],
291 ),
292 ]
293}
294
295fn checked_tree_walk_shape(
296 node_count: u32,
297 out_cap: u32,
298 op_id: &'static str,
299) -> Result<(u32, u32, u32), String> {
300 let stride = NODE_STRIDE_U32 as u32;
301 let node_words = checked_node_words(node_count, stride, op_id)?;
302 let out_words = checked_out_words(out_cap, op_id)?;
303 Ok((stride, node_words, out_words))
304}
305
306fn valid_node_expr(expr: Expr, node_count: u32) -> Expr {
307 Expr::and(
308 Expr::ne(expr.clone(), Expr::u32(SENTINEL)),
309 Expr::lt(expr, Expr::u32(node_count)),
310 )
311}
312
313fn descend_to_leftmost_leaf_node(nodes_name: &str, node_count: u32, stride: u32) -> Node {
314 Node::loop_for(
315 "descend",
316 Expr::u32(0),
317 Expr::u32(node_count),
318 vec![Node::if_then(
319 valid_node_expr(Expr::var("n"), node_count),
320 vec![
321 Node::let_bind(
322 "fc_idx",
323 Expr::add(Expr::mul(Expr::var("n"), Expr::u32(stride)), Expr::u32(2)),
324 ),
325 Node::let_bind("fc", Expr::load(nodes_name, Expr::var("fc_idx"))),
326 Node::if_then(
327 valid_node_expr(Expr::var("fc"), node_count),
328 vec![Node::assign("n", Expr::var("fc"))],
329 ),
330 ],
331 )],
332 )
333}
334
335fn checked_node_words(node_count: u32, stride: u32, op_id: &'static str) -> Result<u32, String> {
336 if node_count == 0 {
337 return Ok(1);
338 }
339 node_count.checked_mul(stride).ok_or_else(|| {
340 format!(
341 "{op_id} node_count={node_count} stride={stride} overflows VAST node buffer words. Fix: shard the tree before GPU dispatch."
342 )
343 })
344}
345
346fn checked_out_words(out_cap: u32, op_id: &'static str) -> Result<u32, String> {
347 if out_cap == 0 {
348 Err(format!(
349 "{op_id} requires out_cap > 0. Fix: allocate traversal output capacity before GPU dispatch."
350 ))
351 } else {
352 Ok(out_cap)
353 }
354}
355
356fn tree_walk_program(
357 op_id: &'static str,
358 nodes: &str,
359 out: &str,
360 node_words: u32,
361 out_words: u32,
362 body: Vec<Node>,
363) -> Program {
364 Program::wrapped(
365 vec![
366 BufferDecl::storage(nodes, 0, BufferAccess::ReadOnly, DataType::U32)
367 .with_count(node_words),
368 BufferDecl::storage(out, 1, BufferAccess::ReadWrite, DataType::U32)
369 .with_count(out_words),
370 ],
371 [1, 1, 1],
372 vec![Node::Region {
373 generator: Ident::from(op_id),
374 source_region: None,
375 body: Arc::new(body),
376 }],
377 )
378}
379
380#[cfg(test)]
381mod tests {
382 use super::*;
383
384 #[test]
385 fn checked_preorder_rejects_zero_output_capacity() {
386 let error = try_ast_walk_preorder("nodes", "out", 1, 0)
387 .expect_err("checked preorder builder must reject zero output capacity");
388
389 assert!(
390 error.contains("out_cap > 0"),
391 "error should describe the launch-shape fix: {error}"
392 );
393 }
394
395 #[test]
396 fn checked_postorder_rejects_node_word_overflow() {
397 let error = try_ast_walk_postorder("nodes", "out", u32::MAX, 1)
398 .expect_err("checked postorder builder must reject node buffer overflow");
399
400 assert!(
401 error.contains("overflows VAST node buffer words"),
402 "error should describe the VAST buffer overflow: {error}"
403 );
404 }
405
406 #[test]
407 fn checked_plan_builds_both_orders_from_primitive_authority() {
408 let plan = try_ast_walk_plan("nodes", "pre", "post", 3, 3)
409 .expect("Fix: primitive VAST plan should build both traversal orders");
410
411 assert_eq!(plan.preorder.workgroup_size(), [1, 1, 1]);
412 assert_eq!(plan.postorder.workgroup_size(), [1, 1, 1]);
413 assert_eq!(plan.preorder.buffers().len(), 2);
414 assert_eq!(plan.postorder.buffers().len(), 2);
415 }
416
417 #[test]
418 fn checked_plan_rejects_shape_before_building_partial_facade_state() {
419 let error = try_ast_walk_plan("nodes", "pre", "post", 3, 0)
420 .expect_err("Fix: primitive VAST plan should reject invalid shared output capacity");
421
422 assert!(
423 error.contains("out_cap > 0"),
424 "Fix: VAST plan diagnostic should come from the primitive output-capacity contract: {error}"
425 );
426 }
427
428 #[test]
429 fn legacy_vast_walk_builders_fail_fast_on_invalid_shape() {
430 let preorder_panic = std::panic::catch_unwind(|| {
431 let _ = ast_walk_preorder("nodes", "out", 1, 0);
432 })
433 .expect_err("legacy preorder builder must fail fast on zero output capacity");
434 let postorder_panic = std::panic::catch_unwind(|| {
435 let _ = ast_walk_postorder("nodes", "out", u32::MAX, 1);
436 })
437 .expect_err("legacy postorder builder must fail fast on node_count overflow");
438
439 let preorder_message = panic_payload_message(preorder_panic);
440 let postorder_message = panic_payload_message(postorder_panic);
441 assert!(
442 preorder_message.contains("out_cap > 0"),
443 "error should describe the launch-shape fix: {preorder_message}"
444 );
445 assert!(
446 postorder_message.contains("node_count"),
447 "error should describe the node_count overflow: {postorder_message}"
448 );
449 }
450
451 fn panic_payload_message(payload: Box<dyn std::any::Any + Send>) -> String {
452 if let Some(message) = payload.downcast_ref::<&str>() {
453 message.to_string()
454 } else if let Some(message) = payload.downcast_ref::<String>() {
455 message.clone()
456 } else {
457 format!("{payload:?}")
458 }
459 }
460
461 fn fixture_tree() -> Vec<u32> {
466 vec![
467 1, SENTINEL, 1, SENTINEL, 0, 0, 0, 0, 0,
468 0, 2, 0, SENTINEL, 2, 0, 0, 0, 0, 0, 0, 3, 0, SENTINEL, SENTINEL, 0, 0, 0, 0, 0,
471 0, ]
473 }
474
475 fn valid(idx: u32, node_count: u32) -> bool {
476 idx != SENTINEL && idx < node_count
477 }
478
479 fn cpu_preorder(nodes: &[u32], node_count: u32) -> Vec<u32> {
480 if node_count == 0 {
481 return Vec::new();
482 }
483 let stride = NODE_STRIDE_U32 as u32;
484 let mut out = Vec::new();
485 let mut n: u32 = 0;
486 for _ in 0..node_count {
487 if !valid(n, node_count) {
488 break;
489 }
490 out.push(n);
491 let base = (n * stride) as usize;
492 let fc = nodes[base + 2];
493 if valid(fc, node_count) {
494 n = fc;
495 } else {
496 let mut walk = n;
498 let mut next = SENTINEL;
499 while valid(walk, node_count) && next == SENTINEL {
500 let wb = (walk * stride) as usize;
501 let sib = nodes[wb + 3];
502 if valid(sib, node_count) {
503 next = sib;
504 } else {
505 walk = nodes[wb + 1]; }
507 }
508 if next == SENTINEL {
509 break;
510 }
511 n = next;
512 }
513 }
514 out
515 }
516
517 fn cpu_postorder(nodes: &[u32], node_count: u32) -> Vec<u32> {
518 if node_count == 0 {
519 return Vec::new();
520 }
521 let stride = NODE_STRIDE_U32 as u32;
522 let mut out = Vec::new();
523 let mut n: u32 = 0;
525 loop {
526 let base = (n * stride) as usize;
527 let fc = nodes[base + 2];
528 if valid(fc, node_count) {
529 n = fc;
530 } else {
531 break;
532 }
533 }
534 for _ in 0..node_count {
535 if !valid(n, node_count) {
536 break;
537 }
538 out.push(n);
539 if n == 0 {
540 break;
541 } let base = (n * stride) as usize;
543 let sib = nodes[base + 3];
544 if valid(sib, node_count) {
545 n = sib;
547 loop {
548 let sb = (n * stride) as usize;
549 let fc = nodes[sb + 2];
550 if valid(fc, node_count) {
551 n = fc;
552 } else {
553 break;
554 }
555 }
556 } else {
557 n = nodes[base + 1]; }
559 }
560 out
561 }
562
563 #[test]
564 fn cpu_preorder_matches_inventory() {
565 let tree = fixture_tree();
566 let result = cpu_preorder(&tree, 3);
567 assert_eq!(result, vec![0, 1, 2]);
568 }
569
570 #[test]
571 fn cpu_postorder_matches_inventory() {
572 let tree = fixture_tree();
573 let result = cpu_postorder(&tree, 3);
574 assert_eq!(result, vec![1, 2, 0]);
575 }
576
577 #[test]
578 fn cpu_preorder_single_node() {
579 let tree = vec![42u32, SENTINEL, SENTINEL, SENTINEL, 0, 0, 0, 0, 0, 0];
580 assert_eq!(cpu_preorder(&tree, 1), vec![0]);
581 }
582
583 #[test]
584 fn cpu_postorder_single_node() {
585 let tree = vec![42u32, SENTINEL, SENTINEL, SENTINEL, 0, 0, 0, 0, 0, 0];
586 assert_eq!(cpu_postorder(&tree, 1), vec![0]);
587 }
588
589 #[test]
590 fn cpu_preorder_empty() {
591 assert_eq!(cpu_preorder(&[], 0), Vec::<u32>::new());
592 }
593
594 #[test]
595 fn cpu_postorder_empty() {
596 assert_eq!(cpu_postorder(&[], 0), Vec::<u32>::new());
597 }
598
599 fn generated_parent(seed: u32, child: u32) -> u32 {
600 seed.wrapping_mul(1_664_525)
601 .wrapping_add(child.wrapping_mul(1_013_904_223))
602 .rotate_left(child % 31)
603 % child
604 }
605
606 fn generated_valid_tree(seed: u32, node_count: u32) -> Vec<u32> {
607 let stride = NODE_STRIDE_U32;
608 let mut nodes = vec![0u32; node_count as usize * stride];
609 for node in 0..node_count {
610 let base = node as usize * stride;
611 nodes[base] = seed ^ node;
612 nodes[base + 1] = SENTINEL;
613 nodes[base + 2] = SENTINEL;
614 nodes[base + 3] = SENTINEL;
615 }
616
617 for child in 1..node_count {
618 let parent = generated_parent(seed, child);
619 let child_base = child as usize * stride;
620 let parent_base = parent as usize * stride;
621 nodes[child_base + 1] = parent;
622
623 if nodes[parent_base + 2] == SENTINEL {
624 nodes[parent_base + 2] = child;
625 continue;
626 }
627
628 let mut sibling = nodes[parent_base + 2];
629 loop {
630 let sibling_next = sibling as usize * stride + 3;
631 if nodes[sibling_next] == SENTINEL {
632 nodes[sibling_next] = child;
633 break;
634 }
635 sibling = nodes[sibling_next];
636 }
637 }
638
639 nodes
640 }
641
642 fn positions(order: &[u32], node_count: u32) -> Vec<u32> {
643 let mut positions = vec![SENTINEL; node_count as usize];
644 for (pos, node) in order.iter().copied().enumerate() {
645 assert!(
646 valid(node, node_count),
647 "generated VAST traversal emitted invalid node {node}"
648 );
649 assert_eq!(
650 positions[node as usize], SENTINEL,
651 "generated VAST traversal emitted node {node} twice"
652 );
653 positions[node as usize] = pos as u32;
654 }
655 assert!(
656 positions.iter().all(|pos| *pos != SENTINEL),
657 "generated VAST traversal missed at least one node"
658 );
659 positions
660 }
661
662 #[test]
663 fn generated_vast_walk_orders_match_tree_order_contracts() {
664 for seed in 0..2048u32 {
665 let node_count = seed % 37 + 1;
666 let tree = generated_valid_tree(seed, node_count);
667 let preorder = cpu_preorder(&tree, node_count);
668 let postorder = cpu_postorder(&tree, node_count);
669
670 assert_eq!(
671 preorder.len(),
672 node_count as usize,
673 "preorder must emit every generated VAST node exactly once for seed {seed}"
674 );
675 assert_eq!(
676 postorder.len(),
677 node_count as usize,
678 "postorder must emit every generated VAST node exactly once for seed {seed}"
679 );
680
681 let preorder_positions = positions(&preorder, node_count);
682 let postorder_positions = positions(&postorder, node_count);
683 let stride = NODE_STRIDE_U32;
684
685 for child in 1..node_count {
686 let parent = tree[child as usize * stride + 1];
687 assert!(
688 preorder_positions[parent as usize] < preorder_positions[child as usize],
689 "preorder must emit parent {parent} before child {child} for seed {seed}"
690 );
691 assert!(
692 postorder_positions[child as usize] < postorder_positions[parent as usize],
693 "postorder must emit child {child} before parent {parent} for seed {seed}"
694 );
695 }
696 }
697 }
698}
699
700#[cfg(feature = "inventory-registry")]
701fn fixture_u32(words: &[u32]) -> Vec<u8> {
702 crate::wire::pack_u32_slice(words)
703}
704
705#[cfg(feature = "inventory-registry")]
706fn fixture_tree_words() -> Vec<u32> {
707 vec![
708 1, SENTINEL, 1, SENTINEL, 0, 0, 0, 0, 0, 0, 2, 0, SENTINEL, 2, 0, 0, 0, 0, 0, 0, 3, 0, SENTINEL, SENTINEL, 0, 0, 0, 0, 0, 0, ]
712}
713
714#[cfg(feature = "inventory-registry")]
715inventory::submit! {
716 vyre_foundation::operation::OperationRegistration::primitive(
717 PREORDER_OP_ID,
718 || ast_walk_preorder("nodes", "out", 3, 3),
719 Some(|| vec![vec![
720 fixture_u32(&fixture_tree_words()),
721 fixture_u32(&[SENTINEL, SENTINEL, SENTINEL]),
722 ]]),
723 Some(|| vec![vec![fixture_u32(&[0, 1, 2])]]),
724 )
725}
726
727#[cfg(feature = "inventory-registry")]
728inventory::submit! {
729 vyre_foundation::operation::OperationRegistration::primitive(
730 POSTORDER_OP_ID,
731 || ast_walk_postorder("nodes", "out", 3, 3),
732 Some(|| vec![vec![
733 fixture_u32(&fixture_tree_words()),
734 fixture_u32(&[SENTINEL, SENTINEL, SENTINEL]),
735 ]]),
736 Some(|| vec![vec![fixture_u32(&[1, 2, 0])]]),
737 )
738}