Skip to main content

lc_langgraph/
subgraph.rs

1// crates/lc-langgraph/src/subgraph.rs
2//! Subgraph Support for LangGraph
3//!
4//! Subgraphs allow nesting compiled graphs within nodes of a parent graph.
5//! This enables composition of complex workflows from simpler components.
6//!
7//! # Example
8//!
9//! ```rust,ignore
10//! use lc_langgraph::{StateGraph, CompiledGraph, SubgraphNode, START, END};
11//!
12//! // Create subgraph
13//! let subgraph = GraphBuilder::<AgentState>::new()
14//!     .add_node_fn("process", |state| Ok(StateUpdate::full(state.clone())))
15//!     .add_node_fn("output", |state| {
16//!         let mut s = state.clone();
17//!         s.set_output("done");
18//!         Ok(StateUpdate::full(s))
19//!     })
20//!     .add_edge(START, "process")
21//!     .add_edge("process", "output")
22//!     .add_edge("output", END)
23//!     .compile()?;
24//!
25//! // Add as subgraph node in parent graph
26//! let parent = GraphBuilder::<AgentState>::new()
27//!     .add_subgraph("subworkflow", subgraph,
28//!         |parent_state| parent_state.clone(),  // input mapper
29//!         |sub_state, parent_state| *parent_state = sub_state.clone()  // output mapper
30//!     )
31//!     .add_edge(START, "subworkflow")
32//!     .add_edge("subworkflow", END)
33//!     .compile()?;
34//! ```
35
36use crate::compiled::CompiledGraph;
37use crate::errors::{GraphError, GraphResult};
38use crate::node::{GraphNode, NodeConfig, NodeResult};
39use crate::state::{StateSchema, StateUpdate};
40use async_trait::async_trait;
41use std::marker::PhantomData;
42use std::sync::Arc;
43
44/// Subgraph Node - A node that executes a nested compiled graph
45///
46/// This allows composition of graphs by embedding one graph inside another.
47/// The subgraph receives mapped input state and returns mapped output state.
48#[allow(clippy::type_complexity)]
49pub struct SubgraphNode<S: StateSchema, SubS: StateSchema> {
50    name: String,
51    subgraph: CompiledGraph<SubS>,
52    input_mapper: Arc<dyn Fn(&S) -> SubS + Send + Sync>,
53    output_mapper: Arc<dyn Fn(&SubS, &mut S) + Send + Sync>,
54    _parent_marker: PhantomData<S>,
55    _sub_marker: PhantomData<SubS>,
56}
57
58impl<S: StateSchema, SubS: StateSchema> SubgraphNode<S, SubS> {
59    pub fn new(
60        name: impl Into<String>,
61        subgraph: CompiledGraph<SubS>,
62        input_mapper: impl Fn(&S) -> SubS + Send + Sync + 'static,
63        output_mapper: impl Fn(&SubS, &mut S) + Send + Sync + 'static,
64    ) -> Self {
65        Self {
66            name: name.into(),
67            subgraph,
68            input_mapper: Arc::new(input_mapper),
69            output_mapper: Arc::new(output_mapper),
70            _parent_marker: PhantomData,
71            _sub_marker: PhantomData,
72        }
73    }
74}
75
76impl<S: StateSchema + Clone> SubgraphNode<S, S> {
77    pub fn same_state(name: impl Into<String>, subgraph: CompiledGraph<S>) -> Self {
78        Self::new(
79            name,
80            subgraph,
81            |s| s.clone(),
82            |sub_s, parent_s| *parent_s = sub_s.clone(),
83        )
84    }
85}
86
87#[async_trait]
88impl<S: StateSchema + 'static, SubS: StateSchema + 'static> GraphNode<S> for SubgraphNode<S, SubS> {
89    async fn execute(&self, state: &S, _config: Option<NodeConfig>) -> NodeResult<S> {
90        // Map parent state to subgraph input
91        let sub_input = (self.input_mapper)(state);
92
93        // Execute subgraph
94        let sub_result = self.subgraph.invoke(sub_input).await.map_err(|e| {
95            GraphError::ExecutionError(format!("Subgraph '{}' execution failed: {}", self.name, e))
96        })?;
97
98        // Map subgraph output back to parent state
99        let mut parent_output = state.clone();
100        (self.output_mapper)(&sub_result.final_state, &mut parent_output);
101
102        // Include subgraph steps in metadata
103        let mut metadata = std::collections::HashMap::new();
104        metadata.insert(
105            "subgraph_steps".to_string(),
106            serde_json::json!(sub_result.steps.len()),
107        );
108        metadata.insert(
109            "subgraph_recursion".to_string(),
110            serde_json::json!(sub_result.recursion_count),
111        );
112
113        Ok(StateUpdate::with_metadata(parent_output, metadata))
114    }
115
116    fn name(&self) -> &str {
117        &self.name
118    }
119}
120
121/// Builder for creating subgraph nodes with fluent API
122#[allow(clippy::type_complexity)]
123pub struct SubgraphBuilder<S: StateSchema, SubS: StateSchema> {
124    name: String,
125    subgraph: Option<CompiledGraph<SubS>>,
126    input_mapper: Option<Arc<dyn Fn(&S) -> SubS + Send + Sync>>,
127    output_mapper: Option<Arc<dyn Fn(&SubS, &mut S) + Send + Sync>>,
128    _parent_marker: PhantomData<S>,
129    _sub_marker: PhantomData<SubS>,
130}
131
132impl<S: StateSchema, SubS: StateSchema> SubgraphBuilder<S, SubS> {
133    pub fn new(name: impl Into<String>) -> Self {
134        Self {
135            name: name.into(),
136            subgraph: None,
137            input_mapper: None,
138            output_mapper: None,
139            _parent_marker: PhantomData,
140            _sub_marker: PhantomData,
141        }
142    }
143
144    pub fn subgraph(mut self, graph: CompiledGraph<SubS>) -> Self {
145        self.subgraph = Some(graph);
146        self
147    }
148
149    pub fn input_mapper(mut self, mapper: impl Fn(&S) -> SubS + Send + Sync + 'static) -> Self {
150        self.input_mapper = Some(Arc::new(mapper));
151        self
152    }
153
154    pub fn output_mapper(mut self, mapper: impl Fn(&SubS, &mut S) + Send + Sync + 'static) -> Self {
155        self.output_mapper = Some(Arc::new(mapper));
156        self
157    }
158
159    pub fn build(self) -> GraphResult<SubgraphNode<S, SubS>> {
160        let subgraph = self
161            .subgraph
162            .ok_or_else(|| GraphError::ValidationError("Subgraph not set".to_string()))?;
163        let input_mapper = self
164            .input_mapper
165            .ok_or_else(|| GraphError::ValidationError("Input mapper not set".to_string()))?;
166        let output_mapper = self
167            .output_mapper
168            .ok_or_else(|| GraphError::ValidationError("Output mapper not set".to_string()))?;
169
170        Ok(SubgraphNode {
171            name: self.name,
172            subgraph,
173            input_mapper,
174            output_mapper,
175            _parent_marker: PhantomData,
176            _sub_marker: PhantomData,
177        })
178    }
179}
180
181#[cfg(test)]
182mod tests {
183    use super::*;
184    use crate::graph::GraphBuilder;
185    use crate::state::AgentState;
186    use crate::{END, START};
187
188    #[tokio::test]
189    async fn test_subgraph_same_state() {
190        // Create simple subgraph
191        let subgraph = GraphBuilder::<AgentState>::new()
192            .add_node_fn("sub_process", |state| {
193                let mut s = state.clone();
194                s.set_output("subgraph_output".to_string());
195                Ok(StateUpdate::full(s))
196            })
197            .add_edge(START, "sub_process")
198            .add_edge("sub_process", END)
199            .compile()
200            .unwrap();
201
202        // Create parent graph with subgraph
203        let parent = GraphBuilder::<AgentState>::new()
204            .add_subgraph_same_state("subworkflow", subgraph)
205            .add_edge(START, "subworkflow")
206            .add_edge("subworkflow", END)
207            .compile()
208            .unwrap();
209
210        let input = AgentState::new("test".to_string());
211        let result = parent.invoke(input).await.unwrap();
212
213        assert!(result.final_state.output.is_some());
214        assert_eq!(result.final_state.output.unwrap(), "subgraph_output");
215    }
216
217    #[tokio::test]
218    async fn test_nested_subgraphs() {
219        // Create innermost subgraph
220        let inner = GraphBuilder::<AgentState>::new()
221            .add_node_fn("inner_node", |state| {
222                let mut s = state.clone();
223                s.input = format!("inner:{}", s.input);
224                Ok(StateUpdate::full(s))
225            })
226            .add_edge(START, "inner_node")
227            .add_edge("inner_node", END)
228            .compile()
229            .unwrap();
230
231        // Create middle subgraph containing inner
232        let middle = GraphBuilder::<AgentState>::new()
233            .add_subgraph_same_state("inner_workflow", inner)
234            .add_node_fn("middle_node", |state| {
235                let mut s = state.clone();
236                s.input = format!("middle:{}", s.input);
237                Ok(StateUpdate::full(s))
238            })
239            .add_edge(START, "inner_workflow")
240            .add_edge("inner_workflow", "middle_node")
241            .add_edge("middle_node", END)
242            .compile()
243            .unwrap();
244
245        // Create outer graph with middle subgraph
246        let outer = GraphBuilder::<AgentState>::new()
247            .add_node_fn("outer_node", |state| {
248                let mut s = state.clone();
249                s.input = format!("outer:{}", s.input);
250                Ok(StateUpdate::full(s))
251            })
252            .add_subgraph_same_state("middle_workflow", middle)
253            .add_edge(START, "outer_node")
254            .add_edge("outer_node", "middle_workflow")
255            .add_edge("middle_workflow", END)
256            .compile()
257            .unwrap();
258
259        let input = AgentState::new("test".to_string());
260        let result = outer.invoke(input).await.unwrap();
261
262        // Verify nested execution: outer:middle:inner:test
263        assert!(result.final_state.input.contains("outer"));
264        assert!(result.final_state.input.contains("middle"));
265        assert!(result.final_state.input.contains("inner"));
266    }
267}