Skip to main content

lumen_engine/
graph.rs

1//! Directed graph model, validation, and topological evaluation ordering.
2
3use std::collections::{HashMap, VecDeque};
4
5use crate::{
6    error::{GraphValidationError, LumenError},
7    node::{Node, NodeId, NodeKind},
8};
9
10#[derive(Debug, Clone, PartialEq, Eq)]
11pub struct Connection {
12    pub from_node: NodeId,
13    pub from_port: String,
14    pub to_node: NodeId,
15    pub to_port: String,
16}
17
18#[derive(Default, Debug)]
19pub struct Graph {
20    pub nodes: HashMap<NodeId, NodeKind>,
21    pub connections: Vec<Connection>,
22    outgoing_connection_counts: HashMap<NodeId, usize>,
23}
24
25unsafe impl Sync for Graph {}
26unsafe impl Send for Graph {}
27
28impl Graph {
29    pub fn new() -> Self {
30        Self {
31            nodes: HashMap::new(),
32            connections: Vec::new(),
33            outgoing_connection_counts: HashMap::new(),
34        }
35    }
36
37    pub fn connect(&mut self, connection: Connection) -> crate::Result<()> {
38        if !self.nodes.contains_key(&connection.from_node) {
39            return Err(GraphValidationError::MissingSourceNode {
40                node_id: connection.from_node,
41            }
42            .into());
43        }
44
45        if !self.nodes.contains_key(&connection.to_node) {
46            return Err(GraphValidationError::MissingTargetNode {
47                node_id: connection.to_node,
48            }
49            .into());
50        }
51
52        *self
53            .outgoing_connection_counts
54            .entry(connection.from_node)
55            .or_default() += 1;
56        self.connections.push(connection);
57        Ok(())
58    }
59
60    pub fn outgoing_connection_count(&self, node_id: NodeId) -> usize {
61        self.outgoing_connection_counts
62            .get(&node_id)
63            .copied()
64            .unwrap_or_default()
65    }
66
67    pub fn validate(&self) -> Result<(), Vec<LumenError>> {
68        let mut errors = Vec::new();
69
70        let media_output_count = self
71            .nodes
72            .values()
73            .filter(|node| matches!(node, NodeKind::MediaOutput(_)))
74            .count();
75        if media_output_count == 0 {
76            errors.push(GraphValidationError::MissingMediaOutput.into());
77        } else if media_output_count > 1 {
78            errors.push(
79                GraphValidationError::MultipleMediaOutputs {
80                    count: media_output_count,
81                }
82                .into(),
83            );
84        }
85
86        for connection in &self.connections {
87            let Some(from_node) = self.nodes.get(&connection.from_node) else {
88                errors.push(
89                    GraphValidationError::MissingSourceNode {
90                        node_id: connection.from_node,
91                    }
92                    .into(),
93                );
94                continue;
95            };
96            let Some(to_node) = self.nodes.get(&connection.to_node) else {
97                errors.push(
98                    GraphValidationError::MissingTargetNode {
99                        node_id: connection.to_node,
100                    }
101                    .into(),
102                );
103                continue;
104            };
105
106            let Some(output_def) = from_node
107                .output_port_defs()
108                .iter()
109                .find(|def| def.name == connection.from_port)
110            else {
111                errors.push(
112                    GraphValidationError::MissingSourcePort {
113                        node_id: connection.from_node,
114                        port: connection.from_port.clone(),
115                    }
116                    .into(),
117                );
118                continue;
119            };
120            let Some(input_def) = to_node
121                .input_port_defs()
122                .iter()
123                .find(|def| def.name == connection.to_port)
124            else {
125                errors.push(
126                    GraphValidationError::MissingTargetPort {
127                        node_id: connection.to_node,
128                        port: connection.to_port.clone(),
129                    }
130                    .into(),
131                );
132                continue;
133            };
134
135            if output_def.kind != input_def.kind {
136                errors.push(
137                    GraphValidationError::PortKindMismatch {
138                        from_node: connection.from_node,
139                        from_port: output_def.name.into(),
140                        from_kind: output_def.kind,
141                        to_node: connection.to_node,
142                        to_port: input_def.name.into(),
143                        expected_kind: input_def.kind,
144                    }
145                    .into(),
146                );
147            }
148        }
149
150        for node in self.nodes.values() {
151            for input in node.input_port_defs() {
152                if input.optional {
153                    continue;
154                }
155
156                let connected = self
157                    .connections
158                    .iter()
159                    .any(|edge| edge.to_node == node.id() && edge.to_port == input.name);
160                if !connected {
161                    errors.push(
162                        GraphValidationError::MissingRequiredInput {
163                            node_id: node.id(),
164                            port: input.name.to_string(),
165                        }
166                        .into(),
167                    );
168                }
169            }
170        }
171
172        if let Err(cycle_error) = self.validate_no_cycle() {
173            errors.push(cycle_error.into());
174        }
175
176        if errors.is_empty() {
177            Ok(())
178        } else {
179            Err(errors)
180        }
181    }
182
183    fn validate_no_cycle(&self) -> Result<(), GraphValidationError> {
184        let mut indegree: HashMap<NodeId, usize> =
185            self.nodes.keys().copied().map(|id| (id, 0)).collect();
186
187        for edge in &self.connections {
188            if let Some(entry) = indegree.get_mut(&edge.to_node) {
189                *entry += 1;
190            }
191        }
192
193        let mut queue: VecDeque<NodeId> = indegree
194            .iter()
195            .filter_map(|(node_id, degree)| (*degree == 0).then_some(*node_id))
196            .collect();
197
198        let mut visited = 0_usize;
199        while let Some(node_id) = queue.pop_front() {
200            visited += 1;
201            for edge in self
202                .connections
203                .iter()
204                .filter(|edge| edge.from_node == node_id)
205            {
206                if let Some(entry) = indegree.get_mut(&edge.to_node) {
207                    *entry -= 1;
208                    if *entry == 0 {
209                        queue.push_back(edge.to_node);
210                    }
211                }
212            }
213        }
214
215        if visited != self.nodes.len() {
216            let cycle_nodes = indegree
217                .into_iter()
218                .filter_map(|(node_id, degree)| (degree > 0).then_some(node_id))
219                .collect();
220            return Err(GraphValidationError::Cycle { path: cycle_nodes });
221        }
222
223        Ok(())
224    }
225}