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(
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 let sub_input = (self.input_mapper)(state);
92
93 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 let mut parent_output = state.clone();
100 (self.output_mapper)(&sub_result.final_state, &mut parent_output);
101
102 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#[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 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 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 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 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 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 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}