use super::agent::LLMAgent;
use super::types::{LLMError, LLMResult};
use std::collections::HashMap;
use std::sync::Arc;
#[derive(Debug, Clone)]
pub enum TeamPattern {
Chain,
Parallel,
Debate {
max_rounds: usize,
},
Supervised,
MapReduce,
Custom,
}
#[derive(Debug, Clone)]
pub struct AgentRole {
pub id: String,
pub name: String,
pub description: String,
pub prompt_template: Option<String>,
}
impl AgentRole {
pub fn new(id: impl Into<String>, name: impl Into<String>) -> Self {
Self {
id: id.into(),
name: name.into(),
description: String::new(),
prompt_template: None,
}
}
pub fn with_description(mut self, desc: impl Into<String>) -> Self {
self.description = desc.into();
self
}
pub fn with_template(mut self, template: impl Into<String>) -> Self {
self.prompt_template = Some(template.into());
self
}
}
pub struct AgentMember {
pub role: AgentRole,
pub agent: Arc<LLMAgent>,
}
impl AgentMember {
pub fn new(id: impl Into<String>, agent: Arc<LLMAgent>) -> Self {
let id = id.into();
Self {
role: AgentRole::new(&id, &id),
agent,
}
}
pub fn with_role(mut self, role: AgentRole) -> Self {
self.role = role;
self
}
pub async fn execute(&self, input: &str, context: Option<&str>) -> LLMResult<String> {
let prompt = if let Some(ref template) = self.role.prompt_template {
let mut p = template.replace("{input}", input);
if let Some(ctx) = context {
p = p.replace("{context}", ctx);
}
p
} else if let Some(ctx) = context {
format!("Context:\n{}\n\nTask:\n{}", ctx, input)
} else {
input.to_string()
};
self.agent.ask(&prompt).await
}
}
pub struct AgentTeam {
pub id: String,
pub name: String,
members: Vec<AgentMember>,
member_map: HashMap<String, usize>,
pattern: TeamPattern,
supervisor_id: Option<String>,
aggregate_prompt: Option<String>,
}
impl AgentTeam {
pub fn new(id: impl Into<String>) -> AgentTeamBuilder {
AgentTeamBuilder::new(id)
}
async fn run_chain(&self, input: &str) -> LLMResult<String> {
let mut current_output = input.to_string();
for member in &self.members {
current_output = member.execute(¤t_output, None).await?;
}
Ok(current_output)
}
async fn run_parallel(&self, input: &str) -> LLMResult<String> {
let mut results = Vec::new();
for member in &self.members {
let result = member.execute(input, None).await?;
results.push((member.role.id.clone(), result));
}
let aggregated = results
.iter()
.map(|(id, result)| format!("=== {} ===\n{}", id, result))
.collect::<Vec<_>>()
.join("\n\n");
if let Some(ref aggregate_prompt) = self.aggregate_prompt
&& let Some(first_member) = self.members.first()
{
let prompt = aggregate_prompt
.replace("{results}", &aggregated)
.replace("{input}", input);
return first_member.agent.ask(&prompt).await;
}
Ok(aggregated)
}
async fn run_debate(&self, input: &str, max_rounds: usize) -> LLMResult<String> {
if self.members.len() < 2 {
return Err(LLMError::Other(
"Debate requires at least 2 agents".to_string(),
));
}
let mut context = format!("Initial topic: {}\n\n", input);
let mut last_response = String::new();
for round in 0..max_rounds {
for (i, member) in self.members.iter().enumerate() {
let prompt = format!(
"Round {}, Speaker {}: {}\n\n\
Previous discussion:\n{}\n\n\
Please provide your perspective. Be constructive and build on previous points.",
round + 1,
i + 1,
member.role.name,
context
);
let response = member.execute(&prompt, None).await?;
context.push_str(&format!(
"[{} - Round {}]:\n{}\n\n",
member.role.name,
round + 1,
response
));
last_response = response;
}
}
if let Some(first_member) = self.members.first() {
let summary_prompt = format!(
"Based on the following debate, provide a concise summary of the key points \
and conclusions:\n\n{}",
context
);
first_member.agent.ask(&summary_prompt).await
} else {
Ok(last_response)
}
}
async fn run_supervised(&self, input: &str) -> LLMResult<String> {
let supervisor_id = self.supervisor_id.as_ref().ok_or_else(|| {
LLMError::Other("Supervisor not specified for Supervised pattern".to_string())
})?;
let supervisor_idx = self
.member_map
.get(supervisor_id)
.ok_or_else(|| LLMError::Other(format!("Supervisor '{}' not found", supervisor_id)))?;
let mut worker_results = Vec::new();
for (i, member) in self.members.iter().enumerate() {
if i != *supervisor_idx {
let result = member.execute(input, None).await?;
worker_results.push((member.role.id.clone(), member.role.name.clone(), result));
}
}
let results_text = worker_results
.iter()
.map(|(id, name, result)| format!("=== {} ({}) ===\n{}", name, id, result))
.collect::<Vec<_>>()
.join("\n\n");
let supervisor = &self.members[*supervisor_idx];
let eval_prompt = format!(
"You are the supervisor. Evaluate the following responses to the task: \"{}\"\n\n\
Responses:\n{}\n\n\
Please provide:\n\
1. An evaluation of each response\n\
2. The best response or a synthesized improved response\n\
3. Suggestions for improvement",
input, results_text
);
supervisor.agent.ask(&eval_prompt).await
}
async fn run_map_reduce(&self, input: &str) -> LLMResult<String> {
let mut mapped_results = Vec::new();
for member in &self.members {
let result = member.execute(input, None).await?;
mapped_results.push((member.role.id.clone(), result));
}
let reduce_input = mapped_results
.iter()
.map(|(id, result)| format!("[{}]: {}", id, result))
.collect::<Vec<_>>()
.join("\n\n");
let reduce_prompt = if let Some(ref aggregate_prompt) = self.aggregate_prompt {
aggregate_prompt
.replace("{results}", &reduce_input)
.replace("{input}", input)
} else {
format!(
"Synthesize the following results into a coherent response:\n\n{}\n\n\
Original task: {}",
reduce_input, input
)
};
if let Some(first_member) = self.members.first() {
first_member.agent.ask(&reduce_prompt).await
} else {
Ok(reduce_input)
}
}
pub async fn run(&self, input: impl Into<String>) -> LLMResult<String> {
let input = input.into();
match &self.pattern {
TeamPattern::Chain => self.run_chain(&input).await,
TeamPattern::Parallel => self.run_parallel(&input).await,
TeamPattern::Debate { max_rounds } => self.run_debate(&input, *max_rounds).await,
TeamPattern::Supervised => self.run_supervised(&input).await,
TeamPattern::MapReduce => self.run_map_reduce(&input).await,
TeamPattern::Custom => {
self.run_chain(&input).await
}
}
}
pub fn get_member(&self, id: &str) -> Option<&AgentMember> {
self.member_map.get(id).map(|idx| &self.members[*idx])
}
pub fn member_ids(&self) -> Vec<&str> {
self.members.iter().map(|m| m.role.id.as_str()).collect()
}
}
pub struct AgentTeamBuilder {
id: String,
name: String,
members: Vec<AgentMember>,
pattern: TeamPattern,
supervisor_id: Option<String>,
aggregate_prompt: Option<String>,
}
impl AgentTeamBuilder {
pub fn new(id: impl Into<String>) -> Self {
let id = id.into();
Self {
name: id.clone(),
id,
members: Vec::new(),
pattern: TeamPattern::Chain,
supervisor_id: None,
aggregate_prompt: None,
}
}
pub fn with_name(mut self, name: impl Into<String>) -> Self {
self.name = name.into();
self
}
pub fn add_member(mut self, id: impl Into<String>, agent: Arc<LLMAgent>) -> Self {
self.members.push(AgentMember::new(id, agent));
self
}
pub fn add_member_with_role(mut self, agent: Arc<LLMAgent>, role: AgentRole) -> Self {
let member = AgentMember::new(&role.id, agent).with_role(role);
self.members.push(member);
self
}
pub fn with_pattern(mut self, pattern: TeamPattern) -> Self {
self.pattern = pattern;
self
}
pub fn with_supervisor(mut self, supervisor_id: impl Into<String>) -> Self {
self.supervisor_id = Some(supervisor_id.into());
self.pattern = TeamPattern::Supervised;
self
}
pub fn with_aggregate_prompt(mut self, prompt: impl Into<String>) -> Self {
self.aggregate_prompt = Some(prompt.into());
self
}
pub fn build(self) -> AgentTeam {
let member_map: HashMap<String, usize> = self
.members
.iter()
.enumerate()
.map(|(i, m)| (m.role.id.clone(), i))
.collect();
AgentTeam {
id: self.id,
name: self.name,
members: self.members,
member_map,
pattern: self.pattern,
supervisor_id: self.supervisor_id,
aggregate_prompt: self.aggregate_prompt,
}
}
}
pub fn content_creation_team(
researcher: Arc<LLMAgent>,
writer: Arc<LLMAgent>,
editor: Arc<LLMAgent>,
) -> AgentTeam {
AgentTeamBuilder::new("content-creation")
.with_name("Content Creation Team")
.add_member_with_role(
researcher,
AgentRole::new("researcher", "Researcher")
.with_description("Research and gather information on the topic")
.with_template(
"Research the following topic thoroughly and provide key findings:\n\n{input}",
),
)
.add_member_with_role(
writer,
AgentRole::new("writer", "Writer")
.with_description("Write engaging content based on research")
.with_template(
"Based on the following research, write an engaging article:\n\n{input}",
),
)
.add_member_with_role(
editor,
AgentRole::new("editor", "Editor")
.with_description("Edit and polish the content")
.with_template(
"Edit and improve the following article for clarity and engagement:\n\n{input}",
),
)
.with_pattern(TeamPattern::Chain)
.build()
}
pub fn code_review_team(
security_reviewer: Arc<LLMAgent>,
performance_reviewer: Arc<LLMAgent>,
style_reviewer: Arc<LLMAgent>,
supervisor: Arc<LLMAgent>,
) -> AgentTeam {
AgentTeamBuilder::new("code-review")
.with_name("Code Review Team")
.add_member_with_role(
security_reviewer,
AgentRole::new("security", "Security Reviewer")
.with_description("Review code for security vulnerabilities"),
)
.add_member_with_role(
performance_reviewer,
AgentRole::new("performance", "Performance Reviewer")
.with_description("Review code for performance issues"),
)
.add_member_with_role(
style_reviewer,
AgentRole::new("style", "Style Reviewer")
.with_description("Review code for style and best practices"),
)
.add_member_with_role(
supervisor,
AgentRole::new("supervisor", "Lead Reviewer")
.with_description("Synthesize reviews and provide final feedback"),
)
.with_supervisor("supervisor")
.build()
}
pub fn debate_team(agent1: Arc<LLMAgent>, agent2: Arc<LLMAgent>, max_rounds: usize) -> AgentTeam {
AgentTeamBuilder::new("debate")
.with_name("Debate Team")
.add_member_with_role(
agent1,
AgentRole::new("debater1", "Debater 1")
.with_description("Present and defend your position"),
)
.add_member_with_role(
agent2,
AgentRole::new("debater2", "Debater 2")
.with_description("Present an alternative perspective"),
)
.with_pattern(TeamPattern::Debate { max_rounds })
.build()
}
pub fn analysis_team(analysts: Vec<(impl Into<String>, Arc<LLMAgent>)>) -> AgentTeam {
let mut builder = AgentTeamBuilder::new("analysis")
.with_name("Analysis Team")
.with_pattern(TeamPattern::MapReduce)
.with_aggregate_prompt(
"Synthesize the following analyses into a comprehensive report:\n\n{results}\n\n\
Original question: {input}",
);
for (id, agent) in analysts {
builder = builder.add_member(id, agent);
}
builder.build()
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_team_builder() {
let builder = AgentTeamBuilder::new("test-team")
.with_name("Test Team")
.with_pattern(TeamPattern::Chain);
assert_eq!(builder.id, "test-team");
assert_eq!(builder.name, "Test Team");
}
#[test]
fn test_agent_role() {
let role = AgentRole::new("researcher", "Researcher")
.with_description("Research topics")
.with_template("{input}");
assert_eq!(role.id, "researcher");
assert_eq!(role.name, "Researcher");
assert_eq!(role.description, "Research topics");
assert!(role.prompt_template.is_some());
}
}