1use 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#[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(
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 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 let sub_input = (self.input_mapper)(state);
94
95 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 let mut parent_output = state.clone();
102 (self.output_mapper)(&sub_result.final_state, &mut parent_output);
103
104 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#[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 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 pub fn subgraph(mut self, graph: CompiledGraph<SubS>) -> Self {
149 self.subgraph = Some(graph);
150 self
151 }
152
153 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 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 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 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 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 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 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 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 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}