use anyhow::{Context, Result};
use std::sync::Arc;
use std::time::{Duration, Instant};
use tokio::sync::{mpsc, Mutex, RwLock, Semaphore};
use tokio::task::JoinSet;
use crate::api::types::{Message, Usage};
use crate::api::{ApiClient, ThinkingMode};
use crate::config::Config;
use super::config::{MultiAgentConfig, MultiAgentFailurePolicy};
use super::types::{
AgentInstance, AgentResult, AgentStatus, MultiAgentEvent, MAX_CONCURRENT_AGENTS,
};
pub struct MultiAgentChat {
pub(super) config: MultiAgentConfig,
pub(super) client: Arc<ApiClient>,
pub(super) semaphore: Arc<Semaphore>,
pub(super) agents: Arc<RwLock<Vec<AgentInstance>>>,
pub(super) results: Arc<Mutex<Vec<AgentResult>>>,
pub(super) event_tx: Option<mpsc::Sender<MultiAgentEvent>>,
}
impl MultiAgentChat {
pub const HEARTBEAT_TIMEOUT: Duration = Duration::from_secs(300);
pub fn new(api_config: &Config, agent_config: MultiAgentConfig) -> Result<Self> {
let client = ApiClient::new(api_config).context("Failed to create API client")?;
let concurrency = agent_config.max_concurrency.clamp(1, MAX_CONCURRENT_AGENTS);
Ok(Self {
config: agent_config,
client: Arc::new(client),
semaphore: Arc::new(Semaphore::new(concurrency)),
agents: Arc::new(RwLock::new(Vec::new())),
results: Arc::new(Mutex::new(Vec::new())),
event_tx: None,
})
}
pub fn with_events(mut self, tx: mpsc::Sender<MultiAgentEvent>) -> Self {
self.event_tx = Some(tx);
self
}
pub async fn initialize_agents(&self) -> Result<()> {
let mut agents = self.agents.write().await;
agents.clear();
for (i, role) in self.config.roles.iter().enumerate() {
let agent = AgentInstance {
id: i,
role: *role,
name: format!("Agent-{}-{}", i, role.name()),
messages: vec![Message::system(role.system_prompt())],
status: AgentStatus::Idle,
last_heartbeat: Instant::now(),
};
agents.push(agent);
}
Ok(())
}
fn emit(&self, event: MultiAgentEvent) {
if let Some(ref tx) = self.event_tx {
if let Err(mpsc::error::TrySendError::Full(_)) = tx.try_send(event) {
tracing::warn!("MultiAgent event channel full, dropping event");
}
}
}
pub async fn run_task(&self, task: &str) -> Result<Vec<AgentResult>> {
let start = Instant::now();
{
let agents = self.agents.read().await;
if agents.is_empty() {
drop(agents);
self.initialize_agents().await?;
}
}
{
let mut results = self.results.lock().await;
results.clear();
}
let agent_count = {
let agents = self.agents.read().await;
agents.len()
};
let base_config = self.client.config();
let resolved_max_tokens = self.config.max_tokens.unwrap_or(base_config.max_tokens);
let timeout = Duration::from_secs(
self.config
.timeout_secs
.unwrap_or(base_config.agent.step_timeout_secs),
);
let limits = BudgetLimits {
max_budget_tokens: base_config.agent.max_budget_tokens,
max_cost_usd: base_config.agent.max_cost_usd,
};
if let Some(0) = limits.max_budget_tokens {
anyhow::bail!(
"multi-chat: --max-budget-tokens is 0; no agent calls can be made. \
Raise the budget or drop the flag."
);
}
if let Some(max) = limits.max_cost_usd {
anyhow::ensure!(
max.is_finite() && max > 0.0,
"multi-chat: --max-cost-usd must be a positive number (got {max}); \
no agent calls can be made."
);
}
let budget = BudgetGuard::new(limits, resolved_max_tokens);
let client = if self.config.temperature.is_some() || self.config.max_tokens.is_some() {
let mut per_agent_config = base_config.clone();
if let Some(t) = self.config.temperature {
per_agent_config.temperature = t;
}
if let Some(mt) = self.config.max_tokens {
per_agent_config.max_tokens = mt;
}
Arc::new(
ApiClient::new(&per_agent_config)
.context("Failed to create per-agent API client")?,
)
} else {
Arc::clone(&self.client)
};
let cancelled = Arc::new(tokio::sync::Notify::new());
let mut join_set = JoinSet::new();
for agent_id in 0..agent_count {
let semaphore = Arc::clone(&self.semaphore);
let agents = Arc::clone(&self.agents);
let results = Arc::clone(&self.results);
let task = task.to_string();
let event_tx = self.event_tx.clone();
let failure_policy = self.config.failure_policy;
let cancelled = Arc::clone(&cancelled);
let client = Arc::clone(&client);
let budget = budget.clone();
join_set.spawn(async move {
tokio::select! {
_ = cancelled.notified() => {
Ok(())
}
res = Self::run_single_agent(
agent_id, task, client, budget, semaphore, agents, results, timeout, event_tx,
) => {
if failure_policy == MultiAgentFailurePolicy::FailFast && res.is_err() {
cancelled.notify_waiters();
}
res
}
}
});
}
while let Some(result) = join_set.join_next().await {
match result {
Ok(Ok(_)) => {
}
Ok(Err(e)) => {
eprintln!("Agent-specific error: {}", e);
if self.config.failure_policy == MultiAgentFailurePolicy::FailFast {
cancelled.notify_waiters();
join_set.abort_all();
while join_set.join_next().await.is_some() {}
break;
}
}
Err(e) if e.is_cancelled() => {
tracing::debug!("Agent task cancelled: {}", e);
}
Err(e) => {
tracing::error!("Agent task panicked: {}", e);
eprintln!("Agent task panicked: {}", e);
if self.config.failure_policy == MultiAgentFailurePolicy::FailFast {
cancelled.notify_waiters();
join_set.abort_all();
while join_set.join_next().await.is_some() {}
break;
}
}
}
}
let total_duration = start.elapsed();
let results = {
let results = self.results.lock().await;
results.clone()
};
self.emit(MultiAgentEvent::AllCompleted {
results: results.clone(),
total_duration,
});
Ok(results)
}
#[allow(clippy::too_many_arguments)]
async fn run_single_agent(
agent_id: usize,
task: String,
client: Arc<ApiClient>,
budget: BudgetGuard,
semaphore: Arc<Semaphore>,
agents: Arc<RwLock<Vec<AgentInstance>>>,
results: Arc<Mutex<Vec<AgentResult>>>,
timeout: Duration,
event_tx: Option<mpsc::Sender<MultiAgentEvent>>,
) -> Result<()> {
let _permit = semaphore.acquire().await?;
let start = Instant::now();
if let Err(reason) = budget.try_reserve() {
let (agent_name, role) = {
let agents = agents.read().await;
match agents.get(agent_id) {
Some(a) => (a.name.clone(), a.role),
None => return Ok(()),
}
};
tracing::info!("multi-chat: agent {} not launched: {}", agent_id, reason);
if let Some(ref tx) = event_tx {
let _ = tx.try_send(MultiAgentEvent::AgentFailed {
agent_id,
error: reason.clone(),
});
}
let mut results = results.lock().await;
results.push(AgentResult {
agent_id,
agent_name,
role,
content: String::new(),
usage: None,
duration: start.elapsed(),
success: false,
error: Some(reason),
});
return Ok(());
}
let (agent_name, role, mut messages) = {
let mut agents = agents.write().await;
if let Some(agent) = agents.get_mut(agent_id) {
agent.status = AgentStatus::Working;
agent.last_heartbeat = Instant::now();
(agent.name.clone(), agent.role, agent.messages.clone())
} else {
budget.settle(None);
return Ok(());
}
};
if let Some(ref tx) = event_tx {
let _ = tx.try_send(MultiAgentEvent::AgentStarted {
agent_id,
name: agent_name.clone(),
task: task.clone(),
});
}
messages.push(Message::user(&task));
let result =
tokio::time::timeout(timeout, client.chat(messages, None, ThinkingMode::Disabled))
.await;
budget.settle(match &result {
Ok(Ok(response)) => Some(&response.usage),
_ => None,
});
let duration = start.elapsed();
let agent_result = match result {
Ok(Ok(response)) => {
let content = response
.choices
.first()
.map(|c| c.message.content.text().to_string())
.unwrap_or_default();
AgentResult {
agent_id,
agent_name: agent_name.clone(),
role,
content,
usage: Some(response.usage),
duration,
success: true,
error: None,
}
}
Ok(Err(e)) => {
if let Some(ref tx) = event_tx {
let _ = tx.try_send(MultiAgentEvent::AgentFailed {
agent_id,
error: e.to_string(),
});
}
AgentResult {
agent_id,
agent_name: agent_name.clone(),
role,
content: String::new(),
usage: None,
duration,
success: false,
error: Some(e.to_string()),
}
}
Err(_) => {
let error = "Request timed out".to_string();
if let Some(ref tx) = event_tx {
let _ = tx.try_send(MultiAgentEvent::AgentFailed {
agent_id,
error: error.clone(),
});
}
AgentResult {
agent_id,
agent_name: agent_name.clone(),
role,
content: String::new(),
usage: None,
duration,
success: false,
error: Some(error),
}
}
};
{
let mut agents = agents.write().await;
if let Some(agent) = agents.get_mut(agent_id) {
agent.status = if agent_result.success {
AgentStatus::Completed
} else {
AgentStatus::Failed
};
agent.last_heartbeat = Instant::now();
if agent_result.success {
agent.messages.push(Message::user(&task));
agent
.messages
.push(Message::assistant(&agent_result.content));
}
}
}
if let Some(ref tx) = event_tx {
let _ = tx.try_send(MultiAgentEvent::AgentCompleted {
agent_id,
result: agent_result.clone(),
});
}
let agent_failed = !agent_result.success;
{
let mut results = results.lock().await;
results.push(agent_result);
}
if agent_failed {
Err(anyhow::anyhow!("Agent {} failed", agent_id))
} else {
Ok(())
}
}
pub async fn is_agent_healthy(&self, agent_id: usize) -> bool {
let agents = self.agents.read().await;
if let Some(agent) = agents.get(agent_id) {
agent.status != AgentStatus::Failed
&& agent.last_heartbeat.elapsed() < Self::HEARTBEAT_TIMEOUT
} else {
false
}
}
pub fn aggregate_results(results: &[AgentResult]) -> String {
let mut summary = String::new();
summary.push_str("## Multi-Agent Summary\n\n");
for result in results {
if result.success {
summary.push_str(&format!(
"### {} ({})\n",
result.agent_name,
result.role.name()
));
summary.push_str(&result.content);
summary.push_str("\n\n");
} else if let Some(error) = &result.error {
summary.push_str(&format!(
"### {} (FAILED)\nError: {}\n\n",
result.agent_name, error
));
}
}
summary
}
pub fn total_usage(results: &[AgentResult]) -> Usage {
let mut total = Usage {
prompt_tokens: 0,
completion_tokens: 0,
total_tokens: 0,
cost: None,
};
for result in results {
if let Some(usage) = &result.usage {
total.prompt_tokens += usage.prompt_tokens;
total.completion_tokens += usage.completion_tokens;
total.total_tokens += usage.total_tokens;
if let Some(cost) = usage.cost {
*total.cost.get_or_insert(0.0) += cost;
}
}
}
total
}
}
#[derive(Debug, Clone, Copy, Default)]
struct BudgetLimits {
max_budget_tokens: Option<usize>,
max_cost_usd: Option<f64>,
}
#[derive(Debug, Default)]
struct BudgetTracker {
actual_tokens: usize,
actual_cost: f64,
reserved_tokens: usize,
}
#[derive(Debug, Clone)]
struct BudgetGuard {
tracker: Arc<std::sync::Mutex<BudgetTracker>>,
limits: BudgetLimits,
estimate: usize,
}
impl BudgetGuard {
fn new(limits: BudgetLimits, estimate: usize) -> Self {
Self {
tracker: Arc::new(std::sync::Mutex::new(BudgetTracker::default())),
limits,
estimate,
}
}
fn try_reserve(&self) -> Result<(), String> {
let mut tracker = self.tracker.lock().unwrap();
if let Some(max) = self.limits.max_budget_tokens {
let committed = tracker.actual_tokens + tracker.reserved_tokens;
if committed + self.estimate > max {
return Err(format!(
"skipped to stay within --max-budget-tokens={max}: \
{committed} tokens used/reserved + ~{} estimated for this call",
self.estimate
));
}
}
if let Some(max) = self.limits.max_cost_usd {
if tracker.actual_cost >= max {
return Err(format!(
"skipped to stay within --max-cost-usd=${max}: \
${:.6} already spent",
tracker.actual_cost
));
}
}
tracker.reserved_tokens += self.estimate;
Ok(())
}
fn settle(&self, usage: Option<&Usage>) {
let mut tracker = self.tracker.lock().unwrap();
tracker.reserved_tokens = tracker.reserved_tokens.saturating_sub(self.estimate);
if let Some(usage) = usage {
tracker.actual_tokens += usage.total_tokens;
if let Some(cost) = usage.cost {
tracker.actual_cost += cost;
}
}
}
}