Skip to main content

adk_agent/team/
templates.rs

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/// Portable LLM-directed team architecture presets.
13///
14/// Presets lower to ordinary [`TeamSpec`] values; they are not new atomic
15/// agent kinds and remain inspectable and editable after lowering.
16#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize, JsonSchema)]
17#[serde(rename_all = "camelCase", tag = "architecture")]
18pub enum TeamArchitectureTemplate {
19    /// One coordinator invokes or transfers to a flat specialist roster.
20    Supervisor {
21        /// Team root name.
22        name: String,
23        /// Coordinator member name.
24        coordinator: String,
25        /// Specialist member names.
26        specialists: Vec<String>,
27        /// Relationship semantics for each specialist edge.
28        relationship: RelationshipKind,
29        /// Runtime policy.
30        #[serde(default)]
31        policy: TeamPolicy,
32    },
33    /// One router dispatches to exact route targets.
34    Router {
35        /// Team root name.
36        name: String,
37        /// Router member name.
38        router: String,
39        /// Exact route targets.
40        routes: Vec<String>,
41        /// Whether dispatch delegates and returns or hands off control.
42        relationship: RelationshipKind,
43        /// Runtime policy.
44        #[serde(default)]
45        policy: TeamPolicy,
46    },
47    /// Root manager delegates through named branch managers to their workers.
48    Hierarchical {
49        /// Team root name.
50        name: String,
51        /// Root manager member name.
52        root: String,
53        /// Manager branches and their exact workers.
54        branches: Vec<TeamManagerBranch>,
55        /// Runtime policy.
56        #[serde(default)]
57        policy: TeamPolicy,
58    },
59}
60
61/// One manager branch in a hierarchical architecture.
62#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize, JsonSchema)]
63#[serde(rename_all = "camelCase", deny_unknown_fields)]
64pub struct TeamManagerBranch {
65    /// Manager member name.
66    pub manager: String,
67    /// Exact worker member names.
68    pub workers: Vec<String>,
69}
70
71impl TeamArchitectureTemplate {
72    /// Lowers this preset to a validated, fully explicit [`TeamSpec`].
73    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/// Portable deterministic workflow presets built from existing agent types.
141#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize, JsonSchema)]
142#[serde(rename_all = "camelCase", tag = "architecture")]
143pub enum WorkflowArchitectureTemplate {
144    /// Execute members once in order.
145    Sequential {
146        /// Workflow root name.
147        name: String,
148        /// Ordered member names.
149        steps: Vec<String>,
150    },
151    /// Execute members concurrently.
152    Parallel {
153        /// Workflow root name.
154        name: String,
155        /// Concurrent member names.
156        members: Vec<String>,
157        /// Share a concurrency-safe blackboard between members.
158        #[serde(default)]
159        shared_state: bool,
160    },
161    /// Fan out to workers in parallel, then run an aggregator.
162    FanOutFanIn {
163        /// Workflow root name.
164        name: String,
165        /// Parallel worker names.
166        workers: Vec<String>,
167        /// Aggregator member name.
168        aggregator: String,
169        /// Share state between fan-out workers.
170        #[serde(default)]
171        shared_state: bool,
172    },
173    /// Alternate producer and reviewer for a bounded number of iterations.
174    ReviewLoop {
175        /// Workflow root name.
176        name: String,
177        /// Producer member name.
178        producer: String,
179        /// Reviewer member name.
180        reviewer: String,
181        /// Hard iteration bound.
182        max_iterations: u32,
183    },
184}
185
186impl WorkflowArchitectureTemplate {
187    /// Binds member names and compiles to existing deterministic agent primitives.
188    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}