Skip to main content

lc_langgraph/compiled/
stream.rs

1// crates/lc-langgraph/src/compiled/stream.rs
2//! CompiledGraph stream, find_next_node, find_fan_out_targets,
3//! find_fan_in_target, and execute_parallel_branches methods
4
5use super::graph::CompiledGraph;
6use super::types::{GraphInvocation, StreamEvent};
7use crate::edge::GraphEdge;
8use crate::errors::{GraphError, GraphResult};
9use crate::graph::END;
10use crate::node::NodeConfig;
11use crate::state::StateSchema;
12use futures_util::future::join_all;
13use std::collections::HashMap;
14
15impl<S: StateSchema> CompiledGraph<S> {
16    pub async fn stream(&self, input: S) -> GraphResult<Vec<StreamEvent<S>>> {
17        let mut events = Vec::new();
18        let mut state = input;
19        let mut current_node = self.entry_point.clone();
20        let mut recursion_count = 0;
21
22        events.push(StreamEvent::start(state.clone()));
23
24        if let Some(ref checkpointer) = self.checkpointer {
25            let _checkpoint_id = checkpointer.lock().await.save(&state).await?;
26        }
27
28        while current_node != END && recursion_count < self.recursion_limit {
29            if self.interrupt_before.contains(&current_node) {
30                return Err(GraphError::ExecutionInterrupted(current_node.clone()));
31            }
32
33            recursion_count += 1;
34
35            events.push(StreamEvent::enter_node(current_node.clone(), state.clone()));
36
37            let node = self.get_node(&current_node).await?;
38
39            let config = NodeConfig {
40                recursion_limit: self.recursion_limit,
41                debug: false,
42                metadata: HashMap::new(),
43            };
44
45            let update = node.execute(&state, Some(config)).await?;
46
47            events.push(StreamEvent::node_complete(
48                current_node.clone(),
49                update.clone(),
50            ));
51
52            if let Some(new_state) = update.update {
53                state = self.default_reducer.reduce(&state, &new_state);
54                events.push(StreamEvent::state_update(state.clone()));
55            }
56
57            if self.interrupt_after.contains(&current_node) {
58                if let Some(ref checkpointer) = self.checkpointer {
59                    let _checkpoint_id = checkpointer.lock().await.save(&state).await?;
60                }
61                return Err(GraphError::ExecutionInterrupted(format!(
62                    "after_{}",
63                    current_node
64                )));
65            }
66
67            let next_node = self.find_next_node(&current_node, &state).await?;
68
69            if let Some(ref checkpointer) = self.checkpointer {
70                let _checkpoint_id = checkpointer.lock().await.save(&state).await?;
71            }
72
73            current_node = next_node;
74        }
75
76        events.push(StreamEvent::end(state.clone()));
77        Ok(events)
78    }
79
80    pub(super) async fn find_next_node(&self, current: &str, state: &S) -> GraphResult<String> {
81        // 1. Check runtime edges — edge is cloned out of the guard before any .await
82        'rt: {
83            let edge = {
84                let re = self.runtime_edges.read().await;
85                match re.iter().find(|e| e.source() == current) {
86                    Some(e) => e.clone(),
87                    None => break 'rt,
88                }
89            }; // re dropped here
90            match edge {
91                GraphEdge::Fixed { target, .. } => return Ok(target),
92                GraphEdge::Conditional {
93                    router_name,
94                    targets,
95                    default_target,
96                    ..
97                } => {
98                    let router = self
99                        .conditional_routers
100                        .get(&router_name)
101                        .cloned()
102                        .or_else(|| {
103                            self.runtime_conditional_routers
104                                .try_read()
105                                .ok()
106                                .and_then(|guard| guard.get(&router_name).cloned())
107                        })
108                        .ok_or_else(|| {
109                            GraphError::ExecutionError(format!(
110                                "Router '{}' not found (runtime)",
111                                router_name
112                            ))
113                        })?;
114                    let route_key = router.route(state).await?;
115                    let target = targets
116                        .get(&route_key)
117                        .or(default_target.as_ref())
118                        .ok_or_else(|| {
119                            GraphError::RoutingError(format!(
120                                "No target for route '{}' (runtime)",
121                                route_key
122                            ))
123                        })?;
124                    return Ok(target.clone());
125                }
126                GraphEdge::FanOut { targets, .. } => {
127                    if targets.is_empty() {
128                        return Err(GraphError::RoutingError(
129                            "FanOut has no targets (runtime)".to_string(),
130                        ));
131                    }
132                    return Ok(targets[0].clone());
133                }
134                GraphEdge::FanIn { .. } => {}
135            }
136        }
137
138        // 2. Check static edges
139        for edge in &self.edges {
140            if edge.source() == current {
141                match edge {
142                    GraphEdge::Fixed { target, .. } => {
143                        return Ok(target.clone());
144                    }
145                    GraphEdge::Conditional {
146                        router_name,
147                        targets,
148                        default_target,
149                        ..
150                    } => {
151                        let router = self
152                            .conditional_routers
153                            .get(router_name)
154                            .cloned()
155                            .or_else(|| {
156                                self.runtime_conditional_routers
157                                    .try_read()
158                                    .ok()
159                                    .and_then(|guard| guard.get(router_name).cloned())
160                            })
161                            .ok_or_else(|| {
162                                GraphError::ExecutionError(format!(
163                                    "Router '{}' not found",
164                                    router_name
165                                ))
166                            })?;
167
168                        let route_key = router.route(state).await?;
169
170                        let target = targets
171                            .get(&route_key)
172                            .or(default_target.as_ref())
173                            .ok_or_else(|| {
174                                GraphError::RoutingError(format!(
175                                    "No target for route '{}'",
176                                    route_key
177                                ))
178                            })?;
179
180                        return Ok(target.clone());
181                    }
182                    GraphEdge::FanOut { targets, .. } => {
183                        if targets.is_empty() {
184                            return Err(GraphError::RoutingError(
185                                "FanOut has no targets".to_string(),
186                            ));
187                        }
188                        return Ok(targets[0].clone());
189                    }
190                    GraphEdge::FanIn { .. } => {
191                        continue;
192                    }
193                }
194            }
195        }
196
197        if current == self.entry_point && self.nodes.len() == 1 {
198            return Ok(END.to_string());
199        }
200
201        Err(GraphError::RoutingError(format!(
202            "No outgoing edge from node '{}'",
203            current
204        )))
205    }
206
207    pub(super) async fn find_fan_out_targets(&self, current: &str) -> Option<Vec<String>> {
208        {
209            let re = self.runtime_edges.read().await;
210            if let Some(GraphEdge::FanOut { targets, .. }) =
211                re.iter().find(|e| e.source() == current)
212            {
213                return Some(targets.clone());
214            }
215        }
216        for edge in &self.edges {
217            if edge.source() == current {
218                if let GraphEdge::FanOut { targets, .. } = edge {
219                    return Some(targets.clone());
220                }
221            }
222        }
223        None
224    }
225
226    pub(super) async fn find_fan_in_target(&self, sources: &[String]) -> Option<String> {
227        {
228            let re = self.runtime_edges.read().await;
229            if let Some(GraphEdge::FanIn {
230                sources: edge_sources,
231                target,
232            }) = re.iter().find(|e| matches!(e, GraphEdge::FanIn { .. }))
233            {
234                if edge_sources.iter().all(|s| sources.contains(s)) {
235                    return Some(target.clone());
236                }
237            }
238        }
239        for edge in &self.edges {
240            if let GraphEdge::FanIn {
241                sources: edge_sources,
242                target,
243            } = edge
244            {
245                if edge_sources.iter().all(|s| sources.contains(s)) {
246                    return Some(target.clone());
247                }
248            }
249        }
250        None
251    }
252
253    pub(super) async fn execute_parallel_branches(
254        &self,
255        targets: &[String],
256        state: &S,
257    ) -> GraphResult<Vec<(String, GraphInvocation<S>)>> {
258        let futures: Vec<_> = targets
259            .iter()
260            .filter(|t| *t != END)
261            .map(|target| {
262                let target = target.clone();
263                let state_clone = state.clone();
264                async move {
265                    let result = self.invoke_from_node(target.clone(), state_clone).await;
266                    result.map(|inv| (target, inv))
267                }
268            })
269            .collect();
270
271        let results = join_all(futures).await;
272
273        let mut successful = Vec::new();
274        for result in results {
275            match result {
276                Ok((name, inv)) => successful.push((name, inv)),
277                Err(e) => return Err(e),
278            }
279        }
280
281        Ok(successful)
282    }
283}