diff --git a/Cargo.lock b/Cargo.lock index 0986b55a6..0fdc7afc9 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -1889,6 +1889,7 @@ dependencies = [ "tempfile", "thiserror 2.0.18", "tokio", + "tokio-tungstenite 0.26.2", "tokio-util", "toml 0.8.23", "toml_edit", @@ -2408,6 +2409,7 @@ dependencies = [ "thiserror 2.0.18", "tokio", "tokio-stream", + "tokio-tungstenite 0.26.2", "tokio-util", "toml 0.8.23", "toml_edit", diff --git a/lib/crates/fabro-cli/Cargo.toml b/lib/crates/fabro-cli/Cargo.toml index a4d15b5b5..d0b8b7373 100644 --- a/lib/crates/fabro-cli/Cargo.toml +++ b/lib/crates/fabro-cli/Cargo.toml @@ -61,6 +61,7 @@ anyhow.workspace = true miette.workspace = true dotenvy.workspace = true tokio.workspace = true +tokio-tungstenite.workspace = true tracing.workspace = true tracing-subscriber.workspace = true tracing-appender.workspace = true @@ -128,3 +129,4 @@ fabro-test = { workspace = true } fabro-macros = { path = "../fabro-macros" } hkdf.workspace = true reqwest = { workspace = true, features = ["cookies"] } +tokio = { workspace = true, features = ["test-util", "macros"] } diff --git a/lib/crates/fabro-cli/src/commands/run/runner.rs b/lib/crates/fabro-cli/src/commands/run/runner.rs index fb3b31d43..ac1a2bbd0 100644 --- a/lib/crates/fabro-cli/src/commands/run/runner.rs +++ b/lib/crates/fabro-cli/src/commands/run/runner.rs @@ -1,10 +1,4 @@ -#![expect( - clippy::disallowed_types, - reason = "sync CLI `run` subprocess wrapper: reads server subprocess stdout line-by-line via \ - std::io::BufReader; not on a Tokio path" -)] - -use std::io::{BufRead as StdBufRead, BufReader as StdBufReader}; +use std::collections::{HashSet, VecDeque}; use std::path::{Path, PathBuf}; use std::sync::Arc; use std::time::Duration; @@ -12,10 +6,14 @@ use std::time::Duration; use anyhow::{Context, Result, anyhow}; use async_trait::async_trait; use fabro_api::types::RunManifest; +use fabro_client::ServerTarget; use fabro_config::user::active_settings_path; use fabro_config::{ServerSettingsBuilder, Storage, load_llm_catalog_settings}; use fabro_interview::{ - AnswerSubmission, ControlInterviewer, WorkerControlEnvelope, WorkerControlMessage, + AnswerSubmission, ControlInterviewer, WORKER_CONTROL_INVALID_CURSOR_REASON, + WORKER_CONTROL_PONG_TIMEOUT_REASON, WORKER_CONTROL_WS_LIVENESS_TIMEOUT, + WORKER_CONTROL_WS_PING_INTERVAL, WorkerControlDeliveryFrame, WorkerControlEnvelope, + WorkerControlMessage, }; use fabro_model::Catalog; use fabro_server::run_tool_manifest; @@ -34,11 +32,23 @@ use fabro_workflow::operations::{self, StartServices}; use fabro_workflow::run_control::RunControlState; use fabro_workflow::runtime_store::{RunStoreBackend, RunStoreHandle}; use fabro_workflow::services::FabroRunToolServices; +use futures::{SinkExt, StreamExt}; use jsonwebtoken::dangerous::insecure_decode; +#[cfg(test)] +use tokio::io::DuplexStream; +use tokio::net::TcpStream; +#[cfg(unix)] +use tokio::net::UnixStream; #[cfg(unix)] use tokio::signal::unix::{SignalKind, signal}; -use tokio::sync::{Mutex, RwLock as AsyncRwLock, mpsc}; -use tokio::time::sleep; +use tokio::sync::{Mutex, RwLock as AsyncRwLock, oneshot}; +use tokio::task::JoinHandle; +use tokio::time::{self, Instant, MissedTickBehavior}; +use tokio_tungstenite::tungstenite::client::IntoClientRequest; +use tokio_tungstenite::tungstenite::http::{HeaderValue, Request as WebSocketRequest, header}; +use tokio_tungstenite::tungstenite::protocol::frame::coding::CloseCode; +use tokio_tungstenite::tungstenite::protocol::{self, Message as WebSocketMessage}; +use tokio_tungstenite::{MaybeTlsStream, WebSocketStream, connect_async, tungstenite}; use tokio_util::sync::CancellationToken; use crate::args::RunWorkerMode; @@ -75,7 +85,7 @@ pub(crate) async fn execute( let _ = fabro_proc::title_init(); set_worker_title(&run_id, initial_worker_title_phase(mode)); - let target = server.parse::()?; + let target = server.parse::()?; let client = server_client::connect_server_target_with_bearer(&target, worker_token).await?; let run_store = HttpRunStore::connect(run_id, client.clone_for_reuse()).await?; let run_state = run_store @@ -110,13 +120,24 @@ pub(crate) async fn execute( let cancel_token = CancellationToken::new(); let emitter = Arc::new(Emitter::new(run_id)); let steering_hub = Arc::new(fabro_workflow::SteeringHub::new(Arc::clone(&emitter))); - spawn_worker_control_stream( - Arc::clone(&interviewer), - cancel_token.clone(), - Arc::clone(&steering_hub), - )?; let run_control = RunControlState::new(); install_signal_handlers(Arc::clone(&run_control), cancel_token.clone())?; + let mut control_manager = if run_state.status.is_terminal() { + None + } else { + Some(spawn_worker_control_manager( + target.clone(), + run_id, + worker_token.to_owned(), + Arc::clone(&interviewer), + cancel_token.clone(), + Arc::clone(&steering_hub), + Arc::clone(&run_control), + )) + }; + if let Some(control_manager) = &mut control_manager { + control_manager.wait_for_first_connection().await?; + } let vault = load_worker_vault(storage_dir.as_deref())?; let github_app = { let vault_guard = match &vault { @@ -158,13 +179,26 @@ pub(crate) async fn execute( fabro_run_tools, }; - match mode { - RunWorkerMode::Start => { - operations::start(&run_dir, services).await?; + let execution = async { + match mode { + RunWorkerMode::Start => operations::start(&run_dir, services).await, + RunWorkerMode::Resume => operations::resume(&run_dir, services).await, } - RunWorkerMode::Resume => { - operations::resume(&run_dir, services).await?; + }; + + if let Some(mut control_manager) = control_manager { + tokio::select! { + result = execution => { + control_manager.finish(); + result?; + } + fatal = control_manager.fatal_control_loss() => { + control_manager.finish(); + return Err(fatal); + } } + } else { + execution.await?; } Ok(()) @@ -254,96 +288,510 @@ fn load_worker_vault(storage_dir: Option<&Path>) -> Result, + order: VecDeque, + recent: HashSet, } -#[expect( - clippy::disallowed_methods, - reason = "Worker control reads blocking stdin on a dedicated OS thread and forwards lines into Tokio." -)] -fn spawn_worker_control_stream( +impl AppliedWorkerControlDeliveryIds { + fn last_applied_id(&self) -> Option<&str> { + self.last.as_deref() + } + + fn contains(&self, id: &str) -> bool { + self.recent.contains(id) + } + + fn record(&mut self, id: String) { + if !self.recent.insert(id.clone()) { + self.last = Some(id); + return; + } + self.order.push_back(id.clone()); + self.last = Some(id); + while self.order.len() > WORKER_CONTROL_APPLIED_ID_DEDUPE_CAPACITY { + if let Some(evicted) = self.order.pop_front() { + self.recent.remove(&evicted); + } + } + } +} + +struct WorkerControlManagerHandle { + first_connection: Option>>, + fatal: Option>, + done: CancellationToken, + task: JoinHandle<()>, +} + +impl WorkerControlManagerHandle { + async fn wait_for_first_connection(&mut self) -> Result<()> { + let receiver = self + .first_connection + .take() + .context("worker control first-connection receiver missing")?; + receiver + .await + .context("worker control manager stopped before first connection")? + } + + async fn fatal_control_loss(&mut self) -> anyhow::Error { + let Some(receiver) = self.fatal.take() else { + return anyhow!("worker control fatal receiver missing"); + }; + receiver + .await + .unwrap_or_else(|_| anyhow!("worker control manager stopped before workflow completed")) + } + + fn finish(&self) { + self.done.cancel(); + self.task.abort(); + } +} + +#[derive(Debug)] +struct WorkerControlStreamConnectRequest { + request: WebSocketRequest<()>, + unix_socket_path: Option, +} + +impl WorkerControlStreamConnectRequest { + fn request_for_tungstenite(&self) -> WebSocketRequest<()> { + self.request.clone() + } + + #[cfg(test)] + fn uri(&self) -> String { + self.request.uri().to_string() + } + + #[cfg(test)] + fn authorization(&self) -> Option<&str> { + self.request + .headers() + .get(header::AUTHORIZATION) + .and_then(|value| value.to_str().ok()) + } +} + +enum WorkerControlSocket { + Tcp(Box>>), + #[cfg(unix)] + Unix(Box>), + #[cfg(test)] + Test(Box>), +} + +impl WorkerControlSocket { + async fn send(&mut self, message: WebSocketMessage) -> Result<(), tungstenite::Error> { + match self { + Self::Tcp(socket) => socket.send(message).await, + #[cfg(unix)] + Self::Unix(socket) => socket.send(message).await, + #[cfg(test)] + Self::Test(socket) => socket.send(message).await, + } + } + + async fn next(&mut self) -> Option> { + match self { + Self::Tcp(socket) => socket.next().await, + #[cfg(unix)] + Self::Unix(socket) => socket.next().await, + #[cfg(test)] + Self::Test(socket) => socket.next().await, + } + } +} + +#[derive(Debug)] +enum WorkerControlConnectError { + InvalidCursor, + Other(anyhow::Error), +} + +fn spawn_worker_control_manager( + target: ServerTarget, + run_id: RunId, + worker_token: String, interviewer: Arc, cancel_token: CancellationToken, steering_hub: Arc, -) -> Result<()> { - let (event_tx, event_rx) = mpsc::unbounded_channel(); - tokio::spawn(handle_worker_control_stream_events( - interviewer, - cancel_token, - steering_hub, - event_rx, - )); - std::thread::Builder::new() - .name("fabro-worker-control".to_string()) - .spawn(move || { - read_worker_control_stream_blocking(StdBufReader::new(std::io::stdin()), &event_tx); - }) - .context("failed to spawn worker control reader thread")?; - Ok(()) + run_control: Arc, +) -> WorkerControlManagerHandle { + let (first_tx, first_rx) = oneshot::channel(); + let (fatal_tx, fatal_rx) = oneshot::channel(); + let done = CancellationToken::new(); + let task_done = done.clone(); + let task = tokio::spawn(async move { + run_worker_control_manager( + target, + run_id, + worker_token, + interviewer, + cancel_token, + steering_hub, + run_control, + task_done, + first_tx, + fatal_tx, + ) + .await; + }); + WorkerControlManagerHandle { + first_connection: Some(first_rx), + fatal: Some(fatal_rx), + done, + task, + } } -fn read_worker_control_stream_blocking( - mut reader: R, - event_tx: &mpsc::UnboundedSender, -) where - R: StdBufRead, -{ - let mut line = String::new(); - loop { - line.clear(); - match reader.read_line(&mut line) { - Ok(0) | Err(_) => { - let _ = event_tx.send(WorkerControlStreamEvent::Eof); - break; +#[allow( + clippy::too_many_arguments, + reason = "Worker control manager owns the worker-side control dependencies." +)] +async fn run_worker_control_manager( + target: ServerTarget, + run_id: RunId, + worker_token: String, + interviewer: Arc, + cancel_token: CancellationToken, + steering_hub: Arc, + run_control: Arc, + done: CancellationToken, + first_tx: oneshot::Sender>, + fatal_tx: oneshot::Sender, +) { + let mut first_tx = Some(first_tx); + let mut fatal_tx = Some(fatal_tx); + let mut backoff = WORKER_CONTROL_RECONNECT_INITIAL_BACKOFF; + let mut applied_ids = AppliedWorkerControlDeliveryIds::default(); + + while !done.is_cancelled() { + let request = match build_worker_control_stream_request( + &target, + &run_id, + &worker_token, + applied_ids.last_applied_id(), + ) { + Ok(request) => request, + Err(err) => { + report_fatal_control_loss( + &interviewer, + &cancel_token, + &mut first_tx, + &mut fatal_tx, + format!("failed to build worker control stream request: {err:#}"), + ) + .await; + return; } - Ok(_) => { - let line = line.trim_end_matches(['\r', '\n']).to_string(); - if event_tx.send(WorkerControlStreamEvent::Line(line)).is_err() { - break; + }; + + match connect_worker_control_stream(request).await { + Ok(mut socket) => { + if let Some(first_tx) = first_tx.take() { + let _ = first_tx.send(Ok(())); + } + backoff = WORKER_CONTROL_RECONNECT_INITIAL_BACKOFF; + match handle_worker_control_socket( + &mut socket, + &interviewer, + &cancel_token, + &steering_hub, + &run_control, + &mut applied_ids, + &done, + ) + .await + { + Ok(()) => {} + Err(WorkerControlConnectError::InvalidCursor) => { + report_fatal_control_loss( + &interviewer, + &cancel_token, + &mut first_tx, + &mut fatal_tx, + "worker control stream replay cursor is invalid".to_string(), + ) + .await; + return; + } + Err(WorkerControlConnectError::Other(err)) => { + tracing::debug!(error = %err, "Worker control stream disconnected"); + } + } + } + Err(WorkerControlConnectError::InvalidCursor) => { + report_fatal_control_loss( + &interviewer, + &cancel_token, + &mut first_tx, + &mut fatal_tx, + "worker control stream replay cursor is invalid".to_string(), + ) + .await; + return; + } + Err(WorkerControlConnectError::Other(err)) => { + tracing::debug!(error = %err, "Worker control stream connection failed"); + } + } + + sleep_or_done(&done, backoff).await; + backoff = next_worker_control_reconnect_backoff(backoff); + } +} + +async fn report_fatal_control_loss( + interviewer: &ControlInterviewer, + cancel_token: &CancellationToken, + first_tx: &mut Option>>, + fatal_tx: &mut Option>, + detail: String, +) { + let message = format!("worker control channel lost: {detail}"); + interviewer.interrupt_all().await; + if let Some(first_tx) = first_tx.take() { + let _ = first_tx.send(Err(anyhow!(message.clone()))); + } + if let Some(fatal_tx) = fatal_tx.take() { + let _ = fatal_tx.send(anyhow!(message)); + } + cancel_token.cancel(); +} + +async fn sleep_or_done(done: &CancellationToken, delay: Duration) { + tokio::select! { + () = done.cancelled() => {} + () = time::sleep(delay) => {} + } +} + +fn next_worker_control_reconnect_backoff(current: Duration) -> Duration { + current + .saturating_mul(2) + .min(WORKER_CONTROL_RECONNECT_MAX_BACKOFF) +} + +fn build_worker_control_stream_request( + target: &ServerTarget, + run_id: &RunId, + worker_token: &str, + after: Option<&str>, +) -> Result { + let (url, unix_socket_path) = match target { + ServerTarget::HttpUrl(_) => { + let base = target + .as_http_url() + .context("HTTP server target missing URL")?; + let websocket_base = if let Some(rest) = base.strip_prefix("http://") { + format!("ws://{rest}") + } else if let Some(rest) = base.strip_prefix("https://") { + format!("wss://{rest}") + } else { + anyhow::bail!("unsupported server URL scheme"); + }; + ( + worker_control_stream_url(&websocket_base, run_id, after), + None, + ) + } + ServerTarget::UnixSocket(path) => { + let url = worker_control_stream_url("ws://fabro", run_id, after); + (url, Some(path.as_path().to_path_buf())) + } + }; + let mut request = url + .as_str() + .into_client_request() + .context("failed to build worker control stream request")?; + request.headers_mut().insert( + header::AUTHORIZATION, + HeaderValue::from_str(&format!("Bearer {worker_token}")) + .context("failed to build worker control stream authorization header")?, + ); + Ok(WorkerControlStreamConnectRequest { + request, + unix_socket_path, + }) +} + +fn worker_control_stream_url(base: &str, run_id: &RunId, after: Option<&str>) -> String { + let mut url = format!("{base}/api/v1/runs/{run_id}/worker/control-stream"); + if let Some(after) = after { + url.push_str("?after="); + url.push_str(after); + } + url +} + +async fn connect_worker_control_stream( + request: WorkerControlStreamConnectRequest, +) -> Result { + if let Some(path) = request.unix_socket_path.as_ref() { + #[cfg(unix)] + { + let ws_request = request.request_for_tungstenite(); + let stream = UnixStream::connect(path) + .await + .map_err(|err| WorkerControlConnectError::Other(anyhow::Error::new(err)))?; + let (socket, _) = tokio_tungstenite::client_async(ws_request, stream) + .await + .map_err(classify_tungstenite_error)?; + Ok(WorkerControlSocket::Unix(Box::new(socket))) + } + #[cfg(not(unix))] + { + let _ = path; + Err(WorkerControlConnectError::Other(anyhow!( + "Unix-socket worker control stream is not supported on this platform" + ))) + } + } else { + let (socket, _) = connect_async(request.request) + .await + .map_err(classify_tungstenite_error)?; + Ok(WorkerControlSocket::Tcp(Box::new(socket))) + } +} + +fn classify_tungstenite_error(error: tungstenite::Error) -> WorkerControlConnectError { + if let tungstenite::Error::Http(response) = &error { + if response.status().as_u16() == 410 { + return WorkerControlConnectError::InvalidCursor; + } + } + WorkerControlConnectError::Other(anyhow::Error::new(error)) +} + +async fn handle_worker_control_socket( + socket: &mut WorkerControlSocket, + interviewer: &ControlInterviewer, + cancel_token: &CancellationToken, + steering_hub: &fabro_workflow::SteeringHub, + run_control: &RunControlState, + applied_ids: &mut AppliedWorkerControlDeliveryIds, + done: &CancellationToken, +) -> Result<(), WorkerControlConnectError> { + let mut ping_interval = time::interval(WORKER_CONTROL_WS_PING_INTERVAL); + ping_interval.set_missed_tick_behavior(MissedTickBehavior::Delay); + let mut last_liveness = Instant::now(); + let liveness_timeout = time::sleep_until(last_liveness + WORKER_CONTROL_WS_LIVENESS_TIMEOUT); + tokio::pin!(liveness_timeout); + + loop { + liveness_timeout + .as_mut() + .reset(last_liveness + WORKER_CONTROL_WS_LIVENESS_TIMEOUT); + + tokio::select! { + () = done.cancelled() => return Ok(()), + _ = ping_interval.tick() => { + socket + .send(WebSocketMessage::Ping(Vec::new().into())) + .await + .map_err(|err| WorkerControlConnectError::Other(anyhow::Error::new(err)))?; + } + () = &mut liveness_timeout => { + let _ = socket + .send(WebSocketMessage::Close(Some(protocol::CloseFrame { + code: CloseCode::Away, + reason: WORKER_CONTROL_PONG_TIMEOUT_REASON.into(), + }))) + .await; + return Err(WorkerControlConnectError::Other(anyhow!( + "worker control WebSocket liveness timed out" + ))); + } + message = socket.next() => { + let Some(message) = message else { + return Ok(()); + }; + match message { + Ok(WebSocketMessage::Text(text)) => { + last_liveness = Instant::now(); + let frame = serde_json::from_str::(text.as_str()) + .map_err(|err| WorkerControlConnectError::Other(anyhow::Error::new(err)))?; + apply_worker_control_delivery_frame( + interviewer, + cancel_token, + steering_hub, + run_control, + applied_ids, + frame, + ) + .await; + } + Ok(WebSocketMessage::Ping(payload)) => { + last_liveness = Instant::now(); + socket + .send(WebSocketMessage::Pong(payload)) + .await + .map_err(|err| WorkerControlConnectError::Other(anyhow::Error::new(err)))?; + } + Ok(WebSocketMessage::Pong(_) | WebSocketMessage::Binary(_)) => { + last_liveness = Instant::now(); + } + Ok(WebSocketMessage::Close(frame)) => { + if frame.as_ref().is_some_and(|frame| { + frame.reason.as_str() == WORKER_CONTROL_INVALID_CURSOR_REASON + }) { + return Err(WorkerControlConnectError::InvalidCursor); + } + return Ok(()); + } + Ok(WebSocketMessage::Frame(_)) => {} + Err(err) => { + return Err(WorkerControlConnectError::Other(anyhow::Error::new(err))); + } } } } } } -async fn handle_worker_control_stream_events( - interviewer: Arc, - cancel_token: CancellationToken, - steering_hub: Arc, - mut event_rx: mpsc::UnboundedReceiver, -) { - while let Some(event) = event_rx.recv().await { - match event { - WorkerControlStreamEvent::Line(line) => { - apply_worker_control_line(&interviewer, &cancel_token, &steering_hub, &line).await; - } - WorkerControlStreamEvent::Eof => { - interviewer.interrupt_all().await; - return; - } - } - } - - interviewer.interrupt_all().await; -} - -async fn apply_worker_control_line( +async fn apply_worker_control_delivery_frame( interviewer: &ControlInterviewer, cancel_token: &CancellationToken, steering_hub: &fabro_workflow::SteeringHub, - line: &str, -) { - if line.trim().is_empty() { - return; + run_control: &RunControlState, + applied_ids: &mut AppliedWorkerControlDeliveryIds, + frame: WorkerControlDeliveryFrame, +) -> bool { + // Duplicate ids cannot reach us under normal operation: the server replays + // strictly after the last applied id. Guard against a server-side bug or + // reconnect race by ignoring recently-applied delivery ids. + if applied_ids.contains(&frame.id) { + return false; } + let frame_id = frame.id; + apply_worker_control_message( + interviewer, + cancel_token, + steering_hub, + run_control, + frame.envelope, + ) + .await; + applied_ids.record(frame_id); + true +} - let Ok(message) = serde_json::from_str::(line) else { - return; - }; - +async fn apply_worker_control_message( + interviewer: &ControlInterviewer, + cancel_token: &CancellationToken, + steering_hub: &fabro_workflow::SteeringHub, + run_control: &RunControlState, + message: WorkerControlEnvelope, +) { match message.message { WorkerControlMessage::InterviewAnswer { qid, answer, actor } => { let _ = interviewer @@ -354,6 +802,12 @@ async fn apply_worker_control_line( cancel_token.cancel(); interviewer.interrupt_all().await; } + WorkerControlMessage::RunPause => { + run_control.request_pause(); + } + WorkerControlMessage::RunUnpause => { + run_control.request_unpause(); + } WorkerControlMessage::Steer { text, actor } => { steering_hub.deliver_steer(text, Some(actor)); } @@ -485,7 +939,7 @@ impl HttpRunStore { Err(err) => last_error = Some(err), } if let Some(delay) = RUN_STORE_RETRY_DELAYS.get(attempt) { - sleep(*delay).await; + time::sleep(*delay).await; } } Err(last_error @@ -757,10 +1211,14 @@ fn install_signal_handlers( )] mod tests { use std::sync::Arc; + use std::time::Duration; use chrono::Utc; + use fabro_client::ServerTarget; use fabro_config::Storage; - use fabro_interview::{AnswerValue, ControlInterviewer, Interviewer, Question}; + use fabro_interview::{ + AnswerValue, ControlInterviewer, Interviewer, Question, WorkerControlEnvelope, + }; use fabro_types::run_event::{ InterviewCompletedProps, InterviewStartedProps, RunCompletedProps, RunControlEffectProps, RunFailedProps, RunStatusTransitionProps, @@ -771,12 +1229,17 @@ mod tests { }; use fabro_vault::{SecretType, Vault}; use fabro_workflow::event::RunEventSink; + use fabro_workflow::run_control::RunControlState; + use tokio::time; + use tokio_tungstenite::tungstenite::protocol::{Message as TestWebSocketMessage, Role}; use tokio_util::sync::CancellationToken; use super::{ - WorkerControlStreamEvent, WorkerTitlePhase, apply_worker_control_line, - handle_worker_control_stream_events, initial_worker_title_phase, load_worker_vault, - read_worker_control_stream_blocking, stamp_system_worker, worker_title, + AppliedWorkerControlDeliveryIds, WorkerControlConnectError, WorkerControlSocket, + WorkerTitlePhase, apply_worker_control_delivery_frame, apply_worker_control_message, + build_worker_control_stream_request, connect_worker_control_stream, + handle_worker_control_socket, initial_worker_title_phase, load_worker_vault, + next_worker_control_reconnect_backoff, stamp_system_worker, worker_title, worker_title_phase_for_event, }; use crate::args::RunWorkerMode; @@ -1018,20 +1481,28 @@ mod tests { } #[tokio::test] - async fn worker_control_line_routes_answer_by_question_id() { + async fn worker_control_routes_answer_by_question_id() { let interviewer = Arc::new(ControlInterviewer::new()); let cancel_token = CancellationToken::new(); + let run_control = RunControlState::new(); let mut question = Question::new("Approve?", QuestionType::YesNo); question.id = "q-1".to_string(); let ask_interviewer = Arc::clone(&interviewer); let answer_task = tokio::spawn(async move { ask_interviewer.ask(question).await }); let hub = test_steering_hub(); - apply_worker_control_line( + apply_worker_control_message( &interviewer, &cancel_token, &hub, - r#"{"v":1,"type":"interview.answer","qid":"q-1","answer":{"kind":"yes"},"actor":{"kind":"system","system_kind":"engine"}}"#, + &run_control, + WorkerControlEnvelope::interview_answer( + "q-1", + fabro_interview::AnswerSubmission::system( + fabro_interview::Answer::yes(), + fabro_types::SystemActorKind::Engine, + ), + ), ) .await; @@ -1041,9 +1512,10 @@ mod tests { } #[tokio::test] - async fn worker_control_line_cancel_sets_cancel_token_and_interrupts_pending_interviews() { + async fn worker_control_cancel_sets_cancel_token_and_interrupts_pending_interviews() { let interviewer = Arc::new(ControlInterviewer::new()); let cancel_token = CancellationToken::new(); + let run_control = RunControlState::new(); let mut question = Question::new("Approve?", QuestionType::YesNo); question.id = "q-1".to_string(); let ask_interviewer = Arc::clone(&interviewer); @@ -1051,11 +1523,12 @@ mod tests { tokio::task::yield_now().await; let hub = test_steering_hub(); - apply_worker_control_line( + apply_worker_control_message( &interviewer, &cancel_token, &hub, - r#"{"v":1,"type":"run.cancel"}"#, + &run_control, + WorkerControlEnvelope::cancel_run(), ) .await; @@ -1065,57 +1538,206 @@ mod tests { } #[tokio::test] - async fn blocking_worker_control_stream_emits_lines_and_eof() { - let (event_tx, mut event_rx) = tokio::sync::mpsc::unbounded_channel(); + async fn worker_control_pause_and_unpause_route_to_run_control() { + let interviewer = Arc::new(ControlInterviewer::new()); + let cancel_token = CancellationToken::new(); + let run_control = RunControlState::new(); + let hub = test_steering_hub(); - read_worker_control_stream_blocking( - std::io::Cursor::new( - b"{\"v\":1,\"type\":\"run.cancel\"}\n{\"v\":1,\"type\":\"interview.answer\",\"qid\":\"q-1\",\"answer\":{\"kind\":\"yes\"},\"actor\":{\"kind\":\"system\",\"system_kind\":\"engine\"}}\n", - ), - &event_tx, - ); + apply_worker_control_message( + &interviewer, + &cancel_token, + &hub, + &run_control, + WorkerControlEnvelope::pause_run(), + ) + .await; + assert!(run_control.pause_requested()); - assert_eq!( - event_rx.try_recv(), - Ok(WorkerControlStreamEvent::Line( - r#"{"v":1,"type":"run.cancel"}"#.to_string() - )) - ); - assert_eq!( - event_rx.try_recv(), - Ok(WorkerControlStreamEvent::Line( - r#"{"v":1,"type":"interview.answer","qid":"q-1","answer":{"kind":"yes"},"actor":{"kind":"system","system_kind":"engine"}}"# - .to_string() - )) - ); - assert_eq!(event_rx.try_recv(), Ok(WorkerControlStreamEvent::Eof)); + apply_worker_control_message( + &interviewer, + &cancel_token, + &hub, + &run_control, + WorkerControlEnvelope::unpause_run(), + ) + .await; + assert!(!run_control.pause_requested()); } #[tokio::test] - async fn worker_control_event_loop_eof_interrupts_pending_interviews() { + async fn duplicate_delivery_ids_are_not_applied_twice() { let interviewer = Arc::new(ControlInterviewer::new()); let cancel_token = CancellationToken::new(); - let mut question = Question::new("Approve?", QuestionType::YesNo); - question.id = "q-1".to_string(); - let ask_interviewer = Arc::clone(&interviewer); - let answer_task = tokio::spawn(async move { ask_interviewer.ask(question).await }); - let (event_tx, event_rx) = tokio::sync::mpsc::unbounded_channel(); - - event_tx.send(WorkerControlStreamEvent::Eof).unwrap(); - drop(event_tx); - + let run_control = RunControlState::new(); let hub = test_steering_hub(); - handle_worker_control_stream_events( - Arc::clone(&interviewer), - cancel_token.clone(), - hub, - event_rx, - ) - .await; + let mut applied_ids = AppliedWorkerControlDeliveryIds::default(); + let frame = fabro_interview::WorkerControlDeliveryFrame { + id: "local:1".to_string(), + envelope: WorkerControlEnvelope::pause_run(), + }; - let answer = answer_task.await.unwrap().answer; - assert_eq!(answer.value, AnswerValue::Interrupted); - assert!(!cancel_token.is_cancelled()); + assert!( + apply_worker_control_delivery_frame( + &interviewer, + &cancel_token, + &hub, + &run_control, + &mut applied_ids, + frame.clone(), + ) + .await + ); + assert!( + !apply_worker_control_delivery_frame( + &interviewer, + &cancel_token, + &hub, + &run_control, + &mut applied_ids, + frame, + ) + .await + ); + + assert_eq!(applied_ids.last_applied_id(), Some("local:1")); + } + + #[test] + fn worker_control_request_construction_for_http_targets() { + let run_id = fixtures::RUN_1; + let request = build_worker_control_stream_request( + &ServerTarget::http_url("http://example.com:3000").unwrap(), + &run_id, + "worker-token", + None, + ) + .unwrap(); + assert_eq!( + request.uri(), + format!("ws://example.com:3000/api/v1/runs/{run_id}/worker/control-stream") + ); + assert_eq!(request.authorization(), Some("Bearer worker-token")); + + let reconnect = build_worker_control_stream_request( + &ServerTarget::http_url("https://example.com").unwrap(), + &run_id, + "worker-token", + Some("local:42"), + ) + .unwrap(); + assert_eq!( + reconnect.uri(), + format!("wss://example.com/api/v1/runs/{run_id}/worker/control-stream?after=local:42") + ); + } + + #[cfg(unix)] + #[test] + fn worker_control_request_construction_for_unix_socket_targets() { + let run_id = fixtures::RUN_1; + let request = build_worker_control_stream_request( + &ServerTarget::unix_socket_path("/tmp/fabro.sock").unwrap(), + &run_id, + "worker-token", + Some("local:42"), + ) + .unwrap(); + + assert_eq!( + request.uri(), + format!("ws://fabro/api/v1/runs/{run_id}/worker/control-stream?after=local:42") + ); + assert_eq!( + request.unix_socket_path.as_deref(), + Some(std::path::Path::new("/tmp/fabro.sock")) + ); + } + + #[test] + fn worker_control_reconnect_backoff_is_bounded() { + assert_eq!( + next_worker_control_reconnect_backoff(Duration::from_millis(100)), + Duration::from_millis(200) + ); + assert_eq!( + next_worker_control_reconnect_backoff(Duration::from_secs(4)), + Duration::from_secs(5) + ); + assert_eq!( + next_worker_control_reconnect_backoff(Duration::from_secs(5)), + Duration::from_secs(5) + ); + } + + #[tokio::test(start_paused = true)] + async fn worker_control_socket_times_out_without_liveness() { + let (worker_io, _server_io) = tokio::io::duplex(1024); + let worker_ws = + tokio_tungstenite::WebSocketStream::from_raw_socket(worker_io, Role::Client, None) + .await; + let mut socket = WorkerControlSocket::Test(Box::new(worker_ws)); + let interviewer = Arc::new(ControlInterviewer::new()); + let cancel_token = CancellationToken::new(); + let hub = test_steering_hub(); + let run_control = RunControlState::new(); + let mut applied_ids = AppliedWorkerControlDeliveryIds::default(); + let done = CancellationToken::new(); + + let task = tokio::spawn(async move { + handle_worker_control_socket( + &mut socket, + &interviewer, + &cancel_token, + &hub, + &run_control, + &mut applied_ids, + &done, + ) + .await + }); + + tokio::task::yield_now().await; + time::advance(Duration::from_secs(46)).await; + let result = task.await.unwrap(); + assert!(matches!(result, Err(WorkerControlConnectError::Other(_)))); + } + + #[cfg(unix)] + #[tokio::test] + async fn worker_control_unix_socket_handshake_completes() { + let temp = tempfile::tempdir().unwrap(); + let socket_path = temp.path().join("fabro.sock"); + let listener = tokio::net::UnixListener::bind(&socket_path).unwrap(); + let server = tokio::spawn(async move { + let (stream, _) = listener.accept().await.unwrap(); + let mut socket = tokio_tungstenite::accept_async(stream).await.unwrap(); + while let Some(message) = futures::StreamExt::next(&mut socket).await { + match message.unwrap() { + TestWebSocketMessage::Close(_) => break, + TestWebSocketMessage::Ping(payload) => { + futures::SinkExt::send(&mut socket, TestWebSocketMessage::Pong(payload)) + .await + .unwrap(); + } + _ => {} + } + } + }); + let request = build_worker_control_stream_request( + &ServerTarget::unix_socket_path(&socket_path).unwrap(), + &fixtures::RUN_1, + "worker-token", + None, + ) + .unwrap(); + + let mut socket = connect_worker_control_stream(request).await.unwrap(); + socket + .send(TestWebSocketMessage::Close(None)) + .await + .unwrap(); + server.await.unwrap(); } #[tokio::test] diff --git a/lib/crates/fabro-cli/tests/it/cmd/runner.rs b/lib/crates/fabro-cli/tests/it/cmd/runner.rs index d7d87fb93..d5b307d35 100644 --- a/lib/crates/fabro-cli/tests/it/cmd/runner.rs +++ b/lib/crates/fabro-cli/tests/it/cmd/runner.rs @@ -794,6 +794,66 @@ fn detached_run_answers_pending_question_without_interview_scratch_files() { ))); } +#[test] +fn detached_run_cancel_reaches_worker_over_control_websocket() { + let context = auth_context(); + let run_id = unique_run_id(); + let workflow_path = context.temp_dir.join("cancel-over-control-websocket.fabro"); + let _gate = write_gated_workflow( + &workflow_path, + "cancel_over_control_websocket", + "Wait for cancellation", + ); + + let output = context + .command() + .args([ + "run", + "--detach", + "--run-id", + run_id.as_str(), + "--environment", + "local", + workflow_path.to_str().unwrap(), + ]) + .timeout(SHARED_DAEMON_TIMEOUT) + .output() + .expect("detached run should execute"); + assert!( + output.status.success(), + "detached run failed:\nstdout:\n{}\nstderr:\n{}", + String::from_utf8_lossy(&output.stdout), + String::from_utf8_lossy(&output.stderr) + ); + + let run_dir = context.find_run_dir(&run_id); + wait_for_event_names(&run_dir, &["run.running"]); + tokio::runtime::Runtime::new() + .expect("test runtime should build") + .block_on(async { + let (client, base_url) = + server_endpoint(&context.storage_dir).expect("server endpoint should exist"); + let response = client + .post(format!("{base_url}/api/v1/runs/{run_id}/cancel")) + .send() + .await + .expect("cancel request should succeed"); + assert_reqwest_status( + response, + fabro_http::StatusCode::OK, + format!("POST /api/v1/runs/{run_id}/cancel"), + ) + .await; + }); + + wait_for_status(&run_dir, &["failed"]); + let events = stored_worker_events(&run_dir); + assert!(events.iter().any(|event| matches!( + &event.body, + EventBody::RunFailed(props) if props.failure.reason == FailureReason::Cancelled + ))); +} + #[cfg(unix)] #[test] fn worker_exits_after_sigterm_cancel_even_when_stdin_stays_open() { diff --git a/lib/crates/fabro-interview/src/control_protocol.rs b/lib/crates/fabro-interview/src/control_protocol.rs index cdc820216..f76b61121 100644 --- a/lib/crates/fabro-interview/src/control_protocol.rs +++ b/lib/crates/fabro-interview/src/control_protocol.rs @@ -1,3 +1,5 @@ +use std::time::Duration; + use fabro_types::{PairId, PairMessageId, PairTarget, Principal, RunId}; use serde::{Deserialize, Serialize}; @@ -5,6 +7,25 @@ use crate::{Answer, AnswerSubmission, AnswerValue}; pub const WORKER_CONTROL_PROTOCOL_VERSION: u8 = 1; +/// Interval between worker-control WebSocket ping frames. +/// +/// Server and worker both initiate pings at this cadence; either side that +/// fails to observe inbound traffic for [`WORKER_CONTROL_WS_LIVENESS_TIMEOUT`] +/// closes the WebSocket. +pub const WORKER_CONTROL_WS_PING_INTERVAL: Duration = Duration::from_secs(15); + +/// Maximum quiet time allowed on a worker-control WebSocket before either side +/// declares the connection dead. +pub const WORKER_CONTROL_WS_LIVENESS_TIMEOUT: Duration = Duration::from_secs(45); + +/// WebSocket close-frame reason used when the server can no longer prove +/// replay correctness for the requested cursor. Workers must treat this as +/// fatal control-channel loss. +pub const WORKER_CONTROL_INVALID_CURSOR_REASON: &str = "invalid_cursor"; + +/// WebSocket close-frame reason used when the ping/pong watchdog fires. +pub const WORKER_CONTROL_PONG_TIMEOUT_REASON: &str = "pong_timeout"; + #[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] pub struct WorkerControlEnvelope { pub v: u8, @@ -33,6 +54,22 @@ impl WorkerControlEnvelope { } } + #[must_use] + pub fn pause_run() -> Self { + Self { + v: WORKER_CONTROL_PROTOCOL_VERSION, + message: WorkerControlMessage::RunPause, + } + } + + #[must_use] + pub fn unpause_run() -> Self { + Self { + v: WORKER_CONTROL_PROTOCOL_VERSION, + message: WorkerControlMessage::RunUnpause, + } + } + #[must_use] pub fn steer(text: impl Into, actor: Principal) -> Self { Self { @@ -121,6 +158,10 @@ pub enum WorkerControlMessage { }, #[serde(rename = "run.cancel")] RunCancel, + #[serde(rename = "run.pause")] + RunPause, + #[serde(rename = "run.unpause")] + RunUnpause, #[serde(rename = "run.steer")] Steer { text: String, actor: Principal }, #[serde(rename = "run.interrupt")] @@ -147,6 +188,12 @@ pub enum WorkerControlMessage { PairEnd { pair_id: PairId, actor: Principal }, } +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub struct WorkerControlDeliveryFrame { + pub id: String, + pub envelope: WorkerControlEnvelope, +} + #[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] #[serde(tag = "kind", rename_all = "snake_case")] pub enum WorkerControlAnswer { @@ -229,6 +276,26 @@ mod tests { assert_eq!(parsed, envelope); } + #[test] + fn pause_run_round_trips_through_json() { + let envelope = WorkerControlEnvelope::pause_run(); + let json = serde_json::to_string(&envelope).unwrap(); + assert_eq!(json, r#"{"v":1,"type":"run.pause"}"#); + + let parsed: WorkerControlEnvelope = serde_json::from_str(&json).unwrap(); + assert_eq!(parsed, envelope); + } + + #[test] + fn unpause_run_round_trips_through_json() { + let envelope = WorkerControlEnvelope::unpause_run(); + let json = serde_json::to_string(&envelope).unwrap(); + assert_eq!(json, r#"{"v":1,"type":"run.unpause"}"#); + + let parsed: WorkerControlEnvelope = serde_json::from_str(&json).unwrap(); + assert_eq!(parsed, envelope); + } + #[test] fn steer_append_round_trips_through_json() { let envelope = WorkerControlEnvelope::steer("try again", Principal::System { @@ -326,4 +393,20 @@ mod tests { serde_json::from_str(&serde_json::to_string(&end).unwrap()).unwrap(); assert_eq!(parsed, end); } + + #[test] + fn delivery_frame_round_trips_through_json() { + let frame = WorkerControlDeliveryFrame { + id: "local:42".to_string(), + envelope: WorkerControlEnvelope::cancel_run(), + }; + let json = serde_json::to_string(&frame).unwrap(); + assert_eq!( + json, + r#"{"id":"local:42","envelope":{"v":1,"type":"run.cancel"}}"# + ); + + let parsed: WorkerControlDeliveryFrame = serde_json::from_str(&json).unwrap(); + assert_eq!(parsed, frame); + } } diff --git a/lib/crates/fabro-interview/src/lib.rs b/lib/crates/fabro-interview/src/lib.rs index 4c050c4f7..6d8c414a1 100644 --- a/lib/crates/fabro-interview/src/lib.rs +++ b/lib/crates/fabro-interview/src/lib.rs @@ -222,8 +222,10 @@ pub use callback::CallbackInterviewer; pub use console::ConsoleInterviewer; pub use control::{ControlInterviewer, SubmitError}; pub use control_protocol::{ - WORKER_CONTROL_PROTOCOL_VERSION, WorkerControlAnswer, WorkerControlEnvelope, - WorkerControlMessage, + WORKER_CONTROL_INVALID_CURSOR_REASON, WORKER_CONTROL_PONG_TIMEOUT_REASON, + WORKER_CONTROL_PROTOCOL_VERSION, WORKER_CONTROL_WS_LIVENESS_TIMEOUT, + WORKER_CONTROL_WS_PING_INTERVAL, WorkerControlAnswer, WorkerControlDeliveryFrame, + WorkerControlEnvelope, WorkerControlMessage, }; pub use queue::QueueInterviewer; pub use recording::RecordingInterviewer; diff --git a/lib/crates/fabro-server/Cargo.toml b/lib/crates/fabro-server/Cargo.toml index 2b9568ce1..879cea7ce 100644 --- a/lib/crates/fabro-server/Cargo.toml +++ b/lib/crates/fabro-server/Cargo.toml @@ -112,6 +112,7 @@ httpmock = "0.8" serde_yaml = "0.9" tracing-subscriber.workspace = true tokio-util.workspace = true +tokio-tungstenite.workspace = true fabro-macros = { path = "../fabro-macros" } fabro-sandbox = { path = "../fabro-sandbox", features = ["test-support"] } fabro-test = { workspace = true } diff --git a/lib/crates/fabro-server/src/lib.rs b/lib/crates/fabro-server/src/lib.rs index 0e591e26d..d2a449732 100644 --- a/lib/crates/fabro-server/src/lib.rs +++ b/lib/crates/fabro-server/src/lib.rs @@ -49,6 +49,7 @@ pub mod static_files; #[cfg(any(test, feature = "test-support"))] pub mod test_support; pub mod web_auth; +mod worker_control; mod worker_token; pub use error::{ApiError, Error, Result}; diff --git a/lib/crates/fabro-server/src/principal_middleware.rs b/lib/crates/fabro-server/src/principal_middleware.rs index acc9eb250..1b037a6ab 100644 --- a/lib/crates/fabro-server/src/principal_middleware.rs +++ b/lib/crates/fabro-server/src/principal_middleware.rs @@ -58,6 +58,7 @@ pub(crate) struct RequestAuth(pub(crate) AuthContextSlot); pub(crate) struct RequiredUser(pub(crate) UserPrincipal); pub(crate) struct RequiredRunManagementActor(pub(crate) Principal); pub(crate) struct RequireRunScoped(pub(crate) RunId); +pub(crate) struct RequireWorkerRunScoped(pub(crate) RunId); pub(crate) struct RequireRunManagementTarget(pub(crate) RunId, pub(crate) Principal); pub(crate) struct RequireRunBlob(pub(crate) RunId, pub(crate) RunBlobId); pub(crate) struct RequireRunStageScoped(pub(crate) RunId, pub(crate) String); @@ -245,6 +246,23 @@ impl FromRequestParts> for RequireRunScoped { } } +impl FromRequestParts> for RequireWorkerRunScoped { + type Rejection = Response; + + async fn from_request_parts( + parts: &mut Parts, + state: &Arc, + ) -> Result { + let Path(id): Path = Path::from_request_parts(parts, state) + .await + .map_err(IntoResponse::into_response)?; + let run_id = parse_run_id_path(&id)?; + require_worker_for_run(&auth_slot_from_parts(parts), &run_id) + .map_err(IntoResponse::into_response)?; + Ok(Self(run_id)) + } +} + impl FromRequestParts> for RequireRunManagementTarget { type Rejection = Response; @@ -418,6 +436,15 @@ fn require_worker_or_user_for_run( } } +fn require_worker_for_run(slot: &AuthContextSlot, route_run_id: &RunId) -> Result<(), ApiError> { + let context = slot.0.lock().expect("auth context lock poisoned"); + match &context.principal { + Principal::Worker { run_id } if run_id == route_run_id => Ok(()), + Principal::Worker { .. } | Principal::User(_) => Err(ApiError::forbidden()), + _ => Err(auth_rejection(context.auth_status, context.auth_error_code)), + } +} + fn require_run_management_target( slot: &AuthContextSlot, route_run_id: &RunId, diff --git a/lib/crates/fabro-server/src/serve.rs b/lib/crates/fabro-server/src/serve.rs index b2eb6d91b..422906eeb 100644 --- a/lib/crates/fabro-server/src/serve.rs +++ b/lib/crates/fabro-server/src/serve.rs @@ -814,6 +814,8 @@ where http_client: None, sandbox_provider_registry: None, shutdown: shutdown.clone(), + #[cfg(test)] + worker_control_bus: None, #[cfg(any(test, feature = "test-support"))] automation_materializer_override: None, })?; diff --git a/lib/crates/fabro-server/src/server.rs b/lib/crates/fabro-server/src/server.rs index 12e05edbc..caa030797 100644 --- a/lib/crates/fabro-server/src/server.rs +++ b/lib/crates/fabro-server/src/server.rs @@ -124,7 +124,7 @@ use sha2::{Digest, Sha256}; use tempfile::NamedTempFile; use tokio::fs; use tokio::io::{AsyncBufReadExt, AsyncWriteExt, BufReader}; -use tokio::process::{ChildStderr, ChildStdin, Command}; +use tokio::process::{ChildStderr, Command}; use tokio::sync::broadcast::error::RecvError; use tokio::sync::{ Mutex as AsyncMutex, Notify, OwnedMutexGuard, RwLock as AsyncRwLock, Semaphore, broadcast, @@ -153,14 +153,15 @@ use crate::ip_allowlist::{IpAllowlistConfig, ip_allowlist_middleware}; use crate::jwt_auth::{self, AuthMode}; use crate::principal_middleware::{ AuthContextSlot, RequestAuth, RequestAuthContext, RequireRunBlob, RequireRunManagementTarget, - RequireRunScoped, RequireRunStageScoped, RequireStageArtifact, RequiredUser, - principal_middleware, + RequireRunScoped, RequireRunStageScoped, RequireStageArtifact, RequireWorkerRunScoped, + RequiredUser, principal_middleware, }; use crate::request_id::{self, RequestId}; use crate::run_files::{FilesInFlight, new_files_in_flight}; use crate::server_secrets::{LlmClientResult, ServerSecrets}; use crate::spawn_env::{apply_render_graph_env, apply_worker_env}; use crate::startup::load_startup_vault; +use crate::worker_control::{LocalWorkerControlBus, WorkerControlBus, WorkerControlBusError}; use crate::worker_token::{WorkerScopeSet, WorkerTokenKeys, issue_worker_token_with_scopes}; use crate::{ canonical_host, demo, diagnostics, run_manifest, security_headers, static_files, web_auth, @@ -272,7 +273,6 @@ enum ExecutionResult { const WORKER_CANCEL_GRACE: Duration = Duration::from_secs(5); const TERMINAL_DELETE_WORKER_GRACE: Duration = Duration::from_millis(50); -const WORKER_CONTROL_QUEUE_CAPACITY: usize = 8; const WORKER_CONTROL_ENQUEUE_TIMEOUT: Duration = Duration::from_secs(1); /// Per-model billing totals. #[derive(Default)] @@ -294,8 +294,9 @@ pub(crate) type RegistryFactoryOverride = #[derive(Clone)] enum RunAnswerTransport { - Subprocess { - control_tx: mpsc::Sender, + Worker { + run_id: RunId, + bus: Arc, }, InProcess { interviewer: Arc, @@ -317,18 +318,46 @@ enum PairTransportError { } impl RunAnswerTransport { + async fn publish_worker_control( + run_id: RunId, + bus: &Arc, + message: WorkerControlEnvelope, + ) -> Result<(), WorkerControlBusError> { + timeout(WORKER_CONTROL_ENQUEUE_TIMEOUT, bus.publish(run_id, message)) + .await + .map_err(|_| WorkerControlBusError::PublishTimeout)? + .map(|_| ()) + } + + fn answer_error_from_bus(error: &WorkerControlBusError) -> AnswerTransportError { + match error { + WorkerControlBusError::PublishTimeout => AnswerTransportError::Timeout, + WorkerControlBusError::Closed + | WorkerControlBusError::Unavailable + | WorkerControlBusError::InvalidCursor { .. } => AnswerTransportError::Closed, + } + } + + fn pair_error_from_bus(error: &WorkerControlBusError) -> PairTransportError { + match error { + WorkerControlBusError::PublishTimeout => PairTransportError::Timeout, + WorkerControlBusError::Closed + | WorkerControlBusError::Unavailable + | WorkerControlBusError::InvalidCursor { .. } => PairTransportError::Closed, + } + } + async fn submit( &self, qid: &str, submission: AnswerSubmission, ) -> Result<(), AnswerTransportError> { match self { - Self::Subprocess { control_tx } => { + Self::Worker { run_id, bus } => { let message = WorkerControlEnvelope::interview_answer(qid.to_string(), submission); - timeout(WORKER_CONTROL_ENQUEUE_TIMEOUT, control_tx.send(message)) + Self::publish_worker_control(*run_id, bus, message) .await - .map_err(|_| AnswerTransportError::Timeout)? - .map_err(|_| AnswerTransportError::Closed) + .map_err(|err| Self::answer_error_from_bus(&err)) } Self::InProcess { interviewer, .. } => interviewer .submit(qid, submission) @@ -339,12 +368,11 @@ impl RunAnswerTransport { async fn cancel_run(&self) -> Result<(), AnswerTransportError> { match self { - Self::Subprocess { control_tx } => { + Self::Worker { run_id, bus } => { let message = WorkerControlEnvelope::cancel_run(); - timeout(WORKER_CONTROL_ENQUEUE_TIMEOUT, control_tx.send(message)) + Self::publish_worker_control(*run_id, bus, message) .await - .map_err(|_| AnswerTransportError::Timeout)? - .map_err(|_| AnswerTransportError::Closed) + .map_err(|err| Self::answer_error_from_bus(&err)) } Self::InProcess { interviewer, .. } => { interviewer.cancel_all().await; @@ -357,12 +385,11 @@ impl RunAnswerTransport { /// in-process steering hub. async fn steer(&self, text: String, actor: Principal) -> Result<(), AnswerTransportError> { match self { - Self::Subprocess { control_tx } => { + Self::Worker { run_id, bus } => { let message = WorkerControlEnvelope::steer(text, actor); - timeout(WORKER_CONTROL_ENQUEUE_TIMEOUT, control_tx.send(message)) + Self::publish_worker_control(*run_id, bus, message) .await - .map_err(|_| AnswerTransportError::Timeout)? - .map_err(|_| AnswerTransportError::Closed) + .map_err(|err| Self::answer_error_from_bus(&err)) } Self::InProcess { steering_hub, .. } => { steering_hub.deliver_steer(text, Some(actor)); @@ -373,12 +400,11 @@ impl RunAnswerTransport { async fn interrupt(&self, actor: Principal) -> Result<(), AnswerTransportError> { match self { - Self::Subprocess { control_tx } => { + Self::Worker { run_id, bus } => { let message = WorkerControlEnvelope::interrupt(actor); - timeout(WORKER_CONTROL_ENQUEUE_TIMEOUT, control_tx.send(message)) + Self::publish_worker_control(*run_id, bus, message) .await - .map_err(|_| AnswerTransportError::Timeout)? - .map_err(|_| AnswerTransportError::Closed) + .map_err(|err| Self::answer_error_from_bus(&err)) } Self::InProcess { steering_hub, .. } => { steering_hub.interrupt(Some(&actor)); @@ -393,12 +419,11 @@ impl RunAnswerTransport { actor: Principal, ) -> Result<(), AnswerTransportError> { match self { - Self::Subprocess { control_tx } => { + Self::Worker { run_id, bus } => { let message = WorkerControlEnvelope::interrupt_then_steer(text, actor); - timeout(WORKER_CONTROL_ENQUEUE_TIMEOUT, control_tx.send(message)) + Self::publish_worker_control(*run_id, bus, message) .await - .map_err(|_| AnswerTransportError::Timeout)? - .map_err(|_| AnswerTransportError::Closed) + .map_err(|err| Self::answer_error_from_bus(&err)) } Self::InProcess { steering_hub, .. } => { steering_hub.interrupt_then_steer(&text, Some(&actor)); @@ -415,12 +440,14 @@ impl RunAnswerTransport { actor: Principal, ) -> Result<(), PairTransportError> { match self { - Self::Subprocess { control_tx } => { + Self::Worker { + run_id: worker_run_id, + bus, + } => { let message = WorkerControlEnvelope::start_pair(run_id, pair_id, target, actor); - timeout(WORKER_CONTROL_ENQUEUE_TIMEOUT, control_tx.send(message)) + Self::publish_worker_control(*worker_run_id, bus, message) .await - .map_err(|_| PairTransportError::Timeout)? - .map_err(|_| PairTransportError::Closed) + .map_err(|err| Self::pair_error_from_bus(&err)) } Self::InProcess { steering_hub, .. } => steering_hub .start_pair(run_id, pair_id, target, Some(actor)) @@ -438,7 +465,7 @@ impl RunAnswerTransport { actor: Principal, ) -> Result<(), PairTransportError> { match self { - Self::Subprocess { control_tx } => { + Self::Worker { run_id, bus } => { let message = WorkerControlEnvelope::pair_message( pair_id, message_id, @@ -446,10 +473,9 @@ impl RunAnswerTransport { client_message_id.clone(), actor, ); - timeout(WORKER_CONTROL_ENQUEUE_TIMEOUT, control_tx.send(message)) + Self::publish_worker_control(*run_id, bus, message) .await - .map_err(|_| PairTransportError::Timeout)? - .map_err(|_| PairTransportError::Closed) + .map_err(|err| Self::pair_error_from_bus(&err)) } Self::InProcess { steering_hub, .. } => steering_hub .send_pair_message(pair_id, message_id, text, client_message_id, Some(actor)) @@ -460,12 +486,11 @@ impl RunAnswerTransport { async fn end_pair(&self, pair_id: PairId, actor: Principal) -> Result<(), PairTransportError> { match self { - Self::Subprocess { control_tx } => { + Self::Worker { run_id, bus } => { let message = WorkerControlEnvelope::end_pair(pair_id, actor); - timeout(WORKER_CONTROL_ENQUEUE_TIMEOUT, control_tx.send(message)) + Self::publish_worker_control(*run_id, bus, message) .await - .map_err(|_| PairTransportError::Timeout)? - .map_err(|_| PairTransportError::Closed) + .map_err(|err| Self::pair_error_from_bus(&err)) } Self::InProcess { steering_hub, .. } => steering_hub .end_pair(pair_id, Some(actor)) @@ -473,6 +498,30 @@ impl RunAnswerTransport { .map_err(PairTransportError::Control), } } + + async fn pause_run(&self) -> Result<(), AnswerTransportError> { + match self { + Self::Worker { run_id, bus } => { + let message = WorkerControlEnvelope::pause_run(); + Self::publish_worker_control(*run_id, bus, message) + .await + .map_err(|err| Self::answer_error_from_bus(&err)) + } + Self::InProcess { .. } => Err(AnswerTransportError::Closed), + } + } + + async fn unpause_run(&self) -> Result<(), AnswerTransportError> { + match self { + Self::Worker { run_id, bus } => { + let message = WorkerControlEnvelope::unpause_run(); + Self::publish_worker_control(*run_id, bus, message) + .await + .map_err(|err| Self::answer_error_from_bus(&err)) + } + Self::InProcess { .. } => Err(AnswerTransportError::Closed), + } + } } #[derive(Debug, Clone)] @@ -1019,6 +1068,7 @@ pub struct AppState { started_at: Instant, resource_sampler: resource_sampler::ResourceSampler, max_concurrent_runs: usize, + pub(crate) worker_control_bus: Arc, scheduler_notify: Notify, global_event_tx: broadcast::Sender, /// Per-run coalescing registry for `GET /runs/{id}/files`. Concurrent @@ -1178,6 +1228,8 @@ pub(crate) struct AppStateConfig { pub(crate) http_client: Option, pub(crate) sandbox_provider_registry: Option, pub(crate) shutdown: CancellationToken, + #[cfg(test)] + pub(crate) worker_control_bus: Option>, #[cfg(any(test, feature = "test-support"))] pub(crate) automation_materializer_override: Option>, } @@ -2230,6 +2282,8 @@ pub(crate) fn build_app_state(config: AppStateConfig) -> anyhow::Result anyhow::Result = { + #[cfg(test)] + { + worker_control_bus.unwrap_or_else(|| Arc::new(LocalWorkerControlBus::new())) + } + #[cfg(not(test))] + { + Arc::new(LocalWorkerControlBus::new()) + } + }; Ok(Arc::new(AppState { runs: Mutex::new(HashMap::new()), aggregate_billing: Mutex::new(BillingAccumulator::default()), @@ -2328,6 +2392,7 @@ pub(crate) fn build_app_state(config: AppStateConfig) -> anyhow::Result { @@ -3023,6 +3095,7 @@ async fn finish_cancelled_run_before_execution(state: &Arc, run_id: Ru clear_live_run_state(managed_run); } drop(runs); + cleanup_worker_control_bus_for_run(state.as_ref(), run_id); state.scheduler_notify.notify_one(); } @@ -3158,6 +3231,7 @@ fn fail_managed_run(state: &Arc, run_id: RunId, reason: FailureReason, managed_run.error = Some(message); clear_live_run_state(managed_run); } + cleanup_worker_control_bus_for_run(state.as_ref(), run_id); } fn update_live_run_from_event(state: &AppState, run_id: RunId, event: &RunEvent) { @@ -3233,6 +3307,7 @@ fn update_live_run_from_event(state: &AppState, run_id: RunId, event: &RunEvent) managed_run.active_api_targets.clear(); managed_run.active_steerable_stages.clear(); managed_run.active_non_steerable_stages.clear(); + cleanup_worker_control_bus_for_run(state, run_id); } EventBody::RunFailed(props) => { managed_run.status = RunStatus::Failed { @@ -3245,6 +3320,7 @@ fn update_live_run_from_event(state: &AppState, run_id: RunId, event: &RunEvent) managed_run.active_api_targets.clear(); managed_run.active_steerable_stages.clear(); managed_run.active_non_steerable_stages.clear(); + cleanup_worker_control_bus_for_run(state, run_id); } // Track active agent sessions by steerability. Activated/deactivated // are leased by session id so stale deactivations cannot clear a newer @@ -3329,20 +3405,6 @@ async fn drain_worker_stderr(run_id: RunId, stderr: ChildStderr) -> anyhow::Resu Ok(()) } -async fn pump_worker_control_jsonl( - mut stdin: ChildStdin, - mut control_rx: mpsc::Receiver, -) -> anyhow::Result<()> { - while let Some(message) = control_rx.recv().await { - let mut line = serde_json::to_vec(&message)?; - line.push(b'\n'); - stdin.write_all(&line).await?; - stdin.flush().await?; - } - - Ok(()) -} - async fn append_worker_exit_failure( run_store: &fabro_store::RunDatabase, run_id: RunId, @@ -3426,7 +3488,7 @@ fn worker_command( .arg(run_id.to_string()) .arg("--mode") .arg(worker_mode_arg(mode)) - .stdin(Stdio::piped()) + .stdin(Stdio::null()) .stdout(worker_stdout) .stderr(Stdio::piped()); @@ -4105,25 +4167,6 @@ async fn execute_run_subprocess(state: Arc, run_id: RunId) { } } - let Some(stdin) = child.stdin.take() else { - let message = "Worker stdin pipe was unavailable".to_string(); - tracing::error!(run_id = %run_id, "{message}"); - let _ = child.start_kill(); - let failure_event = workflow_event::Event::workflow_run_failed_from_error( - &WorkflowError::engine(message.clone()), - fabro_types::RunTiming::default(), - FailureReason::LaunchFailed, - None, - None, - None, - None, - ); - let _ = workflow_event::append_event(&run_store, &run_id, &failure_event).await; - fail_managed_run(&state, run_id, FailureReason::LaunchFailed, message); - state.scheduler_notify.notify_one(); - return; - }; - let Some(stderr) = child.stderr.take() else { let message = "Worker stderr pipe was unavailable".to_string(); tracing::error!(run_id = %run_id, "{message}"); @@ -4143,15 +4186,16 @@ async fn execute_run_subprocess(state: Arc, run_id: RunId) { return; }; - let (control_tx, control_rx) = mpsc::channel(WORKER_CONTROL_QUEUE_CAPACITY); { let mut runs = state.runs.lock().expect("runs lock poisoned"); if let Some(managed_run) = runs.get_mut(&run_id) { - managed_run.answer_transport = Some(RunAnswerTransport::Subprocess { control_tx }); + managed_run.answer_transport = Some(RunAnswerTransport::Worker { + run_id, + bus: Arc::clone(&state.worker_control_bus), + }); } } - let control_task = tokio::spawn(pump_worker_control_jsonl(stdin, control_rx)); let stderr_task = tokio::spawn(drain_worker_stderr(run_id, stderr)); let wait_status = match child.wait().await { @@ -4176,9 +4220,6 @@ async fn execute_run_subprocess(state: Arc, run_id: RunId) { } }; - control_task.abort(); - let _ = control_task.await; - match stderr_task.await { Ok(Ok(())) => {} Ok(Err(err)) => { diff --git a/lib/crates/fabro-server/src/server/handler/lifecycle.rs b/lib/crates/fabro-server/src/server/handler/lifecycle.rs index 0c2697133..d8a2589ac 100644 --- a/lib/crates/fabro-server/src/server/handler/lifecycle.rs +++ b/lib/crates/fabro-server/src/server/handler/lifecycle.rs @@ -387,10 +387,6 @@ async fn cancel_run( | RunStatus::Running | RunStatus::Blocked { .. } | RunStatus::Paused { .. } => { - let use_cancel_signal = !matches!( - managed_run.answer_transport, - Some(RunAnswerTransport::InProcess { .. }) - ); let persist_cancelled_status = matches!( managed_run.status, RunStatus::Submitted | RunStatus::Pending { .. } | RunStatus::Runnable @@ -404,9 +400,7 @@ async fn cancel_run( persist_cancelled_status, managed_run.answer_transport.clone(), managed_run.cancel_token.clone(), - use_cancel_signal - .then(|| managed_run.cancel_tx.take()) - .flatten(), + managed_run.cancel_tx.take(), managed_run.worker_pid, )) } @@ -436,22 +430,29 @@ async fn cancel_run( if let Some(token) = &cancel_token { token.cancel(); } - let sent_cancel_signal = if let Some(cancel_tx) = cancel_tx { + let sent_in_process_cancel = if let Some(cancel_tx) = cancel_tx { let _ = cancel_tx.send(()); true } else { false }; - if let Some(answer_transport) = answer_transport { - if !(sent_cancel_signal && matches!(answer_transport, RunAnswerTransport::InProcess { .. })) + let delivered_control = if let Some(answer_transport) = answer_transport { + if sent_in_process_cancel + && matches!(answer_transport, RunAnswerTransport::InProcess { .. }) { - let _ = answer_transport.cancel_run().await; + true + } else { + answer_transport.cancel_run().await.is_ok() + } + } else { + false + }; + if !delivered_control { + if let Some(worker_pid) = worker_pid { + #[cfg(unix)] + fabro_proc::sigterm(worker_pid); + schedule_worker_kill(Arc::clone(&state), id, worker_pid); } - } - if let Some(worker_pid) = worker_pid { - #[cfg(unix)] - fabro_proc::sigterm(worker_pid); - schedule_worker_kill(Arc::clone(&state), id, worker_pid); } if persist_cancelled_status { @@ -504,9 +505,9 @@ async fn unmanaged_cancel_response( /// How `pause_run` should enact the transition, chosen from the current run /// status. enum PauseMode { - /// Worker is running; ask it to pause via SIGUSR1. Status flips to - /// `Paused` once the worker acknowledges. - Signal { worker_pid: u32 }, + /// Worker is running; ask it to pause via the worker control bus. Status + /// flips to `Paused` once the worker acknowledges. + Transport { transport: RunAnswerTransport }, /// Worker is blocked on a human gate; flip to `Paused` directly by /// appending `RunPaused` ourselves. AppendEvent, @@ -514,8 +515,9 @@ enum PauseMode { /// How `unpause_run` should enact the transition. enum UnpauseMode { - /// No outstanding block; ask the worker to resume via SIGUSR2. - Signal { worker_pid: u32 }, + /// No outstanding block; ask the worker to resume via the worker control + /// bus. + Transport { transport: RunAnswerTransport }, /// Was paused while blocked; append `RunUnpaused` and let the reducer /// restore the underlying blocked state from `Paused { prior_block }`. AppendEvent, @@ -544,11 +546,11 @@ async fn pause_run( let runs = state.runs.lock().expect("runs lock poisoned"); match runs.get(&id) { Some(managed_run) if managed_run.status == RunStatus::Running => { - let Some(worker_pid) = managed_run.worker_pid else { + let Some(transport) = managed_run.answer_transport.clone() else { return ApiError::new(StatusCode::CONFLICT, "Run worker is not available.") .into_response(); }; - PauseMode::Signal { worker_pid } + PauseMode::Transport { transport } } Some(managed_run) if matches!(managed_run.status, RunStatus::Blocked { .. }) => { PauseMode::AppendEvent @@ -578,11 +580,14 @@ async fn pause_run( return ApiError::new(StatusCode::INTERNAL_SERVER_ERROR, err.to_string()).into_response(); } match mode { - PauseMode::Signal { worker_pid } => { - #[cfg(unix)] - fabro_proc::sigusr1(worker_pid); - #[cfg(not(unix))] - let _ = worker_pid; + PauseMode::Transport { transport } => { + if transport.pause_run().await.is_err() { + return ApiError::new( + StatusCode::SERVICE_UNAVAILABLE, + "Failed to deliver pause request to the active run.", + ) + .into_response(); + } } PauseMode::AppendEvent => { if let Some(response) = synchronous_transition(state.as_ref(), id, |events| { @@ -625,11 +630,11 @@ async fn unpause_run( prior_block: Some(_), } => UnpauseMode::AppendEvent, RunStatus::Paused { prior_block: None } => { - let Some(worker_pid) = managed_run.worker_pid else { + let Some(transport) = managed_run.answer_transport.clone() else { return ApiError::new(StatusCode::CONFLICT, "Run worker is not available.") .into_response(); }; - UnpauseMode::Signal { worker_pid } + UnpauseMode::Transport { transport } } _ => { return ApiError::new(StatusCode::CONFLICT, "Run is not paused.") @@ -658,11 +663,14 @@ async fn unpause_run( return ApiError::new(StatusCode::INTERNAL_SERVER_ERROR, err.to_string()).into_response(); } match mode { - UnpauseMode::Signal { worker_pid } => { - #[cfg(unix)] - fabro_proc::sigusr2(worker_pid); - #[cfg(not(unix))] - let _ = worker_pid; + UnpauseMode::Transport { transport } => { + if transport.unpause_run().await.is_err() { + return ApiError::new( + StatusCode::SERVICE_UNAVAILABLE, + "Failed to deliver unpause request to the active run.", + ) + .into_response(); + } } UnpauseMode::AppendEvent => { if let Some(response) = synchronous_transition(state.as_ref(), id, |events| { diff --git a/lib/crates/fabro-server/src/server/handler/mod.rs b/lib/crates/fabro-server/src/server/handler/mod.rs index 3b7cc1ec1..33c6aa8eb 100644 --- a/lib/crates/fabro-server/src/server/handler/mod.rs +++ b/lib/crates/fabro-server/src/server/handler/mod.rs @@ -23,6 +23,7 @@ mod sessions; mod steer; pub(in crate::server) mod system; mod variables; +mod worker_control; pub(super) use system::{health, openapi_spec}; @@ -167,6 +168,7 @@ pub(super) fn real_routes() -> Router> { .merge(models::routes()) .merge(secrets::routes()) .merge(variables::routes()) + .merge(worker_control::routes()) .merge(sessions::routes()) .merge(system::routes()) .merge(completions::routes()) diff --git a/lib/crates/fabro-server/src/server/handler/worker_control.rs b/lib/crates/fabro-server/src/server/handler/worker_control.rs new file mode 100644 index 000000000..d643a904e --- /dev/null +++ b/lib/crates/fabro-server/src/server/handler/worker_control.rs @@ -0,0 +1,166 @@ +use std::sync::Arc; + +use axum::extract::ws::{ + CloseFrame, Message as WsMessage, WebSocket, WebSocketUpgrade, close_code, +}; +use fabro_interview::{ + WORKER_CONTROL_INVALID_CURSOR_REASON, WORKER_CONTROL_PONG_TIMEOUT_REASON, + WORKER_CONTROL_WS_LIVENESS_TIMEOUT, WORKER_CONTROL_WS_PING_INTERVAL, + WorkerControlDeliveryFrame, +}; +use futures_util::{SinkExt, StreamExt}; +use tokio::time::{self, Instant, MissedTickBehavior}; + +use super::super::{ + ApiError, AppState, IntoResponse, Query, RequireWorkerRunScoped, Response, Router, State, + StatusCode, get, +}; +use crate::worker_control::{WorkerControlBusError, WorkerControlCursor, WorkerControlReceiver}; + +#[derive(Debug, serde::Deserialize)] +struct WorkerControlStreamQuery { + after: Option, +} + +pub(super) fn routes() -> Router> { + Router::new().route( + "/runs/{id}/worker/control-stream", + get(worker_control_stream), + ) +} + +async fn worker_control_stream( + RequireWorkerRunScoped(id): RequireWorkerRunScoped, + State(state): State>, + Query(query): Query, + ws: WebSocketUpgrade, +) -> Response { + let cached = match state.store.get_cached_run(&id).await { + Ok(Some(cached)) => cached, + Ok(None) => return ApiError::not_found("Run not found.").into_response(), + Err(err) => { + return ApiError::new(StatusCode::INTERNAL_SERVER_ERROR, err.to_string()) + .into_response(); + } + }; + if cached.projection.archived_at.is_some() { + return ApiError::new(StatusCode::CONFLICT, "Run is archived.").into_response(); + } + if cached.projection.status.is_terminal() { + return ApiError::new( + StatusCode::CONFLICT, + "Run is terminal and cannot accept worker control streams.", + ) + .into_response(); + } + + let cursor = match WorkerControlCursor::from_after_query(query.after.as_deref()) { + Ok(cursor) => cursor, + Err(err @ WorkerControlBusError::InvalidCursor { .. }) => { + return ApiError::new(StatusCode::GONE, err.to_string()).into_response(); + } + Err(err) => return worker_control_bus_error_response(&err), + }; + let receiver = match state.worker_control_bus.subscribe(id, cursor).await { + Ok(receiver) => receiver, + Err(err @ WorkerControlBusError::InvalidCursor { .. }) => { + return ApiError::new(StatusCode::GONE, err.to_string()).into_response(); + } + Err(err) => return worker_control_bus_error_response(&err), + }; + + ws.on_upgrade(move |socket| worker_control_websocket(socket, receiver)) +} + +fn worker_control_bus_error_response(err: &WorkerControlBusError) -> Response { + let status = match err { + WorkerControlBusError::Closed | WorkerControlBusError::Unavailable => { + StatusCode::SERVICE_UNAVAILABLE + } + WorkerControlBusError::PublishTimeout => StatusCode::SERVICE_UNAVAILABLE, + WorkerControlBusError::InvalidCursor { .. } => StatusCode::GONE, + }; + ApiError::new(status, err.to_string()).into_response() +} + +async fn worker_control_websocket(socket: WebSocket, mut receiver: WorkerControlReceiver) { + let (mut sender, mut receiver_ws) = socket.split(); + let mut ping_interval = time::interval(WORKER_CONTROL_WS_PING_INTERVAL); + ping_interval.set_missed_tick_behavior(MissedTickBehavior::Delay); + let mut last_liveness = Instant::now(); + let liveness_timeout = time::sleep_until(last_liveness + WORKER_CONTROL_WS_LIVENESS_TIMEOUT); + tokio::pin!(liveness_timeout); + + loop { + liveness_timeout + .as_mut() + .reset(last_liveness + WORKER_CONTROL_WS_LIVENESS_TIMEOUT); + + tokio::select! { + delivery = receiver.recv() => { + let Some(delivery) = delivery else { + let _ = sender.send(WsMessage::Close(None)).await; + return; + }; + let delivery = match delivery { + Ok(delivery) => delivery, + Err(WorkerControlBusError::InvalidCursor { .. }) => { + let _ = sender.send(invalid_cursor_close_message()).await; + return; + } + Err(_) => { + let _ = sender.send(WsMessage::Close(None)).await; + return; + } + }; + let frame = WorkerControlDeliveryFrame { + id: delivery.id.to_string(), + envelope: delivery.envelope, + }; + let Ok(text) = serde_json::to_string(&frame) else { + let _ = sender.send(WsMessage::Close(None)).await; + return; + }; + if sender.send(WsMessage::Text(text.into())).await.is_err() { + return; + } + } + message = receiver_ws.next() => { + let Some(message) = message else { + return; + }; + match message { + Ok(WsMessage::Ping(payload)) => { + last_liveness = Instant::now(); + if sender.send(WsMessage::Pong(payload)).await.is_err() { + return; + } + } + Ok(WsMessage::Pong(_) | WsMessage::Text(_) | WsMessage::Binary(_)) => { + last_liveness = Instant::now(); + } + Ok(WsMessage::Close(_)) | Err(_) => return, + } + } + _ = ping_interval.tick() => { + if sender.send(WsMessage::Ping(Vec::new().into())).await.is_err() { + return; + } + } + () = &mut liveness_timeout => { + let _ = sender.send(WsMessage::Close(Some(CloseFrame { + code: close_code::AWAY, + reason: WORKER_CONTROL_PONG_TIMEOUT_REASON.into(), + }))).await; + return; + } + } + } +} + +fn invalid_cursor_close_message() -> WsMessage { + WsMessage::Close(Some(CloseFrame { + code: close_code::POLICY, + reason: WORKER_CONTROL_INVALID_CURSOR_REASON.into(), + })) +} diff --git a/lib/crates/fabro-server/src/server/tests.rs b/lib/crates/fabro-server/src/server/tests.rs index a422fa05b..fbb10bbe2 100644 --- a/lib/crates/fabro-server/src/server/tests.rs +++ b/lib/crates/fabro-server/src/server/tests.rs @@ -13,7 +13,8 @@ use fabro_automation::{AutomationId, AutomationTarget}; use fabro_config::ServerSettingsBuilder; use fabro_config::bind::Bind; use fabro_interview::{ - AnswerValue, ControlInterviewer, Interviewer, Question, WorkerControlMessage, + AnswerValue, ControlInterviewer, Interviewer, Question, WorkerControlDeliveryFrame, + WorkerControlEnvelope, WorkerControlMessage, }; use fabro_llm::types::{Message as LlmMessage, Request as LlmRequest, TokenCounts}; use fabro_model::catalog::LlmCatalogSettings; @@ -33,6 +34,8 @@ use httpmock::Method::{GET, POST}; use httpmock::MockServer; use serde_json::json; use tokio_stream::StreamExt as _; +use tokio_tungstenite::tungstenite::client::IntoClientRequest; +use tokio_tungstenite::tungstenite::protocol::Message as WebSocketMessage; use tower::ServiceExt; use tracing::field::{Field, Visit}; use tracing::{Event as TracingEvent, Subscriber, subscriber}; @@ -45,6 +48,9 @@ use crate::automation_materializer::AutomationRunMaterializeInput; use crate::github_webhooks::compute_signature; use crate::jwt_auth::{AuthMode, ConfiguredAuth}; use crate::test_support::*; +use crate::worker_control::{ + LocalWorkerControlBus, WorkerControlBus, WorkerControlCursor, WorkerControlReceiver, +}; const MINIMAL_DOT: &str = r#"digraph Test { graph [goal="Test"] @@ -669,6 +675,292 @@ fn bearer_request(method: Method, path: &str, bearer: &str, body: Body) -> Reque .unwrap() } +struct WorkerControlWsTestServer { + base_url: String, + task: tokio::task::JoinHandle<()>, +} + +impl WorkerControlWsTestServer { + async fn spawn(app: Router) -> Self { + let listener = tokio::net::TcpListener::bind("127.0.0.1:0") + .await + .expect("test WebSocket listener should bind"); + let addr = listener + .local_addr() + .expect("test WebSocket listener should have a local address"); + let task = tokio::spawn(async move { + let result = axum::serve( + listener, + app.into_make_service_with_connect_info::(), + ) + .await; + if let Err(err) = result { + tracing::debug!(error = %err, "test WebSocket server stopped"); + } + }); + Self { + base_url: format!("ws://{addr}"), + task, + } + } + + fn worker_control_url(&self, run_id: RunId, after: Option<&str>) -> String { + let mut url = format!( + "{}/api/v1/runs/{run_id}/worker/control-stream", + self.base_url + ); + if let Some(after) = after { + url.push_str("?after="); + url.push_str(after); + } + url + } +} + +impl Drop for WorkerControlWsTestServer { + fn drop(&mut self) { + self.task.abort(); + } +} + +fn worker_control_ws_request( + server: &WorkerControlWsTestServer, + run_id: RunId, + bearer: Option<&str>, + after: Option<&str>, +) -> Request<()> { + let mut request = server + .worker_control_url(run_id, after) + .into_client_request() + .expect("test worker-control WebSocket request should build"); + if let Some(bearer) = bearer { + request.headers_mut().insert( + header::AUTHORIZATION, + format!("Bearer {bearer}") + .parse() + .expect("test bearer header should parse"), + ); + } + request +} + +async fn connect_worker_control_ws( + server: &WorkerControlWsTestServer, + run_id: RunId, + bearer: &str, + after: Option<&str>, +) -> tokio_tungstenite::WebSocketStream> { + let request = worker_control_ws_request(server, run_id, Some(bearer), after); + let (socket, _) = tokio_tungstenite::connect_async(request) + .await + .expect("worker-control WebSocket should connect"); + socket +} + +async fn assert_worker_control_ws_rejected( + server: &WorkerControlWsTestServer, + run_id: RunId, + bearer: Option<&str>, + after: Option<&str>, + expected: StatusCode, +) { + let request = worker_control_ws_request(server, run_id, bearer, after); + let error = tokio_tungstenite::connect_async(request) + .await + .expect_err("worker-control WebSocket should be rejected"); + match error { + tokio_tungstenite::tungstenite::Error::Http(response) => { + assert_eq!(response.status(), expected); + } + other => panic!("expected HTTP rejection {expected}, got {other:#}"), + } +} + +async fn next_worker_control_frame( + socket: &mut tokio_tungstenite::WebSocketStream< + tokio_tungstenite::MaybeTlsStream, + >, +) -> WorkerControlDeliveryFrame { + let message = tokio::time::timeout(std::time::Duration::from_secs(2), async { + loop { + let message = futures_util::StreamExt::next(socket) + .await + .expect("worker-control WebSocket should remain open") + .expect("worker-control WebSocket frame should be ok"); + match message { + WebSocketMessage::Text(text) => return text, + WebSocketMessage::Ping(payload) => { + futures_util::SinkExt::send(socket, WebSocketMessage::Pong(payload)) + .await + .expect("test worker-control pong should send"); + } + WebSocketMessage::Pong(_) + | WebSocketMessage::Binary(_) + | WebSocketMessage::Frame(_) => {} + WebSocketMessage::Close(frame) => { + panic!("worker-control WebSocket closed before text frame: {frame:?}"); + } + } + } + }) + .await + .expect("worker-control frame should arrive"); + serde_json::from_str(message.as_str()).expect("worker-control delivery frame should parse") +} + +#[tokio::test(flavor = "current_thread")] +async fn worker_control_stream_rejects_missing_user_and_cross_run_auth() { + let (_state, app) = jwt_auth_app(); + let user_bearer = issue_test_user_jwt(); + let run_id = create_run_with_bearer(&app, &user_bearer).await; + let worker_bearer = issue_test_worker_token(&run_id); + let other_run_id = create_run_with_bearer(&app, &user_bearer).await; + let other_worker_bearer = issue_test_worker_token(&other_run_id); + let server = WorkerControlWsTestServer::spawn(app).await; + + assert_worker_control_ws_rejected(&server, run_id, None, None, StatusCode::UNAUTHORIZED).await; + assert_worker_control_ws_rejected( + &server, + run_id, + Some(&user_bearer), + None, + StatusCode::FORBIDDEN, + ) + .await; + assert_worker_control_ws_rejected( + &server, + run_id, + Some(&other_worker_bearer), + None, + StatusCode::FORBIDDEN, + ) + .await; + + let mut socket = connect_worker_control_ws(&server, run_id, &worker_bearer, None).await; + futures_util::SinkExt::send(&mut socket, WebSocketMessage::Close(None)) + .await + .unwrap(); +} + +#[tokio::test(flavor = "current_thread")] +async fn worker_control_stream_start_subscription_delivers_frames() { + let (state, app) = jwt_auth_app(); + let user_bearer = issue_test_user_jwt(); + let run_id = create_run_with_bearer(&app, &user_bearer).await; + let worker_bearer = issue_test_worker_token(&run_id); + let server = WorkerControlWsTestServer::spawn(app).await; + let mut socket = connect_worker_control_ws(&server, run_id, &worker_bearer, None).await; + + let expected = WorkerControlEnvelope::cancel_run(); + let id = state + .worker_control_bus + .publish(run_id, expected.clone()) + .await + .unwrap(); + let frame = next_worker_control_frame(&mut socket).await; + + assert_eq!(frame.id, id.to_string()); + assert_eq!(frame.envelope, expected); +} + +#[tokio::test(flavor = "current_thread")] +async fn worker_control_stream_after_subscription_delivers_only_later_frames() { + let (state, app) = jwt_auth_app(); + let user_bearer = issue_test_user_jwt(); + let run_id = create_run_with_bearer(&app, &user_bearer).await; + let worker_bearer = issue_test_worker_token(&run_id); + let first = state + .worker_control_bus + .publish(run_id, WorkerControlEnvelope::cancel_run()) + .await + .unwrap(); + let expected = WorkerControlEnvelope::pause_run(); + let second = state + .worker_control_bus + .publish(run_id, expected.clone()) + .await + .unwrap(); + let server = WorkerControlWsTestServer::spawn(app).await; + + let mut socket = + connect_worker_control_ws(&server, run_id, &worker_bearer, Some(first.as_str())).await; + let frame = next_worker_control_frame(&mut socket).await; + + assert_eq!(frame.id, second.to_string()); + assert_eq!(frame.envelope, expected); +} + +#[tokio::test(flavor = "current_thread")] +async fn worker_control_stream_invalid_cursor_is_http_gone_before_upgrade() { + let (_state, app) = jwt_auth_app(); + let user_bearer = issue_test_user_jwt(); + let run_id = create_run_with_bearer(&app, &user_bearer).await; + let worker_bearer = issue_test_worker_token(&run_id); + let server = WorkerControlWsTestServer::spawn(app).await; + + assert_worker_control_ws_rejected( + &server, + run_id, + Some(&worker_bearer), + Some("local:999"), + StatusCode::GONE, + ) + .await; +} + +#[tokio::test(flavor = "current_thread")] +async fn worker_control_stream_rejects_missing_terminal_and_archived_runs() { + let (state, app) = jwt_auth_app(); + let user_bearer = issue_test_user_jwt(); + let missing_run_id = RunId::new(); + let missing_worker_bearer = issue_test_worker_token(&missing_run_id); + let terminal_run_id = RunId::new(); + create_succeeded_run(&state, terminal_run_id).await; + let terminal_worker_bearer = issue_test_worker_token(&terminal_run_id); + let archived_run_id = RunId::new(); + create_succeeded_run(&state, archived_run_id).await; + let archived_worker_bearer = issue_test_worker_token(&archived_run_id); + let archive_response = app + .clone() + .oneshot( + Request::builder() + .method("POST") + .uri(api(&format!("/runs/{archived_run_id}/archive"))) + .header(header::AUTHORIZATION, format!("Bearer {user_bearer}")) + .body(Body::empty()) + .unwrap(), + ) + .await + .unwrap(); + response_json!(archive_response, StatusCode::OK).await; + let server = WorkerControlWsTestServer::spawn(app).await; + + assert_worker_control_ws_rejected( + &server, + missing_run_id, + Some(&missing_worker_bearer), + None, + StatusCode::NOT_FOUND, + ) + .await; + assert_worker_control_ws_rejected( + &server, + terminal_run_id, + Some(&terminal_worker_bearer), + None, + StatusCode::CONFLICT, + ) + .await; + assert_worker_control_ws_rejected( + &server, + archived_run_id, + Some(&archived_worker_bearer), + None, + StatusCode::CONFLICT, + ) + .await; +} + fn json_bearer_request( method: Method, path: &str, @@ -1664,6 +1956,7 @@ fn slack_app_state_with_secret_sources( http_client: Some(fabro_http::test_http_client().expect("test HTTP client should build")), sandbox_provider_registry: None, shutdown: tokio_util::sync::CancellationToken::new(), + worker_control_bus: None, automation_materializer_override: None, }) .expect("slack test app state should build") @@ -1764,6 +2057,7 @@ fn slack_service_respects_disabled_server_config_even_with_vault_tokens() { http_client: Some(fabro_http::test_http_client().expect("test HTTP client should build")), sandbox_provider_registry: None, shutdown: tokio_util::sync::CancellationToken::new(), + worker_control_bus: None, automation_materializer_override: None, }) .expect("slack disabled test app state should build"); @@ -1771,6 +2065,23 @@ fn slack_service_respects_disabled_server_config_even_with_vault_tokens() { assert!(state.slack_service.is_none()); } +#[cfg(unix)] +#[test] +fn worker_command_uses_null_stdin_and_token_env() { + let storage_dir = tempfile::tempdir().unwrap(); + let state = worker_command_test_state(storage_dir.path(), &["dev-token"], Some(TEST_DEV_TOKEN)); + let cmd = worker_command( + state.as_ref(), + RunId::new(), + RunExecutionMode::Start, + storage_dir.path(), + false, + ) + .unwrap(); + + assert_worker_command_passes_token_only_by_env(&cmd); +} + #[cfg(unix)] #[test] fn worker_command_default_token_omits_agent_run_tools_scope() { @@ -2069,6 +2380,7 @@ methods = ["dev-token"] http_client: Some(fabro_http::test_http_client().expect("test HTTP client should build")), sandbox_provider_registry: None, shutdown: tokio_util::sync::CancellationToken::new(), + worker_control_bus: None, automation_materializer_override: None, }) else { panic!("build_app_state should require SESSION_SECRET") @@ -2197,6 +2509,7 @@ fn build_test_app_state_with_vault_path(vault_path: &Path) -> anyhow::Result crate::worker_token:: .claims } +async fn worker_transport_with_receiver( + run_id: RunId, +) -> (RunAnswerTransport, WorkerControlReceiver) { + let bus = StdArc::new(LocalWorkerControlBus::new()); + let receiver = bus + .subscribe(run_id, WorkerControlCursor::Start) + .await + .expect("test worker bus should subscribe"); + // Ensure the subscription task is waiting before the test publishes. + tokio::task::yield_now().await; + let bus: StdArc = bus; + let transport = RunAnswerTransport::Worker { run_id, bus }; + (transport, receiver) +} + +async fn recv_worker_control_envelope( + receiver: &mut WorkerControlReceiver, +) -> WorkerControlEnvelope { + receiver + .recv() + .await + .expect("test worker control receiver should stay open") + .expect("test worker control delivery should succeed") + .envelope +} + #[tokio::test] -async fn subprocess_answer_transport_cancel_run_enqueues_cancel_message() { - let (control_tx, mut control_rx) = tokio::sync::mpsc::channel(1); - let transport = RunAnswerTransport::Subprocess { control_tx }; +async fn worker_answer_transport_cancel_run_publishes_cancel_message() { + let (transport, mut control_rx) = worker_transport_with_receiver(fixtures::RUN_1).await; transport.cancel_run().await.unwrap(); assert_eq!( - control_rx.recv().await, - Some(WorkerControlEnvelope::cancel_run()) + recv_worker_control_envelope(&mut control_rx).await, + WorkerControlEnvelope::cancel_run() ); } #[tokio::test] -async fn subprocess_answer_transport_steer_enqueues_plain_steer_message() { - let (control_tx, mut control_rx) = tokio::sync::mpsc::channel(1); - let transport = RunAnswerTransport::Subprocess { control_tx }; +async fn worker_answer_transport_steer_publishes_plain_steer_message() { + let (transport, mut control_rx) = worker_transport_with_receiver(fixtures::RUN_1).await; let actor = Principal::System { system_kind: SystemActorKind::Engine, }; @@ -2404,15 +2741,14 @@ async fn subprocess_answer_transport_steer_enqueues_plain_steer_message() { .unwrap(); assert_eq!( - control_rx.recv().await, - Some(WorkerControlEnvelope::steer("try again", actor)) + recv_worker_control_envelope(&mut control_rx).await, + WorkerControlEnvelope::steer("try again", actor) ); } #[tokio::test] -async fn subprocess_answer_transport_interrupt_enqueues_interrupt_message() { - let (control_tx, mut control_rx) = tokio::sync::mpsc::channel(1); - let transport = RunAnswerTransport::Subprocess { control_tx }; +async fn worker_answer_transport_interrupt_publishes_interrupt_message() { + let (transport, mut control_rx) = worker_transport_with_receiver(fixtures::RUN_1).await; let actor = Principal::System { system_kind: SystemActorKind::Engine, }; @@ -2420,15 +2756,14 @@ async fn subprocess_answer_transport_interrupt_enqueues_interrupt_message() { transport.interrupt(actor.clone()).await.unwrap(); assert_eq!( - control_rx.recv().await, - Some(WorkerControlEnvelope::interrupt(actor)) + recv_worker_control_envelope(&mut control_rx).await, + WorkerControlEnvelope::interrupt(actor) ); } #[tokio::test] -async fn subprocess_answer_transport_interrupt_then_steer_enqueues_single_combined_message() { - let (control_tx, mut control_rx) = tokio::sync::mpsc::channel(1); - let transport = RunAnswerTransport::Subprocess { control_tx }; +async fn worker_answer_transport_interrupt_then_steer_publishes_single_combined_message() { + let (transport, mut control_rx) = worker_transport_with_receiver(fixtures::RUN_1).await; let actor = Principal::System { system_kind: SystemActorKind::Engine, }; @@ -2439,19 +2774,15 @@ async fn subprocess_answer_transport_interrupt_then_steer_enqueues_single_combin .unwrap(); assert_eq!( - control_rx.recv().await, - Some(WorkerControlEnvelope::interrupt_then_steer( - "try again", - actor - )) + recv_worker_control_envelope(&mut control_rx).await, + WorkerControlEnvelope::interrupt_then_steer("try again", actor) ); } #[tokio::test] -async fn subprocess_answer_transport_pair_commands_enqueue_control_messages() { - let (control_tx, mut control_rx) = tokio::sync::mpsc::channel(3); - let transport = RunAnswerTransport::Subprocess { control_tx }; +async fn worker_answer_transport_pair_commands_publish_control_messages() { let run_id = fixtures::RUN_1; + let (transport, mut control_rx) = worker_transport_with_receiver(run_id).await; let pair_id = "01HZX6M29F1CD5YYMHT1F5D7WQ".parse().unwrap(); let message_id = "01HZX6M4D7Y1QW0Q0P6V8Z4DR5".parse().unwrap(); let actor = Principal::System { @@ -2476,27 +2807,39 @@ async fn subprocess_answer_transport_pair_commands_enqueue_control_messages() { transport.end_pair(pair_id, actor.clone()).await.unwrap(); assert_eq!( - control_rx.recv().await, - Some(WorkerControlEnvelope::start_pair( - run_id, - pair_id, - target, - actor.clone() - )) + recv_worker_control_envelope(&mut control_rx).await, + WorkerControlEnvelope::start_pair(run_id, pair_id, target, actor.clone()) ); assert_eq!( - control_rx.recv().await, - Some(WorkerControlEnvelope::pair_message( + recv_worker_control_envelope(&mut control_rx).await, + WorkerControlEnvelope::pair_message( pair_id, message_id, "inspect this", Some("client-1".to_string()), actor.clone() - )) + ) ); assert_eq!( - control_rx.recv().await, - Some(WorkerControlEnvelope::end_pair(pair_id, actor)) + recv_worker_control_envelope(&mut control_rx).await, + WorkerControlEnvelope::end_pair(pair_id, actor) + ); +} + +#[tokio::test] +async fn worker_answer_transport_pause_and_unpause_publish_control_messages() { + let (transport, mut control_rx) = worker_transport_with_receiver(fixtures::RUN_1).await; + + transport.pause_run().await.unwrap(); + transport.unpause_run().await.unwrap(); + + assert_eq!( + recv_worker_control_envelope(&mut control_rx).await, + WorkerControlEnvelope::pause_run() + ); + assert_eq!( + recv_worker_control_envelope(&mut control_rx).await, + WorkerControlEnvelope::unpause_run() ); } @@ -5444,6 +5787,7 @@ fn create_github_token_app_state_with_env_lookup_and_llm_catalog_settings( http_client: Some(fabro_http::test_http_client().expect("test HTTP client should build")), sandbox_provider_registry: None, shutdown: tokio_util::sync::CancellationToken::new(), + worker_control_bus: None, automation_materializer_override: None, }; let state = build_app_state(config).expect("test app state should build"); @@ -10847,12 +11191,8 @@ async fn steer_without_active_steerable_session_forwards_plain_steer_for_bufferi let state = test_app_state(); let app = crate::test_support::build_test_router(Arc::clone(&state)); let run_id = fixtures::RUN_1; - let (control_tx, mut control_rx) = tokio::sync::mpsc::channel(1); - let _temp_dir = insert_running_control_run( - &state, - run_id, - Some(RunAnswerTransport::Subprocess { control_tx }), - ); + let (transport, mut control_rx) = worker_transport_with_receiver(run_id).await; + let _temp_dir = insert_running_control_run(&state, run_id, Some(transport)); let req = Request::builder() .method("POST") @@ -10863,7 +11203,7 @@ async fn steer_without_active_steerable_session_forwards_plain_steer_for_bufferi let response = app.oneshot(req).await.unwrap(); assert_status!(response, StatusCode::ACCEPTED).await; - let envelope = control_rx.recv().await.unwrap(); + let envelope = recv_worker_control_envelope(&mut control_rx).await; assert!(matches!( envelope.message, WorkerControlMessage::Steer { ref text, .. } if text == "try again" @@ -10876,12 +11216,8 @@ async fn steer_with_active_non_steerable_session_returns_conflict() { let app = crate::test_support::build_test_router(Arc::clone(&state)); let run_id = fixtures::RUN_1; let stage_id = StageId::new("agent", 1); - let (control_tx, _control_rx) = tokio::sync::mpsc::channel(1); - let _temp_dir = insert_running_control_run( - &state, - run_id, - Some(RunAnswerTransport::Subprocess { control_tx }), - ); + let (transport, _control_rx) = worker_transport_with_receiver(run_id).await; + let _temp_dir = insert_running_control_run(&state, run_id, Some(transport)); { let mut runs = state.runs.lock().expect("runs lock poisoned"); runs.get_mut(&run_id) @@ -10908,12 +11244,8 @@ async fn steer_interrupt_without_active_steerable_session_returns_conflict() { let state = test_app_state(); let app = crate::test_support::build_test_router(Arc::clone(&state)); let run_id = fixtures::RUN_1; - let (control_tx, _control_rx) = tokio::sync::mpsc::channel(1); - let _temp_dir = insert_running_control_run( - &state, - run_id, - Some(RunAnswerTransport::Subprocess { control_tx }), - ); + let (transport, _control_rx) = worker_transport_with_receiver(run_id).await; + let _temp_dir = insert_running_control_run(&state, run_id, Some(transport)); let req = Request::builder() .method("POST") @@ -10934,12 +11266,8 @@ async fn interrupt_with_active_steerable_session_forwards_interrupt() { let app = crate::test_support::build_test_router(Arc::clone(&state)); let run_id = fixtures::RUN_1; let stage_id = StageId::new("agent", 1); - let (control_tx, mut control_rx) = tokio::sync::mpsc::channel(1); - let _temp_dir = insert_running_control_run( - &state, - run_id, - Some(RunAnswerTransport::Subprocess { control_tx }), - ); + let (transport, mut control_rx) = worker_transport_with_receiver(run_id).await; + let _temp_dir = insert_running_control_run(&state, run_id, Some(transport)); { let mut runs = state.runs.lock().expect("runs lock poisoned"); runs.get_mut(&run_id) @@ -10956,7 +11284,7 @@ async fn interrupt_with_active_steerable_session_forwards_interrupt() { let response = app.oneshot(req).await.unwrap(); assert_status!(response, StatusCode::ACCEPTED).await; - let envelope = control_rx.recv().await.unwrap(); + let envelope = recv_worker_control_envelope(&mut control_rx).await; assert!(matches!( envelope.message, WorkerControlMessage::Interrupt { @@ -10971,12 +11299,8 @@ async fn steer_interrupt_with_active_steerable_session_forwards_combined_control let app = crate::test_support::build_test_router(Arc::clone(&state)); let run_id = fixtures::RUN_1; let stage_id = StageId::new("agent", 1); - let (control_tx, mut control_rx) = tokio::sync::mpsc::channel(1); - let _temp_dir = insert_running_control_run( - &state, - run_id, - Some(RunAnswerTransport::Subprocess { control_tx }), - ); + let (transport, mut control_rx) = worker_transport_with_receiver(run_id).await; + let _temp_dir = insert_running_control_run(&state, run_id, Some(transport)); { let mut runs = state.runs.lock().expect("runs lock poisoned"); runs.get_mut(&run_id) @@ -10994,7 +11318,7 @@ async fn steer_interrupt_with_active_steerable_session_forwards_combined_control let response = app.oneshot(req).await.unwrap(); assert_status!(response, StatusCode::ACCEPTED).await; - let envelope = control_rx.recv().await.unwrap(); + let envelope = recv_worker_control_envelope(&mut control_rx).await; assert!(matches!( envelope.message, WorkerControlMessage::InterruptThenSteer { ref text, .. } if text == "try again" @@ -11160,12 +11484,8 @@ async fn steer_with_active_acp_session_forwards_to_worker() { let state = test_app_state(); let app = crate::test_support::build_test_router(Arc::clone(&state)); let run_id = fixtures::RUN_1; - let (control_tx, mut control_rx) = tokio::sync::mpsc::channel(1); - let _temp_dir = insert_running_control_run( - &state, - run_id, - Some(RunAnswerTransport::Subprocess { control_tx }), - ); + let (transport, mut control_rx) = worker_transport_with_receiver(run_id).await; + let _temp_dir = insert_running_control_run(&state, run_id, Some(transport)); let started = acp_event_for_stage(&run_id, &workflow_event::Event::AgentAcpStarted { node_id: "agent".to_string(), @@ -11198,7 +11518,7 @@ async fn steer_with_active_acp_session_forwards_to_worker() { let response = app.oneshot(req).await.unwrap(); assert_status!(response, StatusCode::ACCEPTED).await; - let envelope = control_rx.recv().await.unwrap(); + let envelope = recv_worker_control_envelope(&mut control_rx).await; assert!(matches!( envelope.message, WorkerControlMessage::Steer { ref text, .. } if text == "try again" @@ -11265,12 +11585,8 @@ async fn active_acp_steerable_marker_clears_on_terminal_paths() { let state = test_app_state(); let app = crate::test_support::build_test_router(Arc::clone(&state)); let run_id = fixtures::RUN_1; - let (control_tx, _control_rx) = tokio::sync::mpsc::channel(1); - let _temp_dir = insert_running_control_run( - &state, - run_id, - Some(RunAnswerTransport::Subprocess { control_tx }), - ); + let (transport, _control_rx) = worker_transport_with_receiver(run_id).await; + let _temp_dir = insert_running_control_run(&state, run_id, Some(transport)); let started = acp_event_for_stage(&run_id, &workflow_event::Event::AgentAcpStarted { node_id: "agent".to_string(), visit: 1, @@ -13200,11 +13516,13 @@ async fn pause_run_sets_pending_control_on_board_response() { let run_id_str = create_and_start_run(&app, MINIMAL_DOT).await; let run_id = run_id_str.parse::().unwrap(); + let (transport, _control_rx) = worker_transport_with_receiver(run_id).await; { let mut runs = state.runs.lock().expect("runs lock poisoned"); let managed_run = runs.get_mut(&run_id).expect("run should exist"); managed_run.status = RunStatus::Running; managed_run.worker_pid = Some(u32::MAX); + managed_run.answer_transport = Some(transport); } let req = Request::builder() @@ -13320,11 +13638,13 @@ async fn unpause_run_sets_pending_control() { let run_id_str = create_and_start_run(&app, MINIMAL_DOT).await; let run_id = run_id_str.parse::().unwrap(); + let (transport, _control_rx) = worker_transport_with_receiver(run_id).await; { let mut runs = state.runs.lock().expect("runs lock poisoned"); let managed_run = runs.get_mut(&run_id).expect("run should exist"); managed_run.status = RunStatus::Paused { prior_block: None }; managed_run.worker_pid = Some(u32::MAX); + managed_run.answer_transport = Some(transport); } let req = Request::builder() diff --git a/lib/crates/fabro-server/src/test_support.rs b/lib/crates/fabro-server/src/test_support.rs index 11f80e0e7..585aa36e9 100644 --- a/lib/crates/fabro-server/src/test_support.rs +++ b/lib/crates/fabro-server/src/test_support.rs @@ -249,6 +249,8 @@ impl TestAppStateBuilder { ), sandbox_provider_registry: self.sandbox_provider_registry, shutdown: CancellationToken::new(), + #[cfg(test)] + worker_control_bus: None, automation_materializer_override: self.automation_materializer, }) } diff --git a/lib/crates/fabro-server/src/worker_control/bus.rs b/lib/crates/fabro-server/src/worker_control/bus.rs new file mode 100644 index 000000000..b6fadedb5 --- /dev/null +++ b/lib/crates/fabro-server/src/worker_control/bus.rs @@ -0,0 +1,155 @@ +use std::fmt; + +use fabro_interview::WorkerControlEnvelope; +use fabro_types::RunId; +use futures_util::future::BoxFuture; +use tokio::sync::mpsc; + +#[derive(Clone, PartialEq, Eq)] +pub(crate) struct WorkerControlMessageId(String); + +impl WorkerControlMessageId { + #[must_use] + pub(crate) fn new(value: impl Into) -> Self { + Self(value.into()) + } + + #[must_use] + pub(crate) fn as_str(&self) -> &str { + &self.0 + } +} + +impl fmt::Debug for WorkerControlMessageId { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.debug_tuple("WorkerControlMessageId") + .field(&self.0) + .finish() + } +} + +impl fmt::Display for WorkerControlMessageId { + fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result { + f.write_str(&self.0) + } +} + +impl From for WorkerControlMessageId { + fn from(value: String) -> Self { + Self::new(value) + } +} + +impl From<&str> for WorkerControlMessageId { + fn from(value: &str) -> Self { + Self::new(value) + } +} + +/// Cursor for replaying a run's worker-control stream. +/// +/// Future Redis Streams mapping: +/// - `Start` maps to `XREAD ... STREAMS fabro:run:{run_id}:control 0-0`. +/// - `After(id)` maps to `XREAD ... STREAMS fabro:run:{run_id}:control {id}`. +#[derive(Debug, Clone, PartialEq, Eq)] +pub(crate) enum WorkerControlCursor { + Start, + After(WorkerControlMessageId), +} + +impl WorkerControlCursor { + pub(crate) fn from_after_query(after: Option<&str>) -> Result { + match after { + None => Ok(Self::Start), + Some(value) if value.trim().is_empty() => Err(WorkerControlBusError::InvalidCursor { + cursor: value.to_string(), + reason: "cursor id must not be empty".to_string(), + }), + Some(value) => Ok(Self::After(WorkerControlMessageId::from(value))), + } + } +} + +#[derive(Debug, Clone, PartialEq, Eq)] +pub(crate) struct WorkerControlDelivery { + pub(crate) id: WorkerControlMessageId, + pub(crate) envelope: WorkerControlEnvelope, +} + +#[allow( + dead_code, + reason = "The local backend does not construct every cross-backend error variant." +)] +#[derive(Debug, Clone, thiserror::Error, PartialEq, Eq)] +pub(crate) enum WorkerControlBusError { + #[error("worker control backend is closed")] + Closed, + #[error("worker control backend is unavailable")] + Unavailable, + #[error("worker control cursor `{cursor}` is invalid: {reason}")] + InvalidCursor { cursor: String, reason: String }, + #[error("timed out publishing worker control message")] + PublishTimeout, +} + +impl WorkerControlBusError { + #[must_use] + pub(crate) fn invalid_cursor(cursor: impl Into, reason: impl Into) -> Self { + Self::InvalidCursor { + cursor: cursor.into(), + reason: reason.into(), + } + } +} + +pub(crate) type WorkerControlReceiver = + mpsc::Receiver>; + +pub(crate) trait WorkerControlBus: Send + Sync { + fn publish( + &self, + run_id: RunId, + envelope: WorkerControlEnvelope, + ) -> BoxFuture<'_, Result>; + + fn subscribe( + &self, + run_id: RunId, + cursor: WorkerControlCursor, + ) -> BoxFuture<'_, Result>; + + fn cleanup_run(&self, run_id: RunId) -> BoxFuture<'_, ()>; +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn message_id_equality_and_debug_are_opaque() { + let first = WorkerControlMessageId::from("local:42"); + let same = WorkerControlMessageId::from("local:42"); + let different = WorkerControlMessageId::from("local:43"); + + assert_eq!(first, same); + assert_ne!(first, different); + assert_eq!(format!("{first}"), "local:42"); + assert_eq!(format!("{first:?}"), "WorkerControlMessageId(\"local:42\")"); + } + + #[test] + fn absent_after_query_parses_as_start_cursor() { + assert_eq!( + WorkerControlCursor::from_after_query(None).unwrap(), + WorkerControlCursor::Start + ); + } + + #[test] + fn present_after_query_parses_as_after_cursor() { + assert_eq!( + WorkerControlCursor::from_after_query(Some("local:42")).unwrap(), + WorkerControlCursor::After(WorkerControlMessageId::from("local:42")) + ); + } +} diff --git a/lib/crates/fabro-server/src/worker_control/local.rs b/lib/crates/fabro-server/src/worker_control/local.rs new file mode 100644 index 000000000..219ff18ec --- /dev/null +++ b/lib/crates/fabro-server/src/worker_control/local.rs @@ -0,0 +1,458 @@ +use std::collections::{HashMap, VecDeque}; +use std::sync::atomic::{AtomicU64, Ordering}; +use std::sync::{Arc, Mutex}; + +use fabro_interview::WorkerControlEnvelope; +use fabro_types::RunId; +use futures_util::FutureExt; +use futures_util::future::BoxFuture; +use tokio::sync::{Notify, mpsc}; + +use super::{ + WorkerControlBus, WorkerControlBusError, WorkerControlCursor, WorkerControlDelivery, + WorkerControlMessageId, WorkerControlReceiver, +}; + +pub(crate) const LOCAL_WORKER_CONTROL_RETAINED_MESSAGES_PER_RUN: usize = 1024; +const LOCAL_WORKER_CONTROL_SUBSCRIBER_BUFFER: usize = 64; + +#[derive(Default)] +pub(crate) struct LocalWorkerControlBus { + streams: Arc>>, + next_sequence: Arc, +} + +struct LocalRunControlStream { + messages: VecDeque, + notify: Arc, + has_trimmed: bool, +} + +#[derive(Clone)] +struct LocalMessage { + sequence: u64, + delivery: WorkerControlDelivery, +} + +impl LocalWorkerControlBus { + #[must_use] + pub(crate) fn new() -> Self { + Self { + streams: Arc::new(Mutex::new(HashMap::new())), + next_sequence: Arc::new(AtomicU64::new(1)), + } + } + + fn next_message_id(&self) -> (u64, WorkerControlMessageId) { + let sequence = self.next_sequence.fetch_add(1, Ordering::Relaxed); + ( + sequence, + WorkerControlMessageId::new(format!("local:{sequence}")), + ) + } + + fn subscribe_inner( + &self, + run_id: RunId, + cursor: &WorkerControlCursor, + ) -> Result { + let (next_sequence, notify) = { + let mut streams = self + .streams + .lock() + .expect("worker control streams poisoned"); + let stream = match cursor { + WorkerControlCursor::Start => streams + .entry(run_id) + .or_insert_with(LocalRunControlStream::new), + WorkerControlCursor::After(id) => streams.get_mut(&run_id).ok_or_else(|| { + WorkerControlBusError::invalid_cursor( + id.as_str(), + "message id is not retained for this run", + ) + })?, + }; + let next_sequence = stream.next_sequence_for_cursor(cursor)?; + (next_sequence, Arc::clone(&stream.notify)) + }; + + let (tx, rx) = mpsc::channel(LOCAL_WORKER_CONTROL_SUBSCRIBER_BUFFER); + let streams = Arc::clone(&self.streams); + tokio::spawn(async move { + local_subscription_task(streams, run_id, notify, next_sequence, tx).await; + }); + Ok(rx) + } + + #[cfg(test)] + pub(crate) fn retained_len(&self, run_id: &RunId) -> usize { + let streams = self + .streams + .lock() + .expect("worker control streams poisoned"); + streams + .get(run_id) + .map_or(0, |stream| stream.messages.len()) + } +} + +impl LocalRunControlStream { + fn new() -> Self { + Self { + messages: VecDeque::new(), + notify: Arc::new(Notify::new()), + has_trimmed: false, + } + } + + fn next_sequence_for_cursor( + &self, + cursor: &WorkerControlCursor, + ) -> Result, WorkerControlBusError> { + match cursor { + WorkerControlCursor::Start if self.has_trimmed => Err( + WorkerControlBusError::invalid_cursor("start", "retained local stream is trimmed"), + ), + WorkerControlCursor::Start => Ok(self.messages.front().map(|message| message.sequence)), + WorkerControlCursor::After(id) => { + let requested_sequence = parse_local_sequence(id)?; + let Some(position) = self + .messages + .iter() + .position(|message| message.delivery.id == *id) + else { + return Err(WorkerControlBusError::invalid_cursor( + id.as_str(), + "message id is not retained for this run", + )); + }; + Ok(self + .messages + .get(position + 1) + .map(|message| message.sequence) + .or(Some(requested_sequence.saturating_add(1)))) + } + } + } + + fn trim_retained(&mut self) { + while self.messages.len() > LOCAL_WORKER_CONTROL_RETAINED_MESSAGES_PER_RUN { + self.messages.pop_front(); + self.has_trimmed = true; + } + } +} + +fn parse_local_sequence(id: &WorkerControlMessageId) -> Result { + let Some(raw) = id.as_str().strip_prefix("local:") else { + return Err(WorkerControlBusError::invalid_cursor( + id.as_str(), + "local backend only understands local message ids", + )); + }; + raw.parse::().map_err(|_| { + WorkerControlBusError::invalid_cursor(id.as_str(), "local message id sequence is invalid") + }) +} + +async fn local_subscription_task( + streams: Arc>>, + run_id: RunId, + notify: Arc, + mut next_sequence: Option, + tx: mpsc::Sender>, +) { + loop { + // Register interest *before* inspecting the stream so that a publish + // racing with this read does not cause a lost wakeup. `notify_waiters` + // does not leave a permit for future `notified()` calls. + let notified = notify.notified(); + tokio::pin!(notified); + notified.as_mut().enable(); + + let collected = { + let streams_guard = streams.lock().expect("worker control streams poisoned"); + match streams_guard.get(&run_id) { + None => None, + Some(stream) => { + // A `Start` subscriber that joined before any publish lazily + // adopts the first retained message as its cursor. + if next_sequence.is_none() && stream.has_trimmed { + Some(Err(WorkerControlBusError::invalid_cursor( + "start", + "retained local stream is trimmed", + ))) + } else { + if next_sequence.is_none() { + next_sequence = stream.messages.front().map(|message| message.sequence); + } + match next_sequence { + None => Some(Ok(Vec::new())), + Some(next) => Some(collect_messages_from(&stream.messages, next)), + } + } + } + } + }; + + let messages = match collected { + None => return, + Some(Err(err)) => { + let _ = tx.send(Err(err)).await; + return; + } + Some(Ok(messages)) => messages, + }; + + if messages.is_empty() { + notified.await; + continue; + } + + for message in messages { + next_sequence = Some(message.sequence.saturating_add(1)); + if tx.send(Ok(message.delivery)).await.is_err() { + return; + } + } + } +} + +/// Returns the retained messages with `sequence >= next`. Cheaper than scanning +/// the whole deque: `partition_point` is O(log N) and we only clone the tail. +fn collect_messages_from( + messages: &VecDeque, + next: u64, +) -> Result, WorkerControlBusError> { + let Some(first_sequence) = messages.front().map(|message| message.sequence) else { + return Ok(Vec::new()); + }; + if next < first_sequence { + return Err(WorkerControlBusError::invalid_cursor( + format!("local:{next}"), + "subscriber fell behind retained local messages", + )); + } + let start = messages.partition_point(|message| message.sequence < next); + Ok(messages.iter().skip(start).cloned().collect()) +} + +impl WorkerControlBus for LocalWorkerControlBus { + fn publish( + &self, + run_id: RunId, + envelope: WorkerControlEnvelope, + ) -> BoxFuture<'_, Result> { + async move { + let (sequence, id) = self.next_message_id(); + let delivery = WorkerControlDelivery { + id: id.clone(), + envelope, + }; + let notify = { + let mut streams = self + .streams + .lock() + .expect("worker control streams poisoned"); + let stream = streams + .entry(run_id) + .or_insert_with(LocalRunControlStream::new); + stream + .messages + .push_back(LocalMessage { sequence, delivery }); + stream.trim_retained(); + Arc::clone(&stream.notify) + }; + notify.notify_waiters(); + Ok(id) + } + .boxed() + } + + fn subscribe( + &self, + run_id: RunId, + cursor: WorkerControlCursor, + ) -> BoxFuture<'_, Result> { + async move { self.subscribe_inner(run_id, &cursor) }.boxed() + } + + fn cleanup_run(&self, run_id: RunId) -> BoxFuture<'_, ()> { + async move { + let notify = { + let mut streams = self + .streams + .lock() + .expect("worker control streams poisoned"); + streams.remove(&run_id).map(|stream| stream.notify) + }; + if let Some(notify) = notify { + notify.notify_waiters(); + } + } + .boxed() + } +} + +#[cfg(test)] +mod tests { + use std::time::Duration; + + use fabro_interview::WorkerControlMessage; + use fabro_types::fixtures; + + use super::*; + + async fn recv_delivery(receiver: &mut WorkerControlReceiver) -> WorkerControlDelivery { + tokio::time::timeout(Duration::from_secs(1), receiver.recv()) + .await + .expect("delivery should arrive") + .expect("subscription should remain open") + .expect("delivery should be ok") + } + + #[tokio::test] + async fn messages_publish_in_order() { + let bus = LocalWorkerControlBus::new(); + bus.publish(fixtures::RUN_1, WorkerControlEnvelope::cancel_run()) + .await + .unwrap(); + bus.publish(fixtures::RUN_1, WorkerControlEnvelope::pause_run()) + .await + .unwrap(); + + let mut receiver = bus + .subscribe(fixtures::RUN_1, WorkerControlCursor::Start) + .await + .unwrap(); + + assert!(matches!( + recv_delivery(&mut receiver).await.envelope.message, + WorkerControlMessage::RunCancel + )); + assert!(matches!( + recv_delivery(&mut receiver).await.envelope.message, + WorkerControlMessage::RunPause + )); + } + + #[tokio::test] + async fn active_subscriber_receives_message_published_after_subscription() { + let bus = LocalWorkerControlBus::new(); + let mut receiver = bus + .subscribe(fixtures::RUN_1, WorkerControlCursor::Start) + .await + .unwrap(); + + bus.publish(fixtures::RUN_1, WorkerControlEnvelope::cancel_run()) + .await + .unwrap(); + + assert!(matches!( + recv_delivery(&mut receiver).await.envelope.message, + WorkerControlMessage::RunCancel + )); + } + + #[tokio::test] + async fn messages_published_before_subscription_replay_from_start() { + let bus = LocalWorkerControlBus::new(); + bus.publish(fixtures::RUN_1, WorkerControlEnvelope::cancel_run()) + .await + .unwrap(); + + let mut receiver = bus + .subscribe(fixtures::RUN_1, WorkerControlCursor::Start) + .await + .unwrap(); + + assert!(matches!( + recv_delivery(&mut receiver).await.envelope.message, + WorkerControlMessage::RunCancel + )); + } + + #[tokio::test] + async fn after_cursor_receives_only_later_messages() { + let bus = LocalWorkerControlBus::new(); + let first_id = bus + .publish(fixtures::RUN_1, WorkerControlEnvelope::cancel_run()) + .await + .unwrap(); + bus.publish(fixtures::RUN_1, WorkerControlEnvelope::pause_run()) + .await + .unwrap(); + + let mut receiver = bus + .subscribe( + fixtures::RUN_1, + WorkerControlCursor::After(first_id.clone()), + ) + .await + .unwrap(); + + assert!(matches!( + recv_delivery(&mut receiver).await.envelope.message, + WorkerControlMessage::RunPause + )); + assert!(receiver.try_recv().is_err()); + } + + #[tokio::test] + async fn trimming_bounds_retained_messages_and_invalidates_old_cursor() { + let bus = LocalWorkerControlBus::new(); + let old_id = bus + .publish(fixtures::RUN_1, WorkerControlEnvelope::cancel_run()) + .await + .unwrap(); + for _ in 0..LOCAL_WORKER_CONTROL_RETAINED_MESSAGES_PER_RUN { + bus.publish(fixtures::RUN_1, WorkerControlEnvelope::pause_run()) + .await + .unwrap(); + } + + assert_eq!( + bus.retained_len(&fixtures::RUN_1), + LOCAL_WORKER_CONTROL_RETAINED_MESSAGES_PER_RUN + ); + let err = bus + .subscribe(fixtures::RUN_1, WorkerControlCursor::After(old_id)) + .await + .unwrap_err(); + assert!(matches!(err, WorkerControlBusError::InvalidCursor { .. })); + + let err = bus + .subscribe(fixtures::RUN_1, WorkerControlCursor::Start) + .await + .unwrap_err(); + assert!(matches!(err, WorkerControlBusError::InvalidCursor { .. })); + } + + #[tokio::test] + async fn after_cursor_for_unknown_run_does_not_create_stream() { + let bus = LocalWorkerControlBus::new(); + let err = bus + .subscribe( + fixtures::RUN_1, + WorkerControlCursor::After(WorkerControlMessageId::new("local:1")), + ) + .await + .unwrap_err(); + + assert!(matches!(err, WorkerControlBusError::InvalidCursor { .. })); + assert_eq!(bus.retained_len(&fixtures::RUN_1), 0); + } + + #[tokio::test] + async fn cleanup_removes_retained_messages() { + let bus = LocalWorkerControlBus::new(); + bus.publish(fixtures::RUN_1, WorkerControlEnvelope::cancel_run()) + .await + .unwrap(); + assert_eq!(bus.retained_len(&fixtures::RUN_1), 1); + + bus.cleanup_run(fixtures::RUN_1).await; + + assert_eq!(bus.retained_len(&fixtures::RUN_1), 0); + } +} diff --git a/lib/crates/fabro-server/src/worker_control/mod.rs b/lib/crates/fabro-server/src/worker_control/mod.rs new file mode 100644 index 000000000..6b6f70896 --- /dev/null +++ b/lib/crates/fabro-server/src/worker_control/mod.rs @@ -0,0 +1,8 @@ +mod bus; +mod local; + +pub(crate) use bus::{ + WorkerControlBus, WorkerControlBusError, WorkerControlCursor, WorkerControlDelivery, + WorkerControlMessageId, WorkerControlReceiver, +}; +pub(crate) use local::LocalWorkerControlBus; diff --git a/lib/crates/fabro-workflow/src/operations/start.rs b/lib/crates/fabro-workflow/src/operations/start.rs index 9d3f08d1c..23c90f3ff 100644 --- a/lib/crates/fabro-workflow/src/operations/start.rs +++ b/lib/crates/fabro-workflow/src/operations/start.rs @@ -26,6 +26,7 @@ use fabro_types::{ManifestPath, RunId, RunRunnableSource, SandboxProviderKind}; use fabro_vault::Vault; use tokio::runtime::Handle; use tokio::sync::RwLock as AsyncRwLock; +use tokio::time; use tokio_util::sync::CancellationToken; use crate::artifact_upload::ArtifactSink; @@ -987,6 +988,24 @@ impl Drop for DetachedRunBootstrapGuard { const POSTRUN_INTERRUPTED_MESSAGE: &str = "Run interrupted before post-run finalization completed."; const POSTRUN_CANCELLED_MESSAGE: &str = "Run cancelled before post-run finalization completed."; +const DETACHED_COMPLETION_GUARD_TERMINAL_GRACE: Duration = Duration::from_millis(25); + +async fn run_store_reaches_terminal(run_store: &RunStoreHandle, timeout: Duration) -> bool { + let start = Instant::now(); + loop { + if run_store + .state() + .await + .is_ok_and(|state| state.status.is_terminal()) + { + return true; + } + if start.elapsed() >= timeout { + return false; + } + time::sleep(Duration::from_millis(10)).await; + } +} struct DetachedRunCompletionGuard { event_sink: RunEventSink, @@ -1044,6 +1063,11 @@ impl Drop for DetachedRunCompletionGuard { let run_store = self.run_store.clone(); if let Ok(handle) = Handle::try_current() { handle.spawn(async move { + if run_store_reaches_terminal(&run_store, DETACHED_COMPLETION_GUARD_TERMINAL_GRACE) + .await + { + return; + } emit_workflow_run_failed( run_id, &run_store,