diff --git a/Cargo.lock b/Cargo.lock index b99dd902..7ca68ba7 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -23,6 +23,7 @@ dependencies = [ "indicatif", "infer", "log", + "nix", "pin-project-lite", "proptest", "rstest", diff --git a/Cargo.toml b/Cargo.toml index c2effee2..27466710 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -27,6 +27,7 @@ indicatif = "0.18" infer = { version = "0.19", default-features = false } crc32fast = "1" log = "0.4.21" +nix = { version = "0.31", features = ["process", "signal"] } pin-project-lite = "0.2.16" same-file = "1.0.6" serde = { version = "1.0.185", features = ["derive"] } diff --git a/src/command/worker.rs b/src/command/worker.rs index 61f8e706..62539bc4 100644 --- a/src/command/worker.rs +++ b/src/command/worker.rs @@ -1,9 +1,9 @@ use crate::command::worker_protocol::{ AnnouncePayload, CRF_SEARCH_TOPIC, CancelPayload, Capabilities, ClientEvent, ClientFrame, - CrfSearchProgressPayload, CrfSearchResultPayload, ErrorReplyPayload, FailureReportPayload, - HeartbeatPayload, PullWorkPayload, ReplyBody, ServerPushFrame, ServerReply, - TransferFailurePayload, TransferProgressPayload, TransferStage, TransferStartedPayload, - WorkStatus, + ControlAction, ControlPayload, ControlState, ControlStatePayload, CrfSearchProgressPayload, + CrfSearchResultPayload, ErrorReplyPayload, FailureReportPayload, HeartbeatPayload, + PullWorkPayload, ReplyBody, ServerPushFrame, ServerReply, TransferFailurePayload, + TransferProgressPayload, TransferStage, TransferStartedPayload, WorkStatus, }; use crate::command::worker_transfer::{Chunk, ChunkReceiver}; use crate::command::{crf_search, sample_encode}; @@ -239,6 +239,19 @@ struct WorkerJobReportState { crf_results: Vec, } +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum WorkerJobOutcome { + Completed, + Stopped, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum WorkerControlState { + Running, + Paused, + Stopped, +} + #[cfg_attr(not(test), allow(dead_code))] #[derive(Debug)] struct PendingJob { @@ -428,7 +441,7 @@ async fn run_worker_job_with_reporting( job: WorkerJob, probe: Arc, worker: &mut Option, -) -> Result { +) -> Result { let crf_config = job.crf_search_config()?; let mut run = std::pin::pin!(crf_search::run(crf_config, probe)); let mut heartbeat = tokio::time::interval(HEARTBEAT_INTERVAL); @@ -437,6 +450,7 @@ async fn run_worker_job_with_reporting( let mut reconnect = tokio::time::interval(Duration::from_secs(5)); reconnect.set_missed_tick_behavior(tokio::time::MissedTickBehavior::Delay); let mut completed_best = None; + let mut paused = false; loop { match worker.as_mut() { @@ -488,6 +502,37 @@ async fn run_worker_job_with_reporting( cancel.job_id, cancel.reason )); } + Some(WorkerPush::Control(control)) + if control.video_id.is_none() + || control.video_id == Some(job.assignment.video_id) => + { + match control.action { + ControlAction::Pause => { + crate::process::managed::pause_active_processes()?; + paused = true; + current_worker.send_control_state( + ControlState::Paused, + Some(job.assignment.video_id), + ).await?; + } + ControlAction::Resume | ControlAction::Start => { + crate::process::managed::resume_active_processes()?; + paused = false; + current_worker.send_control_state( + ControlState::Running, + Some(job.assignment.video_id), + ).await?; + } + ControlAction::Stop => { + crate::process::managed::resume_active_processes()?; + current_worker.send_control_state( + ControlState::Stopped, + None, + ).await?; + return Ok(WorkerJobOutcome::Stopped); + } + } + } _ => {} } } @@ -506,7 +551,7 @@ async fn run_worker_job_with_reporting( } } } - update = run.next() => { + update = run.next(), if !paused => { let (best, disconnected) = handle_crf_update( &job, &mut state, @@ -520,7 +565,7 @@ async fn run_worker_job_with_reporting( if disconnected || worker.is_none() { completed_best = Some(best); } else { - return Ok(best); + return Ok(WorkerJobOutcome::Completed); } } } @@ -535,9 +580,15 @@ async fn run_worker_job_with_reporting( match ConnectedWorker::connect(config).await { Ok(mut reconnected) => { if replay_worker_state(&mut reconnected, &state).await { + if paused { + reconnected.send_control_state( + ControlState::Paused, + Some(job.assignment.video_id), + ).await?; + } *worker = Some(reconnected); - if let Some(best) = completed_best.take() { - return Ok(best); + if completed_best.take().is_some() { + return Ok(WorkerJobOutcome::Completed); } } } @@ -546,7 +597,7 @@ async fn run_worker_job_with_reporting( } } } - update = run.next() => { + update = run.next(), if !paused => { let (best, _) = handle_crf_update( &job, &mut state, @@ -758,8 +809,8 @@ impl Default for WorkerRuntime { fn default() -> Self { Self { idle_delay: Duration::from_secs(5), - reconnect_base_delay: Duration::ZERO, - reconnect_max_delay: Duration::ZERO, + reconnect_base_delay: Duration::from_secs(1), + reconnect_max_delay: Duration::from_secs(30), max_pulls: None, } } @@ -844,11 +895,14 @@ enum PendingJobOutcome { Waiting, Ready, Canceled, + Paused, + Stopped, } #[derive(Debug)] enum WorkerPush { Cancel(CancelPayload), + Control(ControlPayload), Started(TransferStartedPayload), } @@ -871,6 +925,7 @@ struct ConnectedWorker { negotiated_protocol_version: u64, next_ref: u64, socket: WorkerSocket, + pending_control: Option, } impl ConnectedWorker { @@ -905,6 +960,7 @@ impl ConnectedWorker { 2, ClientEvent::Announce(AnnouncePayload { worker_id: config.worker_id.clone(), + hostname: local_hostname(), protocol_version: config.protocol_version, version: config.version.clone(), capabilities: Capabilities { crf_search: true }, @@ -922,6 +978,7 @@ impl ConnectedWorker { negotiated_protocol_version: announce.protocol_version, next_ref: 3, socket, + pending_control: None, }) } @@ -940,7 +997,35 @@ impl ConnectedWorker { ); send_json(&mut self.socket, frame).await?; - expect_reply(&mut self.socket, &request_ref.to_string(), "pull_work").await + let expected_ref = request_ref.to_string(); + while let Some(message) = self.socket.next().await { + match message.context("read websocket message")? { + Message::Text(text) => { + if let Some(WorkerPush::Control(control)) = decode_worker_push(&text)? { + self.pending_control = Some(control.action); + continue; + } + if let Some(reply) = decode_expected_reply(&text, &expected_ref, "pull_work")? { + return reply; + } + } + Message::Ping(payload) => { + self.socket + .send(Message::Pong(payload)) + .await + .context("send websocket pong")?; + } + Message::Close(frame) => { + bail!("websocket closed while waiting for work: {frame:?}") + } + Message::Pong(_) | Message::Binary(_) | Message::Frame(_) => {} + } + } + bail!("websocket ended while waiting for work") + } + + fn take_pending_control(&mut self) -> Option { + self.pending_control.take() } async fn send_event(&mut self, event: ClientEvent) -> Result<()> { @@ -949,6 +1034,18 @@ impl ConnectedWorker { send_json(&mut self.socket, ClientFrame::new(request_ref, event)).await } + async fn send_control_state( + &mut self, + state: ControlState, + active_video_id: Option, + ) -> Result<()> { + self.send_event(ClientEvent::ControlState(ControlStatePayload { + state, + active_video_id, + })) + .await + } + async fn send_transfer_progress(&mut self, payload: TransferProgressPayload) -> Result<()> { let throughput = format_bytes_per_second(payload.bytes_per_second); debug!( @@ -1014,6 +1111,26 @@ impl ConnectedWorker { ); Ok(PendingJobOutcome::Waiting) } + Some(WorkerPush::Control(control)) => match control.action { + ControlAction::Stop => { + self.send_control_state(ControlState::Stopped, None).await?; + Ok(PendingJobOutcome::Stopped) + } + ControlAction::Pause => { + self.send_control_state( + ControlState::Paused, + Some(pending_job.job.assignment.video_id), + ).await?; + Ok(PendingJobOutcome::Paused) + } + ControlAction::Resume | ControlAction::Start => { + self.send_control_state( + ControlState::Running, + Some(pending_job.job.assignment.video_id), + ).await?; + Ok(PendingJobOutcome::Waiting) + } + }, Some(_) => Ok(PendingJobOutcome::Waiting), None => Ok(PendingJobOutcome::Waiting), } @@ -1092,6 +1209,61 @@ impl ConnectedWorker { } } } + + async fn wait_until_running( + &mut self, + control_state: &mut WorkerControlState, + idle_delay: Duration, + ) -> Result { + let mut stopped = *control_state == WorkerControlState::Stopped; + loop { + tokio::select! { + frame = self.socket.next() => match frame { + Some(Ok(Message::Ping(payload))) => { + self.socket.send(Message::Pong(payload)).await.context("send websocket pong")?; + } + Some(Ok(Message::Pong(_))) + | Some(Ok(Message::Binary(_))) + | Some(Ok(Message::Frame(_))) => {} + Some(Ok(Message::Text(text))) => { + if let Some(WorkerPush::Control(control)) = decode_worker_push(&text)? { + match control.action { + ControlAction::Start | ControlAction::Resume => { + self.send_control_state(ControlState::Running, None).await?; + *control_state = WorkerControlState::Running; + return Ok(stopped); + } + ControlAction::Pause if !stopped => { + self.send_control_state(ControlState::Paused, None).await?; + *control_state = WorkerControlState::Paused; + } + ControlAction::Pause | ControlAction::Stop => { + self.send_control_state(ControlState::Stopped, None).await?; + *control_state = WorkerControlState::Stopped; + stopped = true; + } + } + } + } + Some(Ok(Message::Close(frame))) => bail!("websocket closed while worker stopped: {frame:?}"), + Some(Err(error)) => return Err(error).context("read websocket message while worker stopped"), + None => bail!("websocket ended while worker stopped"), + }, + _ = tokio::time::sleep(idle_delay) => { + self.send_event(ClientEvent::Heartbeat(heartbeat_payload(Path::new("."), None))).await?; + } + } + } + } +} + +fn local_hostname() -> Option { + std::env::var("HOSTNAME").ok().or_else(|| { + std::fs::read_to_string("/proc/sys/kernel/hostname") + .ok() + .map(|hostname| hostname.trim().to_owned()) + .filter(|hostname| !hostname.is_empty()) + }) } fn format_bytes_per_second(bytes_per_second: u64) -> String { @@ -1103,7 +1275,7 @@ async fn run_worker_job_and_publish( config: &WorkerConfig, worker: &mut Option, job: &WorkerJob, -) -> Result<()> { +) -> Result { debug!( job_id = %job.assignment.job_id, input = %job.input_path().display(), @@ -1111,36 +1283,40 @@ async fn run_worker_job_and_publish( ); let probe = Arc::new(crate::ffprobe::probe(job.input_path())); debug!(job_id = %job.assignment.job_id, "probe complete, running crf search"); - match run_worker_job_with_reporting(config, job.clone(), probe, worker).await { - Ok(_) => {} + let result = run_worker_job_with_reporting(config, job.clone(), probe, worker).await; + crate::temporary::clean_all().await; + remove_worker_input(job)?; + let outcome = match result { + Ok(outcome) => outcome, Err(error) => { publish_worker_failure(worker, job, &error).await; return Err(error); } - } - remove_completed_worker_input(config, job)?; - Ok(()) + }; + Ok(outcome) } -fn remove_completed_worker_input(config: &WorkerConfig, job: &WorkerJob) -> Result<()> { - if config.local_path.is_some() { +fn remove_worker_input(job: &WorkerJob) -> Result<()> { + if job.input_dir != worker_job_input_dir(&job.assignment.job_id) { return Ok(()); } debug!( job_id = %job.assignment.job_id, input = %job.input_path().display(), - "removing completed worker input" + "removing worker input" ); - fs::remove_file(job.input_path()).with_context(|| { - format!( - "remove completed worker input {}", - job.input_path().display() - ) - })?; + fs::remove_dir_all(&job.input_dir) + .with_context(|| format!("remove worker input directory {}", job.input_dir.display()))?; Ok(()) } +fn remove_pending_worker_input(pending_job: &mut Option) -> Result<()> { + pending_job + .take() + .map_or(Ok(()), |pending| remove_worker_input(&pending.job)) +} + async fn publish_worker_failure( worker: &mut Option, job: &WorkerJob, @@ -1164,6 +1340,16 @@ fn build_worker_job( assignment: crate::command::worker_protocol::JobAssignedPayload, local_path: Option<&Path>, ) -> Result { + if local_path.is_none() + && let Some(input_path) = offered_local_input(&assignment) + { + return Ok(WorkerJob::new( + assignment, + std::env::current_dir().context("current working directory")?, + input_path, + )); + } + let input_dir = worker_job_input_dir(&assignment.job_id); fs::create_dir_all(&input_dir).context("create worker job dir")?; let input_path = local_path @@ -1173,6 +1359,17 @@ fn build_worker_job( Ok(WorkerJob::new(assignment, input_dir, input_path)) } +fn offered_local_input( + assignment: &crate::command::worker_protocol::JobAssignedPayload, +) -> Option { + let path = assignment + .crf_search_args + .windows(2) + .find_map(|args| (args[0] == "--input").then(|| PathBuf::from(&args[1])))?; + let metadata = path.metadata().ok()?; + (metadata.is_file() && metadata.len() == assignment.size_bytes).then_some(path) +} + fn worker_job_input_path( input_dir: &Path, assignment: &crate::command::worker_protocol::JobAssignedPayload, @@ -1414,11 +1611,13 @@ pub async fn worker(config: WorkerConfig) -> Result<()> { async fn run_worker_until(config: &WorkerConfig, runtime: WorkerRuntime) -> Result<()> { let mut completed_pulls = 0usize; + let mut control_state = WorkerControlState::Running; let mut reconnect_backoff = ReconnectBackoff::new(runtime.reconnect_base_delay, runtime.reconnect_max_delay); loop { - match run_connected_worker(config, runtime, &mut completed_pulls).await { + match run_connected_worker(config, runtime, &mut completed_pulls, &mut control_state).await + { Ok(()) => { reconnect_backoff.reset(); return Ok(()); @@ -1519,6 +1718,7 @@ async fn run_connected_worker( config: &WorkerConfig, runtime: WorkerRuntime, completed_pulls: &mut usize, + control_state: &mut WorkerControlState, ) -> Result<()> { debug!( connect = %config.connect, @@ -1530,7 +1730,31 @@ async fn run_connected_worker( let mut worker = Some(ConnectedWorker::connect(config).await?); let mut pending_job: Option = None; + if *control_state != WorkerControlState::Running { + let reported_state = match control_state { + WorkerControlState::Paused => ControlState::Paused, + WorkerControlState::Stopped => ControlState::Stopped, + WorkerControlState::Running => unreachable!(), + }; + worker + .as_mut() + .expect("connected worker") + .send_control_state(reported_state, None) + .await?; + } + loop { + if *control_state != WorkerControlState::Running { + let stopped = worker + .as_mut() + .expect("connected worker") + .wait_until_running(control_state, runtime.idle_delay) + .await?; + if stopped { + remove_pending_worker_input(&mut pending_job)?; + } + } + if pending_job.is_some() { let next = { let job = pending_job.as_mut().expect("pending job"); @@ -1562,7 +1786,16 @@ async fn run_connected_worker( if let Some(job) = pending_job.as_ref() { debug!(job_id = %job.job.assignment.job_id, "pending job canceled"); } - pending_job = None; + remove_pending_worker_input(&mut pending_job)?; + continue; + } + PendingJobOutcome::Paused => { + *control_state = WorkerControlState::Paused; + continue; + } + PendingJobOutcome::Stopped => { + remove_pending_worker_input(&mut pending_job)?; + *control_state = WorkerControlState::Stopped; continue; } PendingJobOutcome::Ready => { @@ -1572,7 +1805,11 @@ async fn run_connected_worker( input = %job.input_path().display(), "pending job input arrived" ); - run_worker_job_and_publish(config, &mut worker, &job.job).await?; + if run_worker_job_and_publish(config, &mut worker, &job.job).await? + == WorkerJobOutcome::Stopped + { + *control_state = WorkerControlState::Stopped; + } continue; } } @@ -1581,6 +1818,29 @@ async fn run_connected_worker( debug!("requesting work"); let worker_ref = worker.as_mut().expect("connected worker"); let work_status = worker_ref.request_work().await?; + if let Some(control) = worker_ref.take_pending_control() { + match control { + ControlAction::Stop => { + worker_ref + .send_control_state(ControlState::Stopped, None) + .await?; + *control_state = WorkerControlState::Stopped; + continue; + } + ControlAction::Pause => { + worker_ref + .send_control_state(ControlState::Paused, None) + .await?; + *control_state = WorkerControlState::Paused; + continue; + } + ControlAction::Resume | ControlAction::Start => { + worker_ref + .send_control_state(ControlState::Running, None) + .await?; + } + } + } *completed_pulls += 1; let status = work_status_label(&work_status); println!( @@ -1617,7 +1877,11 @@ async fn run_connected_worker( phase = ?phase, "input already present, starting job" ); - run_worker_job_and_publish(config, &mut worker, &job).await?; + if run_worker_job_and_publish(config, &mut worker, &job).await? + == WorkerJobOutcome::Stopped + { + *control_state = WorkerControlState::Stopped; + } } WorkerJobPhase::AwaitingInput(delivery) => { let pending = match delivery { @@ -1635,7 +1899,11 @@ async fn run_connected_worker( input = %job.input_path().display(), "downloaded worker input over HTTP, starting job" ); - run_worker_job_and_publish(config, &mut worker, &job).await?; + if run_worker_job_and_publish(config, &mut worker, &job).await? + == WorkerJobOutcome::Stopped + { + *control_state = WorkerControlState::Stopped; + } None } Some(pending) => Some(pending), @@ -1727,6 +1995,10 @@ fn decode_worker_push(text: &str) -> Result> { serde_json::from_value::(payload.clone()) .context("decode cancel push")?, ), + "control" => WorkerPush::Control( + serde_json::from_value::(payload.clone()) + .context("decode control push")?, + ), "transfer_started" => WorkerPush::Started( serde_json::from_value::(payload.clone()) .with_context(|| format!("decode transfer started push event={}", frame.3))?, @@ -1926,53 +2198,9 @@ where while let Some(message) = reader.next().await { match message.context("read websocket message")? { Message::Text(text) => { - let ServerPushFrame(_, msg_ref, topic, event, body): ServerPushFrame = - serde_json::from_str(&text).context("decode phoenix frame")?; - if topic != CRF_SEARCH_TOPIC - || event != "phx_reply" - || msg_ref.as_deref() != Some(expected_ref) - { - continue; + if let Some(reply) = decode_expected_reply(&text, expected_ref, expected_event)? { + return reply; } - debug!( - expected_event, - raw = %text, - "received phoenix reply" - ); - - let ReplyBody { status, response }: ReplyBody = - serde_json::from_value::>(body) - .context("decode phoenix reply body")?; - debug!( - expected_event, - status = %status, - response = %response, - "decoded phoenix reply" - ); - return match status.as_str() { - "ok" => serde_json::from_value(response.clone()).map_err(|error| { - let raw_response = response.to_string(); - anyhow!("decode phoenix ok reply: {error}; raw_response={raw_response}") - }), - "error" => { - let error: ErrorReplyPayload = serde_json::from_value(response) - .context("decode phoenix error reply")?; - let supported_versions = match error.supported_protocol_versions.is_empty() - { - true => String::new(), - false => format!( - " (supported_protocol_versions={:?})", - error.supported_protocol_versions - ), - }; - Err(anyhow!( - "{expected_event} failed: {}{}", - error.reason, - supported_versions - )) - } - other => Err(anyhow!("unexpected phoenix status {other}")), - }; } Message::Close(frame) => { bail!("websocket closed before {expected_event} reply: {frame:?}") @@ -1986,6 +2214,46 @@ where bail!("websocket ended before {expected_event} reply") } +fn decode_expected_reply( + text: &str, + expected_ref: &str, + expected_event: &str, +) -> Result>> +where + T: for<'de> Deserialize<'de>, +{ + let ServerPushFrame(_, msg_ref, topic, event, body): ServerPushFrame = + serde_json::from_str(text).context("decode phoenix frame")?; + if topic != CRF_SEARCH_TOPIC || event != "phx_reply" || msg_ref.as_deref() != Some(expected_ref) + { + return Ok(None); + } + let ReplyBody { status, response }: ReplyBody = + serde_json::from_value(body).context("decode phoenix reply body")?; + Ok(Some(match status.as_str() { + "ok" => serde_json::from_value(response.clone()) + .map_err(|error| anyhow!("decode phoenix ok reply: {error}; raw_response={response}")), + "error" => { + let error: ErrorReplyPayload = + serde_json::from_value(response).context("decode phoenix error reply")?; + let supported_versions = if error.supported_protocol_versions.is_empty() { + String::new() + } else { + format!( + " (supported_protocol_versions={:?})", + error.supported_protocol_versions + ) + }; + Err(anyhow!( + "{expected_event} failed: {}{}", + error.reason, + supported_versions + )) + } + other => Err(anyhow!("unexpected phoenix status {other}")), + })) +} + #[cfg(test)] mod tests { use super::*; @@ -2172,6 +2440,98 @@ mod tests { Ok(()) } + #[tokio::test(flavor = "current_thread")] + async fn stopped_worker_waits_for_start_before_pulling_again() -> Result<()> { + let (listener, address) = FakeCoordinator::bind("127.0.0.1:0").await?; + let server = tokio::spawn(async move { + let (stream, _) = listener.accept().await.expect("accept connection"); + let socket = accept_async(stream).await.expect("accept websocket"); + let (mut writer, mut reader) = socket.split(); + + expect_join(&mut reader).await; + send_join_reply(&mut writer).await; + expect_announce(&mut reader, 1).await; + send_announce_reply(&mut writer).await; + expect_pull_work(&mut reader, 3).await; + send_control_push(&mut writer, ControlAction::Stop, None).await; + send_no_work_reply(&mut writer, 3).await; + assert_eq!( + expect_client_event(&mut reader, 4, "control_state").await, + json!({"state": "stopped"}) + ); + send_control_push(&mut writer, ControlAction::Start, None).await; + assert_eq!( + expect_client_event(&mut reader, 5, "control_state").await, + json!({"state": "running"}) + ); + expect_pull_work(&mut reader, 6).await; + send_no_work_reply(&mut writer, 6).await; + }); + + run_worker_until( + &FakeCoordinator { + address, + server: tokio::spawn(async {}), + } + .worker_config(WorkerTestConfig::continuous()), + WorkerRuntime { + idle_delay: Duration::from_secs(30), + reconnect_base_delay: Duration::from_millis(1), + reconnect_max_delay: Duration::from_millis(1), + max_pulls: Some(1), + }, + ) + .await?; + server.await.expect("server task"); + Ok(()) + } + + #[tokio::test(flavor = "current_thread")] + async fn paused_worker_waits_for_resume_before_pulling_again() -> Result<()> { + let (listener, address) = FakeCoordinator::bind("127.0.0.1:0").await?; + let server = tokio::spawn(async move { + let (stream, _) = listener.accept().await.expect("accept connection"); + let socket = accept_async(stream).await.expect("accept websocket"); + let (mut writer, mut reader) = socket.split(); + + expect_join(&mut reader).await; + send_join_reply(&mut writer).await; + expect_announce(&mut reader, 1).await; + send_announce_reply(&mut writer).await; + expect_pull_work(&mut reader, 3).await; + send_control_push(&mut writer, ControlAction::Pause, None).await; + send_no_work_reply(&mut writer, 3).await; + assert_eq!( + expect_client_event(&mut reader, 4, "control_state").await, + json!({"state": "paused"}) + ); + send_control_push(&mut writer, ControlAction::Resume, None).await; + assert_eq!( + expect_client_event(&mut reader, 5, "control_state").await, + json!({"state": "running"}) + ); + expect_pull_work(&mut reader, 6).await; + send_no_work_reply(&mut writer, 6).await; + }); + + run_worker_until( + &FakeCoordinator { + address, + server: tokio::spawn(async {}), + } + .worker_config(WorkerTestConfig::continuous()), + WorkerRuntime { + idle_delay: Duration::from_secs(30), + reconnect_base_delay: Duration::from_millis(1), + reconnect_max_delay: Duration::from_millis(1), + max_pulls: Some(1), + }, + ) + .await?; + server.await.expect("server task"); + Ok(()) + } + #[tokio::test(flavor = "current_thread")] async fn worker_reconnects_after_disconnect_and_continues_pulling_work() -> Result<()> { let coordinator = FakeCoordinator::with_no_work_replies(1).await?; @@ -2240,6 +2600,7 @@ mod tests { } .worker_config(WorkerTestConfig::continuous()); let mut completed_pulls = 0; + let mut control_state = WorkerControlState::Running; let error = run_connected_worker( &config, WorkerRuntime { @@ -2249,6 +2610,7 @@ mod tests { max_pulls: None, }, &mut completed_pulls, + &mut control_state, ) .await .expect_err("server closes after resend assignment"); @@ -2340,6 +2702,28 @@ mod tests { Ok(()) } + #[test] + fn worker_decodes_pause_control_push() -> Result<()> { + let text = serde_json::to_string(&ServerPushFrame::new( + "control", + crate::command::worker_protocol::ControlPayload { + action: crate::command::worker_protocol::ControlAction::Pause, + video_id: Some(123), + }, + ))?; + + assert!(matches!( + decode_worker_push(&text)?, + Some(WorkerPush::Control( + crate::command::worker_protocol::ControlPayload { + action: crate::command::worker_protocol::ControlAction::Pause, + video_id: Some(123), + } + )) + )); + Ok(()) + } + #[test] fn worker_formats_assigned_job_status_with_job_id() { let status = work_status_label(&ServerReply::JobAssigned(JobAssignedPayload { @@ -2426,6 +2810,36 @@ mod tests { assert_eq!(job.input_path(), local_path.as_path()); } + #[test] + fn build_worker_job_uses_offered_local_file_when_size_matches() -> Result<()> { + let worker_dir = worker_job_input_dir("offered-local"); + let _ = fs::remove_dir_all(&worker_dir); + let root = std::env::temp_dir().join(format!( + "ab-av1-worker-offered-local-{}", + std::process::id() + )); + let local_path = root.join("movie.mkv"); + fs::create_dir_all(&root)?; + fs::write(&local_path, b"data")?; + let assignment = serde_json::from_value(json!({ + "status": "job_assigned", + "job_id": "offered-local", + "video_id": 123, + "source_name": "movie.mkv", + "local_path": local_path, + "size_bytes": 4, + "target_vmaf": 96.5, + "crf_search_args": ["crf-search", "--input", local_path] + }))?; + + let job = build_worker_job(assignment, None)?; + + assert_eq!(job.input_path(), local_path.as_path()); + assert!(!worker_dir.exists()); + fs::remove_dir_all(root)?; + Ok(()) + } + #[test] fn build_worker_job_uses_source_basename_for_worker_file_lookup() -> Result<()> { let job_dir = worker_job_input_dir("job-path-source"); @@ -2578,9 +2992,9 @@ mod tests { } #[test] - fn completed_worker_input_cleanup_respects_local_path() -> Result<()> { - let root = - std::env::temp_dir().join(format!("ab-av1-worker-cleanup-{}", std::process::id())); + fn worker_input_cleanup_removes_owned_directory_only() -> Result<()> { + let root = worker_job_input_dir("cleanup-job"); + let _ = fs::remove_dir_all(&root); let input_path = root.join("movie.mkv"); fs::create_dir_all(&root)?; fs::write(&input_path, b"data")?; @@ -2605,25 +3019,8 @@ mod tests { root.clone(), input_path.clone(), ); - let config = WorkerConfig { - connect: String::new(), - token: String::new(), - worker_id: String::new(), - version: String::new(), - protocol_version: 1, - once: false, - local_path: Some(input_path.clone()), - }; - - remove_completed_worker_input(&config, &job)?; - assert!(input_path.exists()); - let config = WorkerConfig { - local_path: None, - ..config - }; - remove_completed_worker_input(&config, &job)?; - assert!(!input_path.exists()); - fs::remove_dir_all(root)?; + remove_worker_input(&job)?; + assert!(!root.exists()); Ok(()) } @@ -2796,6 +3193,14 @@ mod tests { assert_eq!(backoff.next_delay(), Duration::from_millis(1_000)); } + #[test] + fn worker_runtime_defaults_back_off_reconnects() { + let runtime = WorkerRuntime::default(); + + assert_eq!(runtime.reconnect_base_delay, Duration::from_secs(1)); + assert_eq!(runtime.reconnect_max_delay, Duration::from_secs(30)); + } + #[test] fn reconnect_backoff_resets_after_success() { let mut backoff = @@ -2936,6 +3341,7 @@ mod tests { 2, ClientEvent::Announce(AnnouncePayload { worker_id: "abav1-dev".into(), + hostname: local_hostname(), protocol_version, version: "0.11.4".into(), capabilities: Capabilities { crf_search: true }, @@ -3139,6 +3545,22 @@ mod tests { .expect("send cancel push"); } + async fn send_control_push(writer: &mut W, action: ControlAction, video_id: Option) + where + W: SinkExt + Unpin, + { + writer + .send(Message::Text( + serde_json::to_string(&ServerPushFrame::new( + "control", + ControlPayload { action, video_id }, + )) + .expect("control push json"), + )) + .await + .expect("send control push"); + } + async fn serve_no_work_session(listener: TcpListener, no_work_replies: usize) { let (stream, _) = listener.accept().await.expect("accept connection"); let socket = accept_async(stream).await.expect("accept websocket"); diff --git a/src/command/worker_protocol.rs b/src/command/worker_protocol.rs index 5207a640..c21a0919 100644 --- a/src/command/worker_protocol.rs +++ b/src/command/worker_protocol.rs @@ -24,6 +24,7 @@ pub(crate) enum ClientEvent { Announce(AnnouncePayload), PullWork(PullWorkPayload), Heartbeat(HeartbeatPayload), + ControlState(ControlStatePayload), TransferProgress(TransferProgressPayload), TransferFailure(TransferFailurePayload), CrfSearchProgress(CrfSearchProgressPayload), @@ -38,6 +39,7 @@ impl ClientEvent { Self::Announce(payload) => ("announce", ClientPayload::Announce(payload)), Self::PullWork(payload) => ("pull_work", ClientPayload::PullWork(payload)), Self::Heartbeat(payload) => ("heartbeat", ClientPayload::Heartbeat(payload)), + Self::ControlState(payload) => ("control_state", ClientPayload::ControlState(payload)), Self::TransferProgress(payload) => ( "transfer_progress", ClientPayload::TransferProgress(payload), @@ -61,6 +63,7 @@ enum ClientPayload { Announce(AnnouncePayload), PullWork(PullWorkPayload), Heartbeat(HeartbeatPayload), + ControlState(ControlStatePayload), TransferProgress(TransferProgressPayload), TransferFailure(TransferFailurePayload), Progress(CrfSearchProgressPayload), @@ -92,6 +95,8 @@ fn is_false(value: &bool) -> bool { #[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] pub(crate) struct AnnouncePayload { pub(crate) worker_id: String, + #[serde(skip_serializing_if = "Option::is_none")] + pub(crate) hostname: Option, pub(crate) protocol_version: u64, pub(crate) version: String, pub(crate) capabilities: Capabilities, @@ -350,6 +355,37 @@ pub(crate) struct CancelPayload { pub(crate) reason: String, } +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "snake_case")] +pub(crate) enum ControlAction { + Pause, + Resume, + Start, + Stop, +} + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub(crate) struct ControlPayload { + pub(crate) action: ControlAction, + #[serde(default)] + pub(crate) video_id: Option, +} + +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub(crate) struct ControlStatePayload { + pub(crate) state: ControlState, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub(crate) active_video_id: Option, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "snake_case")] +pub(crate) enum ControlState { + Running, + Paused, + Stopped, +} + #[cfg_attr(not(test), allow(dead_code))] #[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] pub(crate) struct TransferStartedPayload { @@ -453,6 +489,7 @@ mod tests { 2, ClientEvent::Announce(AnnouncePayload { worker_id: "abav1-dev".into(), + hostname: None, protocol_version: 1, version: "0.11.4".into(), capabilities: Capabilities { crf_search: true }, @@ -509,6 +546,28 @@ mod tests { ); } + #[test] + fn control_state_serializes_worker_acknowledgement() { + let frame = ClientFrame::new( + 5, + ClientEvent::ControlState(ControlStatePayload { + state: ControlState::Paused, + active_video_id: Some(123), + }), + ); + + assert_eq!( + serde_json::to_value(frame).expect("serialize control state"), + json!([ + "1", + "5", + "workers:crf_search", + "control_state", + {"state": "paused", "active_video_id": 123} + ]) + ); + } + #[test] fn outbound_float_payloads_round_to_two_places() { assert_eq!( @@ -983,6 +1042,32 @@ mod tests { ); } + #[test] + fn server_push_parses_worker_control_payload() { + let push: ServerPushFrame = serde_json::from_value(json!([ + null, + null, + "workers:crf_search", + "control", + { + "action": "pause", + "video_id": 123 + } + ])) + .expect("parse worker control push"); + + assert_eq!( + push, + ServerPushFrame::new( + "control", + ControlPayload { + action: ControlAction::Pause, + video_id: Some(123), + }, + ) + ); + } + #[test] fn transfer_started_payload_serializes_metadata_side_channel() { let payload = TransferStartedPayload { diff --git a/src/process/managed.rs b/src/process/managed.rs index 3e809dfd..aa729c8e 100644 --- a/src/process/managed.rs +++ b/src/process/managed.rs @@ -4,7 +4,18 @@ )] use anyhow::bail; +#[cfg(unix)] +use nix::{ + sys::signal::{Signal, killpg}, + unistd::Pid, +}; +#[cfg(unix)] +use std::collections::HashSet; +#[cfg(unix)] +use std::io; use std::process::{ExitStatus, Output}; +#[cfg(unix)] +use std::sync::{Mutex, OnceLock}; use std::time::Duration; use tokio::process::Command; use tokio_process_tools::{Chunk, visitors::inspect::InspectChunks}; @@ -21,6 +32,14 @@ const DEFAULT_WAIT_TIMEOUT: Duration = Duration::from_secs(30 * 24 * 60 * 60); const DEFAULT_TERMINATION_GRACE: Duration = Duration::from_millis(25); const DEFAULT_STDERR_LIMIT: usize = 32_768; +#[cfg(unix)] +static ACTIVE_PROCESS_GROUPS: OnceLock>> = OnceLock::new(); + +#[cfg(unix)] +fn active_process_groups() -> &'static Mutex> { + ACTIVE_PROCESS_GROUPS.get_or_init(|| Mutex::new(HashSet::new())) +} + #[derive(Debug, Clone, Copy)] pub struct ManagedProcessOptions { wait_timeout: Duration, @@ -56,6 +75,20 @@ pub struct ManagedProcess { SingleSubscriberOutputStream, >, options: ManagedProcessOptions, + #[cfg(unix)] + process_group: Option, +} + +impl Drop for ManagedProcess { + fn drop(&mut self) { + #[cfg(unix)] + if let Some(process_group) = self.process_group { + active_process_groups() + .lock() + .expect("active process groups lock") + .remove(&process_group); + } + } } /// Process policy for streams that must run through process completion. @@ -186,7 +219,21 @@ impl ManagedProcess { .max_buffered_chunks(DEFAULT_MAX_BUFFERED_CHUNKS) }) .spawn()?; - Ok(Self { handle, options }) + #[cfg(unix)] + let process_group = handle.id().map(|pid| pid as i32); + #[cfg(unix)] + if let Some(process_group) = process_group { + active_process_groups() + .lock() + .expect("active process groups lock") + .insert(process_group); + } + Ok(Self { + handle, + options, + #[cfg(unix)] + process_group, + }) } fn graceful_shutdown_for(options: ManagedProcessOptions) -> GracefulShutdown { @@ -329,6 +376,37 @@ impl ManagedProcess { self.handle.id() } + #[cfg(unix)] + pub fn pause(&mut self) -> anyhow::Result<()> { + self.send_process_group_signal("SIGSTOP", Signal::SIGSTOP) + } + + #[cfg(unix)] + pub fn resume(&mut self) -> anyhow::Result<()> { + self.send_process_group_signal("SIGCONT", Signal::SIGCONT) + } + + #[cfg(unix)] + fn send_process_group_signal( + &mut self, + signal_name: &'static str, + signal: Signal, + ) -> anyhow::Result<()> { + self.handle + .send_signal_with_reaper( + signal_name, + |handle| { + let pid = handle.id().ok_or_else(|| { + io::Error::new(io::ErrorKind::NotFound, "managed process already exited") + })?; + killpg(Pid::from_raw(pid as i32), signal) + .map_err(|error| io::Error::from_raw_os_error(error as i32)) + }, + |_| Ok(None), + ) + .map_err(Into::into) + } + pub async fn terminate_after(mut self, timeout: Duration) -> anyhow::Result { Ok(self .handle @@ -344,6 +422,33 @@ impl ManagedProcess { } } +#[cfg(unix)] +pub fn pause_active_processes() -> anyhow::Result<()> { + signal_active_processes(Signal::SIGSTOP) +} + +#[cfg(unix)] +pub fn resume_active_processes() -> anyhow::Result<()> { + signal_active_processes(Signal::SIGCONT) +} + +#[cfg(unix)] +fn signal_active_processes(signal: Signal) -> anyhow::Result<()> { + let groups = active_process_groups() + .lock() + .expect("active process groups lock") + .iter() + .copied() + .collect::>(); + for group in groups { + match killpg(Pid::from_raw(group), signal) { + Ok(()) | Err(nix::errno::Errno::ESRCH) => {} + Err(error) => return Err(error.into()), + } + } + Ok(()) +} + fn managed_event_from_stream_event(event: StreamEvent) -> anyhow::Result> { Ok(match event { StreamEvent::Chunk(chunk) => Some(ManagedEvent::RawStderr(RawOutputChunk::new( @@ -459,6 +564,7 @@ mod test_support; #[cfg(test)] mod tests { use super::*; + use anyhow::Context; use std::{ env, sync::{Arc, Mutex}, @@ -910,4 +1016,30 @@ mod tests { "timeout termination should return the child terminal status" ); } + + #[cfg(target_os = "linux")] + #[tokio::test] + async fn managed_process_pauses_and_resumes_process_group() -> anyhow::Result<()> { + let mut command = Command::new("sleep"); + command.arg("30"); + let process = ManagedProcess::spawn("pause-resume-fixture", command)?; + let pid = process.id().expect("fixture process id"); + + pause_active_processes()?; + assert_eq!(linux_process_state(pid)?, 'T'); + + resume_active_processes()?; + assert_ne!(linux_process_state(pid)?, 'T'); + + process.terminate_after(Duration::ZERO).await?; + Ok(()) + } + + #[cfg(target_os = "linux")] + fn linux_process_state(pid: u32) -> anyhow::Result { + let stat = std::fs::read_to_string(format!("/proc/{pid}/stat"))?; + stat.rsplit_once(") ") + .and_then(|(_, fields)| fields.chars().next()) + .context("read Linux process state") + } }