mod scheduler;
use mofa_kernel::message::{AgentMessage, TaskRequest, TaskStatus};
use mofa_kernel::{AgentBus, CommunicationMode};
use std::collections::HashMap;
use std::sync::Arc;
use tokio::sync::RwLock;
#[derive(Debug, Clone)]
pub enum CoordinationStrategy {
MasterSlave, PeerToPeer, Pipeline, }
pub struct AgentCoordinator {
bus: Arc<AgentBus>,
strategy: CoordinationStrategy,
role_mapping: Arc<RwLock<HashMap<String, Vec<String>>>>,
task_tracker: Arc<RwLock<HashMap<String, (String, TaskStatus)>>>,
scheduler: scheduler::PriorityScheduler,
}
impl AgentCoordinator {
pub async fn new(bus: Arc<AgentBus>, strategy: CoordinationStrategy) -> Self {
let scheduler = scheduler::PriorityScheduler::new(bus.clone()).await;
Self {
bus,
strategy,
role_mapping: Arc::new(RwLock::new(HashMap::new())),
task_tracker: Arc::new(RwLock::new(HashMap::new())),
scheduler,
}
}
pub async fn register_role(&self, agent_id: &str, role: &str) -> anyhow::Result<()> {
let mut role_map = self.role_mapping.write().await;
role_map
.entry(role.to_string())
.or_default()
.push(agent_id.to_string());
Ok(())
}
pub async fn coordinate_task(&self, task_msg: &AgentMessage) -> anyhow::Result<()> {
match &self.strategy {
CoordinationStrategy::MasterSlave => self.master_slave_coordinate(task_msg).await,
CoordinationStrategy::Pipeline => self.pipeline_coordinate(task_msg).await,
_ => Ok(()),
}
}
pub async fn submit_priority_task(&self, task: TaskRequest) -> anyhow::Result<()> {
self.scheduler.submit_task(task).await
}
async fn master_slave_coordinate(&self, task_msg: &AgentMessage) -> anyhow::Result<()> {
let role_map = self.role_mapping.read().await;
let masters = role_map
.get("master")
.ok_or_else(|| anyhow::anyhow!("No master agent registered"))?;
let master_id = &masters[0];
let workers = role_map
.get("worker")
.ok_or_else(|| anyhow::anyhow!("No worker agents registered"))?;
self.bus
.send_message(master_id, CommunicationMode::Broadcast, task_msg)
.await?;
if let AgentMessage::TaskRequest { task_id, .. } = task_msg {
let mut tracker = self.task_tracker.write().await;
for worker_id in workers {
tracker.insert(task_id.clone(), (worker_id.clone(), TaskStatus::Pending));
}
}
Ok(())
}
async fn pipeline_coordinate(&self, task_msg: &AgentMessage) -> anyhow::Result<()> {
let role_map = self.role_mapping.read().await;
let stages = vec!["stage1", "stage2", "stage3"];
let mut last_output: Option<String> = None;
for stage in stages {
let agents = role_map
.get(stage)
.ok_or_else(|| anyhow::anyhow!("No agent for stage {}", stage))?;
let agent_id = &agents[0];
let current_msg = match last_output {
Some(ref output) => AgentMessage::TaskRequest {
task_id: uuid::Uuid::now_v7().to_string(),
content: output.clone(),
},
None => task_msg.clone(),
};
self.bus
.send_message(
"coordinator",
CommunicationMode::PointToPoint(agent_id.to_string()),
¤t_msg,
)
.await?;
if let Some(AgentMessage::TaskResponse { result, .. }) = self
.bus
.receive_message(
agent_id,
CommunicationMode::PointToPoint("coordinator".to_string()),
)
.await?
{
last_output = Some(result);
}
}
Ok(())
}
}