use std::collections::HashMap;
use std::future::Future;
use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
use std::sync::Arc;
use std::time::Duration;
use tokio::sync::RwLock;
use tokio::task::JoinHandle;
use crate::agent::tui_events::{AgentEvent, BroadcastEmitter, EventEmitter};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum RunStatus {
Running,
Paused,
Completed,
Failed,
Aborted,
}
impl RunStatus {
pub fn is_terminal(self) -> bool {
matches!(
self,
RunStatus::Completed | RunStatus::Failed | RunStatus::Aborted
)
}
fn terminal_event(self) -> AgentEvent {
match self {
RunStatus::Completed => AgentEvent::Completed {
message: "run completed".to_string(),
},
RunStatus::Failed => AgentEvent::Error {
message: "run failed".to_string(),
},
other => AgentEvent::Status {
message: format!("run status: {other:?}"),
},
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, PartialOrd, Ord)]
pub struct RunId(pub u64);
impl std::fmt::Display for RunId {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
write!(f, "run-{}", self.0)
}
}
const EVENT_CHANNEL_CAPACITY: usize = 256;
const STATUS_POLL_INTERVAL: Duration = Duration::from_millis(100);
struct RunHandle {
status: Arc<RwLock<RunStatus>>,
cancel: Arc<AtomicBool>,
join: JoinHandle<()>,
#[allow(dead_code)]
task: String,
events: tokio::sync::broadcast::Sender<AgentEvent>,
}
pub struct RunSupervisor {
runs: RwLock<HashMap<RunId, RunHandle>>,
next_id: AtomicU64,
}
impl Default for RunSupervisor {
fn default() -> Self {
Self::new()
}
}
impl RunSupervisor {
pub fn new() -> Self {
Self {
runs: RwLock::new(HashMap::new()),
next_id: AtomicU64::new(1),
}
}
pub async fn spawn<F>(&self, task: String, run: F) -> RunId
where
F: Future<Output = anyhow::Result<()>> + Send + 'static,
{
let (event_tx, _event_rx) =
tokio::sync::broadcast::channel::<AgentEvent>(EVENT_CHANNEL_CAPACITY);
self.spawn_with_events(task, event_tx, run).await
}
async fn spawn_with_events<F>(
&self,
task: String,
event_tx: tokio::sync::broadcast::Sender<AgentEvent>,
run: F,
) -> RunId
where
F: Future<Output = anyhow::Result<()>> + Send + 'static,
{
let id = RunId(self.next_id.fetch_add(1, Ordering::Relaxed));
let status = Arc::new(RwLock::new(RunStatus::Running));
let cancel = Arc::new(AtomicBool::new(false));
let status_clone = Arc::clone(&status);
let cancel_clone = Arc::clone(&cancel);
let settle_tx = event_tx.clone();
let join = tokio::spawn(async move {
let result = run.await;
let final_status = if cancel_clone.load(Ordering::Relaxed) {
RunStatus::Aborted
} else {
match result {
Ok(()) => RunStatus::Completed,
Err(_) => RunStatus::Failed,
}
};
*status_clone.write().await = final_status;
let _ = settle_tx.send(final_status.terminal_event());
});
let handle = RunHandle {
status,
cancel,
join,
task,
events: event_tx,
};
self.runs.write().await.insert(id, handle);
id
}
pub async fn start(&self, task: String, config: crate::config::Config) -> RunId {
let task_for_run = task.clone();
let (event_tx, _event_rx) =
tokio::sync::broadcast::channel::<AgentEvent>(EVENT_CHANNEL_CAPACITY);
let emitter: Arc<dyn EventEmitter> = Arc::new(BroadcastEmitter::new(event_tx.clone()));
let task_for_spawn = task.clone();
let id = self
.spawn_with_events(task_for_spawn, event_tx, async move {
let mut agent = crate::agent::Agent::new(config).await?;
agent = agent.with_event_emitter(emitter);
agent.run_task(&task_for_run).await
})
.await;
id
}
pub async fn abort(&self, id: &RunId) -> bool {
let runs = self.runs.write().await;
let Some(handle) = runs.get(id) else {
return false;
};
handle.cancel.store(true, Ordering::Relaxed);
handle.join.abort();
let was_terminal = {
let mut st = handle.status.write().await;
let prev = *st;
*st = RunStatus::Aborted;
prev.is_terminal()
};
if !was_terminal {
let _ = handle.events.send(RunStatus::Aborted.terminal_event());
}
true
}
pub async fn status(&self, id: &RunId) -> Option<RunStatus> {
let runs = self.runs.read().await;
match runs.get(id) {
Some(h) => {
let st = *h.status.read().await;
Some(st)
}
None => None,
}
}
pub async fn list(&self) -> Vec<(RunId, RunStatus)> {
let runs = self.runs.read().await;
let mut out = Vec::with_capacity(runs.len());
for (id, handle) in runs.iter() {
let st = *handle.status.read().await;
out.push((*id, st));
}
out
}
pub async fn attach(&self, id: &RunId) -> Option<tokio::sync::broadcast::Receiver<AgentEvent>> {
let runs = self.runs.read().await;
runs.get(id).map(|h| h.events.subscribe())
}
pub async fn wait_for_terminal(
&self,
id: &RunId,
mut rx: tokio::sync::broadcast::Receiver<AgentEvent>,
mut on_event: impl FnMut(AgentEvent),
) -> Option<RunStatus> {
let mut poll = tokio::time::interval(STATUS_POLL_INTERVAL);
poll.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Skip);
loop {
match self.status(id).await {
Some(st) if st.is_terminal() => return Some(st),
Some(_) => {}
None => return None,
}
tokio::select! {
ev = rx.recv() => {
match ev {
Ok(ev) => on_event(ev),
Err(tokio::sync::broadcast::error::RecvError::Closed) => {
return self.status(id).await;
}
Err(tokio::sync::broadcast::error::RecvError::Lagged(_)) => continue,
}
}
_ = poll.tick() => {}
}
}
}
pub async fn pause(&self, _id: &RunId) -> anyhow::Result<()> {
anyhow::bail!("pause not yet implemented (inc2)")
}
pub async fn resume(&self, _id: &RunId) -> anyhow::Result<()> {
anyhow::bail!("resume not yet implemented (inc2)")
}
#[allow(dead_code)]
pub(crate) async fn emit_event(&self, id: &RunId, event: AgentEvent) -> bool {
let runs = self.runs.read().await;
match runs.get(id) {
Some(h) => {
let _ = h.events.send(event);
true
}
None => false,
}
}
}
#[cfg(test)]
#[path = "../../tests/unit/supervision/run_supervisor/run_supervisor_test.rs"]
mod tests;