Skip to main content

vyre_driver/backend/
validation.rs

1//! Backend support validation before dispatch.
2
3use 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
33/// Validate that `backend` supports every operation in `program`.
34pub 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
41/// Default core operation support set for legacy backends.
42pub 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
52/// Default core operation set plus `Node::Trap`.
53///
54/// `Trap` is a structural control-flow node, not a concrete-driver extension:
55/// backends that lower it as lane termination should use this shared set
56/// instead of carrying a backend-local `OnceLock` and literal allocation.
57pub 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        // Leaf nodes and backend-transparent nodes (opaque extensions
103        // validate themselves via `NodeExtension::validate_extension`).
104        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        // `Node` is `#[non_exhaustive]` in vyre-foundation. Future variants
114        // land here as transparent leaves until a dedicated arm is added.
115        _ => {}
116    }
117    Ok(())
118}