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