vyre_driver/backend/
validation.rs1use super::capability::Backend;
4use std::sync::Arc;
5pub use vyre_foundation::ir::model::node::node_op_id;
6use vyre_foundation::ir::model::node::Node;
7use vyre_foundation::ir::{OpId, Program, ValidationError};
8
9const CORE_SUPPORTED_OP_IDS: &[&str] = &[
10 "vyre.node.let",
11 "vyre.node.assign",
12 "vyre.node.store",
13 "vyre.node.if",
14 "vyre.node.loop",
15 "vyre.node.return",
16 "vyre.node.block",
17 "vyre.node.barrier",
18 "vyre.node.indirect_dispatch",
19 "vyre.node.async_load",
20 "vyre.node.async_wait",
21 "vyre.node.region",
22 "vyre.lit_u32",
23 "vyre.lit_i32",
24 "vyre.lit_f32",
25 "vyre.lit_bool",
26 "vyre.var",
27 "vyre.bin_op",
28 "vyre.un_op",
29 "vyre.load",
30 "vyre.store",
31];
32
33pub fn validate_program(program: &Program, backend: &dyn Backend) -> Result<(), ValidationError> {
35 for (index, node) in program.entry().iter().enumerate() {
36 validate_node(node, index, backend.id(), backend.supported_ops())?;
37 }
38 Ok(())
39}
40
41pub fn default_supported_ops() -> &'static std::collections::HashSet<OpId> {
43 static OPS: std::sync::OnceLock<std::collections::HashSet<OpId>> = std::sync::OnceLock::new();
44 OPS.get_or_init(|| {
45 let mut ops = std::collections::HashSet::new();
46 ops.reserve(CORE_SUPPORTED_OP_IDS.len());
47 ops.extend(CORE_SUPPORTED_OP_IDS.iter().copied().map(Arc::<str>::from));
48 ops
49 })
50}
51
52pub fn default_supported_ops_with_trap() -> &'static std::collections::HashSet<OpId> {
58 static OPS: std::sync::OnceLock<std::collections::HashSet<OpId>> = std::sync::OnceLock::new();
59 OPS.get_or_init(|| {
60 let base = default_supported_ops();
61 let reserve = base.len().saturating_add(1);
62 let mut ops = std::collections::HashSet::new();
63 ops.reserve(reserve);
64 ops.extend(base.iter().cloned());
65 ops.insert(Arc::<str>::from("vyre.node.trap"));
66 ops
67 })
68}
69
70fn validate_node(
71 node: &Node,
72 index: usize,
73 backend: &'static str,
74 supported: &std::collections::HashSet<OpId>,
75) -> Result<(), ValidationError> {
76 let op = node_op_id(node);
77 if !supported.contains(op) {
78 let op_id = Arc::<str>::from(op);
79 return Err(ValidationError::unsupported_op(backend, &op_id, index));
80 }
81 match node {
82 Node::If {
83 then, otherwise, ..
84 } => {
85 for (offset, nested) in then.iter().enumerate() {
86 validate_node(nested, offset, backend, supported)?;
87 }
88 for (offset, nested) in otherwise.iter().enumerate() {
89 validate_node(nested, offset, backend, supported)?;
90 }
91 }
92 Node::Loop { body, .. } | Node::Block(body) => {
93 for (offset, nested) in body.iter().enumerate() {
94 validate_node(nested, offset, backend, supported)?;
95 }
96 }
97 Node::Region { body, .. } => {
98 for (offset, nested) in body.iter().enumerate() {
99 validate_node(nested, offset, backend, supported)?;
100 }
101 }
102 Node::Let { .. }
105 | Node::Assign { .. }
106 | Node::Store { .. }
107 | Node::Return
108 | Node::Barrier { .. }
109 | Node::IndirectDispatch { .. }
110 | Node::AsyncLoad { .. }
111 | Node::AsyncWait { .. }
112 | Node::Opaque(_) => {}
113 _ => {}
116 }
117 Ok(())
118}