use anyhow::Result;
pub(super) struct Scheduler {
pub(super) permits: std::sync::Arc<tokio::sync::Semaphore>,
pub(super) running:
std::sync::Mutex<std::collections::HashMap<String, tokio_util::sync::CancellationToken>>,
pub(super) wake: tokio::sync::Notify,
pub(super) task_timeout: Option<std::time::Duration>,
pub(super) streams: std::sync::Mutex<
std::collections::HashMap<String, tokio::sync::broadcast::Sender<mermaid_domain::RunEvent>>,
>,
pub(super) mailboxes: std::sync::Mutex<
std::collections::HashMap<String, crate::engine::EngineHandle<mermaid_domain::RunEvent>>,
>,
}
impl Scheduler {
pub(super) fn stream_for(
&self,
task_id: &str,
) -> tokio::sync::broadcast::Sender<mermaid_domain::RunEvent> {
self.streams
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner())
.entry(task_id.to_string())
.or_insert_with(|| tokio::sync::broadcast::channel(1024).0)
.clone()
}
pub(super) fn register_mailbox(
&self,
task_id: &str,
handle: crate::engine::EngineHandle<mermaid_domain::RunEvent>,
) {
self.mailboxes
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.insert(task_id.to_string(), handle);
}
pub(super) fn mailbox_for(
&self,
task_id: &str,
) -> Option<crate::engine::EngineHandle<mermaid_domain::RunEvent>> {
self.mailboxes
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.get(task_id)
.cloned()
}
pub(super) fn drop_mailbox(&self, task_id: &str) {
self.mailboxes
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner)
.remove(task_id);
}
pub(super) fn drop_stream(&self, task_id: &str) {
self.streams
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner())
.remove(task_id);
}
}
pub(super) static SCHEDULER: std::sync::OnceLock<Scheduler> = std::sync::OnceLock::new();
pub(super) fn scheduler() -> &'static Scheduler {
SCHEDULER
.get()
.expect("scheduler is initialized in main before serving")
}
pub(super) async fn scheduler_drain_loop() {
let sched = scheduler();
loop {
let permit = match sched.permits.clone().acquire_owned().await {
Ok(permit) => permit,
Err(_) => return,
};
let task = loop {
let claimed =
mermaid_runtime::with_shared_store(|store| store.tasks().claim_next_queued());
match claimed {
Ok(Some(task)) => break task,
Ok(None) => {
tokio::select! {
_ = sched.wake.notified() => {},
_ = tokio::time::sleep(std::time::Duration::from_secs(3)) => {},
}
},
Err(error) => {
tracing::warn!(error = %error, "scheduler claim failed; retrying");
tokio::time::sleep(std::time::Duration::from_secs(3)).await;
},
}
};
tokio::spawn(execute_claimed_task(task, permit));
}
}
pub(super) struct RunningGuard {
pub(super) sched: &'static Scheduler,
pub(super) task_id: String,
}
impl Drop for RunningGuard {
fn drop(&mut self) {
self.sched
.running
.lock()
.unwrap_or_else(|poisoned| poisoned.into_inner())
.remove(&self.task_id);
}
}
pub(super) async fn execute_claimed_task(
task: mermaid_runtime::TaskRecord,
permit: tokio::sync::OwnedSemaphorePermit,
) {
let _permit = permit;
let sched = scheduler();
let Some(prompt) = task.prompt.clone() else {
persist_terminal_status(
&task.id,
mermaid_runtime::TaskStatus::Failed,
"task has no persisted prompt",
)
.await;
return;
};
let token = tokio_util::sync::CancellationToken::new();
sched
.running
.lock()
.expect("scheduler running map poisoned")
.insert(task.id.clone(), token.clone());
let _running_guard = RunningGuard {
sched,
task_id: task.id.clone(),
};
let _ = mermaid_runtime::run_plugin_hooks(
"task_start",
&serde_json::json!({
"id": task.id.clone(),
"title": task.title.clone(),
"project_path": task.project_path.clone(),
"model_id": task.model_id.clone(),
}),
);
let config = crate::app::load_config().unwrap_or_default();
let event_tx = sched.stream_for(&task.id);
let _backlink = tokio::spawn(early_backlink(event_tx.subscribe(), task.id.clone()));
let (handle_tx, mut handle_rx) = tokio::sync::mpsc::channel(1);
let _mailbox = tokio::spawn({
let task_id = task.id.clone();
async move {
if let Some(handle) = handle_rx.recv().await {
scheduler().register_mailbox(&task_id, handle);
}
}
});
let result = crate::app::run_non_interactive_with(
config,
std::path::PathBuf::from(&task.project_path),
task.model_id.clone(),
prompt,
crate::app::RunOptions {
task_id: Some(task.id.clone()),
cancel: Some(token.clone()),
deadline: sched.task_timeout,
event_tx: Some(event_tx),
handle_tx: Some(handle_tx),
..crate::app::RunOptions::default()
},
)
.await;
link_completed_session(&task, &result);
let (status, report, hook_status) = classify_run_result(token.is_cancelled(), result);
persist_terminal_status(&task.id, status, &report).await;
sched.drop_mailbox(&task.id);
sched.drop_stream(&task.id);
record_terminal_outcome(&task, status);
let _ = mermaid_runtime::run_plugin_hooks(
"task_stop",
&serde_json::json!({
"id": task.id.clone(),
"status": hook_status,
"final_report": report.clone(),
}),
);
}
pub(super) fn link_completed_session(
task: &mermaid_runtime::TaskRecord,
result: &Result<crate::app::RunResult, anyhow::Error>,
) {
let session_id = result
.as_ref()
.ok()
.map(|run| run.session_id.clone())
.filter(|id| !id.is_empty());
if let Some(session_id) = session_id {
let task_id = task.id.clone();
let _ = mermaid_runtime::with_shared_store(move |store| {
store.tasks().set_conversation(&task_id, &session_id)
});
}
}
pub(super) async fn early_backlink(
mut events: tokio::sync::broadcast::Receiver<mermaid_domain::RunEvent>,
task_id: String,
) {
loop {
match events.recv().await {
Ok(mermaid_domain::RunEvent::SessionStarted { session_id, .. }) => {
if session_id.is_empty() {
return;
}
let owned = task_id.clone();
if let Err(error) = mermaid_runtime::with_shared_store(move |store| {
store.tasks().set_conversation(&owned, &session_id)
}) {
tracing::warn!(task = %task_id, %error, "could not stamp the session backlink at run start");
}
return;
},
Ok(_) | Err(tokio::sync::broadcast::error::RecvError::Lagged(_)) => {},
Err(tokio::sync::broadcast::error::RecvError::Closed) => return,
}
}
}
pub(super) async fn persist_terminal_status(
task_id: &str,
status: mermaid_runtime::TaskStatus,
report: &str,
) {
const MAX_ATTEMPTS: usize = 3;
for attempt in 1..=MAX_ATTEMPTS {
match mermaid_runtime::RuntimeStore::open_default() {
Ok(store) => match store.tasks().update_status(task_id, status, Some(report)) {
Ok(()) => return,
Err(error) => tracing::error!(
task_id,
attempt,
error = %error,
"failed to persist terminal task status; retrying"
),
},
Err(error) => tracing::error!(
task_id,
attempt,
error = %error,
"failed to open store to persist terminal task status; retrying"
),
}
}
tracing::error!(
task_id,
"gave up persisting terminal task status after retries; it may be reconciled as failed on the next daemon restart"
);
}
pub(super) fn classify_run_result<E: std::fmt::Display>(
cancelled: bool,
result: std::result::Result<crate::app::RunResult, E>,
) -> (mermaid_runtime::TaskStatus, String, &'static str) {
use mermaid_runtime::TaskStatus;
if cancelled {
return (
TaskStatus::Cancelled,
"cancelled by user".to_string(),
"cancelled",
);
}
match result {
Ok(run) if run.errors.is_empty() && !run.response.trim().is_empty() => {
(TaskStatus::Completed, run.response, "completed")
},
Ok(run) if run.errors.is_empty() => (
TaskStatus::Failed,
"model returned an empty response".to_string(),
"failed",
),
Ok(run) => (TaskStatus::Failed, run.errors.join("\n"), "failed"),
Err(err) => (TaskStatus::Failed, err.to_string(), "failed"),
}
}
pub(super) fn record_terminal_outcome(
task: &mermaid_runtime::TaskRecord,
status: mermaid_runtime::TaskStatus,
) {
use mermaid_runtime::{
NewOutcome, OUTCOME_LABEL_FAILURE, OUTCOME_LABEL_SUCCESS, OUTCOME_LABEL_UNKNOWN,
OUTCOME_SOURCE_SYSTEM, RuntimeStore, TaskStatus,
};
let (label, reward) = match status {
TaskStatus::Completed => (OUTCOME_LABEL_SUCCESS, 1.0),
TaskStatus::Failed => (OUTCOME_LABEL_FAILURE, -1.0),
_ => (OUTCOME_LABEL_UNKNOWN, 0.0),
};
let detail_json = serde_json::to_string(&serde_json::json!({
"prompt": task.prompt,
"model_id": task.model_id,
"conversation_id": task.conversation_id,
"label": label,
}))
.ok();
let store = match RuntimeStore::open_default() {
Ok(store) => store,
Err(error) => {
tracing::warn!(task_id = %task.id, error = %error, "failed to open store to record terminal outcome");
return;
},
};
if let Err(error) = store.outcomes().record(NewOutcome {
id: None,
task_id: Some(task.id.clone()),
tool_run_id: None,
kind: "task_terminal".to_string(),
label: label.to_string(),
reward: Some(reward),
source: OUTCOME_SOURCE_SYSTEM.to_string(),
detail_json,
}) {
tracing::warn!(task_id = %task.id, error = %error, "failed to record terminal outcome");
}
}
pub(super) fn task_title_from_prompt(prompt: &str) -> String {
let one_line = prompt.split_whitespace().collect::<Vec<_>>().join(" ");
if one_line.is_empty() {
return "daemon task".to_string();
}
if one_line.len() <= 80 {
return one_line;
}
let end = one_line.floor_char_boundary(80);
format!("{}...", &one_line[..end])
}