lc_langgraph/compiled/
stream.rs1use 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 futures_util::Stream;
14use std::collections::HashMap;
15use std::pin::Pin;
16use tokio_stream::wrappers::ReceiverStream;
17
18impl<S: StateSchema + Send + Sync + 'static> CompiledGraph<S> {
19 pub fn stream(
27 &self,
28 input: S,
29 ) -> Pin<Box<dyn Stream<Item = Result<StreamEvent<S>, GraphError>> + Send>> {
30 let (tx, rx) = tokio::sync::mpsc::channel(64);
31
32 let graph = self.clone();
33 tokio::spawn(async move {
34 let mut state = input;
35 let mut current_node = graph.entry_point.clone();
36 let mut recursion_count = 0;
37
38 if tx
39 .send(Ok(StreamEvent::start(state.clone())))
40 .await
41 .is_err()
42 {
43 return;
44 }
45
46 if let Some(ref checkpointer) = graph.checkpointer {
47 match checkpointer.lock().await.save(&state).await {
48 Ok(_) => {}
49 Err(e) => {
50 let _ = tx.send(Err(e)).await;
51 return;
52 }
53 }
54 }
55
56 while current_node != END && recursion_count < graph.recursion_limit {
57 if graph.interrupt_before.contains(¤t_node) {
58 let _ = tx
59 .send(Err(GraphError::ExecutionInterrupted(current_node.clone())))
60 .await;
61 return;
62 }
63
64 recursion_count += 1;
65
66 if tx
67 .send(Ok(StreamEvent::enter_node(
68 current_node.clone(),
69 state.clone(),
70 )))
71 .await
72 .is_err()
73 {
74 return;
75 }
76
77 let node = match graph.get_node(¤t_node).await {
78 Ok(n) => n,
79 Err(e) => {
80 let _ = tx.send(Err(e)).await;
81 return;
82 }
83 };
84
85 let config = NodeConfig {
86 recursion_limit: graph.recursion_limit,
87 debug: false,
88 metadata: HashMap::new(),
89 };
90
91 let update = match node.execute(&state, Some(config)).await {
92 Ok(u) => u,
93 Err(e) => {
94 let _ = tx.send(Err(e)).await;
95 return;
96 }
97 };
98
99 if tx
100 .send(Ok(StreamEvent::node_complete(
101 current_node.clone(),
102 update.clone(),
103 )))
104 .await
105 .is_err()
106 {
107 return;
108 }
109
110 if let Some(new_state) = update.update {
111 state = graph.default_reducer.reduce(&state, &new_state);
112 if tx
113 .send(Ok(StreamEvent::state_update(state.clone())))
114 .await
115 .is_err()
116 {
117 return;
118 }
119 }
120
121 if graph.interrupt_after.contains(¤t_node) {
122 if let Some(ref checkpointer) = graph.checkpointer {
123 let _ = checkpointer.lock().await.save(&state).await;
124 }
125 let _ = tx
126 .send(Err(GraphError::ExecutionInterrupted(format!(
127 "after_{}",
128 current_node
129 ))))
130 .await;
131 return;
132 }
133
134 let next_node = match graph.find_next_node(¤t_node, &state).await {
135 Ok(n) => n,
136 Err(e) => {
137 let _ = tx.send(Err(e)).await;
138 return;
139 }
140 };
141
142 if let Some(ref checkpointer) = graph.checkpointer {
143 let _ = checkpointer.lock().await.save(&state).await;
144 }
145
146 current_node = next_node;
147 }
148
149 let _ = tx.send(Ok(StreamEvent::end(state))).await;
150 });
151
152 Box::pin(ReceiverStream::new(rx))
153 }
154
155 pub async fn stream_collected(&self, input: S) -> GraphResult<Vec<StreamEvent<S>>> {
160 let mut events = Vec::new();
161 let mut stream = self.stream(input);
162 use futures_util::StreamExt;
163 while let Some(event) = stream.next().await {
164 events.push(event?);
165 }
166 Ok(events)
167 }
168
169 pub(super) async fn find_next_node(&self, current: &str, state: &S) -> GraphResult<String> {
170 'rt: {
172 let edge = {
173 let re = self.runtime_edges.read().await;
174 match re.iter().find(|e| e.source() == current) {
175 Some(e) => e.clone(),
176 None => break 'rt,
177 }
178 }; match edge {
180 GraphEdge::Fixed { target, .. } => return Ok(target),
181 GraphEdge::Conditional {
182 router_name,
183 targets,
184 default_target,
185 ..
186 } => {
187 let router = self
188 .conditional_routers
189 .get(&router_name)
190 .cloned()
191 .or_else(|| {
192 self.runtime_conditional_routers
193 .try_read()
194 .ok()
195 .and_then(|guard| guard.get(&router_name).cloned())
196 })
197 .ok_or_else(|| {
198 GraphError::ExecutionError(format!(
199 "Router '{}' not found (runtime)",
200 router_name
201 ))
202 })?;
203 let route_key = router.route(state).await?;
204 let target = targets
205 .get(&route_key)
206 .or(default_target.as_ref())
207 .ok_or_else(|| {
208 GraphError::RoutingError(format!(
209 "No target for route '{}' (runtime)",
210 route_key
211 ))
212 })?;
213 return Ok(target.clone());
214 }
215 GraphEdge::FanOut { targets, .. } => {
216 if targets.is_empty() {
217 return Err(GraphError::RoutingError(
218 "FanOut has no targets (runtime)".to_string(),
219 ));
220 }
221 return Ok(targets[0].clone());
222 }
223 GraphEdge::FanIn { .. } => {}
224 }
225 }
226
227 for edge in &self.edges {
229 if edge.source() == current {
230 match edge {
231 GraphEdge::Fixed { target, .. } => {
232 return Ok(target.clone());
233 }
234 GraphEdge::Conditional {
235 router_name,
236 targets,
237 default_target,
238 ..
239 } => {
240 let router = self
241 .conditional_routers
242 .get(router_name)
243 .cloned()
244 .or_else(|| {
245 self.runtime_conditional_routers
246 .try_read()
247 .ok()
248 .and_then(|guard| guard.get(router_name).cloned())
249 })
250 .ok_or_else(|| {
251 GraphError::ExecutionError(format!(
252 "Router '{}' not found",
253 router_name
254 ))
255 })?;
256
257 let route_key = router.route(state).await?;
258
259 let target = targets
260 .get(&route_key)
261 .or(default_target.as_ref())
262 .ok_or_else(|| {
263 GraphError::RoutingError(format!(
264 "No target for route '{}'",
265 route_key
266 ))
267 })?;
268
269 return Ok(target.clone());
270 }
271 GraphEdge::FanOut { targets, .. } => {
272 if targets.is_empty() {
273 return Err(GraphError::RoutingError(
274 "FanOut has no targets".to_string(),
275 ));
276 }
277 return Ok(targets[0].clone());
278 }
279 GraphEdge::FanIn { .. } => {
280 continue;
281 }
282 }
283 }
284 }
285
286 if current == self.entry_point && self.nodes.len() == 1 {
287 return Ok(END.to_string());
288 }
289
290 Err(GraphError::RoutingError(format!(
291 "No outgoing edge from node '{}'",
292 current
293 )))
294 }
295
296 pub(super) async fn find_fan_out_targets(&self, current: &str) -> Option<Vec<String>> {
297 {
298 let re = self.runtime_edges.read().await;
299 if let Some(GraphEdge::FanOut { targets, .. }) =
300 re.iter().find(|e| e.source() == current)
301 {
302 return Some(targets.clone());
303 }
304 }
305 for edge in &self.edges {
306 if edge.source() == current {
307 if let GraphEdge::FanOut { targets, .. } = edge {
308 return Some(targets.clone());
309 }
310 }
311 }
312 None
313 }
314
315 pub(super) async fn find_fan_in_target(&self, sources: &[String]) -> Option<String> {
316 {
317 let re = self.runtime_edges.read().await;
318 if let Some(GraphEdge::FanIn {
319 sources: edge_sources,
320 target,
321 }) = re.iter().find(|e| matches!(e, GraphEdge::FanIn { .. }))
322 {
323 if edge_sources.iter().all(|s| sources.contains(s)) {
324 return Some(target.clone());
325 }
326 }
327 }
328 for edge in &self.edges {
329 if let GraphEdge::FanIn {
330 sources: edge_sources,
331 target,
332 } = edge
333 {
334 if edge_sources.iter().all(|s| sources.contains(s)) {
335 return Some(target.clone());
336 }
337 }
338 }
339 None
340 }
341
342 pub(super) async fn execute_parallel_branches(
343 &self,
344 targets: &[String],
345 state: &S,
346 ) -> GraphResult<Vec<(String, GraphInvocation<S>)>> {
347 let futures: Vec<_> = targets
348 .iter()
349 .filter(|t| *t != END)
350 .map(|target| {
351 let target = target.clone();
352 let state_clone = state.clone();
353 async move {
354 let result = self.invoke_from_node(target.clone(), state_clone).await;
355 result.map(|inv| (target, inv))
356 }
357 })
358 .collect();
359
360 let results = join_all(futures).await;
361
362 let mut successful = Vec::new();
363 for result in results {
364 match result {
365 Ok((name, inv)) => successful.push((name, inv)),
366 Err(e) => return Err(e),
367 }
368 }
369
370 Ok(successful)
371 }
372}