use mofa_kernel::message::{AgentEvent, AgentMessage, SchedulingStatus, TaskPriority, TaskRequest};
use mofa_kernel::{AgentBus, CommunicationMode};
use std::collections::{BinaryHeap, HashMap};
use std::sync::Arc;
use tokio::sync::RwLock;
#[derive(Debug, Clone, Eq, PartialEq)]
struct PriorityTask {
priority: TaskPriority,
task: TaskRequest,
submit_time: std::time::Instant,
}
impl Ord for PriorityTask {
fn cmp(&self, other: &Self) -> std::cmp::Ordering {
self.priority
.cmp(&other.priority)
.then_with(|| other.submit_time.cmp(&self.submit_time)) }
}
impl PartialOrd for PriorityTask {
fn partial_cmp(&self, other: &Self) -> Option<std::cmp::Ordering> {
Some(self.cmp(other))
}
}
pub struct PriorityScheduler {
task_queue: Arc<RwLock<BinaryHeap<PriorityTask>>>, agent_load: Arc<RwLock<HashMap<String, usize>>>, bus: Arc<AgentBus>,
task_status: Arc<RwLock<HashMap<String, SchedulingStatus>>>, role_mapping: Arc<RwLock<HashMap<String, Vec<String>>>>, }
impl PriorityScheduler {
pub async fn new(bus: Arc<AgentBus>) -> Self {
Self {
task_queue: Arc::new(RwLock::new(BinaryHeap::new())),
agent_load: Arc::new(RwLock::new(HashMap::new())),
bus,
task_status: Arc::new(RwLock::new(HashMap::new())),
role_mapping: Arc::new(RwLock::new(HashMap::new())),
}
}
pub async fn submit_task(&self, task: TaskRequest) -> anyhow::Result<()> {
let priority_task = PriorityTask {
priority: task.priority.clone(),
task: task.clone(),
submit_time: std::time::Instant::now(),
};
self.task_queue.write().await.push(priority_task);
self.task_status
.write()
.await
.insert(task.task_id, SchedulingStatus::Pending);
self.schedule().await?;
Ok(())
}
pub async fn schedule(&self) -> anyhow::Result<()> {
let mut task_queue = self.task_queue.write().await;
let mut agent_load = self.agent_load.write().await;
let mut task_status = self.task_status.write().await;
while let Some(priority_task) = task_queue.pop() {
let task = priority_task.task.clone(); let task_id = task.task_id.clone();
if task_status.get(&task_id) != Some(&SchedulingStatus::Pending) {
continue;
}
let target_agent = self.select_low_load_agent("worker").await?;
if target_agent.is_empty() {
task_queue.push(priority_task);
break;
}
let target_agent = target_agent[0].clone();
self.preempt_low_priority_task(&target_agent, &task).await?;
let task_msg = AgentMessage::TaskRequest {
task_id: task.task_id.clone(),
content: task.content.clone(),
};
self.bus
.send_message(
"scheduler",
CommunicationMode::PointToPoint(target_agent.clone()),
&task_msg,
)
.await?;
task_status.insert(task_id, SchedulingStatus::Running);
*agent_load.entry(target_agent).or_insert(0) += 1;
}
Ok(())
}
async fn select_low_load_agent(&self, role: &str) -> anyhow::Result<Vec<String>> {
let role_map = self.role_mapping.read().await;
let agents = role_map
.get(role)
.ok_or_else(|| anyhow::anyhow!("No agent for role: {}", role))?;
let agent_load = self.agent_load.read().await;
let mut sorted_agents = agents.clone();
sorted_agents.sort_by_key(|agent_id| agent_load.get(agent_id).cloned().unwrap_or(0));
Ok(sorted_agents)
}
async fn preempt_low_priority_task(
&self,
agent_id: &str,
_high_priority_task: &TaskRequest,
) -> anyhow::Result<()> {
let agent_load = self.agent_load.read().await;
let task_status = self.task_status.read().await;
if let Some(&load) = agent_load.get(agent_id)
&& load > 0
{
let low_priority_task_id = task_status
.iter()
.find(|(_, status_ref)| **status_ref == SchedulingStatus::Running)
.map(|(task_id, _)| task_id.clone())
.ok_or_else(|| anyhow::anyhow!("No running task on agent: {}", agent_id))?;
let preempt_msg =
AgentMessage::Event(AgentEvent::TaskPreempted(low_priority_task_id.clone()));
self.bus
.send_message(
"scheduler",
CommunicationMode::PointToPoint(agent_id.to_string()),
&preempt_msg,
)
.await?;
}
Ok(())
}
pub async fn on_task_completed(&self, agent_id: &str, task_id: &str) -> anyhow::Result<()> {
let mut agent_load = self.agent_load.write().await;
let mut task_status = self.task_status.write().await;
agent_load
.entry(agent_id.to_string())
.and_modify(|count| *count -= 1);
task_status.insert(task_id.to_string(), SchedulingStatus::Completed);
self.schedule().await?;
Ok(())
}
}