1use std::collections::{HashMap, HashSet};
2use std::sync::Arc;
3
4use adk_core::Agent;
5use schemars::JsonSchema;
6use serde::{Deserialize, Serialize};
7
8use crate::{LoopAgent, ParallelAgent, SequentialAgent};
9
10use super::{RelationshipKind, TeamError, TeamMemberSpec, TeamPolicy, TeamRelationship, TeamSpec};
11
12#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize, JsonSchema)]
17#[serde(rename_all = "camelCase", tag = "architecture")]
18pub enum TeamArchitectureTemplate {
19 Supervisor {
21 name: String,
23 coordinator: String,
25 specialists: Vec<String>,
27 relationship: RelationshipKind,
29 #[serde(default)]
31 policy: TeamPolicy,
32 },
33 Router {
35 name: String,
37 router: String,
39 routes: Vec<String>,
41 relationship: RelationshipKind,
43 #[serde(default)]
45 policy: TeamPolicy,
46 },
47 Hierarchical {
49 name: String,
51 root: String,
53 branches: Vec<TeamManagerBranch>,
55 #[serde(default)]
57 policy: TeamPolicy,
58 },
59}
60
61#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize, JsonSchema)]
63#[serde(rename_all = "camelCase", deny_unknown_fields)]
64pub struct TeamManagerBranch {
65 pub manager: String,
67 pub workers: Vec<String>,
69}
70
71impl TeamArchitectureTemplate {
72 pub fn lower(&self) -> std::result::Result<TeamSpec, TeamError> {
74 let spec = match self {
75 Self::Supervisor { name, coordinator, specialists, relationship, policy } => TeamSpec {
76 name: name.clone(),
77 description: "Supervisor team compiled from a portable architecture template"
78 .to_string(),
79 coordinator: coordinator.clone(),
80 members: std::iter::once(coordinator)
81 .chain(specialists)
82 .map(|member| TeamMemberSpec::new(member.clone()))
83 .collect(),
84 relationships: specialists
85 .iter()
86 .map(|specialist| TeamRelationship::new(coordinator, specialist, *relationship))
87 .collect(),
88 policy: policy.clone(),
89 },
90 Self::Router { name, router, routes, relationship, policy } => TeamSpec {
91 name: name.clone(),
92 description: "Exact router team compiled from a portable architecture template"
93 .to_string(),
94 coordinator: router.clone(),
95 members: std::iter::once(router)
96 .chain(routes)
97 .map(|member| TeamMemberSpec::new(member.clone()))
98 .collect(),
99 relationships: routes
100 .iter()
101 .map(|route| TeamRelationship::new(router, route, *relationship))
102 .collect(),
103 policy: policy.clone(),
104 },
105 Self::Hierarchical { name, root, branches, policy } => {
106 let mut members = vec![TeamMemberSpec::new(root.clone())];
107 let mut relationships = Vec::new();
108 for branch in branches {
109 members.push(TeamMemberSpec::new(branch.manager.clone()));
110 relationships.push(TeamRelationship::new(
111 root,
112 &branch.manager,
113 RelationshipKind::Delegate,
114 ));
115 for worker in &branch.workers {
116 members.push(TeamMemberSpec::new(worker.clone()));
117 relationships.push(TeamRelationship::new(
118 &branch.manager,
119 worker,
120 RelationshipKind::Delegate,
121 ));
122 }
123 }
124 TeamSpec {
125 name: name.clone(),
126 description: "Hierarchical team compiled from a portable architecture template"
127 .to_string(),
128 coordinator: root.clone(),
129 members,
130 relationships,
131 policy: policy.clone(),
132 }
133 }
134 };
135 spec.validate()?;
136 Ok(spec)
137 }
138}
139
140#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize, JsonSchema)]
142#[serde(rename_all = "camelCase", tag = "architecture")]
143pub enum WorkflowArchitectureTemplate {
144 Sequential {
146 name: String,
148 steps: Vec<String>,
150 },
151 Parallel {
153 name: String,
155 members: Vec<String>,
157 #[serde(default)]
159 shared_state: bool,
160 },
161 FanOutFanIn {
163 name: String,
165 workers: Vec<String>,
167 aggregator: String,
169 #[serde(default)]
171 shared_state: bool,
172 },
173 ReviewLoop {
175 name: String,
177 producer: String,
179 reviewer: String,
181 max_iterations: u32,
183 },
184}
185
186impl WorkflowArchitectureTemplate {
187 pub fn compile(
189 &self,
190 agents: impl IntoIterator<Item = Arc<dyn Agent>>,
191 ) -> std::result::Result<Arc<dyn Agent>, TeamError> {
192 let registry: HashMap<String, Arc<dyn Agent>> =
193 agents.into_iter().map(|agent| (agent.name().to_string(), agent)).collect();
194 let resolve = |names: &[String]| -> std::result::Result<Vec<Arc<dyn Agent>>, TeamError> {
195 names
196 .iter()
197 .map(|name| {
198 registry.get(name).cloned().ok_or_else(|| TeamError::MissingAgent(name.clone()))
199 })
200 .collect()
201 };
202 let ensure_unique = |names: &[String]| -> std::result::Result<(), TeamError> {
203 let mut seen = HashSet::new();
204 for name in names {
205 if !seen.insert(name) {
206 return Err(TeamError::DuplicateMember(name.clone()));
207 }
208 }
209 Ok(())
210 };
211 match self {
212 Self::Sequential { name, steps } => {
213 ensure_unique(steps)?;
214 Ok(Arc::new(SequentialAgent::new(name, resolve(steps)?)))
215 }
216 Self::Parallel { name, members, shared_state } => {
217 ensure_unique(members)?;
218 let parallel = ParallelAgent::new(name, resolve(members)?);
219 Ok(Arc::new(if *shared_state { parallel.with_shared_state() } else { parallel }))
220 }
221 Self::FanOutFanIn { name, workers, aggregator, shared_state } => {
222 ensure_unique(workers)?;
223 if workers.iter().any(|worker| worker == aggregator) {
224 return Err(TeamError::DuplicateMember(aggregator.clone()));
225 }
226 let parallel = ParallelAgent::new(format!("{name}_fan_out"), resolve(workers)?);
227 let parallel = if *shared_state { parallel.with_shared_state() } else { parallel };
228 let aggregator = registry
229 .get(aggregator)
230 .cloned()
231 .ok_or_else(|| TeamError::MissingAgent(aggregator.clone()))?;
232 Ok(Arc::new(SequentialAgent::new(name, vec![Arc::new(parallel), aggregator])))
233 }
234 Self::ReviewLoop { name, producer, reviewer, max_iterations } => {
235 if *max_iterations == 0 {
236 return Err(TeamError::InvalidPolicy("reviewLoop.maxIterations"));
237 }
238 if producer == reviewer {
239 return Err(TeamError::DuplicateMember(producer.clone()));
240 }
241 let members = resolve(&[producer.clone(), reviewer.clone()])?;
242 Ok(Arc::new(LoopAgent::new(name, members).with_max_iterations(*max_iterations)))
243 }
244 }
245 }
246}
247
248#[cfg(test)]
249mod tests {
250 use super::*;
251
252 struct NoopAgent(String);
253
254 #[async_trait::async_trait]
255 impl Agent for NoopAgent {
256 fn name(&self) -> &str {
257 &self.0
258 }
259
260 fn description(&self) -> &str {
261 "workflow template test agent"
262 }
263
264 fn sub_agents(&self) -> &[Arc<dyn Agent>] {
265 &[]
266 }
267
268 async fn run(
269 &self,
270 _ctx: Arc<dyn adk_core::InvocationContext>,
271 ) -> adk_core::Result<adk_core::EventStream> {
272 Ok(Box::pin(futures::stream::empty()))
273 }
274 }
275
276 #[test]
277 fn supervisor_template_lowers_to_exact_edges() {
278 let spec = TeamArchitectureTemplate::Supervisor {
279 name: "support".to_string(),
280 coordinator: "supervisor".to_string(),
281 specialists: vec!["billing".to_string(), "technical".to_string()],
282 relationship: RelationshipKind::Handoff,
283 policy: TeamPolicy::default(),
284 }
285 .lower()
286 .unwrap();
287 assert_eq!(spec.members.len(), 3);
288 assert_eq!(spec.relationships.len(), 2);
289 assert!(spec.relationships.iter().all(|edge| edge.from == "supervisor"));
290 }
291
292 #[test]
293 fn hierarchical_template_rejects_duplicate_workers() {
294 let error = TeamArchitectureTemplate::Hierarchical {
295 name: "org".to_string(),
296 root: "director".to_string(),
297 branches: vec![
298 TeamManagerBranch {
299 manager: "research".to_string(),
300 workers: vec!["analyst".to_string()],
301 },
302 TeamManagerBranch {
303 manager: "review".to_string(),
304 workers: vec!["analyst".to_string()],
305 },
306 ],
307 policy: TeamPolicy::default(),
308 }
309 .lower()
310 .unwrap_err();
311 assert_eq!(error, TeamError::DuplicateMember("analyst".to_string()));
312 }
313
314 #[test]
315 fn workflow_template_compiles_to_existing_primitives() {
316 let root = WorkflowArchitectureTemplate::FanOutFanIn {
317 name: "research".to_string(),
318 workers: vec!["facts".to_string(), "risks".to_string()],
319 aggregator: "reviewer".to_string(),
320 shared_state: true,
321 }
322 .compile([
323 Arc::new(NoopAgent("facts".to_string())) as Arc<dyn Agent>,
324 Arc::new(NoopAgent("risks".to_string())),
325 Arc::new(NoopAgent("reviewer".to_string())),
326 ])
327 .unwrap();
328 assert_eq!(root.name(), "research");
329 assert_eq!(root.sub_agents().len(), 2);
330 assert_eq!(root.sub_agents()[0].name(), "research_fan_out");
331 }
332}