Skip to main content

lc_langgraph/compiled/
invoke.rs

1// crates/lc-langgraph/src/compiled/invoke.rs
2//! CompiledGraph invoke, invoke_with_execution, resume, and invoke_from_node methods
3
4use super::graph::CompiledGraph;
5use super::types::{ExecutionStep, GraphExecution, GraphInvocation, ParallelBranch};
6use crate::errors::{GraphError, GraphResult};
7use crate::node::NodeConfig;
8use crate::state::StateSchema;
9use crate::END;
10use std::collections::HashMap;
11
12impl<S: StateSchema> CompiledGraph<S> {
13    pub async fn invoke(&self, input: S) -> GraphResult<GraphInvocation<S>> {
14        let mut state = input;
15        let mut current_node = self.entry_point.clone();
16        let mut steps: Vec<ExecutionStep> = Vec::new();
17        let mut recursion_count = 0;
18
19        if let Some(ref checkpointer) = self.checkpointer {
20            let checkpoint_id = checkpointer.lock().await.save(&state).await?;
21            steps.push(ExecutionStep::checkpoint(
22                checkpoint_id,
23                current_node.clone(),
24            ));
25        }
26
27        while current_node != END && recursion_count < self.recursion_limit {
28            if self.interrupt_before.contains(&current_node) {
29                return Err(GraphError::ExecutionInterrupted(current_node.clone()));
30            }
31
32            // Check for FanOut edge — execute all branches in parallel and merge
33            let fan_out_targets = self.find_fan_out_targets(&current_node).await;
34            if let Some(targets) = fan_out_targets {
35                recursion_count += 1;
36                let mut parallel_branches: Vec<ParallelBranch<S>> = Vec::new();
37
38                let branch_results = self.execute_parallel_branches(&targets, &state).await?;
39                for (name, inv) in branch_results {
40                    parallel_branches.push(ParallelBranch {
41                        name: name.clone(),
42                        final_state: inv.final_state.clone(),
43                        steps: inv.steps.clone(),
44                    });
45                    steps.push(ExecutionStep::ParallelNode {
46                        branch: name,
47                        metadata: HashMap::new(),
48                    });
49                }
50
51                let merge_target = self.find_fan_in_target(&targets).await;
52                if let Some(merge_node) = merge_target {
53                    state = self.merge_parallel_states(&parallel_branches)?;
54                    current_node = merge_node;
55                } else {
56                    state = self.merge_parallel_states(&parallel_branches)?;
57                    current_node = END.to_string();
58                }
59
60                if let Some(ref checkpointer) = self.checkpointer {
61                    let checkpoint_id = checkpointer.lock().await.save(&state).await?;
62                    steps.push(ExecutionStep::checkpoint(
63                        checkpoint_id,
64                        current_node.clone(),
65                    ));
66                }
67                continue;
68            }
69
70            recursion_count += 1;
71
72            let node = self.get_node(&current_node).await?;
73
74            let config = NodeConfig {
75                recursion_limit: self.recursion_limit,
76                debug: false,
77                metadata: HashMap::new(),
78            };
79
80            let update = node.execute(&state, Some(config)).await?;
81
82            if let Some(new_state) = update.update {
83                state = self.default_reducer.reduce(&state, &new_state);
84            }
85
86            steps.push(ExecutionStep::node(
87                current_node.clone(),
88                update.metadata.clone(),
89            ));
90
91            if self.interrupt_after.contains(&current_node) {
92                return Err(GraphError::ExecutionInterrupted(format!(
93                    "after_{}",
94                    current_node
95                )));
96            }
97
98            let next_node = self.find_next_node(&current_node, &state).await?;
99
100            if let Some(ref checkpointer) = self.checkpointer {
101                let checkpoint_id = checkpointer.lock().await.save(&state).await?;
102                steps.push(ExecutionStep::checkpoint(checkpoint_id, next_node.clone()));
103            }
104
105            current_node = next_node;
106        }
107
108        if recursion_count >= self.recursion_limit {
109            return Err(GraphError::RecursionLimitReached(self.recursion_limit));
110        }
111
112        Ok(GraphInvocation {
113            final_state: state,
114            steps,
115            recursion_count,
116        })
117    }
118
119    pub async fn invoke_with_execution(
120        &self,
121        execution: GraphExecution<S>,
122    ) -> GraphResult<GraphInvocation<S>> {
123        let mut state = execution.state;
124        let mut current_node = if execution.interrupted_at.starts_with("after_") {
125            self.find_next_node(&execution.current_node, &state).await?
126        } else {
127            execution.current_node
128        };
129        let mut steps = execution.steps;
130        let mut recursion_count = execution.recursion_count;
131        let first_node = current_node.clone();
132
133        while current_node != END && recursion_count < self.recursion_limit {
134            if current_node != first_node && self.interrupt_before.contains(&current_node) {
135                return Err(GraphError::ExecutionInterrupted(current_node.clone()));
136            }
137
138            recursion_count += 1;
139
140            let node = self.get_node(&current_node).await?;
141
142            let config = NodeConfig {
143                recursion_limit: self.recursion_limit,
144                debug: false,
145                metadata: HashMap::new(),
146            };
147
148            let update = node.execute(&state, Some(config)).await?;
149
150            if let Some(new_state) = update.update {
151                state = self.default_reducer.reduce(&state, &new_state);
152            }
153
154            steps.push(ExecutionStep::node(
155                current_node.clone(),
156                update.metadata.clone(),
157            ));
158
159            if self.interrupt_after.contains(&current_node) {
160                return Err(GraphError::ExecutionInterrupted(format!(
161                    "after_{}",
162                    current_node
163                )));
164            }
165
166            let next_node = self.find_next_node(&current_node, &state).await?;
167
168            if let Some(ref checkpointer) = self.checkpointer {
169                let checkpoint_id = checkpointer.lock().await.save(&state).await?;
170                steps.push(ExecutionStep::checkpoint(checkpoint_id, next_node.clone()));
171            }
172
173            current_node = next_node;
174        }
175
176        if recursion_count >= self.recursion_limit {
177            return Err(GraphError::RecursionLimitReached(self.recursion_limit));
178        }
179
180        Ok(GraphInvocation {
181            final_state: state,
182            steps,
183            recursion_count,
184        })
185    }
186
187    pub async fn resume(&self, execution: GraphExecution<S>) -> GraphResult<GraphInvocation<S>> {
188        self.invoke_with_execution(execution).await
189    }
190
191    pub async fn invoke_from_node(
192        &self,
193        start_node: String,
194        input: S,
195    ) -> GraphResult<GraphInvocation<S>> {
196        let mut state = input;
197        let mut current_node = start_node;
198        let mut steps: Vec<ExecutionStep> = Vec::new();
199        let mut recursion_count = 0;
200
201        while current_node != END && recursion_count < self.recursion_limit {
202            if self.interrupt_before.contains(&current_node) {
203                return Err(GraphError::ExecutionInterrupted(current_node.clone()));
204            }
205
206            recursion_count += 1;
207
208            let node = self.get_node(&current_node).await?;
209
210            let config = NodeConfig {
211                recursion_limit: self.recursion_limit,
212                debug: false,
213                metadata: HashMap::new(),
214            };
215
216            let update = node.execute(&state, Some(config)).await?;
217
218            if let Some(new_state) = update.update {
219                state = self.default_reducer.reduce(&state, &new_state);
220            }
221
222            steps.push(ExecutionStep::node(
223                current_node.clone(),
224                update.metadata.clone(),
225            ));
226
227            if self.interrupt_after.contains(&current_node) {
228                return Err(GraphError::ExecutionInterrupted(format!(
229                    "after_{}",
230                    current_node
231                )));
232            }
233
234            current_node = self.find_next_node(&current_node, &state).await?;
235        }
236
237        Ok(GraphInvocation {
238            final_state: state,
239            steps,
240            recursion_count,
241        })
242    }
243}