use std::collections::HashMap;
use ai_agents_core::{AgentError, AgentResponse, Result};
use ai_agents_llm::{ChatMessage, LLMProvider};
use ai_agents_observability::{ObservationPurpose, with_observation_purpose};
use ai_agents_state::{
AggregationConfig, AggregationStrategy, TiebreakerStrategy, VoteConfig, VoteMethod,
};
use futures::future::join_all;
use rand::seq::SliceRandom;
use tracing::debug;
use super::types::AgentResult;
pub async fn aggregate(
results: &[AgentResult],
config: &AggregationConfig,
llm: Option<&dyn LLMProvider>,
agent_weights: &HashMap<String, f64>,
vote_parallelism: Option<usize>,
) -> Result<AgentResponse> {
aggregate_with_llms(
results,
config,
super::AggregationLLMs::shared(llm),
agent_weights,
vote_parallelism,
)
.await
}
pub async fn aggregate_with_llms(
results: &[AgentResult],
config: &AggregationConfig,
llms: super::AggregationLLMs<'_>,
agent_weights: &HashMap<String, f64>,
vote_parallelism: Option<usize>,
) -> Result<AgentResponse> {
let successful: Vec<&AgentResult> = results.iter().filter(|r| r.success).collect();
if successful.is_empty() {
return Err(AgentError::Other("All agents failed".into()));
}
match config.strategy {
AggregationStrategy::FirstWins => {
debug!("Aggregation strategy: first_wins");
let first = &successful[0];
Ok(first
.response
.clone()
.unwrap_or_else(|| AgentResponse::new("")))
}
AggregationStrategy::All => {
debug!(count = successful.len(), "Aggregation strategy: all");
let content = successful
.iter()
.filter_map(|r| {
r.response
.as_ref()
.map(|resp| format!("**{}**:\n{}", r.agent_id, resp.content))
})
.collect::<Vec<_>>()
.join("\n\n");
Ok(AgentResponse::new(content))
}
AggregationStrategy::LlmSynthesis => {
debug!("Aggregation strategy: llm_synthesis");
let llm = llms.synthesis.ok_or_else(|| {
AgentError::Config("LLM required for llm_synthesis aggregation".into())
})?;
synthesize_with_llm(llm, &successful, config.synthesizer_prompt.as_deref()).await
}
AggregationStrategy::Voting => {
debug!("Aggregation strategy: voting");
let llm = llms
.vote
.ok_or_else(|| AgentError::Config("LLM required for voting aggregation".into()))?;
let vote_config = config.vote.as_ref();
vote_with_llm(
llm,
llms.tiebreak,
results,
vote_config,
agent_weights,
vote_parallelism,
)
.await
}
}
}
async fn synthesize_with_llm(
llm: &dyn LLMProvider,
results: &[&AgentResult],
custom_prompt: Option<&str>,
) -> Result<AgentResponse> {
let agent_responses = results
.iter()
.filter_map(|r| {
r.response
.as_ref()
.map(|resp| format!("[Agent: {}]\n{}", r.agent_id, resp.content))
})
.collect::<Vec<_>>()
.join("\n\n---\n\n");
let system = custom_prompt.unwrap_or(
"You are a synthesis assistant. \
Multiple agents have provided their analysis. \
Combine their insights into a single coherent response. \
Include the key points from each perspective.",
);
let messages = vec![
ChatMessage::system(system),
ChatMessage::user(format!(
"Synthesize these responses:\n\n{}",
agent_responses
)),
];
let response = with_observation_purpose(
ObservationPurpose::OrchestrationAggregation,
llm.complete(&messages, None),
)
.await
.map_err(|e| AgentError::LLM(format!("Synthesis LLM failed: {}", e)))?;
Ok(AgentResponse::new(response.content))
}
async fn extract_vote(
llm: &dyn LLMProvider,
result: &AgentResult,
vote_prompt: &str,
method: &VoteMethod,
agent_weights: &HashMap<String, f64>,
) -> Result<(usize, String, String, f64)> {
let response = result
.response
.as_ref()
.ok_or_else(|| AgentError::Other(format!("Agent {} has no response", result.agent_id)))?;
let messages = vec![
ChatMessage::system(vote_prompt),
ChatMessage::user(&response.content),
];
let extraction = with_observation_purpose(
ObservationPurpose::OrchestrationAggregation,
llm.complete(&messages, None),
)
.await
.map_err(|e| AgentError::LLM(format!("Vote extraction failed: {}", e)))?;
let weight = match method {
VoteMethod::Weighted => agent_weights.get(&result.agent_id).copied().unwrap_or(1.0),
_ => 1.0,
};
Ok((
result.agent_index,
result.agent_id.clone(),
extraction.content.trim().to_string(),
weight,
))
}
async fn vote_with_llm(
llm: &dyn LLMProvider,
tiebreak_llm: Option<&dyn LLMProvider>,
results: &[AgentResult],
vote_config: Option<&VoteConfig>,
agent_weights: &HashMap<String, f64>,
vote_parallelism: Option<usize>,
) -> Result<AgentResponse> {
let vote_prompt = vote_config
.and_then(|v| v.vote_prompt.as_deref())
.unwrap_or(
"Extract the main recommendation or decision from this response as a single short phrase.",
);
let method = vote_config.map(|v| v.method.clone()).unwrap_or_default();
let tiebreaker = vote_config
.map(|v| v.tiebreaker.clone())
.unwrap_or_default();
let successful: Vec<&AgentResult> = results
.iter()
.filter(|result| result.success && result.response.is_some())
.collect();
let mut votes: Vec<(usize, String, String, f64)> = Vec::new();
if let Some(limit) = vote_parallelism.filter(|limit| *limit > 1) {
for chunk in successful.chunks(limit) {
let extraction_futures = chunk.iter().copied().map(|result| {
let method_for_task = method.clone();
async move {
let response = result.response.as_ref().ok_or_else(|| {
AgentError::Other(format!("Agent {} has no response", result.agent_id))
})?;
let messages = vec![
ChatMessage::system(vote_prompt),
ChatMessage::user(&response.content),
];
let extraction = with_observation_purpose(
ObservationPurpose::OrchestrationAggregation,
llm.complete(&messages, None),
)
.await
.map_err(|e| AgentError::LLM(format!("Vote extraction failed: {}", e)))?;
let weight = match method_for_task {
VoteMethod::Weighted => {
agent_weights.get(&result.agent_id).copied().unwrap_or(1.0)
}
_ => 1.0,
};
Ok::<(usize, String, String, f64), AgentError>((
result.agent_index,
result.agent_id.clone(),
extraction.content.trim().to_string(),
weight,
))
}
});
for extraction in join_all(extraction_futures).await {
votes.push(extraction?);
}
}
} else {
for result in successful {
votes.push(extract_vote(llm, result, vote_prompt, &method, agent_weights).await?);
}
}
votes.sort_by_key(|(agent_index, _, _, _)| *agent_index);
if votes.is_empty() {
return Err(AgentError::Other("No votes extracted".into()));
}
if matches!(method, VoteMethod::Unanimous) {
let first_vote = votes[0].2.to_lowercase();
let all_agree = votes
.iter()
.all(|(_, _, v, _)| v.to_lowercase() == first_vote);
if !all_agree {
let vote_lines = votes
.iter()
.map(|(_, id, v, _)| format!("- {}: {}", id, v))
.collect::<Vec<_>>()
.join("\n");
return Err(AgentError::Other(format!(
"Unanimous vote failed: agents did not agree\n\nVotes:\n{}",
vote_lines
)));
}
return Ok(AgentResponse::new(format!(
"Unanimous decision: {}",
votes[0].2
)));
}
let mut tally: HashMap<String, f64> = HashMap::new();
for (_, _, vote, weight) in &votes {
*tally.entry(vote.clone()).or_default() += weight;
}
let max_score = tally.values().cloned().fold(f64::NEG_INFINITY, f64::max);
let tied: Vec<String> = tally
.iter()
.filter(|(_, v)| (**v - max_score).abs() < f64::EPSILON)
.map(|(k, _)| k.clone())
.collect();
let winner = if tied.len() == 1 {
tied[0].clone()
} else {
match tiebreaker {
TiebreakerStrategy::First => {
votes
.iter()
.find(|(_, _, choice, _)| tied.contains(choice))
.map(|(_, _, choice, _)| choice.clone())
.unwrap_or_else(|| tied[0].clone())
}
TiebreakerStrategy::Random => {
let mut rng = rand::thread_rng();
tied.choose(&mut rng)
.cloned()
.unwrap_or_else(|| tied[0].clone())
}
TiebreakerStrategy::RouterDecides => {
resolve_tie_with_llm(
tiebreak_llm
.ok_or_else(|| AgentError::Config("LLM required for tie-break".into()))?,
&tied,
)
.await?
}
}
};
let vote_lines = votes
.iter()
.map(|(_, id, v, _)| format!("- {}: {}", id, v))
.collect::<Vec<_>>()
.join("\n");
let summary = format!("Vote result: {}\n\nVotes:\n{}", winner, vote_lines);
debug!(winner = %winner, total_votes = votes.len(), "Vote aggregation complete");
Ok(AgentResponse::new(summary))
}
async fn resolve_tie_with_llm(llm: &dyn LLMProvider, tied_choices: &[String]) -> Result<String> {
let choices_list = tied_choices
.iter()
.enumerate()
.map(|(i, c)| format!("{}. {}", i + 1, c))
.collect::<Vec<_>>()
.join("\n");
let messages = vec![
ChatMessage::system(
"You are a tiebreaker. Multiple options received equal votes. \
Pick the single best option. Respond with ONLY the option text.",
),
ChatMessage::user(format!(
"These options are tied:\n{}\n\nPick one.",
choices_list
)),
];
let response = with_observation_purpose(
ObservationPurpose::OrchestrationAggregation,
llm.complete(&messages, None),
)
.await
.map_err(|e| AgentError::LLM(format!("Tiebreaker LLM failed: {}", e)))?;
let raw = response.content.trim().to_string();
for choice in tied_choices {
if raw.contains(choice.as_str()) || choice.contains(raw.as_str()) {
return Ok(choice.clone());
}
}
Ok(tied_choices[0].clone())
}
#[cfg(test)]
mod tests {
use super::*;
use ai_agents_core::{
AgentResponse, FinishReason, LLMChunk, LLMConfig, LLMError, LLMFeature, LLMResponse,
};
use ai_agents_llm::mock::MockLLMProvider;
use ai_agents_state::AggregationStrategy;
use async_trait::async_trait;
use parking_lot::Mutex;
use tokio::sync::Semaphore;
#[derive(Default)]
struct VoteActivity {
started: Vec<usize>,
active: usize,
peak: usize,
}
struct ControlledVoteProvider {
gates: [Semaphore; 3],
activity: Mutex<VoteActivity>,
}
impl ControlledVoteProvider {
fn new() -> Self {
Self {
gates: std::array::from_fn(|_| Semaphore::new(0)),
activity: Mutex::new(VoteActivity::default()),
}
}
fn assert_activity(&self, started: &[usize], active: usize, peak: usize) {
let activity = self.activity.lock();
assert_eq!(activity.started, started);
assert_eq!(activity.active, active);
assert_eq!(activity.peak, peak);
}
}
struct ActiveVote<'a>(&'a Mutex<VoteActivity>);
impl Drop for ActiveVote<'_> {
fn drop(&mut self) {
self.0.lock().active -= 1;
}
}
#[async_trait]
impl LLMProvider for ControlledVoteProvider {
async fn complete(
&self,
messages: &[ChatMessage],
_config: Option<&LLMConfig>,
) -> std::result::Result<LLMResponse, LLMError> {
let id = match messages.last().map(|message| message.content.as_str()) {
Some("first") => 0,
Some("second") => 1,
Some("third") => 2,
other => {
return Err(LLMError::Config(format!(
"Unexpected vote request: {other:?}"
)));
}
};
{
let mut activity = self.activity.lock();
assert!(!activity.started.contains(&id), "Duplicate vote request");
activity.started.push(id);
activity.active += 1;
activity.peak = activity.peak.max(activity.active);
}
let _active = ActiveVote(&self.activity);
self.gates[id].acquire().await.unwrap().forget();
Ok(LLMResponse::new("A", FinishReason::Stop))
}
async fn complete_stream(
&self,
_messages: &[ChatMessage],
_config: Option<&LLMConfig>,
) -> std::result::Result<
Box<dyn futures::Stream<Item = std::result::Result<LLMChunk, LLMError>> + Unpin + Send>,
LLMError,
> {
Err(LLMError::Config("Unexpected streaming vote request".into()))
}
fn provider_name(&self) -> &str {
"controlled-votes"
}
fn supports(&self, _feature: LLMFeature) -> bool {
false
}
}
fn agent_result(index: usize, id: &str, content: &str) -> AgentResult {
AgentResult {
agent_index: index,
agent_id: id.to_string(),
response: Some(AgentResponse::new(content)),
duration_ms: 0,
success: true,
error: None,
}
}
#[tokio::test]
async fn vote_extraction_is_serial_without_parallelism() {
let llm = ControlledVoteProvider::new();
let results = vec![
agent_result(0, "a", "first"),
agent_result(1, "b", "second"),
agent_result(2, "c", "third"),
];
let config = AggregationConfig {
strategy: AggregationStrategy::Voting,
synthesizer_llm: None,
synthesizer_prompt: None,
vote: None,
};
let weights = HashMap::new();
let mut operation = Box::pin(aggregate(&results, &config, Some(&llm), &weights, None));
assert!(futures::poll!(operation.as_mut()).is_pending());
llm.assert_activity(&[0], 1, 1);
llm.gates[0].add_permits(1);
assert!(futures::poll!(operation.as_mut()).is_pending());
llm.assert_activity(&[0, 1], 1, 1);
llm.gates[1].add_permits(1);
assert!(futures::poll!(operation.as_mut()).is_pending());
llm.assert_activity(&[0, 1, 2], 1, 1);
llm.gates[2].add_permits(1);
let response = operation.await.unwrap();
llm.assert_activity(&[0, 1, 2], 0, 1);
assert_eq!(
response.content,
"Vote result: A\n\nVotes:\n- a: A\n- b: A\n- c: A"
);
}
#[tokio::test]
async fn vote_extraction_uses_bounded_parallelism_when_enabled() {
let llm = ControlledVoteProvider::new();
let results = vec![
agent_result(0, "a", "first"),
agent_result(1, "b", "second"),
agent_result(2, "c", "third"),
];
let config = AggregationConfig {
strategy: AggregationStrategy::Voting,
synthesizer_llm: None,
synthesizer_prompt: None,
vote: None,
};
let weights = HashMap::new();
let mut operation = Box::pin(aggregate(&results, &config, Some(&llm), &weights, Some(2)));
assert!(futures::poll!(operation.as_mut()).is_pending());
llm.assert_activity(&[0, 1], 2, 2);
llm.gates[1].add_permits(1);
assert!(futures::poll!(operation.as_mut()).is_pending());
llm.assert_activity(&[0, 1], 1, 2);
llm.gates[0].add_permits(1);
assert!(futures::poll!(operation.as_mut()).is_pending());
llm.assert_activity(&[0, 1, 2], 1, 2);
llm.gates[2].add_permits(1);
let response = operation.await.unwrap();
llm.assert_activity(&[0, 1, 2], 0, 2);
assert_eq!(
response.content,
"Vote result: A\n\nVotes:\n- a: A\n- b: A\n- c: A"
);
}
#[tokio::test]
async fn dropping_vote_extraction_releases_active_requests() {
let llm = ControlledVoteProvider::new();
let results = vec![
agent_result(0, "a", "first"),
agent_result(1, "b", "second"),
];
let config = AggregationConfig {
strategy: AggregationStrategy::Voting,
synthesizer_llm: None,
synthesizer_prompt: None,
vote: None,
};
let weights = HashMap::new();
let mut operation = Box::pin(aggregate(&results, &config, Some(&llm), &weights, Some(2)));
assert!(futures::poll!(operation.as_mut()).is_pending());
llm.assert_activity(&[0, 1], 2, 2);
drop(operation);
llm.assert_activity(&[0, 1], 0, 2);
}
#[tokio::test]
async fn vote_tiebreaker_first_uses_declaration_order() {
let mut llm = MockLLMProvider::new("votes");
llm.set_responses(vec!["B".to_string(), "A".to_string()], false);
let results = vec![
agent_result(1, "b", "second"),
agent_result(0, "a", "first"),
];
let config = AggregationConfig {
strategy: AggregationStrategy::Voting,
synthesizer_llm: None,
synthesizer_prompt: None,
vote: Some(VoteConfig {
method: VoteMethod::Majority,
tiebreaker: TiebreakerStrategy::First,
vote_prompt: None,
}),
};
let response = aggregate(&results, &config, Some(&llm), &HashMap::new(), None)
.await
.unwrap();
assert!(response.content.starts_with("Vote result: A"));
}
}