Skip to main content

lc_langgraph/compiled/
validate.rs

1// crates/lc-langgraph/src/compiled/validate.rs
2//! CompiledGraph validation methods: validate, validate_duplicate_edges,
3//! validate_unreachable_nodes, validate_cycles, and helper graph-traversal methods
4
5use super::graph::CompiledGraph;
6use crate::edge::GraphEdge;
7use crate::errors::{GraphError, GraphResult};
8use crate::state::StateSchema;
9use crate::{END, START};
10
11impl<S: StateSchema> CompiledGraph<S> {
12    /// Validate the graph structure: node references, duplicate edges,
13    /// unreachable nodes, and cycles with no path to `END`.
14    pub fn validate(&self) -> GraphResult<()> {
15        for edge in &self.edges {
16            match edge {
17                GraphEdge::Fixed { source, target } => {
18                    if source != START && !self.nodes.contains_key(source) {
19                        return Err(GraphError::ValidationError(format!(
20                            "Source node '{}' not found",
21                            source
22                        )));
23                    }
24                    if target != END && !self.nodes.contains_key(target) {
25                        return Err(GraphError::ValidationError(format!(
26                            "Target node '{}' not found",
27                            target
28                        )));
29                    }
30                    if target == START {
31                        return Err(GraphError::ValidationError(
32                            "Edge cannot target START node".to_string(),
33                        ));
34                    }
35                }
36                GraphEdge::Conditional {
37                    source,
38                    router_name,
39                    targets,
40                    default_target,
41                } => {
42                    if source != START && !self.nodes.contains_key(source) {
43                        return Err(GraphError::ValidationError(format!(
44                            "Source node '{}' not found",
45                            source
46                        )));
47                    }
48                    if !self.conditional_routers.contains_key(router_name) {
49                        return Err(GraphError::ValidationError(format!(
50                            "Router '{}' not found",
51                            router_name
52                        )));
53                    }
54                    for (route, target) in targets {
55                        if target != END && !self.nodes.contains_key(target) {
56                            return Err(GraphError::ValidationError(format!(
57                                "Target '{}' for route '{}' not found",
58                                target, route
59                            )));
60                        }
61                        if target == START {
62                            return Err(GraphError::ValidationError(
63                                "Conditional edge cannot target START node".to_string(),
64                            ));
65                        }
66                    }
67                    if let Some(default) = default_target {
68                        if default != END && !self.nodes.contains_key(default) {
69                            return Err(GraphError::ValidationError(format!(
70                                "Default target '{}' not found",
71                                default
72                            )));
73                        }
74                    }
75                }
76                GraphEdge::FanOut { source, targets } => {
77                    if source != START && !self.nodes.contains_key(source) {
78                        return Err(GraphError::ValidationError(format!(
79                            "FanOut source node '{}' not found",
80                            source
81                        )));
82                    }
83                    for target in targets {
84                        if target != END && !self.nodes.contains_key(target) {
85                            return Err(GraphError::ValidationError(format!(
86                                "FanOut target node '{}' not found",
87                                target
88                            )));
89                        }
90                    }
91                }
92                GraphEdge::FanIn { sources, target } => {
93                    for source in sources {
94                        if source != START && !self.nodes.contains_key(source) {
95                            return Err(GraphError::ValidationError(format!(
96                                "FanIn source node '{}' not found",
97                                source
98                            )));
99                        }
100                    }
101                    if target != END && !self.nodes.contains_key(target) {
102                        return Err(GraphError::ValidationError(format!(
103                            "FanIn target node '{}' not found",
104                            target
105                        )));
106                    }
107                }
108            }
109        }
110
111        self.validate_duplicate_edges()?;
112        self.validate_unreachable_nodes()?;
113        self.validate_cycles()?;
114
115        Ok(())
116    }
117
118    fn validate_duplicate_edges(&self) -> GraphResult<()> {
119        let mut seen_fixed: std::collections::HashSet<(String, String)> =
120            std::collections::HashSet::new();
121
122        for edge in &self.edges {
123            if let GraphEdge::Fixed { source, target } = edge {
124                let key = (source.clone(), target.clone());
125                if seen_fixed.contains(&key) {
126                    return Err(GraphError::DuplicateEdgeError(format!(
127                        "Duplicate edge: {} -> {}",
128                        source, target
129                    )));
130                }
131                seen_fixed.insert(key);
132            }
133        }
134        Ok(())
135    }
136
137    fn validate_unreachable_nodes(&self) -> GraphResult<()> {
138        let reachable = self.compute_reachable_nodes();
139
140        for node_name in self.nodes.keys() {
141            if !reachable.contains(node_name) {
142                return Err(GraphError::OrphanNodeError(format!(
143                    "Unreachable node: {}",
144                    node_name
145                )));
146            }
147        }
148        Ok(())
149    }
150
151    fn compute_reachable_nodes(&self) -> std::collections::HashSet<String> {
152        let mut reachable: std::collections::HashSet<String> = std::collections::HashSet::new();
153        let mut to_visit: Vec<String> = vec![self.entry_point.clone()];
154
155        while let Some(current) = to_visit.pop() {
156            if reachable.contains(&current) || current == END {
157                continue;
158            }
159            reachable.insert(current.clone());
160
161            for edge in &self.edges {
162                if edge.source() == current {
163                    match edge {
164                        GraphEdge::Fixed { target, .. } => {
165                            if !reachable.contains(target) && target != END {
166                                to_visit.push(target.clone());
167                            }
168                        }
169                        GraphEdge::Conditional {
170                            targets,
171                            default_target,
172                            ..
173                        } => {
174                            for target in targets.values() {
175                                if !reachable.contains(target) && target != END {
176                                    to_visit.push(target.clone());
177                                }
178                            }
179                            if let Some(default) = default_target {
180                                if !reachable.contains(default) && default != END {
181                                    to_visit.push(default.clone());
182                                }
183                            }
184                        }
185                        GraphEdge::FanOut { targets, .. } => {
186                            for target in targets {
187                                if !reachable.contains(target) && target != END {
188                                    to_visit.push(target.clone());
189                                }
190                            }
191                        }
192                        GraphEdge::FanIn { sources, target } => {
193                            if sources.iter().all(|s| reachable.contains(s))
194                                && !reachable.contains(target)
195                                && target != END
196                            {
197                                to_visit.push(target.clone());
198                            }
199                        }
200                    }
201                }
202            }
203        }
204        reachable
205    }
206
207    fn validate_cycles(&self) -> GraphResult<()> {
208        let reachable = self.compute_reachable_nodes();
209        let end_reachable = self.compute_end_reachable_nodes();
210
211        for node in &reachable {
212            if !end_reachable.contains(node) {
213                return Err(GraphError::InfiniteCycleError(format!(
214                    "Node '{}' in cycle with no path to END",
215                    node
216                )));
217            }
218        }
219        Ok(())
220    }
221
222    fn compute_end_reachable_nodes(&self) -> std::collections::HashSet<String> {
223        let mut end_reachable: std::collections::HashSet<String> = std::collections::HashSet::new();
224        end_reachable.insert(END.to_string());
225
226        let mut changed = true;
227        while changed {
228            changed = false;
229            for edge in &self.edges {
230                match edge {
231                    GraphEdge::Fixed { source, target } => {
232                        if end_reachable.contains(target) && !end_reachable.contains(source) {
233                            end_reachable.insert(source.clone());
234                            changed = true;
235                        }
236                    }
237                    GraphEdge::Conditional {
238                        source,
239                        targets,
240                        default_target,
241                        ..
242                    } => {
243                        let any_target_reaches_end =
244                            targets.values().any(|t| end_reachable.contains(t))
245                                || default_target
246                                    .as_ref()
247                                    .is_some_and(|d| end_reachable.contains(d));
248                        if any_target_reaches_end && !end_reachable.contains(source) {
249                            end_reachable.insert(source.clone());
250                            changed = true;
251                        }
252                    }
253                    GraphEdge::FanOut { source, targets } => {
254                        let all_targets_reach_end =
255                            targets.iter().all(|t| end_reachable.contains(t));
256                        if all_targets_reach_end && !end_reachable.contains(source) {
257                            end_reachable.insert(source.clone());
258                            changed = true;
259                        }
260                    }
261                    GraphEdge::FanIn { sources, target } => {
262                        if end_reachable.contains(target) {
263                            for source in sources {
264                                if !end_reachable.contains(source) {
265                                    end_reachable.insert(source.clone());
266                                    changed = true;
267                                }
268                            }
269                        }
270                    }
271                }
272            }
273        }
274        end_reachable
275    }
276}