use anyhow::Result;
use tokio::io::{AsyncRead, AsyncWrite, AsyncWriteExt};
use super::scheduler::scheduler;
pub(super) fn terminal_result_event(
store: &mermaid_runtime::RuntimeStore,
task: &mermaid_runtime::TaskRecord,
) -> mermaid_domain::RunEvent {
let errors = match task.status {
mermaid_runtime::TaskStatus::Completed => Vec::new(),
status => vec![format!("task ended {status}")],
};
let session_id = task.conversation_id.clone().unwrap_or_default();
let total_tokens = (!session_id.is_empty())
.then(|| store.sessions().get(&session_id).ok().flatten())
.flatten()
.and_then(|session| session.total_tokens)
.and_then(|total| u64::try_from(total).ok())
.unwrap_or(0);
mermaid_domain::RunEvent::Result {
response: task.final_report.clone().unwrap_or_default(),
reasoning: None,
total_tokens,
errors,
session_id,
structured_output: None,
}
}
pub(super) const MAX_CATCH_UP_EVENTS: usize = 1_000;
pub(super) fn catch_up_events(task: &mermaid_runtime::TaskRecord) -> Vec<mermaid_domain::RunEvent> {
let Some(session_id) = task.conversation_id.as_deref().filter(|id| !id.is_empty()) else {
return Vec::new();
};
let started = mermaid_domain::RunEvent::SessionStarted {
protocol_version: mermaid_domain::RUN_EVENT_PROTOCOL_VERSION,
cli_version: env!("CARGO_PKG_VERSION").to_string(),
model: task.model_id.clone(),
task_id: Some(task.id.clone()),
session_id: session_id.to_string(),
};
if !std::path::Path::new(&task.project_path).is_dir() {
return vec![started];
}
let events = match crate::session::ConversationManager::new(&task.project_path)
.and_then(|manager| manager.read_session_events(session_id))
{
Ok(Some(events)) => events,
Ok(None) => return vec![started],
Err(error) => {
tracing::warn!(task = %task.id, %error, "session log unreadable; attaching without catch-up");
return vec![started];
},
};
let mut replay = mermaid_domain::RunEvent::catch_up(&events);
if let Some(dropped) = replay
.len()
.checked_sub(MAX_CATCH_UP_EVENTS)
.filter(|n| *n > 0)
{
tracing::info!(task = %task.id, dropped, "catch-up over the cap; replaying the tail");
replay.drain(..dropped);
}
std::iter::once(started).chain(replay).collect()
}
pub(super) async fn handle_subscribe_stream<S>(
mut stream: S,
request: crate::runtime_client::DaemonRequest,
authorized: bool,
) -> Result<()>
where
S: AsyncRead + AsyncWrite + Unpin,
{
const WRITE_TIMEOUT: std::time::Duration = std::time::Duration::from_secs(5);
async fn write_line<S: AsyncWrite + Unpin>(stream: &mut S, line: &str) -> Result<()> {
tokio::time::timeout(WRITE_TIMEOUT, async {
stream.write_all(line.as_bytes()).await?;
stream.write_all(b"\n").await?;
Ok::<(), std::io::Error>(())
})
.await
.map_err(|_| anyhow::anyhow!("subscriber too slow; dropping connection"))??;
Ok(())
}
let crate::runtime_client::DaemonRequest::SubscribeTask { task_id } = request else {
anyhow::bail!("handle_subscribe_stream called with a non-subscribe request");
};
if !authorized {
write_line(
&mut stream,
&serde_json::json!({
"ok": false,
"error": "unauthorized: set MERMAID_DAEMON_TOKEN or include auth.token",
})
.to_string(),
)
.await?;
return Ok(());
}
let store = mermaid_runtime::RuntimeStore::open_default()?;
let Some(task) = store.tasks().get(&task_id)? else {
write_line(
&mut stream,
&serde_json::json!({"ok": false, "error": format!("task not found: {task_id}")})
.to_string(),
)
.await?;
return Ok(());
};
let mut rx = scheduler().stream_for(&task_id).subscribe();
let task = store.tasks().get(&task_id)?.unwrap_or(task);
let catch_up = {
let task = task.clone();
tokio::task::spawn_blocking(move || catch_up_events(&task))
.await
.unwrap_or_default()
};
write_line(
&mut stream,
&serde_json::json!({
"ok": true,
"subscribed": task_id,
"status": task.status.to_string(),
"protocol_version": mermaid_domain::RUN_EVENT_PROTOCOL_VERSION,
"replayed": catch_up.len(),
})
.to_string(),
)
.await?;
for event in &catch_up {
write_line(&mut stream, &serde_json::to_string(event)?).await?;
}
let terminal_status = matches!(
task.status,
mermaid_runtime::TaskStatus::Completed
| mermaid_runtime::TaskStatus::Failed
| mermaid_runtime::TaskStatus::Cancelled
);
if terminal_status {
let event = terminal_result_event(&store, &task);
write_line(&mut stream, &serde_json::to_string(&event)?).await?;
stream.shutdown().await?;
return Ok(());
}
loop {
match rx.recv().await {
Ok(event) => {
let terminal = matches!(event, mermaid_domain::RunEvent::Result { .. });
write_line(&mut stream, &serde_json::to_string(&event)?).await?;
if terminal {
break;
}
},
Err(tokio::sync::broadcast::error::RecvError::Lagged(n)) => {
let event = mermaid_domain::RunEvent::Error {
message: format!("subscriber lagged; {n} events dropped"),
};
write_line(&mut stream, &serde_json::to_string(&event)?).await?;
},
Err(tokio::sync::broadcast::error::RecvError::Closed) => break,
}
}
stream.shutdown().await?;
Ok(())
}