1use 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}