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    /// Create a new subgraph node with the given name, subgraph, and mappers.
60    pub fn new(
61        name: impl Into<String>,
62        subgraph: CompiledGraph<SubS>,
63        input_mapper: impl Fn(&S) -> SubS + Send + Sync + 'static,
64        output_mapper: impl Fn(&SubS, &mut S) + Send + Sync + 'static,
65    ) -> Self {
66        Self {
67            name: name.into(),
68            subgraph,
69            input_mapper: Arc::new(input_mapper),
70            output_mapper: Arc::new(output_mapper),
71            _parent_marker: PhantomData,
72            _sub_marker: PhantomData,
73        }
74    }
75}
76
77impl<S: StateSchema + Clone> SubgraphNode<S, S> {
78    /// Create a subgraph node sharing the same state type, using clone-based mappers.
79    pub fn same_state(name: impl Into<String>, subgraph: CompiledGraph<S>) -> Self {
80        Self::new(
81            name,
82            subgraph,
83            |s| s.clone(),
84            |sub_s, parent_s| *parent_s = sub_s.clone(),
85        )
86    }
87}
88
89#[async_trait]
90impl<S: StateSchema + 'static, SubS: StateSchema + 'static> GraphNode<S> for SubgraphNode<S, SubS> {
91    async fn execute(&self, state: &S, _config: Option<NodeConfig>) -> NodeResult<S> {
92        // Map parent state to subgraph input
93        let sub_input = (self.input_mapper)(state);
94
95        // Execute subgraph
96        let sub_result = self.subgraph.invoke(sub_input).await.map_err(|e| {
97            GraphError::ExecutionError(format!("Subgraph '{}' execution failed: {}", self.name, e))
98        })?;
99
100        // Map subgraph output back to parent state
101        let mut parent_output = state.clone();
102        (self.output_mapper)(&sub_result.final_state, &mut parent_output);
103
104        // Include subgraph steps in metadata
105        let mut metadata = std::collections::HashMap::new();
106        metadata.insert(
107            "subgraph_steps".to_string(),
108            serde_json::json!(sub_result.steps.len()),
109        );
110        metadata.insert(
111            "subgraph_recursion".to_string(),
112            serde_json::json!(sub_result.recursion_count),
113        );
114
115        Ok(StateUpdate::with_metadata(parent_output, metadata))
116    }
117
118    fn name(&self) -> &str {
119        &self.name
120    }
121}
122
123/// Builder for creating subgraph nodes with fluent API
124#[allow(clippy::type_complexity)]
125pub struct SubgraphBuilder<S: StateSchema, SubS: StateSchema> {
126    name: String,
127    subgraph: Option<CompiledGraph<SubS>>,
128    input_mapper: Option<Arc<dyn Fn(&S) -> SubS + Send + Sync>>,
129    output_mapper: Option<Arc<dyn Fn(&SubS, &mut S) + Send + Sync>>,
130    _parent_marker: PhantomData<S>,
131    _sub_marker: PhantomData<SubS>,
132}
133
134impl<S: StateSchema, SubS: StateSchema> SubgraphBuilder<S, SubS> {
135    /// Create a new subgraph builder with the given node name.
136    pub fn new(name: impl Into<String>) -> Self {
137        Self {
138            name: name.into(),
139            subgraph: None,
140            input_mapper: None,
141            output_mapper: None,
142            _parent_marker: PhantomData,
143            _sub_marker: PhantomData,
144        }
145    }
146
147    /// Set the compiled subgraph to execute.
148    pub fn subgraph(mut self, graph: CompiledGraph<SubS>) -> Self {
149        self.subgraph = Some(graph);
150        self
151    }
152
153    /// Set the mapper from parent state to subgraph input.
154    pub fn input_mapper(mut self, mapper: impl Fn(&S) -> SubS + Send + Sync + 'static) -> Self {
155        self.input_mapper = Some(Arc::new(mapper));
156        self
157    }
158
159    /// Set the mapper from subgraph output back to parent state.
160    pub fn output_mapper(mut self, mapper: impl Fn(&SubS, &mut S) + Send + Sync + 'static) -> Self {
161        self.output_mapper = Some(Arc::new(mapper));
162        self
163    }
164
165    /// Build the `SubgraphNode`, failing if any required component is missing.
166    pub fn build(self) -> GraphResult<SubgraphNode<S, SubS>> {
167        let subgraph = self
168            .subgraph
169            .ok_or_else(|| GraphError::ValidationError("Subgraph not set".to_string()))?;
170        let input_mapper = self
171            .input_mapper
172            .ok_or_else(|| GraphError::ValidationError("Input mapper not set".to_string()))?;
173        let output_mapper = self
174            .output_mapper
175            .ok_or_else(|| GraphError::ValidationError("Output mapper not set".to_string()))?;
176
177        Ok(SubgraphNode {
178            name: self.name,
179            subgraph,
180            input_mapper,
181            output_mapper,
182            _parent_marker: PhantomData,
183            _sub_marker: PhantomData,
184        })
185    }
186}
187
188#[cfg(test)]
189mod tests {
190    use super::*;
191    use crate::graph::GraphBuilder;
192    use crate::state::AgentState;
193    use crate::{END, START};
194
195    #[tokio::test]
196    async fn test_subgraph_same_state() {
197        // Create simple subgraph
198        let subgraph = GraphBuilder::<AgentState>::new()
199            .add_node_fn("sub_process", |state| {
200                let mut s = state.clone();
201                s.set_output("subgraph_output".to_string());
202                Ok(StateUpdate::full(s))
203            })
204            .add_edge(START, "sub_process")
205            .add_edge("sub_process", END)
206            .compile()
207            .unwrap();
208
209        // Create parent graph with subgraph
210        let parent = GraphBuilder::<AgentState>::new()
211            .add_subgraph_same_state("subworkflow", subgraph)
212            .add_edge(START, "subworkflow")
213            .add_edge("subworkflow", END)
214            .compile()
215            .unwrap();
216
217        let input = AgentState::new("test".to_string());
218        let result = parent.invoke(input).await.unwrap();
219
220        assert!(result.final_state.output.is_some());
221        assert_eq!(result.final_state.output.unwrap(), "subgraph_output");
222    }
223
224    #[tokio::test]
225    async fn test_nested_subgraphs() {
226        // Create innermost subgraph
227        let inner = GraphBuilder::<AgentState>::new()
228            .add_node_fn("inner_node", |state| {
229                let mut s = state.clone();
230                s.input = format!("inner:{}", s.input);
231                Ok(StateUpdate::full(s))
232            })
233            .add_edge(START, "inner_node")
234            .add_edge("inner_node", END)
235            .compile()
236            .unwrap();
237
238        // Create middle subgraph containing inner
239        let middle = GraphBuilder::<AgentState>::new()
240            .add_subgraph_same_state("inner_workflow", inner)
241            .add_node_fn("middle_node", |state| {
242                let mut s = state.clone();
243                s.input = format!("middle:{}", s.input);
244                Ok(StateUpdate::full(s))
245            })
246            .add_edge(START, "inner_workflow")
247            .add_edge("inner_workflow", "middle_node")
248            .add_edge("middle_node", END)
249            .compile()
250            .unwrap();
251
252        // Create outer graph with middle subgraph
253        let outer = GraphBuilder::<AgentState>::new()
254            .add_node_fn("outer_node", |state| {
255                let mut s = state.clone();
256                s.input = format!("outer:{}", s.input);
257                Ok(StateUpdate::full(s))
258            })
259            .add_subgraph_same_state("middle_workflow", middle)
260            .add_edge(START, "outer_node")
261            .add_edge("outer_node", "middle_workflow")
262            .add_edge("middle_workflow", END)
263            .compile()
264            .unwrap();
265
266        let input = AgentState::new("test".to_string());
267        let result = outer.invoke(input).await.unwrap();
268
269        // Verify nested execution: outer:middle:inner:test
270        assert!(result.final_state.input.contains("outer"));
271        assert!(result.final_state.input.contains("middle"));
272        assert!(result.final_state.input.contains("inner"));
273    }
274}