diff --git a/scripts/smoke_live_handoff_sessions.sh b/scripts/smoke_live_handoff_sessions.sh new file mode 100755 index 00000000..44dba425 --- /dev/null +++ b/scripts/smoke_live_handoff_sessions.sh @@ -0,0 +1,221 @@ +#!/usr/bin/env bash +set -euo pipefail + +ROOT="$(cd "$(dirname "${BASH_SOURCE[0]}")/.." && pwd)" +HERDR_BIN="${HERDR_BIN:-$ROOT/target/debug/herdr}" +BASE="${BASE:-$(mktemp -d /tmp/herdr-handoff-smoke.XXXXXX)}" +CONFIG_HOME="$BASE/config" +RUNTIME_DIR="$BASE/runtime" +STATE_DIR="$BASE/state" + +if [[ "$CONFIG_HOME" == "$HOME/.config" || "$CONFIG_HOME" == "$HOME/.config/"* ]]; then + echo "refusing to run smoke test against $CONFIG_HOME" >&2 + exit 1 +fi + +sessions=("default" "work" "api") +ports=() +server_pids=() + +cleanup() { + set +e + for session in "${sessions[@]}"; do + run_herdr "$session" server stop >/dev/null 2>&1 || true + done + for pid in "${server_pids[@]}"; do + kill "$pid" >/dev/null 2>&1 || true + done +} +trap cleanup EXIT + +run_herdr() { + local session="$1" + shift + local socket + socket="$(api_socket "$session")" + assert_smoke_socket "$socket" + mkdir -p "$(dirname "$socket")" "$RUNTIME_DIR" "$STATE_DIR" + env -u HERDR_SOCKET_PATH \ + -u HERDR_CLIENT_SOCKET_PATH \ + -u HERDR_SESSION \ + XDG_CONFIG_HOME="$CONFIG_HOME" \ + XDG_RUNTIME_DIR="$RUNTIME_DIR" \ + XDG_STATE_HOME="$STATE_DIR" \ + "$HERDR_BIN" --session "$session" "$@" +} + +session_dir() { + local session="$1" + if [[ "$session" == "default" ]]; then + printf '%s/herdr-dev' "$CONFIG_HOME" + else + printf '%s/herdr-dev/sessions/%s' "$CONFIG_HOME" "$session" + fi +} + +api_socket() { + printf '%s/herdr.sock' "$(session_dir "$1")" +} + +client_socket() { + printf '%s/herdr-client.sock' "$(session_dir "$1")" +} + +assert_smoke_socket() { + local socket="$1" + case "$socket" in + "$CONFIG_HOME"/herdr-dev/herdr.sock | "$CONFIG_HOME"/herdr-dev/sessions/*/herdr.sock) + ;; + *) + echo "refusing to use non-smoke socket: $socket" >&2 + exit 1 + ;; + esac +} + +wait_for_socket() { + local socket="$1" + for _ in {1..200}; do + [[ -S "$socket" ]] && return 0 + sleep 0.05 + done + echo "socket did not appear: $socket" >&2 + return 1 +} + +wait_for_http() { + local port="$1" + local expected="$2" + for _ in {1..200}; do + if curl -fsS "http://127.0.0.1:$port/" | grep -q "$expected"; then + return 0 + fi + sleep 0.05 + done + echo "http server on port $port did not return $expected" >&2 + return 1 +} + +json_request() { + local socket="$1" + local body="$2" + python3 - "$socket" "$body" <<'PY' +import socket +import sys + +path, body = sys.argv[1], sys.argv[2] +client = socket.socket(socket.AF_UNIX, socket.SOCK_STREAM) +client.connect(path) +client.sendall(body.encode() + b"\n") +response = b"" +while not response.endswith(b"\n"): + chunk = client.recv(65536) + if not chunk: + break + response += chunk +print(response.decode().strip()) +PY +} + +pane_id_for_session() { + local session="$1" + local socket + socket="$(api_socket "$session")" + json_request "$socket" '{"id":"smoke:workspace:create","method":"workspace.create","params":{"cwd":"/tmp","focus":true}}' \ + | python3 -c 'import json,sys; print(json.load(sys.stdin)["result"]["root_pane"]["pane_id"])' +} + +send_text() { + local session="$1" + local pane="$2" + local text="$3" + local socket + socket="$(api_socket "$session")" + python3 - "$socket" "$pane" "$text" <<'PY' +import json +import socket +import sys + +path, pane, text = sys.argv[1], sys.argv[2], sys.argv[3] +request = { + "id": "smoke:pane:send", + "method": "pane.send_input", + "params": {"pane_id": pane, "text": text, "keys": ["Enter"]}, +} +client = socket.socket(socket.AF_UNIX, socket.SOCK_STREAM) +client.connect(path) +client.sendall(json.dumps(request).encode() + b"\n") +response = b"" +while not response.endswith(b"\n"): + chunk = client.recv(65536) + if not chunk: + break + response += chunk +if b'"error"' in response: + raise SystemExit(response.decode()) +PY +} + +unused_port() { + python3 - <<'PY' +import socket +sock = socket.socket() +sock.bind(("127.0.0.1", 0)) +print(sock.getsockname()[1]) +sock.close() +PY +} + +smoke_http_count() { + local count=0 + local port matches + for port in "${ports[@]}"; do + matches="$(pgrep -fc "python3 -m http.server $port --bind 127.0.0.1" || true)" + if [[ -n "$matches" ]]; then + count=$((count + matches)) + fi + done + printf '%s\n' "$count" +} + +echo "using herdr: $HERDR_BIN" +echo "smoke base: $BASE" + +cargo build --locked --manifest-path "$ROOT/Cargo.toml" >/dev/null +mkdir -p "$CONFIG_HOME/herdr-dev" "$RUNTIME_DIR" "$STATE_DIR" +printf 'onboarding = false\n' > "$CONFIG_HOME/herdr-dev/config.toml" + +for session in "${sessions[@]}"; do + echo "starting smoke session $session at $(api_socket "$session")" + run_herdr "$session" server >/dev/null 2>&1 & + server_pids+=("$!") + wait_for_socket "$(api_socket "$session")" +done + +for session in "${sessions[@]}"; do + port="$(unused_port)" + ports+=("$port") + web="$BASE/web-$session" + mkdir -p "$web" + printf 'hello-from-%s\n' "$session" > "$web/index.html" + pane="$(pane_id_for_session "$session")" + send_text "$session" "$pane" "cd '$web' && python3 -m http.server $port --bind 127.0.0.1" + wait_for_http "$port" "hello-from-$session" +done + +before_count="$(smoke_http_count)" +echo "smoke python http.server process count before handoff: $before_count" + +for session in "${sessions[@]}"; do + socket="$(api_socket "$session")" + json_request "$socket" '{"id":"smoke:handoff","method":"server.live_handoff","params":{}}' >/dev/null + wait_for_socket "$socket" +done + +for i in "${!sessions[@]}"; do + wait_for_http "${ports[$i]}" "hello-from-${sessions[$i]}" +done + +after_count="$(smoke_http_count)" +echo "smoke python http.server process count after handoff: $after_count" +echo "multi-session live handoff smoke passed" diff --git a/src/api/client.rs b/src/api/client.rs index 7929f937..7d966965 100644 --- a/src/api/client.rs +++ b/src/api/client.rs @@ -111,9 +111,14 @@ impl ApiClient { method: Method::Ping(PingParams::default()), })?; match response.result { - ResponseResult::Pong { version, protocol } => Ok(crate::api::RuntimeStatus { + ResponseResult::Pong { + version, + protocol, + capabilities, + } => Ok(crate::api::RuntimeStatus { version: Some(version), protocol: Some(protocol), + capabilities, }), result => Err(ApiClientError::UnexpectedResult(format!("{result:?}"))), } diff --git a/src/api/mod.rs b/src/api/mod.rs index e3ad58c4..1f2ab2bb 100644 --- a/src/api/mod.rs +++ b/src/api/mod.rs @@ -7,7 +7,7 @@ mod subscriptions; mod wait; pub use event_hub::EventHub; -pub use server::start_server; +pub use server::{start_server, start_server_with_capabilities, ServerHandle}; pub use status::{read_runtime_status_at, RuntimeStatus}; use std::path::PathBuf; diff --git a/src/api/schema.rs b/src/api/schema.rs index 907bc2b6..4929d4a7 100644 --- a/src/api/schema.rs +++ b/src/api/schema.rs @@ -14,6 +14,8 @@ pub enum Method { Ping(PingParams), #[serde(rename = "server.stop")] ServerStop(EmptyParams), + #[serde(rename = "server.live_handoff")] + ServerLiveHandoff(ServerLiveHandoffParams), #[serde(rename = "server.reload_config")] ServerReloadConfig(EmptyParams), #[serde(rename = "workspace.create")] @@ -307,6 +309,16 @@ pub struct PaneSendInputParams { pub keys: Vec, } +#[derive(Debug, Clone, Default, PartialEq, Eq, Serialize, Deserialize)] +pub struct ServerLiveHandoffParams { + #[serde(default, skip_serializing_if = "Option::is_none")] + pub import_exe: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub expected_protocol: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub expected_version: Option, +} + #[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] pub struct PaneReadParams { pub pane_id: String, @@ -581,12 +593,19 @@ pub struct ErrorBody { pub message: String, } +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +pub struct ServerCapabilities { + pub live_handoff: bool, +} + #[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] #[serde(tag = "type", rename_all = "snake_case")] pub enum ResponseResult { Pong { version: String, protocol: u32, + #[serde(default)] + capabilities: Option, }, WorkspaceInfo { workspace: WorkspaceInfo, @@ -1225,6 +1244,7 @@ mod tests { result: ResponseResult::Pong { version: "0.1.2".into(), protocol: 6, + capabilities: Some(ServerCapabilities { live_handoff: true }), }, }; diff --git a/src/api/server.rs b/src/api/server.rs index 513262bc..696e1b65 100644 --- a/src/api/server.rs +++ b/src/api/server.rs @@ -1,29 +1,34 @@ -use std::fs; -use std::io::{BufRead, BufReader, Read, Write}; +use std::io::{self, Read, Write}; use std::os::unix::net::{UnixListener, UnixStream}; use std::path::{Path, PathBuf}; use std::sync::atomic::{AtomicBool, Ordering}; use std::sync::Arc; -use std::time::Duration; +use std::time::{Duration, Instant}; use tracing::{debug, error, info, warn}; +#[cfg(test)] +use std::fs; + use crate::api::schema::{ - ErrorBody, ErrorResponse, Method, Request, ResponseResult, SuccessResponse, + ErrorBody, ErrorResponse, Method, Request, ResponseResult, ServerCapabilities, SuccessResponse, }; use crate::api::subscriptions::ActiveSubscription; use crate::api::wait::wait_for_output; use crate::api::{request_changes_ui, socket_path, ApiRequestMessage, ApiRequestSender, EventHub}; +use crate::ipc::{remove_socket_file_if_owned, socket_file_identity, SocketFileIdentity}; const SOCKET_PERMISSION_MODE: u32 = 0o600; pub(super) const CONNECTION_POLL_INTERVAL: Duration = Duration::from_millis(100); pub(super) const APP_RESPONSE_TIMEOUT: Duration = Duration::from_secs(5); const INITIAL_REQUEST_TIMEOUT: Duration = Duration::from_secs(5); const STREAM_WRITE_TIMEOUT: Duration = Duration::from_secs(5); +const MAX_INITIAL_REQUEST_BYTES: usize = 1024 * 1024; pub struct ServerHandle { _thread: std::thread::JoinHandle<()>, path: PathBuf, + identity: SocketFileIdentity, running: Arc, } @@ -31,7 +36,7 @@ impl Drop for ServerHandle { fn drop(&mut self) { self.running.store(false, Ordering::Relaxed); - if let Err(err) = fs::remove_file(&self.path) { + if let Err(err) = self.remove_socket_file_if_owned() { if err.kind() != std::io::ErrorKind::NotFound { warn!(path = %self.path.display(), err = %err, "failed to remove api socket on shutdown"); } @@ -39,15 +44,34 @@ impl Drop for ServerHandle { } } +impl ServerHandle { + pub(crate) fn remove_socket_file_if_owned(&self) -> std::io::Result<()> { + remove_socket_file_if_owned(&self.path, self.identity) + } +} + pub fn start_server( api_tx: ApiRequestSender, event_hub: EventHub, +) -> std::io::Result { + start_server_with_capabilities( + api_tx, + event_hub, + Some(ServerCapabilities { live_handoff: true }), + ) +} + +pub fn start_server_with_capabilities( + api_tx: ApiRequestSender, + event_hub: EventHub, + capabilities: Option, ) -> std::io::Result { let path = socket_path(); prepare_socket_path(&path)?; let listener = UnixListener::bind(&path)?; restrict_socket_permissions(&path)?; + let identity = socket_file_identity(&path)?; info!(path = %path.display(), "api server listening"); let running = Arc::new(AtomicBool::new(true)); @@ -58,11 +82,16 @@ pub fn start_server( Ok(stream) => { let api_tx = api_tx.clone(); let event_hub = event_hub.clone(); + let capabilities = capabilities.clone(); let connection_running = Arc::clone(&listener_running); std::thread::spawn(move || { - if let Err(err) = - handle_connection(stream, &api_tx, &event_hub, &connection_running) - { + if let Err(err) = handle_connection( + stream, + &api_tx, + &event_hub, + &connection_running, + capabilities, + ) { warn!(err = %err, "api connection failed"); } }); @@ -79,6 +108,7 @@ pub fn start_server( Ok(ServerHandle { _thread: thread, path, + identity, running, }) } @@ -101,20 +131,15 @@ fn handle_connection( api_tx: &ApiRequestSender, event_hub: &EventHub, running: &Arc, + capabilities: Option, ) -> std::io::Result<()> { - stream.set_read_timeout(Some(INITIAL_REQUEST_TIMEOUT))?; - stream.set_write_timeout(Some(STREAM_WRITE_TIMEOUT))?; - - let mut line = String::new(); - { - let mut reader = BufReader::new(&stream); - let read = reader.read_line(&mut line)?; - if read == 0 { - return Ok(()); - } + if let Err(err) = stream.set_write_timeout(Some(STREAM_WRITE_TIMEOUT)) { + debug!(err = %err, "api connection write timeout unavailable"); } - stream.set_read_timeout(None)?; + let Some(line) = read_initial_request_line(&mut stream)? else { + return Ok(()); + }; let line = line.trim(); if line.is_empty() { @@ -199,6 +224,7 @@ fn handle_connection( method: method_body, }, api_tx, + capabilities, ); let result = write_text_line_allow_disconnect(&mut stream, &response); match &result { @@ -217,13 +243,18 @@ fn handle_connection( } } -fn handle_request(request: Request, api_tx: &ApiRequestSender) -> String { +fn handle_request( + request: Request, + api_tx: &ApiRequestSender, + capabilities: Option, +) -> String { match request.method { Method::Ping(_) => serde_json::to_string(&SuccessResponse { id: request.id, result: ResponseResult::Pong { version: env!("CARGO_PKG_VERSION").into(), protocol: crate::protocol::PROTOCOL_VERSION, + capabilities, }, }) .unwrap_or_else(|_| { @@ -238,6 +269,7 @@ fn api_method_name(method: &Method) -> &'static str { match method { Method::Ping(_) => "ping", Method::ServerStop(_) => "server.stop", + Method::ServerLiveHandoff(_) => "server.live_handoff", Method::ServerReloadConfig(_) => "server.reload_config", Method::WorkspaceCreate(_) => "workspace.create", Method::WorkspaceList(_) => "workspace.list", @@ -298,6 +330,52 @@ fn api_response_outcome(response: &str) -> &'static str { } } +fn read_initial_request_line(stream: &mut UnixStream) -> std::io::Result> { + stream.set_nonblocking(true)?; + let deadline = Instant::now() + INITIAL_REQUEST_TIMEOUT; + let mut bytes = Vec::new(); + let mut byte = [0u8; 1]; + + loop { + match stream.read(&mut byte) { + Ok(0) => { + stream.set_nonblocking(false)?; + return Ok(None); + } + Ok(_) => { + bytes.push(byte[0]); + if byte[0] == b'\n' { + stream.set_nonblocking(false)?; + return String::from_utf8(bytes) + .map(Some) + .map_err(|err| io::Error::new(io::ErrorKind::InvalidData, err)); + } + if bytes.len() > MAX_INITIAL_REQUEST_BYTES { + stream.set_nonblocking(false)?; + return Err(io::Error::new( + io::ErrorKind::InvalidData, + "api request line is too large", + )); + } + } + Err(err) if err.kind() == io::ErrorKind::WouldBlock => { + if Instant::now() >= deadline { + stream.set_nonblocking(false)?; + return Err(io::Error::new( + io::ErrorKind::TimedOut, + "timed out reading api request", + )); + } + std::thread::sleep(CONNECTION_POLL_INTERVAL); + } + Err(err) => { + stream.set_nonblocking(false)?; + return Err(err); + } + } + } +} + fn stream_subscriptions( mut stream: UnixStream, request_id: String, @@ -496,6 +574,7 @@ fn error_response_json(id: String, code: &str, message: String) -> String { #[cfg(test)] mod tests { use super::*; + use std::io::{BufRead, BufReader}; use std::os::unix::fs::PermissionsExt; use std::sync::{Mutex, OnceLock}; use tokio::sync::mpsc; @@ -610,6 +689,7 @@ mod tests { method: Method::Ping(crate::api::schema::PingParams::default()), }, &tx, + Some(ServerCapabilities { live_handoff: true }), ); let parsed: SuccessResponse = serde_json::from_str(&response).unwrap(); @@ -626,7 +706,7 @@ mod tests { }; let request_for_thread = request.clone(); - let thread = std::thread::spawn(move || handle_request(request_for_thread, &tx)); + let thread = std::thread::spawn(move || handle_request(request_for_thread, &tx, None)); let msg = rx.blocking_recv().unwrap(); assert_eq!(msg.request.id, "req_2"); @@ -692,7 +772,7 @@ mod tests { let event_hub = EventHub::default(); let (done_tx, done_rx) = std::sync::mpsc::channel(); let server_thread = std::thread::spawn(move || { - let result = handle_connection(server, &api_tx, &event_hub, &server_running); + let result = handle_connection(server, &api_tx, &event_hub, &server_running, None); done_tx.send(result).unwrap(); }); @@ -724,7 +804,7 @@ mod tests { let event_hub = EventHub::default(); let (done_tx, done_rx) = std::sync::mpsc::channel(); let server_thread = std::thread::spawn(move || { - let result = handle_connection(server, &api_tx, &event_hub, &server_running); + let result = handle_connection(server, &api_tx, &event_hub, &server_running, None); done_tx.send(result).unwrap(); }); @@ -756,7 +836,7 @@ mod tests { let event_hub = EventHub::default(); let (done_tx, done_rx) = std::sync::mpsc::channel(); let server_thread = std::thread::spawn(move || { - let result = handle_connection(server, &api_tx, &event_hub, &server_running); + let result = handle_connection(server, &api_tx, &event_hub, &server_running, None); done_tx.send(result).unwrap(); }); diff --git a/src/api/status.rs b/src/api/status.rs index f58e69c0..209b93df 100644 --- a/src/api/status.rs +++ b/src/api/status.rs @@ -8,6 +8,7 @@ use crate::api::schema::{Method, Request, ResponseResult}; pub struct RuntimeStatus { pub version: Option, pub protocol: Option, + pub capabilities: Option, } pub fn read_runtime_status_at( @@ -43,9 +44,14 @@ pub fn read_runtime_status_at( Err(err) => return Err(io::Error::other(err)), }; match response.result { - ResponseResult::Pong { version, protocol } => Ok(Some(RuntimeStatus { + ResponseResult::Pong { + version, + protocol, + capabilities, + } => Ok(Some(RuntimeStatus { version: Some(version), protocol: Some(protocol), + capabilities, })), result => Err(io::Error::other(format!( "server status request returned unexpected result: {result:?}" diff --git a/src/app/api.rs b/src/app/api.rs index 6f88f6c3..c51a1bb0 100644 --- a/src/app/api.rs +++ b/src/app/api.rs @@ -330,7 +330,9 @@ impl App { &mut self, request: crate::api::schema::Request, ) -> String { - use crate::api::schema::{Method, ResponseResult, SuccessResponse}; + use crate::api::schema::{ + ErrorBody, ErrorResponse, Method, ResponseResult, SuccessResponse, + }; let response = match request.method { Method::ServerStop(_) => { @@ -340,6 +342,16 @@ impl App { result: ResponseResult::Ok {}, } } + Method::ServerLiveHandoff(_) => { + let response = ErrorResponse { + id: request.id, + error: ErrorBody { + code: "unsupported_in_app_mode".into(), + message: "live handoff is only supported by the headless server".into(), + }, + }; + return serde_json::to_string(&response).unwrap_or_else(|_| "{}".to_string()); + } Method::ServerReloadConfig(_) => { let report = self.reload_config(); SuccessResponse { diff --git a/src/app/mod.rs b/src/app/mod.rs index 1679364b..bcbd40cf 100644 --- a/src/app/mod.rs +++ b/src/app/mod.rs @@ -555,6 +555,70 @@ impl App { } } + #[cfg(unix)] + pub fn new_from_handoff( + config: &Config, + config_diagnostic: Option, + api_rx: tokio::sync::mpsc::UnboundedReceiver, + event_hub: crate::api::EventHub, + snapshot: &crate::persist::SessionSnapshot, + imports: &mut std::collections::HashMap, + ) -> io::Result { + let mut app = Self::new(config, true, config_diagnostic, api_rx, event_hub); + let (workspaces, terminals, runtimes) = crate::persist::restore_handoff( + snapshot, + config.advanced.scrollback_limit_bytes, + &config.terminal.default_shell, + imports, + app.event_tx.clone(), + app.render_notify.clone(), + app.render_dirty.clone(), + )?; + + app.no_session = false; + app.state.detach_exits = false; + app.state.workspaces = workspaces; + app.state.terminals = terminals; + app.terminal_runtimes = runtimes.into(); + app.state.active = snapshot + .active + .filter(|&idx| idx < app.state.workspaces.len()); + app.state.selected = snapshot + .selected + .min(app.state.workspaces.len().saturating_sub(1)); + app.state.agent_panel_scope = snapshot.agent_panel_scope; + if let Some(width) = snapshot.sidebar_width { + app.state.sidebar_width = width; + app.state.sidebar_width_source = state::SidebarWidthSource::Persisted; + } + if let Some(split) = snapshot.sidebar_section_split { + app.state.sidebar_section_split = split; + } + app.state.collapsed_space_keys = snapshot.collapsed_space_keys.clone(); + app.state.mode = if app.state.active.is_some() { + state::Mode::Terminal + } else { + state::Mode::Navigate + }; + app.last_focus = app.state.active.and_then(|idx| { + app.state + .workspaces + .get(idx) + .and_then(|ws| ws.focused_pane_id().map(|pane_id| (idx, pane_id))) + }); + Ok(app) + } + + #[cfg(unix)] + pub fn unpause_handoff_readers(&self) { + self.terminal_runtimes.set_handoff_readers_paused(false); + } + + #[cfg(unix)] + pub fn assume_handoff_ownership(&mut self) { + self.terminal_runtimes.assume_handoff_ownership(); + } + fn request_full_redraw(&mut self) { self.full_redraw_pending = true; } diff --git a/src/cli/server.rs b/src/cli/server.rs index 311cc31a..0102b2d4 100644 --- a/src/cli/server.rs +++ b/src/cli/server.rs @@ -1,4 +1,4 @@ -use crate::api::schema::{EmptyParams, Method, Request}; +use crate::api::schema::{EmptyParams, Method, Request, ServerLiveHandoffParams}; pub(super) fn run_server_command(args: &[String]) -> std::io::Result> { let Some(subcommand) = args.first().map(|arg| arg.as_str()) else { @@ -7,6 +7,8 @@ pub(super) fn run_server_command(args: &[String]) -> std::io::Result match subcommand { "stop" => server_stop(&args[1..]).map(Some), + "live-handoff" => server_live_handoff(&args[1..]).map(Some), + "--handoff-import" => Ok(None), "reload-config" => server_reload_config(&args[1..]).map(Some), "help" | "--help" | "-h" => { print_server_help(); @@ -40,9 +42,92 @@ fn server_reload_config(args: &[String]) -> std::io::Result { })?) } +fn server_live_handoff(args: &[String]) -> std::io::Result { + let Some(params) = parse_live_handoff_params(args) else { + eprintln!( + "usage: herdr server live-handoff [--import-exe ] [--expected-protocol ] [--expected-version ]" + ); + return Ok(2); + }; + + let response = super::send_request(&Request { + id: "cli:server:live-handoff".into(), + method: Method::ServerLiveHandoff(params), + })?; + if response.get("error").is_some() { + let rendered = serde_json::to_string(&response).unwrap_or_else(|err| { + format!( + "{{\"error\":{{\"code\":\"render_failed\",\"message\":\"failed to render error response: {err}\"}}}}" + ) + }); + eprintln!("{rendered}"); + return Ok(1); + } + + eprintln!( + "live handoff complete; server log: {}", + crate::session::data_dir() + .join("herdr-server.log") + .display() + ); + Ok(0) +} + +fn parse_live_handoff_params(args: &[String]) -> Option { + let mut params = ServerLiveHandoffParams::default(); + let mut idx = 0; + while idx < args.len() { + let arg = &args[idx]; + let (flag, value) = if let Some((flag, value)) = arg.split_once('=') { + (flag, Some(value.to_string())) + } else { + let value = args.get(idx + 1).cloned(); + idx += 1; + (arg.as_str(), value) + }; + let value = value?; + match flag { + "--import-exe" => params.import_exe = Some(value), + "--expected-protocol" => { + params.expected_protocol = Some(value.parse().ok()?); + } + "--expected-version" => params.expected_version = Some(value), + _ => return None, + } + idx += 1; + } + Some(params) +} + fn print_server_help() { eprintln!("herdr server commands:"); eprintln!(" herdr server run as headless server"); eprintln!(" herdr server stop stop the running server via the API socket"); + eprintln!(" herdr server live-handoff hand off live panes to a new local server"); eprintln!(" herdr server reload-config reload config.toml in the running server"); } + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn live_handoff_params_parse_remote_update_fields() { + let args = vec![ + "--import-exe".to_string(), + "/home/me/.local/bin/herdr".to_string(), + "--expected-protocol=9".to_string(), + "--expected-version".to_string(), + "0.6.2".to_string(), + ]; + + let params = parse_live_handoff_params(&args).expect("params"); + + assert_eq!( + params.import_exe.as_deref(), + Some("/home/me/.local/bin/herdr") + ); + assert_eq!(params.expected_protocol, Some(9)); + assert_eq!(params.expected_version.as_deref(), Some("0.6.2")); + } +} diff --git a/src/cli/status.rs b/src/cli/status.rs index 9b8a1648..0062367a 100644 --- a/src/cli/status.rs +++ b/src/cli/status.rs @@ -74,6 +74,7 @@ enum ServerRuntimeStatus { Running { version: Option, protocol: Option, + capabilities: Option, }, NotRunning, } @@ -127,7 +128,9 @@ fn print_client_status(json: bool) -> std::io::Result<()> { fn print_server_status_body(server: &ServerRuntimeStatus, indent: &str) { match server { - ServerRuntimeStatus::Running { version, protocol } => { + ServerRuntimeStatus::Running { + version, protocol, .. + } => { println!("{indent}status: running"); println!("{indent}version: {}", option_label(version.as_deref())); println!("{indent}protocol: {}", protocol_label(*protocol)); @@ -146,6 +149,7 @@ fn read_server_runtime_status() -> std::io::Result { Ok(status) => Ok(ServerRuntimeStatus::Running { version: status.version, protocol: status.protocol, + capabilities: status.capabilities, }), Err(ApiClientError::Io(err)) if server_not_running_error(&err) => { Ok(ServerRuntimeStatus::NotRunning) @@ -218,12 +222,18 @@ struct ServerStatusJson { running: bool, version: Option, protocol: Option, + capabilities: Option, compatible: Option, socket: String, session: Option, restart_needed: Option, } +#[derive(Serialize)] +struct ServerCapabilitiesJson { + live_handoff: bool, +} + #[derive(Serialize)] struct UpdateStatusJson { restart_needed: Option, @@ -240,11 +250,20 @@ fn client_status_json() -> ClientStatusJson { fn server_status_json(server: &ServerRuntimeStatus) -> ServerStatusJson { match server { - ServerRuntimeStatus::Running { version, protocol } => ServerStatusJson { + ServerRuntimeStatus::Running { + version, + protocol, + capabilities, + } => ServerStatusJson { status: "running", running: true, version: version.clone(), protocol: *protocol, + capabilities: capabilities + .as_ref() + .map(|capabilities| ServerCapabilitiesJson { + live_handoff: capabilities.live_handoff, + }), compatible: protocol.map(|value| value == crate::protocol::PROTOCOL_VERSION), socket: api::socket_path().display().to_string(), session: crate::session::active_name(), @@ -255,6 +274,7 @@ fn server_status_json(server: &ServerRuntimeStatus) -> ServerStatusJson { running: false, version: None, protocol: None, + capabilities: None, compatible: None, socket: api::socket_path().display().to_string(), session: crate::session::active_name(), diff --git a/src/ipc.rs b/src/ipc.rs index aa0148a1..5f973ccf 100644 --- a/src/ipc.rs +++ b/src/ipc.rs @@ -1,9 +1,16 @@ use std::fs; use std::io; +use std::os::unix::fs::MetadataExt; use std::os::unix::fs::PermissionsExt; use std::os::unix::net::UnixStream; use std::path::Path; +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +pub(crate) struct SocketFileIdentity { + dev: u64, + ino: u64, +} + pub(crate) fn prepare_socket_path( path: &Path, busy_message: impl FnOnce(&Path) -> String, @@ -39,6 +46,35 @@ pub(crate) fn prepare_socket_path( Ok(()) } +pub(crate) fn socket_file_identity(path: &Path) -> io::Result { + let metadata = fs::metadata(path)?; + Ok(SocketFileIdentity { + dev: metadata.dev(), + ino: metadata.ino(), + }) +} + +pub(crate) fn remove_socket_file_if_owned( + path: &Path, + identity: SocketFileIdentity, +) -> io::Result<()> { + let current = match socket_file_identity(path) { + Ok(current) => current, + Err(err) if err.kind() == io::ErrorKind::NotFound => return Ok(()), + Err(err) => return Err(err), + }; + + if current != identity { + return Ok(()); + } + + match fs::remove_file(path) { + Ok(()) => Ok(()), + Err(err) if err.kind() == io::ErrorKind::NotFound => Ok(()), + Err(err) => Err(err), + } +} + pub(crate) fn restrict_socket_permissions(path: &Path, mode: u32) -> io::Result<()> { let mut permissions = fs::metadata(path)?.permissions(); permissions.set_mode(mode); diff --git a/src/main.rs b/src/main.rs index 43737e01..725aed45 100644 --- a/src/main.rs +++ b/src/main.rs @@ -524,7 +524,7 @@ fn main() -> io::Result<()> { let (api_tx, api_rx) = tokio::sync::mpsc::unbounded_channel(); let event_hub = api::EventHub::default(); - let _api_server = match api::start_server(api_tx, event_hub.clone()) { + let _api_server = match api::start_server_with_capabilities(api_tx, event_hub.clone(), None) { Ok(server) => server, Err(err) if err.kind() == io::ErrorKind::AddrInUse => { eprintln!("error: herdr is already running"); diff --git a/src/pane.rs b/src/pane.rs index 9b1bf125..49d7a0e7 100644 --- a/src/pane.rs +++ b/src/pane.rs @@ -1,12 +1,12 @@ use std::cell::Cell; -use std::io::{BufWriter, Write}; +use std::io::{Read, Write}; use std::sync::{ atomic::{AtomicBool, AtomicU16, AtomicU32, Ordering}, Arc, Mutex, }; use bytes::Bytes; -use portable_pty::{native_pty_system, CommandBuilder, PtySize}; +use portable_pty::{native_pty_system, CommandBuilder, MasterPty, PtySize}; use ratatui::{layout::Rect, Frame}; use tokio::sync::{mpsc, watch, Notify}; use tracing::{debug, error, info, warn}; @@ -168,6 +168,104 @@ fn should_publish_detection_update( || (next.visible_idle && previous.visible_idle) } +fn spawn_basic_detection_task( + pane_id: PaneId, + child_pid: Arc, + terminal: Arc, + state_events: mpsc::Sender, +) -> ( + tokio::task::AbortHandle, + Arc, + Arc>>, +) { + let detect_reset_notify = Arc::new(Notify::new()); + let detect_reset = detect_reset_notify.clone(); + let pending_release = Arc::new(Mutex::new(None)); + let pending_release_for_task = pending_release.clone(); + + let handle = tokio::spawn(async move { + let mut agent_presence = AgentDetectionPresence::from_agent(None); + let mut state = AgentState::Unknown; + let mut last_visible_blocker = false; + let mut last_visible_idle = false; + let mut last_visible_working = false; + + loop { + tokio::select! { + _ = tokio::time::sleep(std::time::Duration::from_millis(300)) => {} + _ = detect_reset.notified() => { + agent_presence = AgentDetectionPresence::from_agent(None); + state = AgentState::Unknown; + last_visible_blocker = false; + last_visible_idle = false; + last_visible_working = false; + } + } + + let now = std::time::Instant::now(); + let pid = child_pid.load(Ordering::Acquire); + let mut agent_changed = false; + let mut agent = agent_presence.current_agent(); + + if pid > 0 { + let new_agent = crate::detect::foreground_job(pid).and_then(|job| { + crate::detect::identify_agent_in_job(&job).map(|(agent, _)| agent) + }); + let previous_agent = agent_presence.current_agent(); + if agent_presence.observe_process_probe(new_agent) { + agent = agent_presence.current_agent(); + agent_changed = previous_agent != agent; + } + } + + let content = terminal.detection_text(); + let detection = crate::detect::detect_agent(agent, &content); + let new_state = detection.state; + let visible_blocker = detection.visible_blocker && new_state == AgentState::Blocked; + let visible_idle = detection.visible_idle && new_state == AgentState::Idle; + let visible_working = detection.visible_working && new_state == AgentState::Working; + + if should_publish_detection_update( + DetectionPublishState { + state, + visible_blocker: last_visible_blocker, + visible_idle: last_visible_idle, + visible_working: last_visible_working, + }, + DetectionPublishState { + state: new_state, + visible_blocker, + visible_idle, + visible_working, + }, + agent_changed, + false, + ) { + state = new_state; + last_visible_blocker = visible_blocker; + last_visible_idle = visible_idle; + last_visible_working = visible_working; + publish_state_changed_event( + state_events.clone(), + pane_id, + agent, + new_state, + visible_blocker, + visible_idle, + visible_working, + false, + now, + ) + .await; + } + + let _ = active_pending_release(&pending_release_for_task, now); + } + }); + + (handle.abort_handle(), detect_reset_notify, pending_release) +} + impl AgentDetectionPresence { fn from_agent(current_agent: Option) -> Self { Self { @@ -230,13 +328,34 @@ pub struct PaneRuntime { resize_tx: watch::Sender<(u16, u16, u32, u32)>, current_size: Cell<(u16, u16, u32, u32)>, child_pid: Arc, + pty_master: Option>, + raw_master_fd: Option, + force_resize_fd: Option, + io_stop: Arc, + reader_paused: Arc, + reader_pause_ack: Arc, + reader_stopped_rx: Option>, kitty_keyboard_flags: Arc, detect_reset_notify: Arc, pending_release: Arc>>, + preserve_processes_on_drop: bool, // Task handles for deterministic shutdown detect_handle: tokio::task::AbortHandle, } +#[cfg(unix)] +#[derive(Debug)] +pub struct PaneRuntimeImport { + pub pane_id: PaneId, + pub master_fd: std::os::fd::RawFd, + pub child_pid: u32, + pub rows: u16, + pub cols: u16, + pub cell_width_px: u32, + pub cell_height_px: u32, + pub initial_history_ansi: Option, +} + #[derive(Debug, Clone, Copy, PartialEq, Eq)] pub enum WheelRouting { HostScroll, @@ -250,7 +369,16 @@ impl Drop for PaneRuntime { // Reader/writer/resize tasks shut down naturally via channel close // and PTY EOF when the rest of PaneRuntime is dropped. self.detect_handle.abort(); - shutdown_pane_processes(self.pane_id, self.child_pid.load(Ordering::Acquire)); + self.io_stop.store(true, Ordering::Release); + if !self.preserve_processes_on_drop { + shutdown_pane_processes(self.pane_id, self.child_pid.load(Ordering::Acquire)); + } + if let Some(fd) = self.raw_master_fd.take() { + let _ = unsafe { libc::close(fd) }; + } + if let Some(fd) = self.force_resize_fd.take() { + let _ = unsafe { libc::close(fd) }; + } } } @@ -316,6 +444,22 @@ fn shutdown_pane_processes(pane_id: PaneId, child_pid: u32) { ); } +#[cfg(unix)] +fn truncate_handoff_history(history: String, max_bytes: usize) -> String { + if history.len() <= max_bytes { + return history; + } + let mut start = history.len().saturating_sub(max_bytes); + while !history.is_char_boundary(start) { + start += 1; + } + let Some(newline_offset) = history[start..].find('\n') else { + return String::new(); + }; + start += newline_offset + 1; + history[start..].to_owned() +} + fn pane_shell(configured_shell: &str) -> String { pane_shell_from(configured_shell, std::env::var("SHELL").ok()) } @@ -368,12 +512,243 @@ fn restore_command_builder(agent: &str, fallback_shell: &str, argv: &[String]) - cmd } +#[cfg(unix)] +fn duplicate_fd(fd: std::os::fd::RawFd) -> std::io::Result { + let duplicated = unsafe { libc::dup(fd) }; + if duplicated < 0 { + return Err(std::io::Error::last_os_error()); + } + Ok(duplicated) +} + +#[cfg(unix)] +fn set_cloexec(fd: std::os::fd::RawFd) -> std::io::Result<()> { + let flags = unsafe { libc::fcntl(fd, libc::F_GETFD) }; + if flags < 0 { + return Err(std::io::Error::last_os_error()); + } + if unsafe { libc::fcntl(fd, libc::F_SETFD, flags | libc::FD_CLOEXEC) } < 0 { + return Err(std::io::Error::last_os_error()); + } + Ok(()) +} + +#[cfg(unix)] +fn set_nonblocking(fd: std::os::fd::RawFd) -> std::io::Result<()> { + let flags = unsafe { libc::fcntl(fd, libc::F_GETFL) }; + if flags < 0 { + return Err(std::io::Error::last_os_error()); + } + if unsafe { libc::fcntl(fd, libc::F_SETFL, flags | libc::O_NONBLOCK) } < 0 { + return Err(std::io::Error::last_os_error()); + } + Ok(()) +} + +#[cfg(unix)] +fn duplicate_cloexec_fd(fd: std::os::fd::RawFd) -> std::io::Result { + let duplicated = duplicate_fd(fd)?; + if let Err(err) = set_cloexec(duplicated) { + let _ = unsafe { libc::close(duplicated) }; + return Err(err); + } + Ok(duplicated) +} + +#[cfg(unix)] +fn file_from_duplicated_fd(fd: std::os::fd::RawFd) -> std::io::Result { + use std::os::fd::FromRawFd; + + let duplicated = duplicate_cloexec_fd(fd)?; + Ok(unsafe { std::fs::File::from_raw_fd(duplicated) }) +} + +#[cfg(unix)] +fn poll_read_ready(fd: std::os::fd::RawFd, timeout_ms: i32) -> std::io::Result { + let mut poll_fd = libc::pollfd { + fd, + events: libc::POLLIN, + revents: 0, + }; + loop { + let result = unsafe { libc::poll(&mut poll_fd, 1, timeout_ms) }; + if result < 0 { + let err = std::io::Error::last_os_error(); + if err.kind() == std::io::ErrorKind::Interrupted { + continue; + } + return Err(err); + } + return Ok(result > 0 && (poll_fd.revents & (libc::POLLIN | libc::POLLHUP)) != 0); + } +} + +#[cfg(unix)] +fn poll_write_ready(fd: std::os::fd::RawFd, timeout_ms: i32) -> std::io::Result { + let mut poll_fd = libc::pollfd { + fd, + events: libc::POLLOUT, + revents: 0, + }; + loop { + let result = unsafe { libc::poll(&mut poll_fd, 1, timeout_ms) }; + if result < 0 { + let err = std::io::Error::last_os_error(); + if err.kind() == std::io::ErrorKind::Interrupted { + continue; + } + return Err(err); + } + return Ok(result > 0 && (poll_fd.revents & (libc::POLLOUT | libc::POLLHUP)) != 0); + } +} + +#[cfg(unix)] +fn write_all_nonblocking( + writer: &mut std::fs::File, + fd: std::os::fd::RawFd, + mut bytes: &[u8], + io_stop: &AtomicBool, +) -> std::io::Result<()> { + while !bytes.is_empty() { + if io_stop.load(Ordering::Acquire) { + return Ok(()); + } + match writer.write(bytes) { + Ok(0) => { + return Err(std::io::Error::new( + std::io::ErrorKind::WriteZero, + "pty write returned zero bytes", + )); + } + Ok(written) => bytes = &bytes[written..], + Err(err) if err.kind() == std::io::ErrorKind::WouldBlock => { + let _ = poll_write_ready(fd, 50)?; + } + Err(err) if err.kind() == std::io::ErrorKind::Interrupted => {} + Err(err) => return Err(err), + } + } + Ok(()) +} + +#[cfg(unix)] +fn resize_pty_fd( + fd: std::os::fd::RawFd, + rows: u16, + cols: u16, + cell_width_px: u32, + cell_height_px: u32, +) -> std::io::Result<()> { + let size = libc::winsize { + ws_row: rows, + ws_col: cols, + ws_xpixel: (cols as u32) + .saturating_mul(cell_width_px) + .min(u16::MAX as u32) as u16, + ws_ypixel: (rows as u32) + .saturating_mul(cell_height_px) + .min(u16::MAX as u32) as u16, + }; + if unsafe { libc::ioctl(fd, libc::TIOCSWINSZ, &size) } < 0 { + return Err(std::io::Error::last_os_error()); + } + Ok(()) +} + impl PaneRuntime { + #[cfg(unix)] + fn master_fd(&self) -> Option { + self.raw_master_fd.or_else(|| { + self.pty_master + .as_ref() + .and_then(|master| master.as_raw_fd()) + }) + } + pub fn shutdown(self) { self.detect_handle.abort(); shutdown_pane_processes(self.pane_id, self.child_pid.load(Ordering::Acquire)); } + #[cfg(unix)] + pub fn duplicate_handoff_fd(&self) -> std::io::Result { + let master_fd = self + .master_fd() + .ok_or_else(|| std::io::Error::other("runtime has no PTY master fd"))?; + duplicate_cloexec_fd(master_fd) + } + + #[cfg(unix)] + pub fn preserve_for_handoff(mut self) { + self.io_stop.store(true, Ordering::Release); + if let Some(reader_stopped_rx) = self.reader_stopped_rx.take() { + let _ = reader_stopped_rx.recv_timeout(std::time::Duration::from_millis(500)); + } + self.detect_handle.abort(); + self.preserve_processes_on_drop = true; + } + + #[cfg(unix)] + pub fn assume_handoff_ownership(&mut self) { + self.preserve_processes_on_drop = false; + } + + #[cfg(unix)] + pub fn set_handoff_reader_paused(&self, paused: bool) { + self.reader_paused.store(paused, Ordering::Release); + if !paused { + self.reader_pause_ack.store(false, Ordering::Release); + } + } + + #[cfg(unix)] + pub fn pause_handoff_reader(&self, timeout: std::time::Duration) -> std::io::Result<()> { + self.reader_pause_ack.store(false, Ordering::Release); + self.reader_paused.store(true, Ordering::Release); + let deadline = std::time::Instant::now() + timeout; + while std::time::Instant::now() < deadline { + if self.reader_pause_ack.load(Ordering::Acquire) || self.io_stop.load(Ordering::Acquire) + { + return Ok(()); + } + std::thread::sleep(std::time::Duration::from_millis(5)); + } + Err(std::io::Error::new( + std::io::ErrorKind::TimedOut, + "timed out waiting for pane reader to pause for handoff", + )) + } + + #[cfg(unix)] + pub fn handoff_pane(&self, pane_id: u32) -> crate::server::handoff::HandoffPane { + let child_pid = self.child_pid.load(Ordering::Acquire); + let (rows, cols, cell_width_px, cell_height_px) = self.current_size.get(); + crate::server::handoff::HandoffPane { + pane_id, + child_pid, + rows, + cols, + cell_width_px, + cell_height_px, + initial_history_ansi: None, + } + } + + #[cfg(unix)] + pub fn handoff_history_ansi(&self) -> Option { + if self + .terminal + .input_state() + .is_some_and(|input_state| input_state.alternate_screen) + { + return None; + } + self.snapshot_history().map(|history| { + truncate_handoff_history(history, crate::server::handoff::MAX_REPLAY_BYTES_PER_PANE) + }) + } + pub fn apply_host_terminal_theme(&self, theme: crate::terminal_theme::TerminalTheme) { self.terminal.apply_host_terminal_theme(theme); } @@ -565,6 +940,213 @@ impl PaneRuntime { ) } + #[cfg(unix)] + pub fn from_handoff_fd( + import: PaneRuntimeImport, + scrollback_limit_bytes: usize, + host_terminal_theme: crate::terminal_theme::TerminalTheme, + events: mpsc::Sender, + render_notify: Arc, + render_dirty: Arc, + ) -> std::io::Result { + let PaneRuntimeImport { + pane_id, + master_fd, + child_pid, + rows, + cols, + cell_width_px, + cell_height_px, + initial_history_ansi, + } = import; + use std::os::fd::{AsRawFd, FromRawFd, IntoRawFd}; + + let master_fd = unsafe { std::os::fd::OwnedFd::from_raw_fd(master_fd) }; + set_cloexec(master_fd.as_raw_fd())?; + set_nonblocking(master_fd.as_raw_fd())?; + let reader = file_from_duplicated_fd(master_fd.as_raw_fd())?; + let writer = file_from_duplicated_fd(master_fd.as_raw_fd())?; + let force_resize_fd = duplicate_cloexec_fd(master_fd.as_raw_fd())?; + let resize_fd = unsafe { + std::os::fd::OwnedFd::from_raw_fd(duplicate_cloexec_fd(master_fd.as_raw_fd())?) + }; + let io_stop = Arc::new(AtomicBool::new(false)); + let reader_paused = Arc::new(AtomicBool::new(true)); + let reader_pause_ack = Arc::new(AtomicBool::new(false)); + + let (input_tx, mut input_rx) = mpsc::channel::(32); + let mut terminal = crate::ghostty::Terminal::new(cols, rows, scrollback_limit_bytes) + .map_err(|e| std::io::Error::other(e.to_string()))?; + if crate::kitty_graphics::is_enabled() { + terminal + .enable_kitty_graphics() + .map_err(|e| std::io::Error::other(e.to_string()))?; + } + let pane_terminal = GhosttyPaneTerminal::new(terminal, input_tx.clone())?; + pane_terminal.apply_host_terminal_theme(host_terminal_theme); + if let Some(ansi) = initial_history_ansi.as_deref() { + pane_terminal.seed_history_ansi(ansi); + } + let terminal = Arc::new(PaneTerminal::new(pane_terminal)); + let child_pid = Arc::new(AtomicU32::new(child_pid)); + let kitty_keyboard_flags = Arc::new(AtomicU16::new(0)); + let (reader_stopped_tx, reader_stopped_rx) = std::sync::mpsc::channel(); + + { + use std::os::fd::AsRawFd; + + let mut reader = reader; + let reader_fd = reader.as_raw_fd(); + let terminal = terminal.clone(); + let response_writer = input_tx.clone(); + let render_notify = render_notify.clone(); + let render_dirty = render_dirty.clone(); + let child_pid = child_pid.clone(); + let events = events.clone(); + let io_stop = io_stop.clone(); + let reader_paused = reader_paused.clone(); + let reader_pause_ack = reader_pause_ack.clone(); + let rt = tokio::runtime::Handle::current(); + tokio::task::spawn_blocking(move || { + let mut buf = [0u8; 8192]; + loop { + if io_stop.load(Ordering::Acquire) { + break; + } + if reader_paused.load(Ordering::Acquire) { + reader_pause_ack.store(true, Ordering::Release); + std::thread::sleep(std::time::Duration::from_millis(5)); + continue; + } + reader_pause_ack.store(false, Ordering::Release); + match poll_read_ready(reader_fd, 50) { + Ok(true) => {} + Ok(false) => continue, + Err(e) => { + debug!(pane = pane_id.raw(), err = %e, "handoff pty reader poll failed"); + break; + } + } + match reader.read(&mut buf) { + Ok(0) => break, + Err(e) if e.kind() == std::io::ErrorKind::WouldBlock => continue, + Err(e) => { + debug!(pane = pane_id.raw(), err = %e, "handoff pty reader closed"); + break; + } + Ok(n) => { + let shell_pid = child_pid.load(Ordering::Acquire); + let result = terminal.process_pty_bytes( + pane_id, + shell_pid, + &buf[..n], + &response_writer, + ); + if result.request_render && !render_dirty.swap(true, Ordering::AcqRel) { + render_notify.notify_one(); + } + if let Some(delay) = result.render_delay { + let render_notify = render_notify.clone(); + let render_dirty = render_dirty.clone(); + rt.spawn(async move { + tokio::time::sleep(delay).await; + if !render_dirty.swap(true, Ordering::AcqRel) { + render_notify.notify_one(); + } + }); + } + for content in result.clipboard_writes { + let _ = + rt.block_on(events.send(AppEvent::ClipboardWrite { content })); + } + } + } + } + let _ = reader_stopped_tx.send(()); + let _ = rt.block_on(events.send(AppEvent::PaneDied { pane_id })); + debug!(pane = pane_id.raw(), "handoff reader task exiting"); + }); + } + + { + use std::os::fd::AsRawFd; + + let mut writer = writer; + let writer_fd = writer.as_raw_fd(); + let io_stop = io_stop.clone(); + tokio::task::spawn_blocking(move || { + let rt = tokio::runtime::Handle::current(); + while let Some(bytes) = rt.block_on(input_rx.recv()) { + if io_stop.load(Ordering::Acquire) { + break; + } + if let Err(e) = write_all_nonblocking(&mut writer, writer_fd, &bytes, &io_stop) + { + warn!(pane = pane_id.raw(), err = %e, "handoff pty write failed"); + break; + } + if let Err(e) = writer.flush() { + warn!(pane = pane_id.raw(), err = %e, "handoff pty flush failed"); + break; + } + } + debug!(pane = pane_id.raw(), "handoff writer task exiting"); + }); + } + + let (resize_tx, mut resize_rx) = + watch::channel::<(u16, u16, u32, u32)>((rows, cols, cell_width_px, cell_height_px)); + { + let io_stop = io_stop.clone(); + let resize_fd = resize_fd.into_raw_fd(); + tokio::task::spawn_blocking(move || { + let rt = tokio::runtime::Handle::current(); + let mut last_size = (rows, cols, cell_width_px, cell_height_px); + while rt.block_on(resize_rx.changed()).is_ok() { + if io_stop.load(Ordering::Acquire) { + break; + } + let (rows, cols, cell_width_px, cell_height_px) = + *resize_rx.borrow_and_update(); + if (rows, cols, cell_width_px, cell_height_px) == last_size { + continue; + } + last_size = (rows, cols, cell_width_px, cell_height_px); + if let Err(e) = + resize_pty_fd(resize_fd, rows, cols, cell_width_px, cell_height_px) + { + warn!(pane = pane_id.raw(), err = %e, rows, cols, "handoff pty resize failed"); + } + } + let _ = unsafe { libc::close(resize_fd) }; + }); + } + + let (detect_handle, detect_reset_notify, pending_release) = + spawn_basic_detection_task(pane_id, child_pid.clone(), terminal.clone(), events); + + Ok(Self { + pane_id, + terminal, + sender: input_tx, + resize_tx, + current_size: Cell::new((rows, cols, cell_width_px, cell_height_px)), + child_pid, + pty_master: None, + raw_master_fd: Some(master_fd.into_raw_fd()), + force_resize_fd: Some(force_resize_fd), + io_stop, + reader_paused, + reader_pause_ack, + reader_stopped_rx: Some(reader_stopped_rx), + kitty_keyboard_flags, + detect_reset_notify, + pending_release, + preserve_processes_on_drop: true, + detect_handle, + }) + } + fn spawn_command_builder( pane_id: PaneId, rows: u16, @@ -608,14 +1190,19 @@ impl PaneRuntime { let terminal = Arc::new(PaneTerminal::new(pane_terminal)); let kitty_keyboard_flags = Arc::new(AtomicU16::new(0)); - let reader = pair + let master_fd = pair .master - .try_clone_reader() - .map_err(|e| std::io::Error::other(e.to_string()))?; - let writer = pair - .master - .take_writer() - .map_err(|e| std::io::Error::other(e.to_string()))?; + .as_raw_fd() + .ok_or_else(|| std::io::Error::other("pty master fd is unavailable"))?; + set_nonblocking(master_fd)?; + let reader = file_from_duplicated_fd(master_fd)?; + let writer = file_from_duplicated_fd(master_fd)?; + let force_resize_fd = duplicate_cloexec_fd(master_fd)?; + let resize_fd = duplicate_cloexec_fd(master_fd)?; + let io_stop = Arc::new(AtomicBool::new(false)); + let reader_paused = Arc::new(AtomicBool::new(false)); + let reader_pause_ack = Arc::new(AtomicBool::new(false)); + let (reader_stopped_tx, reader_stopped_rx) = std::sync::mpsc::channel(); // --- Child watcher task --- let child_pid = Arc::new(AtomicU32::new(0)); @@ -652,19 +1239,43 @@ impl PaneRuntime { // --- Reader task: PTY → terminal backend + screen snapshot + terminal query responses --- { + use std::os::fd::AsRawFd; + let mut reader = reader; + let reader_fd = reader.as_raw_fd(); let terminal = terminal.clone(); let response_writer = input_tx.clone(); let render_notify = render_notify.clone(); let render_dirty = render_dirty.clone(); let child_pid = child_pid.clone(); let events = events.clone(); + let io_stop = io_stop.clone(); + let reader_paused = reader_paused.clone(); + let reader_pause_ack = reader_pause_ack.clone(); let rt = tokio::runtime::Handle::current(); tokio::task::spawn_blocking(move || { let mut buf = [0u8; 8192]; loop { + if io_stop.load(Ordering::Acquire) { + break; + } + if reader_paused.load(Ordering::Acquire) { + reader_pause_ack.store(true, Ordering::Release); + std::thread::sleep(std::time::Duration::from_millis(5)); + continue; + } + reader_pause_ack.store(false, Ordering::Release); + match poll_read_ready(reader_fd, 50) { + Ok(true) => {} + Ok(false) => continue, + Err(e) => { + debug!(pane = pane_id.raw(), err = %e, "pty reader poll failed"); + break; + } + } match reader.read(&mut buf) { Ok(0) => break, + Err(e) if e.kind() == std::io::ErrorKind::WouldBlock => continue, Err(e) => { debug!(pane = pane_id.raw(), err = %e, "pty reader closed"); break; @@ -704,6 +1315,7 @@ impl PaneRuntime { } } } + let _ = reader_stopped_tx.send(()); debug!(pane = pane_id.raw(), "reader task exiting"); }); } @@ -969,11 +1581,19 @@ impl PaneRuntime { // --- Writer task: channel → PTY --- { - let mut writer = BufWriter::new(writer); + use std::os::fd::AsRawFd; + + let mut writer = writer; + let writer_fd = writer.as_raw_fd(); + let io_stop = io_stop.clone(); tokio::task::spawn_blocking(move || { let rt = tokio::runtime::Handle::current(); while let Some(bytes) = rt.block_on(input_rx.recv()) { - if let Err(e) = writer.write_all(&bytes) { + if io_stop.load(Ordering::Acquire) { + break; + } + if let Err(e) = write_all_nonblocking(&mut writer, writer_fd, &bytes, &io_stop) + { warn!(pane = pane_id.raw(), err = %e, "pty write failed"); break; } @@ -989,30 +1609,27 @@ impl PaneRuntime { // --- Resize task --- let (resize_tx, mut resize_rx) = watch::channel::<(u16, u16, u32, u32)>((rows, cols, 0, 0)); { - let master = pair.master; + let io_stop = io_stop.clone(); tokio::task::spawn_blocking(move || { let rt = tokio::runtime::Handle::current(); let mut last_size = (rows, cols, 0, 0); while rt.block_on(resize_rx.changed()).is_ok() { + if io_stop.load(Ordering::Acquire) { + break; + } let (rows, cols, cell_width_px, cell_height_px) = *resize_rx.borrow_and_update(); if (rows, cols, cell_width_px, cell_height_px) == last_size { continue; } last_size = (rows, cols, cell_width_px, cell_height_px); - if let Err(e) = master.resize(PtySize { - rows, - cols, - pixel_width: (cols as u32) - .saturating_mul(cell_width_px) - .min(u16::MAX as u32) as u16, - pixel_height: (rows as u32) - .saturating_mul(cell_height_px) - .min(u16::MAX as u32) as u16, - }) { + if let Err(e) = + resize_pty_fd(resize_fd, rows, cols, cell_width_px, cell_height_px) + { warn!(pane = pane_id.raw(), err = %e, rows, cols, "pty resize failed"); } } + let _ = unsafe { libc::close(resize_fd) }; }); } @@ -1023,9 +1640,17 @@ impl PaneRuntime { resize_tx, current_size: Cell::new((rows, cols, 0, 0)), child_pid, + pty_master: Some(pair.master), + raw_master_fd: None, + force_resize_fd: Some(force_resize_fd), + io_stop, + reader_paused, + reader_pause_ack, + reader_stopped_rx: Some(reader_stopped_rx), kitty_keyboard_flags, detect_reset_notify, pending_release, + preserve_processes_on_drop: false, detect_handle, }) } @@ -1060,6 +1685,36 @@ impl PaneRuntime { let _ = self.resize_tx.send(size); } + pub fn nudge_child_redraw_after_handoff(&self) { + let Some(fd) = self.force_resize_fd else { + return; + }; + let (rows, cols, cell_width_px, cell_height_px) = self.current_size.get(); + let nudge = if rows > 2 { + (rows - 1, cols, cell_width_px, cell_height_px) + } else { + ( + rows, + cols.saturating_sub(1).max(4), + cell_width_px, + cell_height_px, + ) + }; + if nudge == (rows, cols, cell_width_px, cell_height_px) { + return; + } + + let Ok(fd) = duplicate_cloexec_fd(fd) else { + return; + }; + std::thread::spawn(move || { + let _ = resize_pty_fd(fd, nudge.0, nudge.1, nudge.2, nudge.3); + std::thread::sleep(std::time::Duration::from_millis(30)); + let _ = resize_pty_fd(fd, rows, cols, cell_width_px, cell_height_px); + let _ = unsafe { libc::close(fd) }; + }); + } + /// Scroll up by N lines (into scrollback history). pub fn scroll_up(&self, lines: usize) { self.terminal.scroll_up(lines); @@ -1322,9 +1977,17 @@ impl PaneRuntime { resize_tx, current_size: Cell::new((rows, cols, 0, 0)), child_pid: Arc::new(AtomicU32::new(0)), + pty_master: None, + raw_master_fd: None, + force_resize_fd: None, + io_stop: Arc::new(AtomicBool::new(false)), + reader_paused: Arc::new(AtomicBool::new(false)), + reader_pause_ack: Arc::new(AtomicBool::new(false)), + reader_stopped_rx: None, kitty_keyboard_flags: Arc::new(AtomicU16::new(0)), detect_reset_notify: Arc::new(Notify::new()), pending_release: Arc::new(Mutex::new(None)), + preserve_processes_on_drop: true, detect_handle: tokio::spawn(async {}).abort_handle(), }, rx, @@ -1410,6 +2073,47 @@ mod tests { assert_eq!(output, "vt100\n24bit\n"); } + #[tokio::test] + async fn handoff_history_ansi_captures_primary_screen() { + let runtime = + PaneRuntime::test_with_scrollback_bytes(40, 5, 4096, b"handoff-primary-history\r\n"); + + let history = runtime.handoff_history_ansi().unwrap(); + + assert!(history.contains("handoff-primary-history")); + } + + #[tokio::test] + async fn handoff_history_ansi_skips_alternate_screen() { + let runtime = PaneRuntime::test_with_scrollback_bytes( + 40, + 5, + 4096, + b"primary\r\n\x1b[?1049halt-screen", + ); + + assert!(runtime.handoff_history_ansi().is_none()); + } + + #[test] + fn truncate_handoff_history_keeps_recent_utf8_boundary() { + let history = format!("old\n{}\nrecent\n", "é".repeat(8)); + + let truncated = truncate_handoff_history(history, 20); + + assert_eq!(truncated, "recent\n"); + assert!(truncated.is_char_boundary(0)); + } + + #[test] + fn truncate_handoff_history_drops_partial_long_line() { + let history = format!("old\n{}", "x".repeat(64)); + + let truncated = truncate_handoff_history(history, 12); + + assert!(truncated.is_empty()); + } + #[test] fn restore_wrapper_falls_back_after_early_resume_failure() { let argv = vec!["/bin/sh".into(), "-c".into(), "exit 7".into()]; @@ -1510,9 +2214,17 @@ mod tests { resize_tx, current_size: Cell::new((80, 24, 0, 0)), child_pid: Arc::new(AtomicU32::new(0)), + pty_master: None, + raw_master_fd: None, + force_resize_fd: None, + io_stop: Arc::new(AtomicBool::new(false)), + reader_paused: Arc::new(AtomicBool::new(false)), + reader_pause_ack: Arc::new(AtomicBool::new(false)), + reader_stopped_rx: None, kitty_keyboard_flags: Arc::new(AtomicU16::new(0)), detect_reset_notify: Arc::new(Notify::new()), pending_release: Arc::new(Mutex::new(None)), + preserve_processes_on_drop: true, detect_handle: tokio::spawn(async {}).abort_handle(), }; @@ -1534,9 +2246,17 @@ mod tests { resize_tx, current_size: Cell::new((80, 24, 0, 0)), child_pid: Arc::new(AtomicU32::new(0)), + pty_master: None, + raw_master_fd: None, + force_resize_fd: None, + io_stop: Arc::new(AtomicBool::new(false)), + reader_paused: Arc::new(AtomicBool::new(false)), + reader_pause_ack: Arc::new(AtomicBool::new(false)), + reader_stopped_rx: None, kitty_keyboard_flags: Arc::new(AtomicU16::new(0)), detect_reset_notify: Arc::new(Notify::new()), pending_release: Arc::new(Mutex::new(None)), + preserve_processes_on_drop: true, detect_handle: tokio::spawn(async {}).abort_handle(), }; diff --git a/src/persist.rs b/src/persist.rs index b2698364..0409511c 100644 --- a/src/persist.rs +++ b/src/persist.rs @@ -9,6 +9,8 @@ mod snapshot; pub use self::io::{clear, clear_history, load, load_history, save}; pub use self::restore::restore; +#[cfg(unix)] +pub use self::restore::{restore_handoff, ImportedPaneRuntime}; pub use self::snapshot::{ capture, capture_history, DirectionSnapshot, LayoutSnapshot, SessionHistorySnapshot, SessionSnapshot, TabSnapshot, WorkspaceSnapshot, diff --git a/src/persist/restore.rs b/src/persist/restore.rs index 03bbf4e3..ef6d39f4 100644 --- a/src/persist/restore.rs +++ b/src/persist/restore.rs @@ -22,6 +22,17 @@ use super::{ WorkspaceSnapshot, }; +#[cfg(unix)] +pub struct ImportedPaneRuntime { + pub master_fd: std::os::fd::RawFd, + pub child_pid: u32, + pub rows: u16, + pub cols: u16, + pub cell_width_px: u32, + pub cell_height_px: u32, + pub initial_history_ansi: Option, +} + struct AgentRestoreState<'a> { enabled: bool, resumed_sessions: &'a mut HashSet, @@ -34,6 +45,32 @@ struct PaneRestoreStartup<'a> { reserved_agent_session: Option, } +struct RestoreRuntimeContext<'a> { + scrollback_limit_bytes: usize, + default_shell: &'a str, + resume_agents_on_restore: bool, + events: mpsc::Sender, + render_notify: Arc, + render_dirty: Arc, +} + +type RestoredSession = ( + Vec, + HashMap, + HashMap, +); +type RestoredWorkspace = ( + Workspace, + Vec, + HashMap, +); +type RestoredTab = ( + crate::workspace::Tab, + Vec, + HashMap, +); +type RestoreFailures = (T, usize); + /// Restore workspaces from a snapshot. Each pane gets a fresh shell in its saved cwd. pub fn restore( snapshot: &SessionSnapshot, @@ -46,29 +83,156 @@ pub fn restore( events: mpsc::Sender, render_notify: Arc, render_dirty: Arc, -) -> ( - Vec, - HashMap, - HashMap, -) { +) -> RestoredSession { + let mut imported_panes = HashMap::new(); + restore_with_imports( + snapshot, + history, + rows, + cols, + scrollback_limit_bytes, + default_shell, + resume_agents_on_restore, + &mut imported_panes, + events, + render_notify, + render_dirty, + ) +} + +#[cfg(unix)] +pub fn restore_handoff( + snapshot: &SessionSnapshot, + scrollback_limit_bytes: usize, + default_shell: &str, + imports: &mut HashMap, + events: mpsc::Sender, + render_notify: Arc, + render_dirty: Arc, +) -> std::io::Result { + restore_with_imports_strict( + snapshot, + None, + 24, + 80, + scrollback_limit_bytes, + default_shell, + false, + imports, + events, + render_notify, + render_dirty, + ) +} + +#[cfg(unix)] +fn restore_with_imports_strict( + snapshot: &SessionSnapshot, + history: Option<&SessionHistorySnapshot>, + rows: u16, + cols: u16, + scrollback_limit_bytes: usize, + default_shell: &str, + resume_agents_on_restore: bool, + imported_panes: &mut HashMap, + events: mpsc::Sender, + render_notify: Arc, + render_dirty: Arc, +) -> std::io::Result { + let (restored, failed_imports) = restore_with_imports_and_failures( + snapshot, + history, + rows, + cols, + scrollback_limit_bytes, + default_shell, + resume_agents_on_restore, + imported_panes, + events, + render_notify, + render_dirty, + ); + if failed_imports > 0 { + return Err(std::io::Error::other(format!( + "handoff failed to restore {failed_imports} imported pane runtime(s)" + ))); + } + if !imported_panes.is_empty() { + return Err(std::io::Error::other(format!( + "handoff import did not consume {} pane runtime(s)", + imported_panes.len() + ))); + } + Ok(restored) +} + +fn restore_with_imports( + snapshot: &SessionSnapshot, + history: Option<&SessionHistorySnapshot>, + rows: u16, + cols: u16, + scrollback_limit_bytes: usize, + default_shell: &str, + resume_agents_on_restore: bool, + imported_panes: &mut HashMap, + events: mpsc::Sender, + render_notify: Arc, + render_dirty: Arc, +) -> RestoredSession { + restore_with_imports_and_failures( + snapshot, + history, + rows, + cols, + scrollback_limit_bytes, + default_shell, + resume_agents_on_restore, + imported_panes, + events, + render_notify, + render_dirty, + ) + .0 +} + +fn restore_with_imports_and_failures( + snapshot: &SessionSnapshot, + history: Option<&SessionHistorySnapshot>, + rows: u16, + cols: u16, + scrollback_limit_bytes: usize, + default_shell: &str, + resume_agents_on_restore: bool, + imported_panes: &mut HashMap, + events: mpsc::Sender, + render_notify: Arc, + render_dirty: Arc, +) -> RestoreFailures { let mut workspaces = Vec::new(); let mut terminals = HashMap::new(); let mut terminal_runtimes = HashMap::new(); let mut resumed_agent_sessions = HashSet::new(); + let mut failed_imports = 0; for (idx, ws_snap) in snapshot.workspaces.iter().enumerate() { - if let Some((workspace, restored_terminals, restored_runtimes)) = restore_workspace( + let runtime_context = RestoreRuntimeContext { + scrollback_limit_bytes, + default_shell, + resume_agents_on_restore, + events: events.clone(), + render_notify: render_notify.clone(), + render_dirty: render_dirty.clone(), + }; + let (restored, workspace_failed_imports) = restore_workspace( ws_snap, history.and_then(|history| history.workspaces.get(idx)), rows, cols, - scrollback_limit_bytes, - default_shell, - resume_agents_on_restore, + &runtime_context, &mut resumed_agent_sessions, - events.clone(), - render_notify.clone(), - render_dirty.clone(), - ) { + imported_panes, + ); + failed_imports += workspace_failed_imports; + if let Some((workspace, restored_terminals, restored_runtimes)) = restored { for terminal in restored_terminals { terminals.insert(terminal.id.clone(), terminal); } @@ -76,7 +240,7 @@ pub fn restore( workspaces.push(workspace); } } - (workspaces, terminals, terminal_runtimes) + ((workspaces, terminals, terminal_runtimes), failed_imports) } fn restore_workspace( @@ -84,42 +248,32 @@ fn restore_workspace( history: Option<&WorkspaceHistorySnapshot>, rows: u16, cols: u16, - scrollback_limit_bytes: usize, - default_shell: &str, - resume_agents_on_restore: bool, + runtime_context: &RestoreRuntimeContext<'_>, resumed_agent_sessions: &mut HashSet, - events: mpsc::Sender, - render_notify: Arc, - render_dirty: Arc, -) -> Option<( - Workspace, - Vec, - HashMap, -)> { + imported_panes: &mut HashMap, +) -> RestoreFailures> { let mut tabs = Vec::new(); let mut terminals = Vec::new(); let mut terminal_runtimes = HashMap::new(); let mut public_pane_numbers = HashMap::new(); let mut next_public_pane_number = 1; - let mut agent_restore = AgentRestoreState { - enabled: resume_agents_on_restore, - resumed_sessions: resumed_agent_sessions, - }; + let mut failed_imports = 0; for (idx, tab_snap) in snap.tabs.iter().enumerate() { - let (tab, restored_terminals, restored_runtimes) = restore_tab( + let (restored_tab, tab_failed_imports) = restore_tab( tab_snap, history.and_then(|history| history.tabs.get(idx)), idx + 1, rows, cols, - scrollback_limit_bytes, - default_shell, - &mut agent_restore, - events.clone(), - render_notify.clone(), - render_dirty.clone(), - )?; + runtime_context, + resumed_agent_sessions, + imported_panes, + ); + failed_imports += tab_failed_imports; + let Some((tab, restored_terminals, restored_runtimes)) = restored_tab else { + continue; + }; for pane_id in tab.layout.pane_ids() { public_pane_numbers.insert(pane_id, next_public_pane_number); next_public_pane_number += 1; @@ -130,13 +284,13 @@ fn restore_workspace( } if tabs.is_empty() { - return None; + return (None, failed_imports); } let worktree_space = restored_worktree_space_membership(snap.worktree_space.clone()); - Some(( - Workspace { + ( + Some(Workspace { id: snap .id .clone() @@ -153,10 +307,10 @@ fn restore_workspace( tabs, #[cfg(test)] test_runtimes: HashMap::new(), - }, - terminals, - terminal_runtimes, - )) + }) + .map(|workspace| (workspace, terminals, terminal_runtimes)), + failed_imports, + ) } fn restored_worktree_space_membership( @@ -175,17 +329,10 @@ fn restore_tab( number: usize, rows: u16, cols: u16, - scrollback_limit_bytes: usize, - default_shell: &str, - agent_restore: &mut AgentRestoreState<'_>, - events: mpsc::Sender, - render_notify: Arc, - render_dirty: Arc, -) -> Option<( - crate::workspace::Tab, - Vec, - HashMap, -)> { + runtime_context: &RestoreRuntimeContext<'_>, + resumed_agent_sessions: &mut HashSet, + imported_panes: &mut HashMap, +) -> RestoreFailures> { let (node, id_map) = restore_node_remapped(&snap.layout); let reverse_id_map: HashMap = id_map .iter() @@ -196,6 +343,7 @@ fn restore_tab( let mut panes = HashMap::new(); let mut terminals = Vec::new(); let mut terminal_runtimes = HashMap::new(); + let mut failed_imports = 0; for id in &pane_ids { let old_id = reverse_id_map.get(id); let saved_pane = old_id.and_then(|old_id| snap.panes.get(old_id)); @@ -225,13 +373,40 @@ fn restore_tab( let saved_agent_session = saved_pane.and_then(|p| p.agent_session.as_ref()); let saved_history = old_id.and_then(|old_id| history.and_then(|history| history.panes.get(old_id))); - let startup = pane_restore_startup(saved_agent_session, saved_history, agent_restore); + let startup = { + let mut agent_restore = AgentRestoreState { + enabled: runtime_context.resume_agents_on_restore, + resumed_sessions: resumed_agent_sessions, + }; + pane_restore_startup(saved_agent_session, saved_history, &mut agent_restore) + }; let initial_restore_agent = startup .restore_plan .as_ref() .and_then(|plan| crate::detect::parse_agent_label(&plan.agent)); - let runtime_result = if let Some(plan) = startup.restore_plan { + let old_pane_id = reverse_id_map.get(id).copied(); + let imported_runtime = old_pane_id.and_then(|old_id| imported_panes.remove(&old_id)); + let was_imported = imported_runtime.is_some(); + let runtime_result = if let Some(imported) = imported_runtime { + TerminalRuntime::from_handoff_fd( + crate::pane::PaneRuntimeImport { + pane_id: *id, + master_fd: imported.master_fd, + child_pid: imported.child_pid, + rows: imported.rows, + cols: imported.cols, + cell_width_px: imported.cell_width_px, + cell_height_px: imported.cell_height_px, + initial_history_ansi: imported.initial_history_ansi, + }, + runtime_context.scrollback_limit_bytes, + crate::terminal_theme::TerminalTheme::default(), + runtime_context.events.clone(), + runtime_context.render_notify.clone(), + runtime_context.render_dirty.clone(), + ) + } else if let Some(plan) = startup.restore_plan { let launch = crate::agent_resume::AgentResumeLaunch { plan: &plan, initial_history_ansi: startup.initial_history_ansi, @@ -242,12 +417,12 @@ fn restore_tab( cols, cwd.clone(), launch, - scrollback_limit_bytes, + runtime_context.scrollback_limit_bytes, crate::terminal_theme::TerminalTheme::default(), - default_shell, - events.clone(), - render_notify.clone(), - render_dirty.clone(), + runtime_context.default_shell, + runtime_context.events.clone(), + runtime_context.render_notify.clone(), + runtime_context.render_dirty.clone(), ) } else { TerminalRuntime::spawn_with_initial_history( @@ -255,13 +430,13 @@ fn restore_tab( rows, cols, cwd.clone(), - scrollback_limit_bytes, + runtime_context.scrollback_limit_bytes, crate::terminal_theme::TerminalTheme::default(), - default_shell, + runtime_context.default_shell, startup.initial_history_ansi, - events.clone(), - render_notify.clone(), - render_dirty.clone(), + runtime_context.events.clone(), + runtime_context.render_notify.clone(), + runtime_context.render_dirty.clone(), ) }; @@ -298,7 +473,16 @@ fn restore_tab( } Err(e) => { if let Some(key) = startup.reserved_agent_session.as_deref() { - agent_restore.resumed_sessions.remove(key); + resumed_agent_sessions.remove(key); + } + if was_imported { + failed_imports += 1; + error!( + tab = ?snap.custom_name, + pane_id = id.raw(), + err = %e, + "failed to restore imported pane" + ); } error!( tab = ?snap.custom_name, @@ -315,7 +499,7 @@ fn restore_tab( tab = ?snap.custom_name, "no panes could be restored for tab, dropping it" ); - return None; + return (None, failed_imports); } let surviving: HashSet = panes.keys().copied().collect(); @@ -324,30 +508,38 @@ fn restore_tab( tab = ?snap.custom_name, "restored tab lost all panes after pruning missing layout nodes" ); - return None; + return (None, failed_imports); }; let pane_ids = collect_pane_ids(&node); - let focus = resolve_restored_pane(snap.focused, &id_map, &surviving, &pane_ids)?; - let root_pane = resolve_restored_pane(snap.root_pane, &id_map, &surviving, &pane_ids)?; + let Some(focus) = resolve_restored_pane(snap.focused, &id_map, &surviving, &pane_ids) else { + return (None, failed_imports); + }; + let Some(root_pane) = resolve_restored_pane(snap.root_pane, &id_map, &surviving, &pane_ids) + else { + return (None, failed_imports); + }; let layout = TileLayout::from_saved(node, focus); - Some(( - crate::workspace::Tab { - custom_name: snap.custom_name.clone(), - number, - root_pane, - layout, - panes, - #[cfg(test)] - runtimes: HashMap::new(), - zoomed: snap.zoomed, - events, - render_notify, - render_dirty, - }, - terminals, - terminal_runtimes, - )) + ( + Some(( + crate::workspace::Tab { + custom_name: snap.custom_name.clone(), + number, + root_pane, + layout, + panes, + #[cfg(test)] + runtimes: HashMap::new(), + zoomed: snap.zoomed, + events: runtime_context.events.clone(), + render_notify: runtime_context.render_notify.clone(), + render_dirty: runtime_context.render_dirty.clone(), + }, + terminals, + terminal_runtimes, + )), + failed_imports, + ) } fn pane_restore_startup<'a>( diff --git a/src/remote.rs b/src/remote.rs index 237f06de..990ffb7b 100644 --- a/src/remote.rs +++ b/src/remote.rs @@ -322,11 +322,15 @@ fn prepare_remote_herdr(target: &str) -> io::Result { let platform = detect_remote_platform(target)?; let remote_herdr = RemoteHerdr::for_platform(platform); let override_binary = remote_binary_override_path()?; + let path_remote_herdr = remote_binary_on_path_any(target, &remote_herdr)?; if override_binary.is_none() { - if let Some(path_remote_herdr) = remote_binary_on_path(target, &remote_herdr)? { + if let Some(path_remote_herdr) = path_remote_herdr + .as_ref() + .filter(|candidate| remote_binary_matches(target, candidate).unwrap_or(false)) + { return Ok(PreparedRemoteHerdr { - remote_herdr: path_remote_herdr, + remote_herdr: path_remote_herdr.clone(), installed_or_replaced: false, }); } @@ -338,6 +342,13 @@ fn prepare_remote_herdr(target: &str) -> io::Result { } } + if let Some(status_probe_herdr) = path_remote_herdr.as_ref().or_else(|| { + remote_binary_exists(target, &remote_herdr) + .ok() + .and_then(|exists| exists.then_some(&remote_herdr)) + }) { + confirm_remote_install_with_running_server(target, status_probe_herdr)?; + } confirm_remote_install( target, &remote_herdr, @@ -381,28 +392,27 @@ fn detect_remote_platform(target: &str) -> io::Result { }) } -fn remote_binary_on_path( +fn remote_binary_on_path_any( target: &str, remote_herdr: &RemoteHerdr, ) -> io::Result> { - let output = ssh_output(target, remote_path_probe_command())?; + let output = ssh_output(target, remote_path_probe_any_command())?; if !output.status.success() { return Ok(None); } let stdout = String::from_utf8_lossy(&output.stdout); - Ok(remote_herdr_from_path_probe(remote_herdr, &stdout)) + Ok(remote_herdr_from_path_probe_any(remote_herdr, &stdout)) } -fn remote_path_probe_command() -> &'static str { +fn remote_path_probe_any_command() -> &'static str { r#"path=$(command -v herdr) || exit 1 test -n "$path" || exit 1 -version=$("$path" --version) || exit 1 -status=$("$path" status client --json) || exit 1 -printf '%s\n%s\n%s\n' "$path" "$version" "$status" +printf '%s\n' "$path" "# } +#[cfg(test)] fn remote_herdr_from_path_probe(remote_herdr: &RemoteHerdr, stdout: &str) -> Option { let mut lines = stdout.lines(); let path = lines.next()?; @@ -419,6 +429,18 @@ fn remote_herdr_from_path_probe(remote_herdr: &RemoteHerdr, stdout: &str) -> Opt Some(remote_herdr.clone().with_shell_path(shell_quote(path))) } +fn remote_herdr_from_path_probe_any( + remote_herdr: &RemoteHerdr, + stdout: &str, +) -> Option { + let mut lines = stdout.lines(); + let path = lines.next()?; + if !path.starts_with('/') { + return None; + } + Some(remote_herdr.clone().with_shell_path(shell_quote(path))) +} + fn remote_binary_matches(target: &str, remote_herdr: &RemoteHerdr) -> io::Result { let command = format!( "test -x {0} && {0} --version && {0} status client --json", @@ -439,6 +461,11 @@ fn remote_binary_matches(target: &str, remote_herdr: &RemoteHerdr) -> io::Result .unwrap_or(false)) } +fn remote_binary_exists(target: &str, remote_herdr: &RemoteHerdr) -> io::Result { + let command = format!("test -x {}", remote_herdr.shell_path); + Ok(ssh_output(target, &command)?.status.success()) +} + fn remote_binary_override_path() -> io::Result> { let Some(value) = std::env::var_os(REMOTE_BINARY_ENV_VAR) else { return Ok(None); @@ -509,6 +536,7 @@ enum RemoteServerStatus { Running { version: Option, protocol: Option, + live_handoff: bool, }, NotRunning, } @@ -526,7 +554,12 @@ fn ensure_remote_server_ready( remote_binary_changed: bool, ) -> io::Result<()> { let status = remote_server_status(target, remote_herdr)?; - let RemoteServerStatus::Running { version, protocol } = status else { + let RemoteServerStatus::Running { + version, + protocol, + live_handoff, + } = status + else { return Ok(()); }; @@ -536,6 +569,17 @@ fn ensure_remote_server_ready( return Ok(()); }; + if live_handoff && confirm_remote_server_handoff(target, version.as_deref(), protocol, reason)? + { + match live_handoff_remote_server(target, remote_herdr) { + Ok(()) => return Ok(()), + Err(err) => { + eprintln!("remote live handoff failed: {err}"); + eprintln!("falling back to remote server restart."); + } + } + } + if confirm_remote_server_stop(target, version.as_deref(), protocol, reason)? { stop_remote_server(target, remote_herdr)?; } @@ -559,6 +603,85 @@ fn remote_server_restart_reason( None } +fn confirm_remote_install_with_running_server( + target: &str, + remote_herdr: &RemoteHerdr, +) -> io::Result<()> { + let status = match remote_server_status(target, remote_herdr) { + Ok(status) => status, + Err(err) => { + if !io::stdin().is_terminal() { + return Err(io::Error::other(format!( + "could not inspect the running remote herdr server on {target} before installing: {err}; run from an interactive terminal to approve updating the remote binary" + ))); + } + eprintln!( + "could not inspect the running remote herdr server on {target} before installing: {err}" + ); + eprint!("continue installing the remote herdr binary? [Y/n] "); + io::stderr().flush()?; + + let mut answer = String::new(); + io::stdin().read_line(&mut answer)?; + let answer = answer.trim().to_ascii_lowercase(); + if answer == "n" || answer == "no" { + return Err(io::Error::new( + io::ErrorKind::Interrupted, + "remote herdr install cancelled", + )); + } + return Ok(()); + } + }; + let RemoteServerStatus::Running { + version, + protocol, + live_handoff, + } = status + else { + return Ok(()); + }; + if live_handoff { + return Ok(()); + } + + if !io::stdin().is_terminal() { + return Err(io::Error::other(format!( + "remote herdr server on {target} is running v{} protocol {}, but it does not advertise live handoff; run from an interactive terminal to approve updating the remote binary", + version_label(version.as_deref()), + protocol_label(protocol) + ))); + } + + eprintln!("remote herdr server on {target} is currently running:"); + eprintln!( + " server: v{} protocol {}", + version_label(version.as_deref()), + protocol_label(protocol) + ); + eprintln!( + "this server does not advertise live handoff, so this attach cannot preserve its running panes during the update." + ); + eprintln!( + "future remote updates can preserve panes after the remote server has run a handoff-capable version once." + ); + eprintln!(); + eprint!("continue installing the remote herdr binary? [Y/n] "); + io::stderr().flush()?; + + let mut answer = String::new(); + io::stdin().read_line(&mut answer)?; + let answer = answer.trim().to_ascii_lowercase(); + if answer == "n" || answer == "no" { + return Err(io::Error::new( + io::ErrorKind::Interrupted, + "remote herdr install cancelled", + )); + } + + Ok(()) +} + fn remote_server_status( target: &str, remote_herdr: &RemoteHerdr, @@ -583,6 +706,12 @@ struct RemoteServerStatusJson { running: bool, version: Option, protocol: Option, + capabilities: Option, +} + +#[derive(Debug, Deserialize)] +struct RemoteServerCapabilitiesJson { + live_handoff: bool, } fn parse_client_status_json(status: &str) -> Option { @@ -602,6 +731,9 @@ fn parse_remote_server_status_json(status: &str) -> io::Result, + protocol: Option, + reason: RemoteServerRestartReason, +) -> io::Result { + if !io::stdin().is_terminal() { + if reason == RemoteServerRestartReason::ProtocolMismatch { + return Err(io::Error::other(format!( + "remote herdr server on {target} is running with protocol {}, but this client needs protocol {CURRENT_PROTOCOL}; run from an interactive terminal to approve live handoff or stopping it", + protocol_label(protocol) + ))); + } + + eprintln!( + "remote herdr server on {target} is still running v{}; it will use v{CURRENT_VERSION} after it restarts.", + version_label(version) + ); + return Ok(false); + } + + eprintln!("remote herdr server on {target} is currently running:"); + eprintln!( + " server: v{} protocol {}", + version_label(version), + protocol_label(protocol) + ); + eprintln!(" prepared binary: v{CURRENT_VERSION} protocol {CURRENT_PROTOCOL}"); + eprintln!(); + + match reason { + RemoteServerRestartReason::ProtocolMismatch => { + eprintln!( + "the remote server protocol does not match this client. herdr will try to hand off live pane processes to the prepared remote server before the old server exits." + ); + } + RemoteServerRestartReason::BinaryUpdated => { + eprintln!( + "the remote herdr binary was installed or replaced. herdr will try to hand off live pane processes to the prepared remote server." + ); + } + RemoteServerRestartReason::VersionMismatch => { + eprintln!( + "the remote server is still running a different herdr version. herdr will try to hand off live pane processes to the prepared remote server." + ); + } + } + + eprint!("live-handoff remote panes to the prepared server? [Y/n] "); + io::stderr().flush()?; + + let mut answer = String::new(); + io::stdin().read_line(&mut answer)?; + let answer = answer.trim().to_ascii_lowercase(); + Ok(answer != "n" && answer != "no") +} + +fn live_handoff_remote_server(target: &str, remote_herdr: &RemoteHerdr) -> io::Result<()> { + let command = format!( + "{} server live-handoff --import-exe {} --expected-protocol {CURRENT_PROTOCOL} --expected-version {CURRENT_VERSION}", + remote_herdr.shell_path, + remote_herdr.shell_path + ); + let output = ssh_output(target, &command)?; + if !output.status.success() { + return Err(command_failed("remote server live handoff failed", &output)); + } + + eprintln!( + "handed off the remote herdr server on {target}; reconnecting to the prepared server." + ); + Ok(()) +} + fn stop_remote_server(target: &str, remote_herdr: &RemoteHerdr) -> io::Result<()> { let command = format!("{} server stop", remote_herdr.shell_path); let output = ssh_output(target, &command)?; @@ -1495,6 +1701,21 @@ mod tests { #[test] fn parse_remote_server_status_json_reads_running_server() { + assert_eq!( + parse_remote_server_status_json( + r#"{"status":"running","running":true,"version":"0.6.0","protocol":8,"capabilities":{"live_handoff":true}}"# + ) + .unwrap(), + RemoteServerStatus::Running { + version: Some("0.6.0".into()), + protocol: Some(8), + live_handoff: true + } + ); + } + + #[test] + fn parse_remote_server_status_json_treats_missing_capability_as_no_handoff() { assert_eq!( parse_remote_server_status_json( r#"{"status":"running","running":true,"version":"0.6.0","protocol":8}"# @@ -1502,7 +1723,8 @@ mod tests { .unwrap(), RemoteServerStatus::Running { version: Some("0.6.0".into()), - protocol: Some(8) + protocol: Some(8), + live_handoff: false } ); } diff --git a/src/server/client_accept.rs b/src/server/client_accept.rs index 61db2eb9..afb2c398 100644 --- a/src/server/client_accept.rs +++ b/src/server/client_accept.rs @@ -48,3 +48,22 @@ pub(crate) fn accept_pending_client_connections( Ok(()) } + +/// Drains pending thin-client connections without starting handshakes. +/// +/// During live handoff the old server must not let clients sit in the Unix +/// listener backlog waiting for a welcome frame that will never be sent. +pub(crate) fn reject_pending_client_connections(listener: &UnixListener) -> io::Result<()> { + loop { + match listener.accept() { + Ok((_stream, _addr)) => {} + Err(ref err) if err.kind() == io::ErrorKind::WouldBlock => break, + Err(err) => { + error!(err = %err, "client listener reject failed"); + break; + } + } + } + + Ok(()) +} diff --git a/src/server/handoff.rs b/src/server/handoff.rs new file mode 100644 index 00000000..d38bdc29 --- /dev/null +++ b/src/server/handoff.rs @@ -0,0 +1,446 @@ +#[cfg(unix)] +use std::io::{self, Read, Write}; +#[cfg(unix)] +use std::os::fd::{AsRawFd, RawFd}; +#[cfg(unix)] +use std::os::unix::net::{UnixListener, UnixStream}; +#[cfg(unix)] +use std::os::unix::process::CommandExt; +#[cfg(unix)] +use std::path::{Path, PathBuf}; +#[cfg(unix)] +use std::process::Command; +#[cfg(unix)] +use std::time::Duration; + +#[cfg(unix)] +use serde::{Deserialize, Serialize}; +#[cfg(unix)] +use tracing::{info, warn}; + +#[cfg(unix)] +const HANDOFF_VERSION: u32 = 1; +#[cfg(unix)] +const READY_TIMEOUT: Duration = Duration::from_secs(30); +#[cfg(unix)] +pub(crate) const MAX_FDS_PER_HANDOFF: usize = 64; +#[cfg(unix)] +pub(crate) const MAX_REPLAY_BYTES_PER_PANE: usize = 8 * 1024; +#[cfg(unix)] +pub(crate) const COMMIT_TIMEOUT: Duration = READY_TIMEOUT; + +#[cfg(unix)] +#[derive(Serialize, Deserialize)] +pub(crate) struct HandoffManifest { + pub version: u32, + pub source_version: String, + pub source_protocol: u32, + pub expected_version: Option, + pub expected_protocol: Option, + pub snapshot: crate::persist::SessionSnapshot, + pub panes: Vec, +} + +#[cfg(unix)] +#[derive(Debug, Serialize, Deserialize)] +pub(crate) struct HandoffPane { + pub pane_id: u32, + pub child_pid: u32, + pub rows: u16, + pub cols: u16, + pub cell_width_px: u32, + pub cell_height_px: u32, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub initial_history_ansi: Option, +} + +#[cfg(unix)] +pub(crate) struct ReceivedHandoff { + pub manifest: HandoffManifest, + pub fds: Vec, + pub stream: UnixStream, +} + +#[cfg(unix)] +pub(crate) fn handoff_socket_path() -> PathBuf { + crate::session::data_dir().join(format!("herdr-handoff-{}.sock", std::process::id())) +} + +#[cfg(unix)] +pub(crate) fn spawn_handoff_import( + import_exe: Option<&Path>, + socket_path: &Path, + token: &str, +) -> io::Result { + let fallback_exe; + let exe = if let Some(import_exe) = import_exe { + import_exe + } else { + fallback_exe = std::env::current_exe().map_err(|err| { + io::Error::new( + err.kind(), + format!("failed to determine herdr executable path: {err}"), + ) + })?; + &fallback_exe + }; + let mut command = Command::new(exe); + command + .arg("server") + .arg("--handoff-import") + .arg(socket_path) + .arg(token) + .process_group(0) + .stdin(std::process::Stdio::null()) + .stdout(std::process::Stdio::null()) + .stderr(std::process::Stdio::null()); + let child = command.spawn().map_err(|err| { + io::Error::new( + err.kind(), + format!( + "failed to spawn handoff import server at {}: {err}", + exe.display() + ), + ) + })?; + Ok(child.id()) +} + +#[cfg(unix)] +pub(crate) fn bind_listener(socket_path: &Path) -> io::Result { + let _ = std::fs::remove_file(socket_path); + let listener = UnixListener::bind(socket_path)?; + listener.set_nonblocking(true)?; + restrict_socket_permissions(socket_path)?; + Ok(listener) +} + +#[cfg(unix)] +pub(crate) fn accept_and_validate_on( + listener: UnixListener, + socket_path: &Path, + token: &str, + manifest: &HandoffManifest, +) -> io::Result { + let (mut stream, _) = accept_with_timeout(&listener, READY_TIMEOUT)?; + stream.set_nonblocking(false)?; + stream.set_read_timeout(Some(READY_TIMEOUT))?; + stream.set_write_timeout(Some(READY_TIMEOUT))?; + let token_line = read_line_unbuffered(&mut stream)?; + if token_line.trim_end() != token { + return Err(io::Error::new( + io::ErrorKind::PermissionDenied, + "handoff import token mismatch", + )); + } + + serde_json::to_writer(&mut stream, manifest).map_err(io::Error::other)?; + stream.write_all(b"\n")?; + stream.flush()?; + + stream.set_read_timeout(Some(READY_TIMEOUT))?; + let validated = read_line_unbuffered(&mut stream)?; + if validated.trim_end() != "validated" { + return Err(io::Error::other("handoff import did not validate manifest")); + } + let _ = std::fs::remove_file(socket_path); + Ok(stream) +} + +#[cfg(unix)] +pub(crate) fn send_fds_and_wait_restored(stream: &mut UnixStream, fds: &[RawFd]) -> io::Result<()> { + if fds.len() > MAX_FDS_PER_HANDOFF { + return Err(io::Error::new( + io::ErrorKind::InvalidInput, + format!("handoff supports at most {MAX_FDS_PER_HANDOFF} pane file descriptors at once"), + )); + } + send_fds(stream, fds)?; + + stream.set_read_timeout(Some(READY_TIMEOUT))?; + let restored = read_line_unbuffered(&mut *stream)?; + if restored.trim_end() != "restored" { + return Err(io::Error::other( + "handoff import did not report restored runtimes", + )); + } + Ok(()) +} + +#[cfg(unix)] +pub(crate) fn wait_ready(stream: &mut UnixStream) -> io::Result<()> { + stream.set_read_timeout(Some(READY_TIMEOUT))?; + let ready = read_line_unbuffered(&mut *stream)?; + if ready.trim_end() != "ready" { + return Err(io::Error::other("handoff import did not report ready")); + } + Ok(()) +} + +#[cfg(unix)] +pub(crate) fn report_committed(stream: &mut UnixStream) -> io::Result<()> { + stream.write_all(b"committed\n")?; + stream.flush() +} + +#[cfg(unix)] +pub(crate) fn wait_owned_ack(stream: &mut UnixStream) { + if let Err(err) = stream.set_read_timeout(Some(READY_TIMEOUT)) { + warn!(err = %err, "failed to set handoff ownership ack timeout"); + return; + } + match read_line_unbuffered(&mut *stream) { + Ok(owned) if owned.trim_end() == "owned" => {} + Ok(other) => { + warn!( + response = %other.trim_end(), + "handoff import sent unexpected ownership ack after commit" + ); + } + Err(err) => { + warn!(err = %err, "handoff import ownership ack was not received after commit"); + } + } +} + +#[cfg(unix)] +pub(crate) fn receive(socket_path: &Path, token: &str) -> io::Result { + let mut stream = UnixStream::connect(socket_path)?; + stream.write_all(token.as_bytes())?; + stream.write_all(b"\n")?; + stream.flush()?; + + let manifest_line = read_line_unbuffered(&mut stream)?; + let manifest: HandoffManifest = + serde_json::from_str(&manifest_line).map_err(io::Error::other)?; + if manifest.version != HANDOFF_VERSION { + return Err(io::Error::other(format!( + "unsupported handoff version {}", + manifest.version + ))); + } + if manifest + .expected_protocol + .is_some_and(|protocol| protocol != crate::protocol::PROTOCOL_VERSION) + { + return Err(io::Error::other(format!( + "handoff expected protocol {}, but this server speaks protocol {}", + manifest.expected_protocol.unwrap_or_default(), + crate::protocol::PROTOCOL_VERSION + ))); + } + if manifest + .expected_version + .as_deref() + .is_some_and(|version| version != env!("CARGO_PKG_VERSION")) + { + return Err(io::Error::other(format!( + "handoff expected herdr v{}, but this server is v{}", + manifest.expected_version.as_deref().unwrap_or("unknown"), + env!("CARGO_PKG_VERSION") + ))); + } + stream.write_all(b"validated\n")?; + stream.flush()?; + let fds = recv_fds(&stream, manifest.panes.len())?; + Ok(ReceivedHandoff { + manifest, + fds, + stream, + }) +} + +#[cfg(unix)] +pub(crate) fn report_restored(stream: &mut UnixStream) -> io::Result<()> { + stream.write_all(b"restored\n")?; + stream.flush() +} + +#[cfg(unix)] +pub(crate) fn report_ready(stream: &mut UnixStream) -> io::Result<()> { + stream.write_all(b"ready\n")?; + stream.flush() +} + +#[cfg(unix)] +pub(crate) fn wait_committed(stream: &mut UnixStream) -> io::Result<()> { + stream.set_read_timeout(Some(READY_TIMEOUT))?; + let committed = read_line_unbuffered(&mut *stream)?; + if committed.trim_end() != "committed" { + return Err(io::Error::other("handoff source did not commit")); + } + Ok(()) +} + +#[cfg(unix)] +pub(crate) fn report_owned(stream: &mut UnixStream) -> io::Result<()> { + stream.write_all(b"owned\n")?; + stream.flush() +} + +#[cfg(unix)] +pub(crate) fn manifest_for( + snapshot: crate::persist::SessionSnapshot, + panes: Vec, + expected_protocol: Option, + expected_version: Option, +) -> HandoffManifest { + HandoffManifest { + version: HANDOFF_VERSION, + source_version: env!("CARGO_PKG_VERSION").to_string(), + source_protocol: crate::protocol::PROTOCOL_VERSION, + expected_version, + expected_protocol, + snapshot, + panes, + } +} + +#[cfg(unix)] +fn restrict_socket_permissions(path: &Path) -> io::Result<()> { + use std::os::unix::fs::PermissionsExt; + + std::fs::set_permissions(path, std::fs::Permissions::from_mode(0o600)) +} + +#[cfg(unix)] +fn accept_with_timeout( + listener: &UnixListener, + timeout: Duration, +) -> io::Result<(UnixStream, std::os::unix::net::SocketAddr)> { + let deadline = std::time::Instant::now() + timeout; + loop { + match listener.accept() { + Ok(accepted) => return Ok(accepted), + Err(err) if err.kind() == io::ErrorKind::WouldBlock => { + if std::time::Instant::now() >= deadline { + return Err(io::Error::new( + io::ErrorKind::TimedOut, + "timed out waiting for handoff import connection", + )); + } + std::thread::sleep(Duration::from_millis(25)); + } + Err(err) if err.kind() == io::ErrorKind::Interrupted => {} + Err(err) => return Err(err), + } + } +} + +#[cfg(unix)] +fn read_line_unbuffered(stream: &mut UnixStream) -> io::Result { + let mut bytes = Vec::new(); + let mut byte = [0u8; 1]; + loop { + let read = stream.read(&mut byte)?; + if read == 0 { + return Err(io::Error::new( + io::ErrorKind::UnexpectedEof, + "handoff stream closed while reading line", + )); + } + bytes.push(byte[0]); + if byte[0] == b'\n' { + return String::from_utf8(bytes) + .map_err(|err| io::Error::new(io::ErrorKind::InvalidData, err)); + } + if bytes.len() > 16 * 1024 * 1024 { + return Err(io::Error::new( + io::ErrorKind::InvalidData, + "handoff line exceeded maximum size", + )); + } + } +} + +#[cfg(unix)] +fn send_fds(stream: &UnixStream, fds: &[RawFd]) -> io::Result<()> { + if fds.is_empty() { + return Ok(()); + } + let byte = [b'F']; + let iov = [libc::iovec { + iov_base: byte.as_ptr() as *mut libc::c_void, + iov_len: byte.len(), + }]; + let fd_bytes = std::mem::size_of_val(fds); + let mut control = vec![0u8; unsafe { libc::CMSG_SPACE(fd_bytes as u32) as usize }]; + let mut msg: libc::msghdr = unsafe { std::mem::zeroed() }; + msg.msg_iov = iov.as_ptr() as *mut libc::iovec; + msg.msg_iovlen = iov.len() as _; + msg.msg_control = control.as_mut_ptr() as *mut libc::c_void; + msg.msg_controllen = control.len() as _; + + unsafe { + let cmsg = libc::CMSG_FIRSTHDR(&msg); + if cmsg.is_null() { + return Err(io::Error::other("failed to allocate fd control message")); + } + (*cmsg).cmsg_level = libc::SOL_SOCKET; + (*cmsg).cmsg_type = libc::SCM_RIGHTS; + (*cmsg).cmsg_len = libc::CMSG_LEN(fd_bytes as u32) as _; + std::ptr::copy_nonoverlapping(fds.as_ptr() as *const u8, libc::CMSG_DATA(cmsg), fd_bytes); + if libc::sendmsg(stream.as_raw_fd(), &msg, 0) < 0 { + return Err(io::Error::last_os_error()); + } + } + Ok(()) +} + +#[cfg(unix)] +fn recv_fds(stream: &UnixStream, expected: usize) -> io::Result> { + if expected == 0 { + return Ok(Vec::new()); + } + let mut byte = [0u8; 1]; + let mut iov = [libc::iovec { + iov_base: byte.as_mut_ptr() as *mut libc::c_void, + iov_len: byte.len(), + }]; + let fd_bytes = expected * std::mem::size_of::(); + let mut control = vec![0u8; unsafe { libc::CMSG_SPACE(fd_bytes as u32) as usize }]; + let mut msg: libc::msghdr = unsafe { std::mem::zeroed() }; + msg.msg_iov = iov.as_mut_ptr(); + msg.msg_iovlen = iov.len() as _; + msg.msg_control = control.as_mut_ptr() as *mut libc::c_void; + msg.msg_controllen = control.len() as _; + + let read = unsafe { libc::recvmsg(stream.as_raw_fd(), &mut msg, 0) }; + if read < 0 { + return Err(io::Error::last_os_error()); + } + if msg.msg_flags & libc::MSG_CTRUNC != 0 { + return Err(io::Error::other("handoff fd control message was truncated")); + } + + let mut out = Vec::new(); + unsafe { + let cmsg = libc::CMSG_FIRSTHDR(&msg); + if cmsg.is_null() + || (*cmsg).cmsg_level != libc::SOL_SOCKET + || (*cmsg).cmsg_type != libc::SCM_RIGHTS + { + return Err(io::Error::other("handoff fd message missing SCM_RIGHTS")); + } + let data_len = ((*cmsg).cmsg_len as usize).saturating_sub(libc::CMSG_LEN(0) as usize); + let count = data_len / std::mem::size_of::(); + let data = libc::CMSG_DATA(cmsg) as *const RawFd; + for idx in 0..count { + out.push(*data.add(idx)); + } + } + if out.len() != expected { + for fd in out { + let _ = unsafe { libc::close(fd) }; + } + return Err(io::Error::other(format!( + "expected {expected} handoff fds, received fewer" + ))); + } + Ok(out) +} + +#[cfg(unix)] +pub(crate) fn log_import_result(panes: usize) { + info!(panes, "handoff import ready"); +} diff --git a/src/server/headless.rs b/src/server/headless.rs index d10aef04..4dafb6d7 100644 --- a/src/server/headless.rs +++ b/src/server/headless.rs @@ -15,7 +15,6 @@ //! and pane spawn failure during restore use std::collections::HashMap; -use std::fs; use std::io; use std::os::unix::net::UnixListener; use std::path::{Path, PathBuf}; @@ -35,11 +34,14 @@ use crate::api; use crate::app; use crate::config; use crate::events::AppEvent; +use crate::ipc::{remove_socket_file_if_owned, socket_file_identity, SocketFileIdentity}; use crate::protocol::{ self, AttachScrollDirection, AttachScrollSource, FrameData, ServerMessage, MAX_FRAME_SIZE, MAX_GRAPHICS_FRAME_SIZE, }; -use crate::server::client_accept::accept_pending_client_connections; +use crate::server::client_accept::{ + accept_pending_client_connections, reject_pending_client_connections, +}; use crate::server::client_transport::ServerEvent; use crate::server::clients::{ events_include_interaction, latest_app_client, render_targets, terminal_attach_client_ids, @@ -58,6 +60,8 @@ use crate::server::terminal_attach::paste_payload_for_runtime; use crate::protocol::RenderEncoding; #[cfg(test)] use crate::server::client_transport::ClientWriter; +#[cfg(test)] +use std::fs; // --------------------------------------------------------------------------- // Loop event enum for the headless server event loop @@ -100,8 +104,11 @@ const CLIENT_ACCEPT_POLL_INTERVAL: Duration = Duration::from_millis(250); /// The headless server — runs the herdr event loop without a real terminal. pub struct HeadlessServer { app: app::App, + api_tx: Option, + api_server: Option, client_listener: UnixListener, client_socket_path: PathBuf, + client_socket_identity: SocketFileIdentity, clients: HashMap, next_client_id: u64, /// The client currently driving the shared pane runtime size, theme, and input keybindings. @@ -121,6 +128,10 @@ pub struct HeadlessServer { effective_size: (u16, u16), /// Flag set when shutdown is initiated. shutting_down: bool, + /// Flag set while exporting live PTYs to a replacement server. + handoff_in_progress: bool, + /// Imported panes get one app-safe resize nudge after the first client attaches. + pending_handoff_repaint_nudge: bool, /// Flag set by Ctrl+C or `server stop` signal. should_quit: Arc, /// Channel for receiving server events from client connection threads. @@ -209,12 +220,18 @@ impl HeadlessServer { /// 1. Prepares the client socket path (cleans up stale sockets) /// 2. Binds the client socket listener /// 3. Returns the server ready to run - pub fn new(app: app::App, config_diagnostics: &[String]) -> io::Result { + pub fn new( + app: app::App, + config_diagnostics: &[String], + api_tx: Option, + api_server: Option, + ) -> io::Result { let client_path = client_socket_path(); prepare_socket_path(&client_path)?; let listener = UnixListener::bind(&client_path)?; restrict_socket_permissions(&client_path)?; + let client_socket_identity = socket_file_identity(&client_path)?; info!(path = %client_path.display(), "client protocol socket listening"); // Set non-blocking on the listener so we can poll it from the event loop. @@ -230,8 +247,11 @@ impl HeadlessServer { Ok(Self { app, + api_tx, + api_server, client_listener: listener, client_socket_path: client_path, + client_socket_identity, clients: HashMap::new(), next_client_id: 1, foreground_client_id: None, @@ -242,6 +262,8 @@ impl HeadlessServer { next_activity_stamp: 1, effective_size: (MIN_COLS, MIN_ROWS), shutting_down: false, + handoff_in_progress: false, + pending_handoff_repaint_nudge: false, should_quit, server_event_rx, server_event_tx, @@ -539,6 +561,217 @@ impl HeadlessServer { } } + #[cfg(unix)] + fn perform_live_handoff( + &mut self, + params: crate::api::schema::ServerLiveHandoffParams, + ) -> io::Result<()> { + info!("starting live handoff"); + let import_exe = params.import_exe.as_deref().map(std::path::PathBuf::from); + let socket_path = crate::server::handoff::handoff_socket_path(); + let token = format!( + "{}-{}", + std::process::id(), + std::time::SystemTime::now() + .duration_since(std::time::UNIX_EPOCH) + .unwrap_or_default() + .as_nanos() + ); + let listener = match crate::server::handoff::bind_listener(&socket_path) { + Ok(listener) => listener, + Err(err) => { + self.handoff_in_progress = false; + return Err(err); + } + }; + + let mut pane_by_terminal = HashMap::new(); + for ws in &self.app.state.workspaces { + for tab in &ws.tabs { + for (pane_id, pane) in &tab.panes { + pane_by_terminal.insert(pane.attached_terminal_id.clone(), pane_id.raw()); + } + } + } + if pane_by_terminal.len() > crate::server::handoff::MAX_FDS_PER_HANDOFF { + let _ = std::fs::remove_file(&socket_path); + return Err(io::Error::new( + io::ErrorKind::InvalidInput, + format!( + "live handoff supports at most {} panes in one update; close panes or restart herdr normally", + crate::server::handoff::MAX_FDS_PER_HANDOFF + ), + )); + } + + self.handoff_in_progress = true; + self.disconnect_all_clients_for_handoff(); + let _ = reject_pending_client_connections(&self.client_listener); + + let mut paused_terminal_ids = Vec::new(); + for terminal_id in pane_by_terminal.keys() { + if let Some(runtime) = self.app.terminal_runtimes.get(terminal_id) { + if let Err(err) = runtime.pause_handoff_reader(Duration::from_secs(2)) { + self.rollback_handoff_before_commit(&socket_path, &paused_terminal_ids); + return Err(err); + } + paused_terminal_ids.push(terminal_id.clone()); + } + } + + let snapshot = crate::persist::capture( + &self.app.state.workspaces, + &self.app.state.terminals, + &self.app.terminal_runtimes, + self.app.state.active, + self.app.state.selected, + self.app.state.agent_panel_scope, + self.app.state.sidebar_width, + self.app.state.sidebar_section_split, + self.app.state.collapsed_space_keys.clone(), + ); + + let mut panes = Vec::new(); + for (terminal_id, runtime) in self.app.terminal_runtimes.iter() { + let Some(pane_id) = pane_by_terminal.get(terminal_id).copied() else { + continue; + }; + let mut handoff_pane = runtime.handoff_pane(pane_id); + let has_agent_session = self + .app + .state + .terminals + .get(terminal_id) + .is_some_and(|terminal| terminal.persisted_agent_session.is_some()); + if !has_agent_session { + handoff_pane.initial_history_ansi = runtime.handoff_history_ansi(); + } + panes.push(handoff_pane); + } + + let manifest = crate::server::handoff::manifest_for( + snapshot, + panes, + params.expected_protocol, + params.expected_version, + ); + let child_pid = match crate::server::handoff::spawn_handoff_import( + import_exe.as_deref(), + &socket_path, + &token, + ) { + Ok(child_pid) => child_pid, + Err(err) => { + self.rollback_handoff_before_commit(&socket_path, &paused_terminal_ids); + return Err(err); + } + }; + info!(pid = child_pid, socket = %socket_path.display(), "spawned handoff import server"); + + let mut fds = Vec::new(); + let duplicate_result = (|| { + for (terminal_id, runtime) in self.app.terminal_runtimes.iter() { + if !pane_by_terminal.contains_key(terminal_id) { + continue; + } + fds.push(runtime.duplicate_handoff_fd()?); + } + Ok::<(), io::Error>(()) + })(); + if let Err(err) = duplicate_result { + for fd in fds { + let _ = unsafe { libc::close(fd) }; + } + self.rollback_handoff_before_commit(&socket_path, &paused_terminal_ids); + return Err(err); + } + + let mut stream = match crate::server::handoff::accept_and_validate_on( + listener, + &socket_path, + &token, + &manifest, + ) { + Ok(stream) => stream, + Err(err) => { + for fd in fds { + let _ = unsafe { libc::close(fd) }; + } + self.rollback_handoff_before_commit(&socket_path, &paused_terminal_ids); + return Err(err); + } + }; + + let send_result = crate::server::handoff::send_fds_and_wait_restored(&mut stream, &fds); + for fd in fds { + let _ = unsafe { libc::close(fd) }; + } + if let Err(err) = send_result { + self.rollback_handoff_before_commit(&socket_path, &paused_terminal_ids); + return Err(err); + } + + if let Some(api_server) = &self.api_server { + let _ = api_server.remove_socket_file_if_owned(); + } else { + let _ = std::fs::remove_file(crate::api::socket_path()); + } + let _ = remove_socket_file_if_owned(&self.client_socket_path, self.client_socket_identity); + if let Err(err) = crate::server::handoff::wait_ready(&mut stream) { + match self.wait_then_restore_public_sockets_after_failed_handoff() { + Ok(()) => { + self.rollback_handoff_before_commit(&socket_path, &paused_terminal_ids); + } + Err(restore_err) => { + self.rollback_handoff_before_commit(&socket_path, &paused_terminal_ids); + return Err(io::Error::other(format!( + "handoff replacement server did not become ready: {err}; old server could not restore public sockets: {restore_err}" + ))); + } + } + return Err(io::Error::other(format!( + "handoff replacement server did not become ready: {err}" + ))); + } + if let Err(err) = crate::server::handoff::report_committed(&mut stream) { + match self.wait_then_restore_public_sockets_after_failed_handoff() { + Ok(()) => { + self.rollback_handoff_before_commit(&socket_path, &paused_terminal_ids); + } + Err(restore_err) => { + self.rollback_handoff_before_commit(&socket_path, &paused_terminal_ids); + return Err(io::Error::other(format!( + "handoff replacement server was ready, but commit failed: {err}; old server could not restore public sockets: {restore_err}" + ))); + } + } + return Err(err); + } + + for (terminal_id, runtime) in self.app.terminal_runtimes.drain_for_handoff() { + if !pane_by_terminal.contains_key(&terminal_id) { + continue; + } + debug!(terminal = %terminal_id, "preserving pane runtime for handoff"); + runtime.preserve_for_handoff(); + } + crate::server::handoff::wait_owned_ack(&mut stream); + + self.shutting_down = true; + self.app.state.should_quit = true; + self.app.no_session = true; + info!("live handoff completed; old server exiting"); + Ok(()) + } + + #[cfg(not(unix))] + fn perform_live_handoff( + &mut self, + _params: crate::api::schema::ServerLiveHandoffParams, + ) -> io::Result<()> { + Err(io::Error::other("live handoff is only supported on Unix")) + } + fn sync_visible_server_config_diagnostic(&mut self, uses_local_keybindings: bool) { let visible = if uses_local_keybindings { &self.server_config_diagnostic_without_keybindings @@ -552,6 +785,64 @@ impl HeadlessServer { } } + #[cfg(unix)] + fn restore_public_sockets_after_failed_handoff(&mut self) -> io::Result<()> { + let api_tx = self + .api_tx + .clone() + .ok_or_else(|| io::Error::other("cannot restore api socket without api sender"))?; + let api_server = api::start_server(api_tx, self.app.event_hub.clone())?; + + let client_path = client_socket_path(); + prepare_socket_path(&client_path)?; + let listener = UnixListener::bind(&client_path)?; + restrict_socket_permissions(&client_path)?; + let client_socket_identity = socket_file_identity(&client_path)?; + listener.set_nonblocking(true)?; + + self.api_server = Some(api_server); + self.client_listener = listener; + self.client_socket_path = client_path; + self.client_socket_identity = client_socket_identity; + Ok(()) + } + + #[cfg(unix)] + fn wait_then_restore_public_sockets_after_failed_handoff(&mut self) -> io::Result<()> { + let timeout = crate::server::handoff::COMMIT_TIMEOUT + Duration::from_secs(2); + wait_for_old_public_sockets_to_close(timeout)?; + self.restore_public_sockets_after_failed_handoff() + } + + #[cfg(unix)] + fn rollback_handoff_before_commit( + &mut self, + socket_path: &Path, + paused_terminal_ids: &[crate::terminal::TerminalId], + ) { + for terminal_id in paused_terminal_ids { + if let Some(runtime) = self.app.terminal_runtimes.get(terminal_id) { + runtime.set_handoff_reader_paused(false); + } + } + self.handoff_in_progress = false; + let _ = std::fs::remove_file(socket_path); + } + + #[cfg(unix)] + fn nudge_handoff_panes_on_first_client_attach(&mut self) { + if !self.pending_handoff_repaint_nudge { + return; + } + self.pending_handoff_repaint_nudge = false; + self.app + .terminal_runtimes + .nudge_child_redraw_after_handoff(); + } + + #[cfg(not(unix))] + fn nudge_handoff_panes_on_first_client_attach(&mut self) {} + fn reload_server_config(&mut self, notify_success: bool) -> crate::config::ConfigReloadReport { let server_keybindings = self.server_keybindings.clone(); apply_keybindings(&mut self.app, &server_keybindings); @@ -706,6 +997,9 @@ impl HeadlessServer { /// Accepts pending client connections from the non-blocking listener. fn accept_client_connections(&mut self) -> io::Result<()> { + if self.handoff_in_progress { + return reject_pending_client_connections(&self.client_listener); + } accept_pending_client_connections( &self.client_listener, &mut self.next_client_id, @@ -1236,6 +1530,28 @@ impl HeadlessServer { } } + fn disconnect_all_clients_for_handoff(&mut self) { + let client_ids = self.clients.keys().copied().collect::>(); + for client_id in client_ids { + self.send_client_graphics_cleanup(client_id); + self.send_to_client( + client_id, + ServerMessage::ServerShutdown { + reason: Some( + "live update in progress; reconnect after handoff completes".to_owned(), + ), + }, + ); + if let Some(client) = self.clients.get_mut(&client_id) { + client.writer = None; + } + let _ = self.remove_client(client_id); + } + self.foreground_client_id = None; + self.sync_foreground_client_state(); + self.resize_shared_runtime_to_effective_size(); + } + fn attach_terminal_client( &mut self, client_id: u64, @@ -1310,6 +1626,10 @@ impl HeadlessServer { /// Handles a server event. Returns true if the event requires a re-render. fn handle_server_event(&mut self, ev: ServerEvent) -> bool { + if self.handoff_in_progress && Self::ignore_client_event_during_handoff(&ev) { + return false; + } + match ev { ServerEvent::ClientConnected { client_id, @@ -1321,6 +1641,19 @@ impl HeadlessServer { writer, render_encoding, } => { + if self.handoff_in_progress { + if let Ok(message) = + Self::frame_server_message(&ServerMessage::ServerShutdown { + reason: Some( + "live update in progress; reconnect after handoff completes" + .to_owned(), + ), + }) + { + let _ = writer.control.send(message); + } + return false; + } info!( client_id, cols, @@ -1351,6 +1684,7 @@ impl HeadlessServer { self.foreground_client_id = Some(client_id); self.sync_foreground_client_state(); self.resize_shared_runtime_to_effective_size(); + self.nudge_handoff_panes_on_first_client_attach(); true } ServerEvent::ClientAttachTerminal { @@ -1370,6 +1704,14 @@ impl HeadlessServer { client_id, source, direction, lines, column, row, modifiers, ), ServerEvent::ClientInput { client_id, data } => { + if self.handoff_in_progress { + debug!( + client_id, + len = data.len(), + "ignored client input during handoff" + ); + return false; + } debug!(client_id, len = data.len(), "client input received"); if let Some(ClientConnection { mode: ClientConnectionMode::TerminalAttach { terminal_id }, @@ -1548,6 +1890,16 @@ impl HeadlessServer { } } + fn ignore_client_event_during_handoff(ev: &ServerEvent) -> bool { + !matches!( + ev, + ServerEvent::ClientConnected { .. } + | ServerEvent::ClientDisconnected { .. } + | ServerEvent::ClientWriterDrained { .. } + | ServerEvent::QuitSignal + ) + } + /// Drains API requests with shutdown awareness. /// /// During shutdown, remaining requests get a `server_unavailable` error. @@ -1583,6 +1935,25 @@ impl HeadlessServer { return false; } + if let api::schema::Method::ServerLiveHandoff(params) = &msg.request.method { + let response = match self.perform_live_handoff(params.clone()) { + Ok(()) => serde_json::to_string(&api::schema::SuccessResponse { + id: msg.request.id, + result: api::schema::ResponseResult::Ok {}, + }), + Err(err) => serde_json::to_string(&api::schema::ErrorResponse { + id: msg.request.id, + error: api::schema::ErrorBody { + code: "handoff_failed".into(), + message: err.to_string(), + }, + }), + } + .unwrap_or_else(|_| "{}".to_string()); + let _ = msg.respond_to.send(response); + return true; + } + let changed = api::request_changes_ui(&msg.request); let changed = self.drain_all_internal_events_with_forwarding() || changed; @@ -2153,7 +2524,9 @@ impl HeadlessServer { /// Removes socket files created by the server. fn cleanup_sockets(&self) -> io::Result<()> { - if let Err(err) = fs::remove_file(&self.client_socket_path) { + if let Err(err) = + remove_socket_file_if_owned(&self.client_socket_path, self.client_socket_identity) + { if err.kind() != io::ErrorKind::NotFound { warn!( path = %self.client_socket_path.display(), @@ -2224,12 +2597,24 @@ fn is_keybinding_config_diagnostic(diagnostic: &str) -> bool { pub fn run_server() -> io::Result<()> { init_logging(); + let args: Vec = std::env::args().collect(); + if args.get(2).map(String::as_str) == Some("--handoff-import") { + let socket_path = args + .get(3) + .map(PathBuf::from) + .ok_or_else(|| io::Error::new(io::ErrorKind::InvalidInput, "missing handoff socket"))?; + let token = args + .get(4) + .ok_or_else(|| io::Error::new(io::ErrorKind::InvalidInput, "missing handoff token"))?; + return run_handoff_import_server(&socket_path, token); + } + let loaded_config = config::Config::load(); let (api_tx, api_rx) = tokio::sync::mpsc::unbounded_channel(); let event_hub = api::EventHub::default(); // Start the JSON API socket server. - let _api_server = match api::start_server(api_tx, event_hub.clone()) { + let _api_server = match api::start_server(api_tx.clone(), event_hub.clone()) { Ok(server) => server, Err(err) if err.kind() == io::ErrorKind::AddrInUse => { eprintln!("error: herdr server is already running"); @@ -2263,7 +2648,12 @@ pub fn run_server() -> io::Result<()> { app.local_terminal_notifications = false; // Create the headless server. - let mut server = match HeadlessServer::new(app, &loaded_config.diagnostics) { + let mut server = match HeadlessServer::new( + app, + &loaded_config.diagnostics, + Some(api_tx.clone()), + Some(_api_server), + ) { Ok(server) => server, Err(err) if err.kind() == io::ErrorKind::AddrInUse => { eprintln!("error: herdr server is already running"); @@ -2288,6 +2678,106 @@ pub fn run_server() -> io::Result<()> { result } +#[cfg(unix)] +fn run_handoff_import_server(socket_path: &Path, token: &str) -> io::Result<()> { + let loaded_config = config::Config::load(); + let mut received = crate::server::handoff::receive(socket_path, token)?; + crate::server::handoff::log_import_result(received.manifest.panes.len()); + + let (api_tx, api_rx) = tokio::sync::mpsc::unbounded_channel(); + let event_hub = api::EventHub::default(); + + let mut imports = HashMap::new(); + for (pane, fd) in received.manifest.panes.iter().zip(received.fds) { + imports.insert( + pane.pane_id, + crate::persist::ImportedPaneRuntime { + master_fd: fd, + child_pid: pane.child_pid, + rows: pane.rows, + cols: pane.cols, + cell_width_px: pane.cell_width_px, + cell_height_px: pane.cell_height_px, + initial_history_ansi: pane.initial_history_ansi.clone(), + }, + ); + } + + let rt = tokio::runtime::Builder::new_multi_thread() + .enable_all() + .build() + .map_err(io::Error::other)?; + + let result = rt.block_on(async { + let mut app = app::App::new_from_handoff( + &loaded_config.config, + config::config_diagnostic_summary(&loaded_config.diagnostics), + api_rx, + event_hub.clone(), + &received.manifest.snapshot, + &mut imports, + )?; + app.state.local_sound_playback = false; + app.local_terminal_notifications = false; + crate::server::handoff::report_restored(&mut received.stream)?; + if std::env::var("HERDR_TEST_HANDOFF_IMPORT_FAIL").as_deref() == Ok("after_restored") { + return Err(io::Error::other( + "test handoff import failure after restored", + )); + } + wait_for_old_public_sockets_to_close(Duration::from_secs(5))?; + + let api_server = api::start_server(api_tx.clone(), event_hub.clone())?; + let mut server = HeadlessServer::new( + app, + &loaded_config.diagnostics, + Some(api_tx.clone()), + Some(api_server), + )?; + crate::server::handoff::report_ready(&mut received.stream)?; + crate::server::handoff::wait_committed(&mut received.stream)?; + server.app.assume_handoff_ownership(); + server.app.unpause_handoff_readers(); + server.pending_handoff_repaint_nudge = true; + if let Err(err) = crate::server::handoff::report_owned(&mut received.stream) { + warn!(err = %err, "failed to report handoff ownership; continuing as owner"); + } + info!("handoff import server started"); + print_ready_message(&api::socket_path(), &client_socket_path()); + server.run().await + }); + + rt.shutdown_timeout(Duration::from_millis(100)); + crate::logging::shutdown("server"); + result +} + +#[cfg(unix)] +fn wait_for_old_public_sockets_to_close(timeout: Duration) -> io::Result<()> { + use std::os::unix::net::UnixStream; + + let deadline = Instant::now() + timeout; + let api_socket = api::socket_path(); + let client_socket = client_socket_path(); + while Instant::now() < deadline { + let api_open = api_socket.exists() && UnixStream::connect(&api_socket).is_ok(); + let client_open = client_socket.exists() && UnixStream::connect(&client_socket).is_ok(); + if !api_open && !client_open { + return Ok(()); + } + std::thread::sleep(Duration::from_millis(50)); + } + Err(io::Error::new( + io::ErrorKind::TimedOut, + "old server sockets did not close before handoff import bind", + )) +} + +#[cfg(not(unix))] +fn run_handoff_import_server(_socket_path: &Path, _token: &str) -> io::Result<()> { + Err(io::Error::other("live handoff is only supported on Unix")) +} + fn print_ready_message(api_socket: &Path, client_socket: &Path) { eprintln!("herdr server running; you can use any herdr CLI command in another terminal."); eprintln!("api socket: {}", api_socket.display()); @@ -2336,6 +2826,8 @@ mod tests { let socket_path = dir.join("client.sock"); let _ = fs::remove_file(&socket_path); let listener = UnixListener::bind(&socket_path).expect("bind test listener"); + let client_socket_identity = + socket_file_identity(&socket_path).expect("test listener socket identity"); listener .set_nonblocking(true) .expect("set listener nonblocking"); @@ -2344,8 +2836,11 @@ mod tests { HeadlessServer { app, + api_tx: None, + api_server: None, client_listener: listener, client_socket_path: socket_path, + client_socket_identity, clients: HashMap::new(), next_client_id: 1, foreground_client_id: None, @@ -2356,6 +2851,8 @@ mod tests { next_activity_stamp: 1, effective_size: (MIN_COLS, MIN_ROWS), shutting_down: false, + handoff_in_progress: false, + pending_handoff_repaint_nudge: false, should_quit: Arc::new(AtomicBool::new(false)), server_event_rx, server_event_tx, diff --git a/src/server/mod.rs b/src/server/mod.rs index b123e111..3386d0d7 100644 --- a/src/server/mod.rs +++ b/src/server/mod.rs @@ -3,6 +3,8 @@ pub(crate) mod client_accept; pub(crate) mod client_transport; pub(crate) mod clients; pub(crate) mod clipboard_image; +#[cfg(unix)] +pub(crate) mod handoff; pub mod headless; pub(crate) mod keybindings; pub(crate) mod notifications; diff --git a/src/session.rs b/src/session.rs index f00797c0..d537a0f9 100644 --- a/src/session.rs +++ b/src/session.rs @@ -100,7 +100,11 @@ pub fn local_attach_command() -> String { } pub fn local_stop_command() -> String { - match active_name() { + stop_command_for(active_name().as_deref()) +} + +pub fn stop_command_for(name: Option<&str>) -> String { + match name { Some(name) => format!("herdr session stop {name}"), None => "herdr server stop".to_string(), } diff --git a/src/terminal/runtime.rs b/src/terminal/runtime.rs index 2aa49b68..13ce75c9 100644 --- a/src/terminal/runtime.rs +++ b/src/terminal/runtime.rs @@ -19,6 +19,61 @@ impl TerminalRuntime { self.0.shutdown(); } + #[cfg(unix)] + pub fn duplicate_handoff_fd(&self) -> std::io::Result { + self.0.duplicate_handoff_fd() + } + + #[cfg(unix)] + pub fn preserve_for_handoff(self) { + self.0.preserve_for_handoff() + } + + #[cfg(unix)] + pub fn assume_handoff_ownership(&mut self) { + self.0.assume_handoff_ownership(); + } + + #[cfg(unix)] + pub fn set_handoff_reader_paused(&self, paused: bool) { + self.0.set_handoff_reader_paused(paused); + } + + #[cfg(unix)] + pub fn pause_handoff_reader(&self, timeout: std::time::Duration) -> std::io::Result<()> { + self.0.pause_handoff_reader(timeout) + } + + #[cfg(unix)] + pub fn handoff_pane(&self, pane_id: u32) -> crate::server::handoff::HandoffPane { + self.0.handoff_pane(pane_id) + } + + #[cfg(unix)] + pub fn handoff_history_ansi(&self) -> Option { + self.0.handoff_history_ansi() + } + + #[cfg(unix)] + pub fn from_handoff_fd( + import: crate::pane::PaneRuntimeImport, + scrollback_limit_bytes: usize, + host_terminal_theme: crate::terminal_theme::TerminalTheme, + events: mpsc::Sender, + render_notify: Arc, + render_dirty: Arc, + ) -> std::io::Result { + crate::pane::PaneRuntime::from_handoff_fd( + import, + scrollback_limit_bytes, + host_terminal_theme, + events, + render_notify, + render_dirty, + ) + .map(Self) + } + pub fn spawn( pane_id: PaneId, rows: u16, @@ -172,6 +227,10 @@ impl TerminalRuntime { self.0.resize(rows, cols, cell_width_px, cell_height_px); } + pub fn nudge_child_redraw_after_handoff(&self) { + self.0.nudge_child_redraw_after_handoff(); + } + pub fn scroll_up(&self, lines: usize) { self.0.scroll_up(lines); } diff --git a/src/terminal/runtime_registry.rs b/src/terminal/runtime_registry.rs index d1aeeae2..a909985e 100644 --- a/src/terminal/runtime_registry.rs +++ b/src/terminal/runtime_registry.rs @@ -37,10 +37,43 @@ impl TerminalRuntimeRegistry { self.runtimes.values() } + #[cfg(unix)] + pub(crate) fn iter(&self) -> impl Iterator { + self.runtimes.iter() + } + + #[cfg(unix)] + pub(crate) fn set_handoff_readers_paused(&self, paused: bool) { + for runtime in self.runtimes.values() { + runtime.set_handoff_reader_paused(paused); + } + } + + #[cfg(unix)] + pub(crate) fn assume_handoff_ownership(&mut self) { + for runtime in self.runtimes.values_mut() { + runtime.assume_handoff_ownership(); + } + } + pub(crate) fn len(&self) -> usize { self.runtimes.len() } + #[cfg(unix)] + pub(crate) fn nudge_child_redraw_after_handoff(&self) { + for runtime in self.runtimes.values() { + runtime.nudge_child_redraw_after_handoff(); + } + } + + #[cfg(unix)] + pub(crate) fn drain_for_handoff( + &mut self, + ) -> impl Iterator + '_ { + self.runtimes.drain() + } + #[cfg(test)] pub(crate) fn drain(&mut self) -> impl Iterator + '_ { self.runtimes.drain() diff --git a/src/update.rs b/src/update.rs index 2a2776fa..27039ac1 100644 --- a/src/update.rs +++ b/src/update.rs @@ -27,7 +27,7 @@ const FAKE_UPDATE_VERSION_ENV: &str = "HERDR_FAKE_UPDATE_VERSION"; const FAKE_UPDATE_NOTES_VERSION_ENV: &str = "HERDR_FAKE_UPDATE_NOTES_VERSION"; const DEFAULT_FAKE_UPDATE_NOTES_VERSION: &str = "0.3.0"; const SERVER_STOP_RESPONSE_TIMEOUT: Duration = Duration::from_secs(5); -const SERVER_SHUTDOWN_CONFIRM_TIMEOUT: Duration = Duration::from_secs(5); +const SERVER_HANDOFF_CONFIRM_TIMEOUT: Duration = Duration::from_secs(30); const SERVER_SHUTDOWN_POLL_INTERVAL: Duration = Duration::from_millis(100); const STAR_PROMPT_REPO: &str = "ogulcancelik/herdr"; const STAR_PROMPT_STATE_FILE: &str = "github-star-prompt.json"; @@ -400,18 +400,6 @@ fn running_inside_herdr() -> bool { running_inside_herdr_env(env::var(crate::HERDR_ENV_VAR).ok().as_deref()) } -fn api_server_is_running_at(socket_path: &Path) -> bool { - if !socket_path.exists() { - return false; - } - - UnixStream::connect(socket_path).is_ok() -} - -fn api_server_is_running() -> bool { - api_server_is_running_at(&crate::api::socket_path()) -} - fn client_protocol_server_is_running_at(socket_path: &Path) -> bool { if !socket_path.exists() { return false; @@ -424,11 +412,6 @@ fn client_protocol_server_is_running() -> bool { client_protocol_server_is_running_at(&crate::server::socket_paths::client_socket_path()) } -fn read_running_server_info() -> Result, String> { - crate::api::read_runtime_status_at(&crate::api::socket_path(), SERVER_STOP_RESPONSE_TIMEOUT) - .map_err(|e| format!("failed to read running server status: {e}")) -} - fn protocol_label(protocol: Option) -> String { protocol .map(|value| value.to_string()) @@ -439,14 +422,21 @@ fn version_label(version: Option<&str>) -> &str { version.unwrap_or("unknown") } -fn update_requires_server_stop(server: &crate::api::RuntimeStatus, release: &ReleaseInfo) -> bool { +fn update_requires_live_handoff(server: &crate::api::RuntimeStatus, release: &ReleaseInfo) -> bool { match (server.protocol, release.target_protocol) { (Some(server_protocol), Some(target_protocol)) => server_protocol != target_protocol, _ => true, } } -fn parse_stop_server_before_update_response(input: &str) -> Option { +fn server_supports_live_handoff(server: &crate::api::RuntimeStatus) -> bool { + server + .capabilities + .as_ref() + .is_some_and(|capabilities| capabilities.live_handoff) +} + +fn parse_live_handoff_before_update_response(input: &str) -> Option { let trimmed = input.trim().to_ascii_lowercase(); match trimmed.as_str() { "" | "y" | "yes" => Some(true), @@ -455,53 +445,391 @@ fn parse_stop_server_before_update_response(input: &str) -> Option { } } -fn prompt_to_stop_server_before_update( - server: &crate::api::RuntimeStatus, - release: &ReleaseInfo, - requires_stop: bool, -) -> Result { - if !io::stdin().is_terminal() { - if requires_stop { - return Err(format!( - "a herdr server is running and updating to v{} requires stopping it; run `{}`, then run `herdr update` again", - release.version, - crate::session::local_stop_command() - )); - } +#[derive(Debug, Clone, PartialEq, Eq)] +struct RunningServerUpdatePlan { + target: RunningUpdateTarget, + server: crate::api::RuntimeStatus, + requires_live_handoff: bool, +} - eprintln!( - "a herdr server is running. updating the binary will not affect that server until it restarts." - ); - return Ok(false); +impl RunningServerUpdatePlan { + fn label(&self) -> &str { + &self.target.label + } + + fn socket_path(&self) -> &Path { + &self.target.socket_path + } + + fn stop_command(&self) -> String { + self.target.stop_command.clone() + } + + fn attach_command(&self) -> Option { + self.target.attach_command.clone() + } + + fn target_noun(&self) -> &'static str { + if self.target.attach_command.is_some() { + "session" + } else { + "server" + } + } +} + +#[derive(Debug, Clone, PartialEq, Eq)] +struct RunningServerUpdateDecision { + plan: RunningServerUpdatePlan, + action: RunningServerUpdateAction, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum RunningServerUpdateAction { + None, + LiveHandoff, + StopOldServer, +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum RunningServerUpdateOutcome { + RestartDeferred, + Stopped, + LiveHandoffComplete, + FailedHandoffOldServerKept, + FailedHandoffOldServerStopped, + FailedHandoffNoServer, + FailedHandoffUnknown, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +struct RunningSessionUpdateOutcome { + session_label: String, + stop_command: String, + attach_command: Option, + outcome: RunningServerUpdateOutcome, +} + +#[derive(Debug, Clone, PartialEq, Eq)] +enum FailedHandoffServerState { + UpdatedServerRunning, + OldServerRunning(crate::api::RuntimeStatus), + NoServerResponding, + Unknown(String), +} + +fn plan_running_server_updates( + release: &ReleaseInfo, +) -> Result, String> { + let targets = running_update_targets()?; + let mut plans = Vec::new(); + + for target in targets { + let server = match crate::api::read_runtime_status_at( + &target.socket_path, + SERVER_STOP_RESPONSE_TIMEOUT, + ) + .map_err(|err| { + format!( + "failed to read status for herdr target {} at {}: {err}. stop it with `{}` and run `herdr update` again", + target.label, + target.socket_path.display(), + target.stop_command + ) + })? { + Some(server) => server, + None if target.must_be_running => { + return Err(format!( + "herdr target {} looked running, but its status API did not respond at {}. stop it with `{}` and run `herdr update` again", + target.label, + target.socket_path.display(), + target.stop_command + )); + } + None if client_protocol_server_is_running_at(&target.client_socket_path) => { + return Err(format!( + "herdr target {} has a client socket, but its status API did not respond at {}. stop it with `{}` and run `herdr update` again", + target.label, + target.socket_path.display(), + target.stop_command + )); + } + None => continue, + }; + + plans.push(RunningServerUpdatePlan { + requires_live_handoff: update_requires_live_handoff(&server, release), + server, + target, + }); + } + + if plans.is_empty() && target_client_protocol_server_is_running()? { + return Err(format!( + "a herdr server is listening, but its status API is unavailable; try `{}`, or stop the old server process manually, then run `herdr update` again", + crate::session::local_stop_command() + )); + } + + Ok(plans) +} + +#[derive(Debug, Clone, PartialEq, Eq)] +struct RunningUpdateTarget { + name: Option, + label: String, + stop_command: String, + attach_command: Option, + socket_path: PathBuf, + client_socket_path: PathBuf, + must_be_running: bool, +} + +fn running_update_targets() -> Result, String> { + if crate::session::explicit_session_requested() { + return Ok(vec![RunningUpdateTarget { + name: crate::session::active_name(), + label: crate::session::active_name() + .unwrap_or_else(|| crate::session::DEFAULT_SESSION_NAME.to_string()), + stop_command: crate::session::local_stop_command(), + attach_command: Some(crate::session::local_attach_command()), + socket_path: crate::api::socket_path(), + client_socket_path: crate::server::socket_paths::client_socket_path(), + must_be_running: false, + }]); + } + + if let Some(socket_path) = std::env::var_os(crate::api::SOCKET_PATH_ENV_VAR) { + let socket_path = PathBuf::from(socket_path); + return Ok(vec![RunningUpdateTarget { + name: None, + label: socket_path.display().to_string(), + stop_command: format!( + "{}={} herdr server stop", + crate::api::SOCKET_PATH_ENV_VAR, + socket_path.display() + ), + attach_command: None, + client_socket_path: crate::server::socket_paths::client_socket_path_from_overrides( + Some(&socket_path.to_string_lossy()), + None, + ), + socket_path, + must_be_running: false, + }]); + } + + let sessions = crate::session::list_sessions() + .map_err(|err| format!("failed to list herdr sessions: {err}"))?; + Ok(sessions + .into_iter() + .map(|session| RunningUpdateTarget { + name: if session.default { + None + } else { + Some(session.name.clone()) + }, + stop_command: crate::session::stop_command_for(if session.default { + None + } else { + Some(&session.name) + }), + attach_command: Some(if session.default { + "herdr".to_string() + } else { + format!("herdr session attach {}", session.name) + }), + label: session.name.clone(), + client_socket_path: crate::session::client_socket_path_for(if session.default { + None + } else { + Some(&session.name) + }), + socket_path: PathBuf::from(session.socket_path), + must_be_running: session.running, + }) + .collect()) +} + +fn target_client_protocol_server_is_running() -> Result { + if crate::session::explicit_session_requested() + || std::env::var_os(crate::api::SOCKET_PATH_ENV_VAR).is_some() + { + return Ok(client_protocol_server_is_running()); + } + + let sessions = crate::session::list_sessions() + .map_err(|err| format!("failed to list herdr sessions: {err}"))?; + Ok(sessions.into_iter().any(|session| { + let client_socket = crate::session::client_socket_path_for(if session.default { + None + } else { + Some(&session.name) + }); + client_protocol_server_is_running_at(&client_socket) + })) +} + +fn prompt_to_stop_old_servers_before_update( + plans: &[RunningServerUpdatePlan], + release: &ReleaseInfo, +) -> Result { + if !io::stdin().is_terminal() { + return Err( + "one or more herdr targets are running and cannot perform live handoff for this update; run `herdr update` from an interactive terminal, or stop those targets and run `herdr update` again" + .to_string(), + ); } - eprintln!("a herdr server is currently running:"); eprintln!( - " server: v{} protocol {}", - version_label(server.version.as_deref()), - protocol_label(server.protocol) + "these running herdr targets are too old to preserve panes during this update to v{}:", + release.version ); + for plan in plans { + eprintln!( + " {}: v{} protocol {}", + plan.label(), + version_label(plan.server.version.as_deref()), + protocol_label(plan.server.protocol) + ); + } + eprintln!( + "herdr can leave them running, or stop them after installing the update. stopping them will exit their pane processes." + ); + + loop { + eprint!("stop these old targets after updating? [y/N] "); + io::stderr() + .flush() + .map_err(|e| format!("failed to flush prompt: {e}"))?; + + let mut input = String::new(); + let read = io::stdin() + .read_line(&mut input) + .map_err(|e| format!("failed to read prompt response: {e}"))?; + if read == 0 { + return Ok(false); + } + + match input.trim().to_ascii_lowercase().as_str() { + "y" | "yes" => return Ok(true), + "" | "n" | "no" => return Ok(false), + _ => eprintln!("please answer y or n"), + } + } +} + +fn confirm_running_server_update_action( + plans: Vec, + release: &ReleaseInfo, +) -> Result, String> { + if plans.is_empty() { + return Ok(Vec::new()); + } + + print_running_session_update_summary(&plans, release); + + let handoff_supported: Vec<&RunningServerUpdatePlan> = plans + .iter() + .filter(|plan| server_supports_live_handoff(&plan.server)) + .collect(); + let handoff_unsupported_requiring_update: Vec<&RunningServerUpdatePlan> = plans + .iter() + .filter(|plan| !server_supports_live_handoff(&plan.server) && plan.requires_live_handoff) + .collect(); + + let live_handoff = if handoff_supported.is_empty() { + false + } else { + prompt_to_live_handoff_sessions_before_update( + &handoff_supported, + release, + plans.iter().any(|plan| plan.requires_live_handoff), + )? + }; + + let stop_unsupported = if handoff_unsupported_requiring_update.is_empty() { + false + } else { + let owned: Vec = handoff_unsupported_requiring_update + .iter() + .map(|plan| (*plan).clone()) + .collect(); + prompt_to_stop_old_servers_before_update(&owned, release)? + }; + + let mut decisions = Vec::new(); + for plan in plans { + let action = if server_supports_live_handoff(&plan.server) && live_handoff { + RunningServerUpdateAction::LiveHandoff + } else if !server_supports_live_handoff(&plan.server) + && plan.requires_live_handoff + && stop_unsupported + { + RunningServerUpdateAction::StopOldServer + } else { + RunningServerUpdateAction::None + }; + decisions.push(RunningServerUpdateDecision { plan, action }); + } + + Ok(decisions) +} + +fn print_running_session_update_summary(plans: &[RunningServerUpdatePlan], release: &ReleaseInfo) { + eprintln!("running herdr targets:"); + for plan in plans { + let capability = if server_supports_live_handoff(&plan.server) { + "handoff supported" + } else { + "too old for handoff" + }; + eprintln!( + " {}: v{} protocol {} ({})", + plan.label(), + version_label(plan.server.version.as_deref()), + protocol_label(plan.server.protocol), + capability + ); + } eprintln!( " update: v{} protocol {}", release.version, protocol_label(release.target_protocol) ); eprintln!(); +} - if requires_stop { +fn prompt_to_live_handoff_sessions_before_update( + plans: &[&RunningServerUpdatePlan], + release: &ReleaseInfo, + requires_live_handoff: bool, +) -> Result { + if !io::stdin().is_terminal() { + if requires_live_handoff { + return Err(format!( + "one or more herdr targets are running and updating to v{} requires live server handoff; run `herdr update` from an interactive terminal, or stop those targets and run `herdr update` again", + release.version + )); + } eprintln!( - "this update changes the herdr client/server protocol. the running server must be stopped before the new client can attach." + "herdr targets are running. updating the binary will not affect them until they restart." ); - eprintln!("stopping the server will end the current herdr session and its panes."); - } else { - eprintln!("updating the binary will not affect the running server until it restarts."); + return Ok(false); } + eprintln!( + "herdr can hand off {} running target{} to the new server so pane processes keep running.", + plans.len(), + if plans.len() == 1 { "" } else { "s" } + ); + eprintln!("connected clients will disconnect during handoff and can attach again afterward."); + loop { - let prompt = if requires_stop { - "stop the server and continue updating? [Y/n] " + let prompt = if requires_live_handoff { + "update and live-handoff supported targets to the new server? [Y/n] " } else { - "stop the server before updating? [Y/n] " + "live-handoff supported running targets after updating? [Y/n] " }; eprint!("{prompt}"); io::stderr() @@ -516,7 +844,7 @@ fn prompt_to_stop_server_before_update( return Ok(false); } - if let Some(answer) = parse_stop_server_before_update_response(&input) { + if let Some(answer) = parse_live_handoff_before_update_response(&input) { return Ok(answer); } @@ -524,107 +852,288 @@ fn prompt_to_stop_server_before_update( } } -#[derive(Debug, Clone, PartialEq, Eq)] -struct RunningServerUpdatePlan { - server: crate::api::RuntimeStatus, - requires_stop: bool, +fn live_handoff_running_server_for_update( + plan: &RunningServerUpdatePlan, + release: &ReleaseInfo, + updated_exe: &Path, +) -> Result<(), String> { + eprintln!( + "asking {} {} to hand off live panes to the updated server...", + plan.target_noun(), + plan.label() + ); + live_handoff_server_via_api_for_update_at(plan.socket_path(), updated_exe, release)?; + wait_for_server_handoff_at(plan.socket_path(), SERVER_HANDOFF_CONFIRM_TIMEOUT, release)?; + eprintln!( + "live handoff complete for {} {}; pane processes should still be running.", + plan.target_noun(), + plan.label() + ); + Ok(()) } -fn plan_running_server_update( +fn runtime_matches_release(status: &crate::api::RuntimeStatus, release: &ReleaseInfo) -> bool { + let protocol_matches = release + .target_protocol + .is_none_or(|protocol| status.protocol == Some(protocol)); + let version_matches = status.version.as_deref() == Some(&release.version.to_string()); + protocol_matches && version_matches +} + +fn classify_failed_live_handoff_state_at( + socket_path: &Path, release: &ReleaseInfo, -) -> Result, String> { - let Some(server) = read_running_server_info()? else { - if client_protocol_server_is_running() { - return Err(format!( - "a herdr server is listening, but its status API is unavailable; try `{}`, or stop the old server process manually, then run `herdr update` again", - crate::session::local_stop_command() - )); +) -> FailedHandoffServerState { + match crate::api::read_runtime_status_at(socket_path, SERVER_STOP_RESPONSE_TIMEOUT) { + Ok(Some(status)) if runtime_matches_release(&status, release) => { + FailedHandoffServerState::UpdatedServerRunning } - return Ok(None); - }; - - let requires_stop = update_requires_server_stop(&server, release); - Ok(Some(RunningServerUpdatePlan { - server, - requires_stop, - })) + Ok(Some(status)) => FailedHandoffServerState::OldServerRunning(status), + Ok(None) => FailedHandoffServerState::NoServerResponding, + Err(err) => FailedHandoffServerState::Unknown(err.to_string()), + } } -fn stop_running_server_for_update( - plan: Option<&RunningServerUpdatePlan>, +fn prompt_to_stop_old_server_after_failed_handoff( + plan: &RunningServerUpdatePlan, release: &ReleaseInfo, + status: &crate::api::RuntimeStatus, ) -> Result { - let Some(plan) = plan else { - return Ok(false); - }; + eprintln!( + "live handoff failed, but {} {} is still running with your panes.", + plan.target_noun(), + plan.label() + ); + eprintln!( + " server: v{} protocol {}", + version_label(status.version.as_deref()), + protocol_label(status.protocol) + ); + eprintln!( + " installed: v{} protocol {}", + release.version, + protocol_label(release.target_protocol) + ); + eprintln!( + "you can keep using the old server, or stop it now so the next `herdr` start uses v{}.", + release.version + ); + eprintln!("stopping the old server will exit its pane processes."); - let stop_server = - prompt_to_stop_server_before_update(&plan.server, release, plan.requires_stop)?; - if !stop_server { - if plan.requires_stop { - return Err(format!( - "update cancelled; stop the running herdr server with `{}`, then run `herdr update` again", - crate::session::local_stop_command() - )); - } + if !io::stdin().is_terminal() { + eprintln!( + "not stopping the old server from a non-interactive update; run `{}` when you are ready.", + plan.stop_command() + ); return Ok(false); } - stop_server_via_api()?; - wait_for_server_shutdown(SERVER_SHUTDOWN_CONFIRM_TIMEOUT)?; - eprintln!("stopped the running herdr server."); - Ok(true) + loop { + eprint!("stop the old server now? [y/N] "); + io::stderr() + .flush() + .map_err(|e| format!("failed to flush prompt: {e}"))?; + + let mut input = String::new(); + let read = io::stdin() + .read_line(&mut input) + .map_err(|e| format!("failed to read prompt response: {e}"))?; + if read == 0 { + return Ok(false); + } + + match input.trim().to_ascii_lowercase().as_str() { + "y" | "yes" => return Ok(true), + "" | "n" | "no" => return Ok(false), + _ => eprintln!("please answer y or n"), + } + } +} + +fn recover_failed_live_handoff_for_update( + plan: &RunningServerUpdatePlan, + release: &ReleaseInfo, + error: &str, +) -> Result { + eprintln!( + "live handoff failed for {} {}: {error}", + plan.target_noun(), + plan.label() + ); + + match classify_failed_live_handoff_state_at(plan.socket_path(), release) { + FailedHandoffServerState::UpdatedServerRunning => { + eprintln!( + "the updated server is running for {} {}.", + plan.target_noun(), + plan.label() + ); + Ok(RunningServerUpdateOutcome::LiveHandoffComplete) + } + FailedHandoffServerState::OldServerRunning(status) => { + if prompt_to_stop_old_server_after_failed_handoff(plan, release, &status)? { + stop_running_server_for_update(plan)?; + Ok(RunningServerUpdateOutcome::FailedHandoffOldServerStopped) + } else { + Ok(RunningServerUpdateOutcome::FailedHandoffOldServerKept) + } + } + FailedHandoffServerState::NoServerResponding => { + if let Some(command) = plan.attach_command() { + eprintln!( + "no herdr server is responding for session {}. the binary was updated; run `{command}` to start v{}.", + plan.label(), + release.version + ); + } else { + eprintln!( + "no herdr server is responding at {}. the binary was updated; restart with the same socket override to use v{}.", + plan.socket_path().display(), + release.version + ); + } + Ok(RunningServerUpdateOutcome::FailedHandoffNoServer) + } + FailedHandoffServerState::Unknown(status_error) => { + eprintln!( + "herdr could not determine server state for {} {} after the failed handoff: {status_error}", + plan.target_noun(), + plan.label() + ); + eprintln!("{}", reconnect_or_stop_guidance(plan)); + Ok(RunningServerUpdateOutcome::FailedHandoffUnknown) + } + } +} + +fn reconnect_or_stop_guidance(plan: &RunningServerUpdatePlan) -> String { + if let Some(command) = plan.attach_command() { + format!( + "if `{command}` does not reconnect cleanly, stop the old server with `{}` and run `{command}` again.", + plan.stop_command() + ) + } else { + format!( + "if reconnecting with the same socket override does not work, stop the old server with `{}`.", + plan.stop_command() + ) + } } fn stop_server_via_api_at(socket_path: &Path, timeout: Duration) -> Result<(), String> { - use crate::api::schema::{EmptyParams, Method, Request}; + use crate::api::schema::{EmptyParams, Method}; + + send_server_update_method_at( + socket_path, + timeout, + "update:server:stop", + Method::ServerStop(EmptyParams::default()), + "server stop", + ) +} + +fn send_server_update_method_at( + socket_path: &Path, + timeout: Duration, + request_id: &str, + method: crate::api::schema::Method, + error_prefix: &str, +) -> Result<(), String> { + use crate::api::schema::Request; let request = Request { - id: "update:server:stop".into(), - method: Method::ServerStop(EmptyParams::default()), + id: request_id.into(), + method, }; let mut stream = UnixStream::connect(socket_path) .map_err(|e| format!("failed to connect to running server: {e}"))?; stream .set_write_timeout(Some(timeout)) - .map_err(|e| format!("failed to set server stop write timeout: {e}"))?; + .map_err(|e| format!("failed to set {error_prefix} write timeout: {e}"))?; stream .set_read_timeout(Some(timeout)) - .map_err(|e| format!("failed to set server stop read timeout: {e}"))?; + .map_err(|e| format!("failed to set {error_prefix} read timeout: {e}"))?; stream .write_all( serde_json::to_string(&request) .map_err(|e| e.to_string())? .as_bytes(), ) - .map_err(|e| format!("failed to send server stop request: {e}"))?; + .map_err(|e| format!("failed to send {error_prefix} request: {e}"))?; stream .write_all(b"\n") - .map_err(|e| format!("failed to finish server stop request: {e}"))?; + .map_err(|e| format!("failed to finish {error_prefix} request: {e}"))?; stream .flush() - .map_err(|e| format!("failed to flush server stop request: {e}"))?; + .map_err(|e| format!("failed to flush {error_prefix} request: {e}"))?; let mut reader = BufReader::new(stream); let mut line = String::new(); let read = reader .read_line(&mut line) - .map_err(|e| format!("failed to read server stop response: {e}"))?; + .map_err(|e| format!("failed to read {error_prefix} response: {e}"))?; if read == 0 || line.trim().is_empty() { - return Err("empty server stop response".into()); + return Err(format!("empty {error_prefix} response")); } let response: serde_json::Value = serde_json::from_str(&line).map_err(|e| format!("invalid server response: {e}"))?; if let Some(error) = response.get("error") { - return Err(format!("server stop failed: {error}")); + return Err(format!("{error_prefix} failed: {error}")); } Ok(()) } -fn stop_server_via_api() -> Result<(), String> { - stop_server_via_api_at(&crate::api::socket_path(), SERVER_STOP_RESPONSE_TIMEOUT) +#[cfg(test)] +fn live_handoff_server_via_api_at(socket_path: &Path, timeout: Duration) -> Result<(), String> { + use crate::api::schema::{Method, ServerLiveHandoffParams}; + + let params = ServerLiveHandoffParams::default(); + + send_server_update_method_at( + socket_path, + timeout, + "update:server:live-handoff", + Method::ServerLiveHandoff(params), + "server live handoff", + ) +} + +fn live_handoff_server_via_api_for_release_at( + socket_path: &Path, + timeout: Duration, + updated_exe: &Path, + release: &ReleaseInfo, +) -> Result<(), String> { + use crate::api::schema::{Method, ServerLiveHandoffParams}; + + let params = ServerLiveHandoffParams { + import_exe: Some(updated_exe.display().to_string()), + expected_protocol: release.target_protocol, + expected_version: Some(release.version.to_string()), + }; + + send_server_update_method_at( + socket_path, + timeout, + "update:server:live-handoff", + Method::ServerLiveHandoff(params), + "server live handoff", + ) +} + +fn live_handoff_server_via_api_for_update_at( + socket_path: &Path, + updated_exe: &Path, + release: &ReleaseInfo, +) -> Result<(), String> { + live_handoff_server_via_api_for_release_at( + socket_path, + SERVER_HANDOFF_CONFIRM_TIMEOUT, + updated_exe, + release, + ) } fn server_shutdown_confirmed_at(socket_path: &Path) -> Result { @@ -668,8 +1177,173 @@ fn wait_for_server_shutdown_at(socket_path: &Path, timeout: Duration) -> Result< } } -fn wait_for_server_shutdown(timeout: Duration) -> Result<(), String> { - wait_for_server_shutdown_at(&crate::api::socket_path(), timeout) +fn stop_running_server_for_update(plan: &RunningServerUpdatePlan) -> Result<(), String> { + eprintln!("stopping herdr {} {}...", plan.target_noun(), plan.label()); + stop_server_via_api_at(plan.socket_path(), SERVER_STOP_RESPONSE_TIMEOUT)?; + wait_for_server_shutdown_at(plan.socket_path(), SERVER_HANDOFF_CONFIRM_TIMEOUT)?; + Ok(()) +} + +fn wait_for_server_handoff_at( + socket_path: &Path, + timeout: Duration, + release: &ReleaseInfo, +) -> Result<(), String> { + wait_for_running_server_protocol_at( + socket_path, + timeout, + release.target_protocol, + Some(&release.version.to_string()), + ) +} + +fn wait_for_running_server_protocol_at( + socket_path: &Path, + timeout: Duration, + expected_protocol: Option, + expected_version: Option<&str>, +) -> Result<(), String> { + let deadline = Instant::now() + timeout; + loop { + if let Some(status) = + crate::api::read_runtime_status_at(socket_path, SERVER_STOP_RESPONSE_TIMEOUT) + .map_err(|e| format!("failed to read server status after handoff: {e}"))? + { + let protocol_matches = + expected_protocol.is_none_or(|protocol| status.protocol == Some(protocol)); + let version_matches = + expected_version.is_none_or(|version| status.version.as_deref() == Some(version)); + if protocol_matches && version_matches { + return Ok(()); + } + } + if Instant::now() >= deadline { + return Err(format!( + "live handoff was requested, but no compatible server responded on {} after {} seconds", + socket_path.display(), + timeout.as_secs() + )); + } + std::thread::sleep(SERVER_SHUTDOWN_POLL_INTERVAL); + } +} + +fn apply_running_session_update_decisions( + release: &ReleaseInfo, + updated_exe: &Path, + decisions: Vec, +) -> Result, String> { + let mut outcomes = Vec::new(); + + for decision in decisions { + let outcome = match decision.action { + RunningServerUpdateAction::None => RunningServerUpdateOutcome::RestartDeferred, + RunningServerUpdateAction::StopOldServer => { + stop_running_server_for_update(&decision.plan)?; + RunningServerUpdateOutcome::Stopped + } + RunningServerUpdateAction::LiveHandoff => { + match live_handoff_running_server_for_update(&decision.plan, release, updated_exe) { + Ok(()) => RunningServerUpdateOutcome::LiveHandoffComplete, + Err(err) => { + recover_failed_live_handoff_for_update(&decision.plan, release, &err)? + } + } + } + }; + + let stop_command = decision.plan.stop_command(); + outcomes.push(RunningSessionUpdateOutcome { + session_label: decision.plan.label().to_string(), + stop_command, + attach_command: decision.plan.attach_command(), + outcome, + }); + } + + Ok(outcomes) +} + +fn print_running_session_update_outcomes( + outcomes: &[RunningSessionUpdateOutcome], + release: &ReleaseInfo, +) { + if outcomes.is_empty() { + eprintln!("run herdr again."); + return; + } + + for outcome in outcomes { + match outcome.outcome { + RunningServerUpdateOutcome::LiveHandoffComplete => { + if let Some(command) = &outcome.attach_command { + eprintln!( + "session {} was replaced; reconnect clients with `{command}`.", + outcome.session_label + ); + } else { + eprintln!( + "server {} was replaced; reconnect using the same socket override.", + outcome.session_label + ); + } + } + RunningServerUpdateOutcome::RestartDeferred => { + if outcome.attach_command.is_some() { + eprintln!( + "session {} will use v{} after it restarts.", + outcome.session_label, release.version + ); + } else { + eprintln!( + "server {} will use v{} after it restarts.", + outcome.session_label, release.version + ); + } + } + RunningServerUpdateOutcome::Stopped + | RunningServerUpdateOutcome::FailedHandoffOldServerStopped + | RunningServerUpdateOutcome::FailedHandoffNoServer => { + if let Some(command) = &outcome.attach_command { + eprintln!( + "session {} is stopped; run `{command}` again.", + outcome.session_label + ); + } else { + eprintln!( + "server {} is stopped; restart it with the same socket override.", + outcome.session_label + ); + } + } + RunningServerUpdateOutcome::FailedHandoffOldServerKept => { + if outcome.attach_command.is_some() { + eprintln!( + "session {} is still running with your panes; stop it later with `{}` to use v{}.", + outcome.session_label, outcome.stop_command, release.version + ); + } else { + eprintln!( + "server {} is still running with your panes; stop it later with `{}` to use v{}.", + outcome.session_label, outcome.stop_command, release.version + ); + } + } + RunningServerUpdateOutcome::FailedHandoffUnknown => { + if let Some(command) = &outcome.attach_command { + eprintln!( + "session {} state is unclear; run `{command}`, or stop the old server with `{}` if reconnect fails.", + outcome.session_label, outcome.stop_command + ); + } else { + eprintln!( + "server {} state is unclear; reconnect with the same socket override, or stop it with `{}` if reconnect fails.", + outcome.session_label, outcome.stop_command + ); + } + } + } + } } // --------------------------------------------------------------------------- @@ -1034,7 +1708,9 @@ pub fn self_update() -> Result { } }; - let running_server_plan = plan_running_server_update(&release)?; + let running_server_plans = plan_running_server_updates(&release)?; + let server_update_decisions = + confirm_running_server_update_action(running_server_plans, &release)?; eprintln!("downloading v{}...", release.version); if let Err(e) = @@ -1044,18 +1720,13 @@ pub fn self_update() -> Result { } let downloaded_update = download_update(&release)?; let updated_exe = downloaded_update.current_exe.clone(); - let stopped_server = stop_running_server_for_update(running_server_plan.as_ref(), &release)?; install_downloaded_update(downloaded_update)?; + let server_update_outcomes = + apply_running_session_update_decisions(&release, &updated_exe, server_update_decisions)?; eprintln!("updated to v{}", release.version); print_outdated_integration_notice_with_updated_binary(&updated_exe); - if stopped_server { - eprintln!("run herdr again to start the updated server."); - } else if api_server_is_running() { - eprintln!("the running herdr server will use the new version after it restarts."); - } else { - eprintln!("run herdr again."); - } + print_running_session_update_outcomes(&server_update_outcomes, &release); maybe_offer_star_after_successful_update(); @@ -1259,6 +1930,46 @@ mod tests { (running, handle) } + fn spawn_status_server_once( + path: &Path, + version: &str, + protocol: u32, + ) -> thread::JoinHandle<()> { + let listener = UnixListener::bind(path).unwrap(); + let version = version.to_string(); + thread::spawn(move || { + let (mut stream, _) = listener.accept().unwrap(); + let mut request = String::new(); + BufReader::new(stream.try_clone().unwrap()) + .read_line(&mut request) + .unwrap(); + assert!(request.contains("\"method\":\"ping\"")); + let response = format!( + r#"{{"id":"runtime:status","result":{{"type":"pong","version":"{version}","protocol":{protocol},"capabilities":{{"live_handoff":true}}}}}}"# + ); + stream.write_all(response.as_bytes()).unwrap(); + stream.write_all(b"\n").unwrap(); + stream.flush().unwrap(); + }) + } + + fn fake_release(version: &str, target_protocol: Option) -> ReleaseInfo { + ReleaseInfo { + version: Version::parse(version).unwrap(), + target_protocol, + download_url: "https://example.com/herdr".to_string(), + notes_body: "### Changed\n- One".to_string(), + } + } + + fn set_test_config_home(name: &str) -> PathBuf { + let dir = std::env::temp_dir().join(format!("herdr-update-{name}-{}", std::process::id())); + let _ = fs::remove_dir_all(&dir); + fs::create_dir_all(&dir).unwrap(); + std::env::set_var("XDG_CONFIG_HOME", &dir); + dir + } + #[test] fn parse_version_basic() { assert_eq!( @@ -1348,6 +2059,7 @@ mod tests { #[test] fn fake_release_notes_default_to_real_large_changelog_section() { + let _guard = env_lock().lock().unwrap(); std::env::remove_var(FAKE_UPDATE_NOTES_VERSION_ENV); let body = fake_release_notes_body("9.4.9"); @@ -1357,6 +2069,7 @@ mod tests { #[test] fn fake_release_notes_fallback_include_version_and_context() { + let _guard = env_lock().lock().unwrap(); std::env::set_var(FAKE_UPDATE_NOTES_VERSION_ENV, "does-not-exist"); let body = fake_release_notes_body("9.4.9"); @@ -1375,21 +2088,22 @@ mod tests { } #[test] - fn parse_stop_server_before_update_response_defaults_yes_for_blank() { - assert_eq!(parse_stop_server_before_update_response(""), Some(true)); - assert_eq!(parse_stop_server_before_update_response("\n"), Some(true)); - assert_eq!(parse_stop_server_before_update_response("y"), Some(true)); - assert_eq!(parse_stop_server_before_update_response("yes"), Some(true)); - assert_eq!(parse_stop_server_before_update_response("n"), Some(false)); - assert_eq!(parse_stop_server_before_update_response("no"), Some(false)); - assert_eq!(parse_stop_server_before_update_response("later"), None); + fn parse_live_handoff_before_update_response_defaults_yes_for_blank() { + assert_eq!(parse_live_handoff_before_update_response(""), Some(true)); + assert_eq!(parse_live_handoff_before_update_response("\n"), Some(true)); + assert_eq!(parse_live_handoff_before_update_response("y"), Some(true)); + assert_eq!(parse_live_handoff_before_update_response("yes"), Some(true)); + assert_eq!(parse_live_handoff_before_update_response("n"), Some(false)); + assert_eq!(parse_live_handoff_before_update_response("no"), Some(false)); + assert_eq!(parse_live_handoff_before_update_response("later"), None); } #[test] - fn update_requires_server_stop_when_target_protocol_differs_or_unknown() { + fn update_requires_live_handoff_when_target_protocol_differs_or_unknown() { let server = crate::api::RuntimeStatus { version: Some("0.5.5".to_string()), protocol: Some(2), + capabilities: None, }; let compatible_release = ReleaseInfo { version: Version::parse("0.5.6").unwrap(), @@ -1406,13 +2120,168 @@ mod tests { ..compatible_release.clone() }; - assert!(!update_requires_server_stop(&server, &compatible_release)); - assert!(update_requires_server_stop(&server, &incompatible_release)); - assert!(update_requires_server_stop(&server, &unknown_release)); + assert!(!update_requires_live_handoff(&server, &compatible_release)); + assert!(update_requires_live_handoff(&server, &incompatible_release)); + assert!(update_requires_live_handoff(&server, &unknown_release)); } #[test] - fn noninteractive_update_requires_stop_names_session_stop_command() { + fn plain_update_targets_all_running_sessions() { + let _guard = env_lock().lock().unwrap(); + let config_home = set_test_config_home("all-sessions"); + std::env::remove_var(crate::api::SOCKET_PATH_ENV_VAR); + std::env::remove_var(crate::session::SESSION_ENV_VAR); + crate::session::clear_explicit_session_for_test(); + + let default_socket = crate::session::api_socket_path_for(None); + let work_socket = crate::session::api_socket_path_for(Some("work")); + fs::create_dir_all(default_socket.parent().unwrap()).unwrap(); + fs::create_dir_all(work_socket.parent().unwrap()).unwrap(); + let default_listener = UnixListener::bind(&default_socket).unwrap(); + let work_listener = UnixListener::bind(&work_socket).unwrap(); + + let mut targets = running_update_targets().unwrap(); + targets.sort_by(|left, right| left.label.cmp(&right.label)); + + drop(default_listener); + drop(work_listener); + let _ = fs::remove_dir_all(config_home); + std::env::remove_var("XDG_CONFIG_HOME"); + + assert_eq!(targets.len(), 2); + assert_eq!(targets[0].label, crate::session::DEFAULT_SESSION_NAME); + assert_eq!(targets[0].name, None); + assert_eq!(targets[1].label, "work"); + assert_eq!(targets[1].name.as_deref(), Some("work")); + } + + #[test] + fn explicit_session_update_targets_only_that_session() { + let _guard = env_lock().lock().unwrap(); + let config_home = set_test_config_home("explicit-session"); + std::env::set_var(crate::api::SOCKET_PATH_ENV_VAR, "/tmp/ignored-herdr.sock"); + std::env::remove_var(crate::session::SESSION_ENV_VAR); + crate::session::clear_explicit_session_for_test(); + let args = vec![ + "herdr".to_string(), + "--session".to_string(), + "work".to_string(), + "update".to_string(), + ]; + let _ = crate::session::configure_from_args(&args).unwrap(); + + let targets = running_update_targets().unwrap(); + + let expected_socket = crate::session::api_socket_path_for(Some("work")); + std::env::remove_var(crate::api::SOCKET_PATH_ENV_VAR); + std::env::remove_var(crate::session::SESSION_ENV_VAR); + std::env::remove_var("XDG_CONFIG_HOME"); + crate::session::clear_explicit_session_for_test(); + let _ = fs::remove_dir_all(config_home); + + assert_eq!(targets.len(), 1); + assert_eq!(targets[0].label, "work"); + assert_eq!(targets[0].name.as_deref(), Some("work")); + assert_eq!(targets[0].socket_path, expected_socket); + } + + #[test] + fn socket_override_update_targets_socket_not_env_session() { + let _guard = env_lock().lock().unwrap(); + std::env::set_var(crate::api::SOCKET_PATH_ENV_VAR, "/tmp/custom-herdr.sock"); + std::env::set_var(crate::session::SESSION_ENV_VAR, "work"); + crate::session::clear_explicit_session_for_test(); + + let targets = running_update_targets().unwrap(); + + std::env::remove_var(crate::api::SOCKET_PATH_ENV_VAR); + std::env::remove_var(crate::session::SESSION_ENV_VAR); + crate::session::clear_explicit_session_for_test(); + + assert_eq!(targets.len(), 1); + assert_eq!(targets[0].name, None); + assert_eq!( + targets[0].socket_path, + PathBuf::from("/tmp/custom-herdr.sock") + ); + assert!(targets[0] + .stop_command + .contains(crate::api::SOCKET_PATH_ENV_VAR)); + } + + #[test] + fn plain_update_errors_when_named_session_has_client_socket_without_status_api() { + let _guard = env_lock().lock().unwrap(); + let config_home = set_test_config_home("client-only-session"); + std::env::remove_var(crate::api::SOCKET_PATH_ENV_VAR); + std::env::remove_var(crate::session::SESSION_ENV_VAR); + crate::session::clear_explicit_session_for_test(); + + let work_client_socket = crate::session::client_socket_path_for(Some("work")); + fs::create_dir_all(work_client_socket.parent().unwrap()).unwrap(); + let work_client_listener = UnixListener::bind(&work_client_socket).unwrap(); + let release = fake_release("9.8.7", Some(77)); + + let err = plan_running_server_updates(&release).unwrap_err(); + + drop(work_client_listener); + let _ = fs::remove_dir_all(config_home); + std::env::remove_var("XDG_CONFIG_HOME"); + + assert!( + err.contains("work") && err.contains("status API did not respond"), + "unexpected error: {err}" + ); + assert!( + err.contains("herdr session stop work"), + "unexpected error: {err}" + ); + } + + #[test] + fn failed_handoff_classification_detects_updated_server() { + let socket_path = unique_test_socket_path("handoff-updated-status"); + let handle = spawn_status_server_once(&socket_path, "9.8.7", 77); + let release = fake_release("9.8.7", Some(77)); + + let state = classify_failed_live_handoff_state_at(&socket_path, &release); + + let _ = handle.join(); + let _ = fs::remove_file(&socket_path); + assert_eq!(state, FailedHandoffServerState::UpdatedServerRunning); + } + + #[test] + fn failed_handoff_classification_detects_old_server() { + let socket_path = unique_test_socket_path("handoff-old-status"); + let handle = spawn_status_server_once(&socket_path, "0.6.2", 76); + let release = fake_release("9.8.7", Some(77)); + + let state = classify_failed_live_handoff_state_at(&socket_path, &release); + + let _ = handle.join(); + let _ = fs::remove_file(&socket_path); + match state { + FailedHandoffServerState::OldServerRunning(status) => { + assert_eq!(status.version.as_deref(), Some("0.6.2")); + assert_eq!(status.protocol, Some(76)); + } + other => panic!("unexpected state: {other:?}"), + } + } + + #[test] + fn failed_handoff_classification_detects_missing_server() { + let socket_path = unique_test_socket_path("handoff-missing-status"); + let release = fake_release("9.8.7", Some(77)); + + let state = classify_failed_live_handoff_state_at(&socket_path, &release); + + assert_eq!(state, FailedHandoffServerState::NoServerResponding); + } + + #[test] + fn noninteractive_update_requires_handoff_fails_before_install() { let _guard = env_lock().lock().unwrap(); assert!( !io::stdin().is_terminal(), @@ -1423,6 +2292,7 @@ mod tests { let server = crate::api::RuntimeStatus { version: Some("0.5.5".to_string()), protocol: Some(2), + capabilities: None, }; let release = ReleaseInfo { version: Version::parse("0.5.6").unwrap(), @@ -1430,15 +2300,29 @@ mod tests { download_url: "https://example.com/herdr".to_string(), notes_body: "### Changed\n- One".to_string(), }; + let plan = RunningServerUpdatePlan { + target: RunningUpdateTarget { + name: Some("work".to_string()), + label: "work".to_string(), + stop_command: "herdr session stop work".to_string(), + attach_command: Some("herdr session attach work".to_string()), + socket_path: crate::session::api_socket_path_for(Some("work")), + client_socket_path: crate::session::client_socket_path_for(Some("work")), + must_be_running: true, + }, + requires_live_handoff: true, + server, + }; - let err = prompt_to_stop_server_before_update(&server, &release, true).unwrap_err(); + let err = + prompt_to_live_handoff_sessions_before_update(&[&plan], &release, true).unwrap_err(); assert!( - err.contains("run `herdr session stop work`"), + err.contains("requires live server handoff"), "unexpected error: {err}" ); assert!( - err.contains("then run `herdr update` again"), + err.contains("run `herdr update` from an interactive terminal"), "unexpected error: {err}" ); std::env::remove_var(crate::session::SESSION_ENV_VAR); @@ -1494,6 +2378,73 @@ mod tests { ); } + #[test] + fn live_handoff_server_via_api_sends_handoff_request() { + let socket_path = unique_test_socket_path("handoff-ok"); + let listener = UnixListener::bind(&socket_path).unwrap(); + let handle = thread::spawn(move || { + let (mut stream, _) = listener.accept().unwrap(); + let mut request = String::new(); + BufReader::new(stream.try_clone().unwrap()) + .read_line(&mut request) + .unwrap(); + assert!(request.contains("server.live_handoff")); + stream + .write_all(b"{\"id\":\"update:server:live-handoff\",\"result\":{}}\n") + .unwrap(); + stream.flush().unwrap(); + }); + + let result = live_handoff_server_via_api_at(&socket_path, Duration::from_millis(200)); + let _ = handle.join(); + let _ = fs::remove_file(&socket_path); + assert!( + result.is_ok(), + "expected handoff request to succeed: {result:?}" + ); + } + + #[test] + fn update_live_handoff_request_names_import_binary_and_expected_release() { + let socket_path = unique_test_socket_path("handoff-update-ok"); + let listener = UnixListener::bind(&socket_path).unwrap(); + let handle = thread::spawn(move || { + let (mut stream, _) = listener.accept().unwrap(); + let mut request = String::new(); + BufReader::new(stream.try_clone().unwrap()) + .read_line(&mut request) + .unwrap(); + let value: serde_json::Value = serde_json::from_str(&request).unwrap(); + assert_eq!(value["method"], "server.live_handoff"); + assert_eq!(value["params"]["import_exe"], "/tmp/herdr-new"); + assert_eq!(value["params"]["expected_protocol"], 77); + assert_eq!(value["params"]["expected_version"], "9.8.7"); + stream + .write_all(b"{\"id\":\"update:server:live-handoff\",\"result\":{}}\n") + .unwrap(); + stream.flush().unwrap(); + }); + let release = ReleaseInfo { + version: Version::parse("9.8.7").unwrap(), + target_protocol: Some(77), + download_url: "https://example.com/herdr".to_string(), + notes_body: "### Changed\n- One".to_string(), + }; + + let result = live_handoff_server_via_api_for_release_at( + &socket_path, + Duration::from_millis(200), + Path::new("/tmp/herdr-new"), + &release, + ); + let _ = handle.join(); + let _ = fs::remove_file(&socket_path); + assert!( + result.is_ok(), + "expected handoff request to succeed: {result:?}" + ); + } + #[test] fn stop_server_via_api_times_out_when_server_never_replies() { let socket_path = unique_test_socket_path("stop-timeout"); diff --git a/tests/live_handoff.rs b/tests/live_handoff.rs new file mode 100644 index 00000000..b8e9ad76 --- /dev/null +++ b/tests/live_handoff.rs @@ -0,0 +1,817 @@ +mod support; + +use std::fs; +use std::io::{BufRead, BufReader, Read, Write}; +use std::net::{TcpListener, TcpStream}; +use std::os::unix::net::UnixStream; +use std::path::{Path, PathBuf}; +use std::sync::atomic::{AtomicUsize, Ordering}; +use std::sync::{Mutex, MutexGuard, OnceLock}; +use std::thread; +use std::time::{Duration, Instant}; + +use portable_pty::{native_pty_system, Child, CommandBuilder, MasterPty, PtySize}; +use support::{ + cleanup_test_base, client_handshake, register_runtime_dir, register_spawned_herdr_pid, + unregister_spawned_herdr_pid, wait_for_disconnect, wait_for_socket, +}; + +struct SpawnedHerdr { + _master: Box, + child: Box, +} + +impl Drop for SpawnedHerdr { + fn drop(&mut self) { + let pid = self.child.process_id(); + let _ = self.child.kill(); + unregister_spawned_herdr_pid(pid); + } +} + +fn test_lock() -> MutexGuard<'static, ()> { + static LOCK: OnceLock> = OnceLock::new(); + LOCK.get_or_init(|| Mutex::new(())) + .lock() + .unwrap_or_else(|poisoned| poisoned.into_inner()) +} + +fn unique_test_dir() -> PathBuf { + static COUNTER: AtomicUsize = AtomicUsize::new(0); + let n = COUNTER.fetch_add(1, Ordering::Relaxed); + PathBuf::from(format!("/tmp/hlh-{}-{n}", std::process::id())) +} + +fn spawn_server(config_home: &Path, runtime_dir: &Path, api_socket: &Path) -> SpawnedHerdr { + spawn_server_with_env(config_home, runtime_dir, api_socket, &[]) +} + +fn spawn_server_with_env( + config_home: &Path, + runtime_dir: &Path, + api_socket: &Path, + extra_env: &[(&str, &str)], +) -> SpawnedHerdr { + fs::create_dir_all(config_home.join("herdr")).unwrap(); + fs::create_dir_all(runtime_dir).unwrap(); + fs::write( + config_home.join("herdr/config.toml"), + "onboarding = false\n", + ) + .unwrap(); + + let pair = native_pty_system() + .openpty(PtySize { + rows: 24, + cols: 80, + pixel_width: 0, + pixel_height: 0, + }) + .unwrap(); + let mut cmd = CommandBuilder::new(env!("CARGO_BIN_EXE_herdr")); + cmd.arg("server"); + cmd.env("XDG_CONFIG_HOME", config_home); + cmd.env("XDG_RUNTIME_DIR", runtime_dir); + cmd.env("HERDR_SOCKET_PATH", api_socket); + cmd.env( + "HERDR_CLIENT_SOCKET_PATH", + runtime_dir.join("herdr-client.sock"), + ); + cmd.env("SHELL", "/bin/sh"); + for (key, value) in extra_env { + cmd.env(key, value); + } + + let child = pair.slave.spawn_command(cmd).unwrap(); + register_spawned_herdr_pid(child.process_id()); + SpawnedHerdr { + _master: pair.master, + child, + } +} + +fn spawn_named_session_server( + config_home: &Path, + runtime_dir: &Path, + session_name: &str, +) -> SpawnedHerdr { + fs::create_dir_all(config_home.join("herdr-dev")).unwrap(); + fs::create_dir_all(runtime_dir).unwrap(); + fs::write( + config_home.join("herdr-dev/config.toml"), + "onboarding = false\n", + ) + .unwrap(); + + let pair = native_pty_system() + .openpty(PtySize { + rows: 24, + cols: 80, + pixel_width: 0, + pixel_height: 0, + }) + .unwrap(); + let mut cmd = CommandBuilder::new(env!("CARGO_BIN_EXE_herdr")); + cmd.arg("server"); + cmd.env("XDG_CONFIG_HOME", config_home); + cmd.env("XDG_RUNTIME_DIR", runtime_dir); + cmd.env("HERDR_SESSION", session_name); + cmd.env_remove("HERDR_SOCKET_PATH"); + cmd.env_remove("HERDR_CLIENT_SOCKET_PATH"); + cmd.env("SHELL", "/bin/sh"); + + let child = pair.slave.spawn_command(cmd).unwrap(); + register_spawned_herdr_pid(child.process_id()); + SpawnedHerdr { + _master: pair.master, + child, + } +} + +fn spawn_default_session_server(config_home: &Path, runtime_dir: &Path) -> SpawnedHerdr { + fs::create_dir_all(config_home.join("herdr-dev")).unwrap(); + fs::create_dir_all(runtime_dir).unwrap(); + fs::write( + config_home.join("herdr-dev/config.toml"), + "onboarding = false\n", + ) + .unwrap(); + + let pair = native_pty_system() + .openpty(PtySize { + rows: 24, + cols: 80, + pixel_width: 0, + pixel_height: 0, + }) + .unwrap(); + let mut cmd = CommandBuilder::new(env!("CARGO_BIN_EXE_herdr")); + cmd.arg("server"); + cmd.env("XDG_CONFIG_HOME", config_home); + cmd.env("XDG_RUNTIME_DIR", runtime_dir); + cmd.env_remove("HERDR_SESSION"); + cmd.env_remove("HERDR_SOCKET_PATH"); + cmd.env_remove("HERDR_CLIENT_SOCKET_PATH"); + cmd.env("SHELL", "/bin/sh"); + + let child = pair.slave.spawn_command(cmd).unwrap(); + register_spawned_herdr_pid(child.process_id()); + SpawnedHerdr { + _master: pair.master, + child, + } +} + +fn request(socket_path: &Path, request: serde_json::Value) -> serde_json::Value { + let mut stream = UnixStream::connect(socket_path).unwrap(); + stream.write_all(request.to_string().as_bytes()).unwrap(); + stream.write_all(b"\n").unwrap(); + stream.flush().unwrap(); + let mut line = String::new(); + BufReader::new(stream).read_line(&mut line).unwrap(); + serde_json::from_str(&line).unwrap() +} + +fn assert_ok(response: serde_json::Value) { + assert!( + response.get("result").is_some(), + "api request failed: {response}" + ); +} + +fn wait_for_api(socket_path: &Path, timeout: Duration) { + let deadline = Instant::now() + timeout; + while Instant::now() < deadline { + if UnixStream::connect(socket_path).is_ok() { + let response = request( + socket_path, + serde_json::json!({"id":"test:ping","method":"ping","params":{}}), + ); + if response.get("result").is_some() { + return; + } + } + thread::sleep(Duration::from_millis(25)); + } + panic!("api did not become ready at {}", socket_path.display()); +} + +fn wait_for_output(socket_path: &Path, pane_id: &str, needle: &str) { + let deadline = Instant::now() + Duration::from_secs(5); + let mut last_text = String::new(); + let mut last_response = serde_json::Value::Null; + while Instant::now() < deadline { + let response = request( + socket_path, + serde_json::json!({ + "id": "test:pane:read", + "method": "pane.read", + "params": { + "pane_id": pane_id, + "source": "visible", + "lines": 20, + "format": "text", + "strip_ansi": true + } + }), + ); + last_response = response.clone(); + let text = response["result"]["read"]["text"] + .as_str() + .unwrap_or_default(); + last_text = text.to_string(); + if text.contains(needle) { + return; + } + thread::sleep(Duration::from_millis(50)); + } + panic!( + "pane output did not contain {needle:?}; last text was {last_text:?}; last response was {last_response}" + ); +} + +fn wait_for_file_contains(path: &Path, needle: &str, timeout: Duration) -> String { + let deadline = Instant::now() + timeout; + let mut last_text = String::new(); + while Instant::now() < deadline { + if let Ok(text) = fs::read_to_string(path) { + last_text = text; + if last_text.contains(needle) { + return last_text; + } + } + thread::sleep(Duration::from_millis(50)); + } + panic!( + "{} did not contain {needle:?}; last text was {last_text:?}", + path.display() + ); +} + +fn unused_local_port() -> u16 { + TcpListener::bind("127.0.0.1:0") + .unwrap() + .local_addr() + .unwrap() + .port() +} + +fn wait_for_http_contains(port: u16, needle: &str, timeout: Duration) -> String { + let deadline = Instant::now() + timeout; + let mut last_response = String::new(); + while Instant::now() < deadline { + if let Ok(mut stream) = TcpStream::connect(("127.0.0.1", port)) { + let _ = + stream.write_all(b"GET / HTTP/1.1\r\nHost: 127.0.0.1\r\nConnection: close\r\n\r\n"); + let mut response = String::new(); + let _ = stream.read_to_string(&mut response); + last_response = response; + if last_response.contains(needle) { + return last_response; + } + } + thread::sleep(Duration::from_millis(50)); + } + panic!( + "http server on port {port} did not return {needle:?}; last response was {last_response:?}" + ); +} + +#[test] +fn live_handoff_preserves_named_session_socket_paths() { + let _lock = test_lock(); + let base = unique_test_dir(); + let config_home = base.join("config"); + let runtime_dir = base.join("runtime"); + let session_dir = config_home.join("herdr-dev/sessions/work"); + let api_socket = session_dir.join("herdr.sock"); + let client_socket = session_dir.join("herdr-client.sock"); + + let spawned = spawn_named_session_server(&config_home, &runtime_dir, "work"); + wait_for_socket(&api_socket, Duration::from_secs(10)); + register_runtime_dir(&runtime_dir); + + assert_ok(request( + &api_socket, + serde_json::json!({"id":"test:handoff","method":"server.live_handoff","params":{}}), + )); + drop(spawned); + wait_for_api(&api_socket, Duration::from_secs(10)); + wait_for_socket(&client_socket, Duration::from_secs(5)); + assert!( + !config_home.join("herdr-dev/herdr.sock").exists(), + "named handoff unexpectedly bound the default session API socket" + ); + + let _ = request( + &api_socket, + serde_json::json!({"id":"test:stop","method":"server.stop","params":{}}), + ); + cleanup_test_base(&base); +} + +#[test] +fn live_handoff_preserves_pane_process_io() { + let _lock = test_lock(); + let base = unique_test_dir(); + let config_home = base.join("config"); + let runtime_dir = base.join("runtime"); + let api_socket = runtime_dir.join("herdr.sock"); + let client_socket = runtime_dir.join("herdr-client.sock"); + let marker = base.join("child.pid"); + let second_marker = base.join("second-child.pid"); + let hup_marker = base.join("hup"); + let second_hup_marker = base.join("second-hup"); + let received_marker = base.join("received"); + let second_received_marker = base.join("second-received"); + + let spawned = spawn_server(&config_home, &runtime_dir, &api_socket); + wait_for_socket(&api_socket, Duration::from_secs(10)); + register_runtime_dir(&runtime_dir); + + let created = request( + &api_socket, + serde_json::json!({ + "id": "test:workspace:create", + "method": "workspace.create", + "params": {"cwd": "/tmp", "focus": true} + }), + ); + let pane_id = created["result"]["root_pane"]["pane_id"] + .as_str() + .unwrap() + .to_string(); + let split = request( + &api_socket, + serde_json::json!({ + "id": "test:pane:split", + "method": "pane.split", + "params": { + "target_pane_id": pane_id, + "direction": "right", + "focus": false + } + }), + ); + assert_ok(split.clone()); + let second_pane_id = split["result"]["pane"]["pane_id"] + .as_str() + .unwrap() + .to_string(); + + let command = format!( + "sh -c 'echo READY $$ > {}; trap \"echo HUP >> {}\" HUP; while read line; do echo got:$line; echo got:$line >> {}; done'", + marker.display(), + hup_marker.display(), + received_marker.display() + ); + let second_command = format!( + "sh -c 'echo SECOND_READY $$ > {}; trap \"echo HUP >> {}\" HUP; while read line; do echo second:$line; echo second:$line >> {}; done'", + second_marker.display(), + second_hup_marker.display(), + second_received_marker.display() + ); + assert_ok(request( + &api_socket, + serde_json::json!({ + "id": "test:pane:run", + "method": "pane.send_input", + "params": {"pane_id": pane_id, "text": command, "keys": ["Enter"]} + }), + )); + assert_ok(request( + &api_socket, + serde_json::json!({ + "id": "test:second-pane:run", + "method": "pane.send_input", + "params": {"pane_id": second_pane_id, "text": second_command, "keys": ["Enter"]} + }), + )); + support::wait_for_file(&marker, Duration::from_secs(5)); + support::wait_for_file(&second_marker, Duration::from_secs(5)); + let pid_text = fs::read_to_string(&marker).unwrap(); + let child_pid: u32 = pid_text.split_whitespace().last().unwrap().parse().unwrap(); + let second_pid_text = fs::read_to_string(&second_marker).unwrap(); + let second_child_pid: u32 = second_pid_text + .split_whitespace() + .last() + .unwrap() + .parse() + .unwrap(); + assert_eq!(unsafe { libc::kill(child_pid as libc::pid_t, 0) }, 0); + assert_eq!(unsafe { libc::kill(second_child_pid as libc::pid_t, 0) }, 0); + + let protocol = request( + &api_socket, + serde_json::json!({"id":"test:protocol","method":"ping","params":{}}), + )["result"]["protocol"] + .as_u64() + .unwrap() as u32; + let mut client_stream = UnixStream::connect(&client_socket).unwrap(); + let (server_protocol, error) = client_handshake(&mut client_stream, protocol, 80, 24).unwrap(); + assert_eq!(server_protocol, protocol); + assert!(error.is_none(), "client handshake failed: {error:?}"); + + assert_ok(request( + &api_socket, + serde_json::json!({ + "id": "test:pane:before-log", + "method": "pane.send_input", + "params": {"pane_id": pane_id, "text": "before_replay", "keys": ["Enter"]} + }), + )); + wait_for_output(&api_socket, &pane_id, "got:before_replay"); + + assert_ok(request( + &api_socket, + serde_json::json!({"id":"test:handoff","method":"server.live_handoff","params":{}}), + )); + drop(spawned); + assert!( + wait_for_disconnect(&mut client_stream, Duration::from_secs(5)).unwrap(), + "connected clients should disconnect during live handoff" + ); + thread::sleep(Duration::from_millis(300)); + wait_for_api(&api_socket, Duration::from_secs(10)); + wait_for_socket(&client_socket, Duration::from_secs(5)); + assert_eq!(unsafe { libc::kill(child_pid as libc::pid_t, 0) }, 0); + assert_eq!(unsafe { libc::kill(second_child_pid as libc::pid_t, 0) }, 0); + assert!( + !hup_marker.exists(), + "pane process received HUP during handoff" + ); + assert!( + !second_hup_marker.exists(), + "second pane process received HUP during handoff" + ); + wait_for_output(&api_socket, &pane_id, "got:before_replay"); + + assert_ok(request( + &api_socket, + serde_json::json!({ + "id": "test:pane:send", + "method": "pane.send_input", + "params": {"pane_id": pane_id, "text": "after-handoff", "keys": ["Enter"]} + }), + )); + wait_for_file_contains( + &received_marker, + "got:after-handoff", + Duration::from_secs(5), + ); + wait_for_output(&api_socket, &pane_id, "got:after-handoff"); + assert_ok(request( + &api_socket, + serde_json::json!({ + "id": "test:second-pane:send", + "method": "pane.send_input", + "params": {"pane_id": second_pane_id, "text": "after-handoff-second", "keys": ["Enter"]} + }), + )); + wait_for_file_contains( + &second_received_marker, + "second:after-handoff-second", + Duration::from_secs(5), + ); + wait_for_output(&api_socket, &second_pane_id, "second:after-handoff-sec"); + + let _ = request( + &api_socket, + serde_json::json!({"id":"test:stop","method":"server.stop","params":{}}), + ); + let _ = client_socket; + cleanup_test_base(&base); +} + +#[test] +fn live_handoff_preserves_python_http_server() { + let _lock = test_lock(); + let base = unique_test_dir(); + let config_home = base.join("config"); + let runtime_dir = base.join("runtime"); + let api_socket = runtime_dir.join("herdr.sock"); + let client_socket = runtime_dir.join("herdr-client.sock"); + let web_root = base.join("web"); + fs::create_dir_all(&web_root).unwrap(); + fs::write( + web_root.join("index.html"), + "hello-from-python-before-and-after", + ) + .unwrap(); + let port = unused_local_port(); + + let spawned = spawn_server(&config_home, &runtime_dir, &api_socket); + wait_for_socket(&api_socket, Duration::from_secs(10)); + register_runtime_dir(&runtime_dir); + + let created = request( + &api_socket, + serde_json::json!({ + "id": "test:workspace:create", + "method": "workspace.create", + "params": {"cwd": web_root, "focus": true} + }), + ); + let pane_id = created["result"]["root_pane"]["pane_id"] + .as_str() + .unwrap() + .to_string(); + + assert_ok(request( + &api_socket, + serde_json::json!({ + "id": "test:pane:run-python", + "method": "pane.send_input", + "params": { + "pane_id": pane_id, + "text": format!("python3 -m http.server {port} --bind 127.0.0.1"), + "keys": ["Enter"] + } + }), + )); + wait_for_http_contains( + port, + "hello-from-python-before-and-after", + Duration::from_secs(10), + ); + + assert_ok(request( + &api_socket, + serde_json::json!({"id":"test:handoff","method":"server.live_handoff","params":{}}), + )); + drop(spawned); + wait_for_api(&api_socket, Duration::from_secs(10)); + wait_for_http_contains( + port, + "hello-from-python-before-and-after", + Duration::from_secs(10), + ); + + let _ = request( + &api_socket, + serde_json::json!({"id":"test:stop","method":"server.stop","params":{}}), + ); + let _ = client_socket; + cleanup_test_base(&base); +} + +#[test] +fn live_handoff_preserves_http_servers_across_multiple_sessions() { + let _lock = test_lock(); + let base = unique_test_dir(); + let config_home = base.join("config"); + let runtime_dir = base.join("runtime"); + let sessions = [ + (None, config_home.join("herdr-dev/herdr.sock")), + ( + Some("work"), + config_home.join("herdr-dev/sessions/work/herdr.sock"), + ), + ]; + let mut spawned = Vec::new(); + let mut ports = Vec::new(); + + for (session_name, api_socket) in &sessions { + let web_root = base.join(format!("web-{}", session_name.unwrap_or("default"))); + fs::create_dir_all(&web_root).unwrap(); + fs::write( + web_root.join("index.html"), + format!("hello-from-{}", session_name.unwrap_or("default")), + ) + .unwrap(); + let port = unused_local_port(); + let server = if let Some(session_name) = session_name { + spawn_named_session_server(&config_home, &runtime_dir, session_name) + } else { + spawn_default_session_server(&config_home, &runtime_dir) + }; + wait_for_socket(api_socket, Duration::from_secs(10)); + let created = request( + api_socket, + serde_json::json!({ + "id": "test:workspace:create", + "method": "workspace.create", + "params": {"cwd": web_root, "focus": true} + }), + ); + let pane_id = created["result"]["root_pane"]["pane_id"] + .as_str() + .unwrap() + .to_string(); + assert_ok(request( + api_socket, + serde_json::json!({ + "id": "test:pane:run-python", + "method": "pane.send_input", + "params": { + "pane_id": pane_id, + "text": format!("python3 -m http.server {port} --bind 127.0.0.1"), + "keys": ["Enter"] + } + }), + )); + wait_for_http_contains( + port, + &format!("hello-from-{}", session_name.unwrap_or("default")), + Duration::from_secs(10), + ); + spawned.push(server); + ports.push((port, session_name.unwrap_or("default").to_string())); + } + register_runtime_dir(&runtime_dir); + + for (_session_name, api_socket) in &sessions { + assert_ok(request( + api_socket, + serde_json::json!({"id":"test:handoff","method":"server.live_handoff","params":{}}), + )); + } + drop(spawned); + + for (_session_name, api_socket) in &sessions { + wait_for_api(api_socket, Duration::from_secs(10)); + } + for (port, label) in ports { + wait_for_http_contains( + port, + &format!("hello-from-{label}"), + Duration::from_secs(10), + ); + } + + for (_session_name, api_socket) in &sessions { + let _ = request( + api_socket, + serde_json::json!({"id":"test:stop","method":"server.stop","params":{}}), + ); + } + cleanup_test_base(&base); +} + +#[test] +fn live_handoff_bad_expected_protocol_rolls_back_old_server() { + let _lock = test_lock(); + let base = unique_test_dir(); + let config_home = base.join("config"); + let runtime_dir = base.join("runtime"); + let api_socket = runtime_dir.join("herdr.sock"); + let marker = base.join("child.pid"); + let received_marker = base.join("received"); + + let spawned = spawn_server(&config_home, &runtime_dir, &api_socket); + wait_for_socket(&api_socket, Duration::from_secs(10)); + register_runtime_dir(&runtime_dir); + + let created = request( + &api_socket, + serde_json::json!({ + "id": "test:workspace:create", + "method": "workspace.create", + "params": {"cwd": "/tmp", "focus": true} + }), + ); + let pane_id = created["result"]["root_pane"]["pane_id"] + .as_str() + .unwrap() + .to_string(); + let command = format!( + "sh -c 'echo READY $$ > {}; while read line; do echo got:$line; echo got:$line >> {}; done'", + marker.display(), + received_marker.display() + ); + assert_ok(request( + &api_socket, + serde_json::json!({ + "id": "test:pane:run", + "method": "pane.send_input", + "params": {"pane_id": pane_id, "text": command, "keys": ["Enter"]} + }), + )); + support::wait_for_file(&marker, Duration::from_secs(5)); + let pid_text = fs::read_to_string(&marker).unwrap(); + let child_pid: u32 = pid_text.split_whitespace().last().unwrap().parse().unwrap(); + + let failed = request( + &api_socket, + serde_json::json!({ + "id": "test:bad-handoff", + "method": "server.live_handoff", + "params": {"expected_protocol": 999999} + }), + ); + assert!( + failed.get("error").is_some(), + "bad protocol handoff should fail: {failed}" + ); + wait_for_api(&api_socket, Duration::from_secs(5)); + assert_eq!(unsafe { libc::kill(child_pid as libc::pid_t, 0) }, 0); + + assert_ok(request( + &api_socket, + serde_json::json!({ + "id": "test:pane:send-after-failed-handoff", + "method": "pane.send_input", + "params": {"pane_id": pane_id, "text": "after-failed-handoff", "keys": ["Enter"]} + }), + )); + wait_for_file_contains( + &received_marker, + "got:after-failed-handoff", + Duration::from_secs(5), + ); + wait_for_output(&api_socket, &pane_id, "got:after-failed-handoff"); + + let _ = request( + &api_socket, + serde_json::json!({"id":"test:stop","method":"server.stop","params":{}}), + ); + drop(spawned); + cleanup_test_base(&base); +} + +fn live_handoff_import_failure_rolls_back_old_server_at(failure_point: &str) { + let _lock = test_lock(); + let base = unique_test_dir(); + let config_home = base.join("config"); + let runtime_dir = base.join("runtime"); + let api_socket = runtime_dir.join("herdr.sock"); + let client_socket = runtime_dir.join("herdr-client.sock"); + let marker = base.join("child.pid"); + let received_marker = base.join("received"); + + let spawned = spawn_server_with_env( + &config_home, + &runtime_dir, + &api_socket, + &[("HERDR_TEST_HANDOFF_IMPORT_FAIL", failure_point)], + ); + wait_for_socket(&api_socket, Duration::from_secs(10)); + register_runtime_dir(&runtime_dir); + + let created = request( + &api_socket, + serde_json::json!({ + "id": "test:workspace:create", + "method": "workspace.create", + "params": {"cwd": "/tmp", "focus": true} + }), + ); + let pane_id = created["result"]["root_pane"]["pane_id"] + .as_str() + .unwrap() + .to_string(); + let command = format!( + "sh -c 'echo READY $$ > {}; while read line; do echo got:$line; echo got:$line >> {}; done'", + marker.display(), + received_marker.display() + ); + assert_ok(request( + &api_socket, + serde_json::json!({ + "id": "test:pane:run", + "method": "pane.send_input", + "params": {"pane_id": pane_id, "text": command, "keys": ["Enter"]} + }), + )); + support::wait_for_file(&marker, Duration::from_secs(5)); + let pid_text = fs::read_to_string(&marker).unwrap(); + let child_pid: u32 = pid_text.split_whitespace().last().unwrap().parse().unwrap(); + + let failed = request( + &api_socket, + serde_json::json!({"id":"test:handoff-fail","method":"server.live_handoff","params":{}}), + ); + assert!( + failed.get("error").is_some(), + "{failure_point} handoff should fail: {failed}" + ); + wait_for_api(&api_socket, Duration::from_secs(10)); + wait_for_socket(&client_socket, Duration::from_secs(5)); + assert_eq!(unsafe { libc::kill(child_pid as libc::pid_t, 0) }, 0); + + assert_ok(request( + &api_socket, + serde_json::json!({ + "id": "test:pane:send-after-import-failure", + "method": "pane.send_input", + "params": {"pane_id": pane_id, "text": failure_point, "keys": ["Enter"]} + }), + )); + wait_for_file_contains( + &received_marker, + &format!("got:{failure_point}"), + Duration::from_secs(5), + ); + + let _ = request( + &api_socket, + serde_json::json!({"id":"test:stop","method":"server.stop","params":{}}), + ); + drop(spawned); + cleanup_test_base(&base); +} + +#[test] +fn live_handoff_after_restored_failure_rolls_back_old_server() { + live_handoff_import_failure_rolls_back_old_server_at("after_restored"); +} diff --git a/tests/support/mod.rs b/tests/support/mod.rs index e810b843..cae140c9 100644 --- a/tests/support/mod.rs +++ b/tests/support/mod.rs @@ -328,16 +328,30 @@ pub fn wait_for_message_variant( } pub fn wait_for_disconnect(stream: &mut UnixStream, timeout: Duration) -> Result { - stream - .set_read_timeout(Some(Duration::from_millis(200))) - .map_err(|e| e.to_string())?; + stream.set_nonblocking(true).map_err(|e| e.to_string())?; let deadline = Instant::now() + timeout; - while Instant::now() < deadline { - if read_server_message(stream).is_err() { - return Ok(true); + let mut idle_since = None; + let result = loop { + match read_server_message(stream) { + Ok(_) => idle_since = None, + Err(err) + if err.to_ascii_lowercase().contains("would block") + || err.contains("Resource temporarily unavailable") => + { + let idle_started = *idle_since.get_or_insert_with(Instant::now); + if idle_started.elapsed() >= Duration::from_millis(200) { + break Ok(true); + } + } + Err(_) => break Ok(true), } - } - Ok(false) + if Instant::now() >= deadline { + break Ok(false); + } + thread::sleep(Duration::from_millis(25)); + }; + let _ = stream.set_nonblocking(false); + result } pub fn cleanup_registered_herdr_pids() {