diff --git a/wire/cli/src/tracing_setup.rs b/wire/cli/src/tracing_setup.rs index f52c7b9..abb4c9a 100644 --- a/wire/cli/src/tracing_setup.rs +++ b/wire/cli/src/tracing_setup.rs @@ -4,7 +4,6 @@ use std::{ collections::VecDeque, io::{self, Stderr, Write, stderr}, - sync::TryLockError, }; use clap_verbosity_flag::{LogLevel, Verbosity}; @@ -56,22 +55,14 @@ impl NonClobberingWriter { impl Write for NonClobberingWriter { fn write(&mut self, buf: &[u8]) -> std::io::Result { - match STDIN_CLOBBER_LOCK.clone().try_lock() { - Ok(_) => { - self.dump_previous().map(|()| 0)?; - - self.stderr.write(buf) - } - Err(e) => match e { - TryLockError::Poisoned(_) => { - panic!("Internal stdout clobber lock is posioned. Please create an issue."); - } - TryLockError::WouldBlock => { - self.queue.push_front(buf.to_vec()); - - Ok(buf.len()) - } - }, + if let 1.. = STDIN_CLOBBER_LOCK.available_permits() { + self.dump_previous().map(|()| 0)?; + + self.stderr.write(buf) + } else { + self.queue.push_front(buf.to_vec()); + + Ok(buf.len()) } } diff --git a/wire/lib/src/commands/common.rs b/wire/lib/src/commands/common.rs index fb3c8b7..1c6dc71 100644 --- a/wire/lib/src/commands/common.rs +++ b/wire/lib/src/commands/common.rs @@ -37,7 +37,8 @@ pub async fn push(context: &Context<'_>, push: Push<'_>) -> Result<(), HiveLibEr .target .create_ssh_opts(context.modifiers, false)?, )]), - )?; + ) + .await?; child .wait_till_success() @@ -95,7 +96,8 @@ pub async fn evaluate_hive_attribute( let child = run_command( &CommandArguments::new(command_string, modifiers) .mode(crate::commands::ChildOutputMode::Nix), - )?; + ) + .await?; child .wait_till_success() diff --git a/wire/lib/src/commands/interactive.rs b/wire/lib/src/commands/interactive.rs deleted file mode 100644 index 1ffbe48..0000000 --- a/wire/lib/src/commands/interactive.rs +++ /dev/null @@ -1,891 +0,0 @@ -// SPDX-License-Identifier: AGPL-3.0-or-later -// Copyright 2024-2025 wire Contributors - -use aho_corasick::{AhoCorasick, PatternID}; -use itertools::Itertools; -use nix::sys::termios::{LocalFlags, SetArg, Termios, tcgetattr, tcsetattr}; -use nix::{ - poll::{PollFd, PollFlags, PollTimeout, poll}, - unistd::{pipe as posix_pipe, read as posix_read, write as posix_write}, -}; -use portable_pty::{CommandBuilder, NativePtySystem, PtyPair, PtySize}; -use rand::distr::Alphabetic; -use std::collections::VecDeque; -use std::sync::mpsc::{self, Sender}; -use std::sync::{Condvar, LazyLock, Mutex}; -use std::thread::JoinHandle; -use std::{ - io::{Read, Write}, - os::fd::{AsFd, OwnedFd}, - sync::Arc, -}; -use tracing::instrument; -use tracing::{Span, debug, error, trace, warn}; - -use crate::commands::CommandArguments; -use crate::commands::interactive_logbuffer::LogBuffer; -use crate::errors::CommandError; -use crate::{STDIN_CLOBBER_LOCK, SubCommandModifiers}; -use crate::{ - commands::{ChildOutputMode, WireCommandChip}, - errors::HiveLibError, - hive::node::Target, -}; - -type MasterWriter = Box; -type MasterReader = Box; -type Child = Box; - -pub(crate) struct InteractiveChildChip { - child: Child, - - cancel_stdin_pipe_w: OwnedFd, - write_stdin_pipe_w: OwnedFd, - - stderr_collection: Arc>>, - stdout_collection: Arc>>, - - original_command: String, - - completion_status: Arc, - stdout_handle: JoinHandle>, -} - -struct StdinTermiosAttrGuard(Termios); - -struct CompletionStatus { - completed: Mutex, - success: Mutex>, - condvar: Condvar, -} - -struct WatchStdoutArguments { - began_tx: Sender<()>, - reader: MasterReader, - succeed_needle: Arc>, - failed_needle: Arc>, - start_needle: Arc>, - output_mode: ChildOutputMode, - stderr_collection: Arc>>, - stdout_collection: Arc>>, - completion_status: Arc, - span: Span, - log_stdout: bool, -} - -#[derive(Debug)] -enum SearchFindings { - None, - Started, - Terminate, -} - -/// the underlying command began -const THREAD_BEGAN_SIGNAL: &[u8; 1] = b"b"; -const THREAD_QUIT_SIGNAL: &[u8; 1] = b"q"; - -static STARTED_PATTERN: LazyLock = LazyLock::new(|| PatternID::must(0)); -static SUCCEEDED_PATTERN: LazyLock = LazyLock::new(|| PatternID::must(1)); -static FAILED_PATTERN: LazyLock = LazyLock::new(|| PatternID::must(2)); - -/// substitutes STDOUT with #$line. stdout is far less common than stderr. -const IO_SUBS: &str = "1> >(while IFS= read -r line; do echo \"#$line\"; done)"; - -fn create_ending_segment>( - arguments: &CommandArguments<'_, S>, - needles: Needles, -) -> String { - let (succeed_needle, failed_needle, start_needle) = needles; - - format!( - "echo -e '{succeed}' || echo '{failed}'", - succeed = if matches!(arguments.output_mode, ChildOutputMode::Interactive) { - format!( - "{start}\\n{succeed}", - start = String::from_utf8_lossy(&start_needle), - succeed = String::from_utf8_lossy(&succeed_needle) - ) - } else { - String::from_utf8_lossy(&succeed_needle).to_string() - }, - failed = String::from_utf8_lossy(&failed_needle) - ) -} - -fn create_starting_segment>( - arguments: &CommandArguments<'_, S>, - start_needle: &Arc>, -) -> String { - if matches!(arguments.output_mode, ChildOutputMode::Interactive) { - String::new() - } else { - format!( - "echo '{start}' && ", - start = String::from_utf8_lossy(start_needle) - ) - } -} - -#[instrument(skip_all, name = "run-int", fields(elevated = %arguments.is_elevated()))] -pub(crate) fn interactive_command_with_env>( - arguments: &CommandArguments, - envs: std::collections::HashMap, -) -> Result { - print_authenticate_warning(arguments)?; - - let (succeed_needle, failed_needle, start_needle) = create_needles(); - - let pty_system = NativePtySystem::default(); - let pty_pair = portable_pty::PtySystem::openpty(&pty_system, PtySize::default()).unwrap(); - setup_master(&pty_pair)?; - - let command_string = &format!( - "{starting}{command} {flags} {IO_SUBS} && {ending}", - command = arguments.command_string.as_ref(), - flags = match arguments.output_mode { - ChildOutputMode::Nix => "--log-format internal-json", - ChildOutputMode::Generic | ChildOutputMode::Interactive => "", - }, - starting = create_starting_segment(arguments, &start_needle), - ending = create_ending_segment( - arguments, - ( - succeed_needle.clone(), - failed_needle.clone(), - start_needle.clone() - ) - ) - ); - - debug!("{command_string}"); - - let mut command = build_command(arguments, command_string)?; - - // give command all env vars - for (key, value) in envs { - command.env(key, value); - } - - let clobber_guard = STDIN_CLOBBER_LOCK.lock().unwrap(); - let _guard = StdinTermiosAttrGuard::new().map_err(HiveLibError::CommandError)?; - let child = pty_pair - .slave - .spawn_command(command) - .map_err(|x| HiveLibError::CommandError(CommandError::PortablePty(x)))?; - - // Release any handles owned by the slave: we don't need it now - // that we've spawned the child. - drop(pty_pair.slave); - - let reader = pty_pair - .master - .try_clone_reader() - .map_err(|x| HiveLibError::CommandError(CommandError::PortablePty(x)))?; - let master_writer = pty_pair - .master - .take_writer() - .map_err(|x| HiveLibError::CommandError(CommandError::PortablePty(x)))?; - - let stderr_collection = Arc::new(Mutex::new(VecDeque::::with_capacity(10))); - let stdout_collection = Arc::new(Mutex::new(VecDeque::::with_capacity(10))); - let (began_tx, began_rx) = mpsc::channel::<()>(); - let completion_status = Arc::new(CompletionStatus::new()); - - let stdout_handle = { - let arguments = WatchStdoutArguments { - began_tx, - reader, - succeed_needle: succeed_needle.clone(), - failed_needle: failed_needle.clone(), - start_needle: start_needle.clone(), - output_mode: arguments.output_mode, - stderr_collection: stderr_collection.clone(), - stdout_collection: stdout_collection.clone(), - completion_status: completion_status.clone(), - span: Span::current(), - log_stdout: arguments.log_stdout, - }; - - std::thread::spawn(move || dynamic_watch_sudo_stdout(arguments)) - }; - - let (write_stdin_pipe_r, write_stdin_pipe_w) = - posix_pipe().map_err(|x| HiveLibError::CommandError(CommandError::PosixPipe(x)))?; - let (cancel_stdin_pipe_r, cancel_stdin_pipe_w) = - posix_pipe().map_err(|x| HiveLibError::CommandError(CommandError::PosixPipe(x)))?; - - std::thread::spawn(move || { - watch_stdin_from_user( - &cancel_stdin_pipe_r, - master_writer, - &write_stdin_pipe_r, - Span::current(), - ) - }); - - debug!("Setup threads"); - - let () = began_rx - .recv() - .map_err(|x| HiveLibError::CommandError(CommandError::RecvError(x)))?; - - drop(clobber_guard); - - if arguments.keep_stdin_open { - trace!("Sending THREAD_BEGAN_SIGNAL"); - - posix_write(&cancel_stdin_pipe_w, THREAD_BEGAN_SIGNAL) - .map_err(|x| HiveLibError::CommandError(CommandError::PosixPipe(x)))?; - } else { - trace!("Sending THREAD_QUIT_SIGNAL"); - - posix_write(&cancel_stdin_pipe_w, THREAD_QUIT_SIGNAL) - .map_err(|x| HiveLibError::CommandError(CommandError::PosixPipe(x)))?; - } - - Ok(InteractiveChildChip { - child, - cancel_stdin_pipe_w, - write_stdin_pipe_w, - stderr_collection, - stdout_collection, - original_command: arguments.command_string.as_ref().to_string(), - completion_status, - stdout_handle, - }) -} - -fn print_authenticate_warning>( - arguments: &CommandArguments, -) -> Result<(), HiveLibError> { - if !arguments.is_elevated() { - return Ok(()); - } - - eprintln!( - "{} | Authenticate for \"sudo {}\":", - arguments - .target - .map_or(Ok("localhost (!)".to_string()), |target| Ok(format!( - "{}@{}:{}", - target.user, - target.get_preferred_host()?, - target.port - )))?, - arguments.command_string.as_ref() - ); - - Ok(()) -} - -type Needles = (Arc>, Arc>, Arc>); - -fn create_needles() -> Needles { - let tmp_prefix = rand::distr::SampleString::sample_string(&Alphabetic, &mut rand::rng(), 5); - - ( - Arc::new(format!("{tmp_prefix}_W_Q").as_bytes().to_vec()), - Arc::new(format!("{tmp_prefix}_W_F").as_bytes().to_vec()), - Arc::new(format!("{tmp_prefix}_W_S").as_bytes().to_vec()), - ) -} - -fn setup_master(pty_pair: &PtyPair) -> Result<(), HiveLibError> { - if let Some(fd) = pty_pair.master.as_raw_fd() { - // convert raw fd to a BorrowedFd - // safe as `fd` is dropped well before `pty_pair.master` - let fd = unsafe { std::os::unix::io::BorrowedFd::borrow_raw(fd) }; - let mut termios = - tcgetattr(fd).map_err(|x| HiveLibError::CommandError(CommandError::TermAttrs(x)))?; - - termios.local_flags &= !LocalFlags::ECHO; - // Key agent does not work well without canonical mode - termios.local_flags &= !LocalFlags::ICANON; - // Actually quit - termios.local_flags &= !LocalFlags::ISIG; - - tcsetattr(fd, SetArg::TCSANOW, &termios) - .map_err(|x| HiveLibError::CommandError(CommandError::TermAttrs(x)))?; - } - - Ok(()) -} - -fn build_command>( - arguments: &CommandArguments<'_, S>, - command_string: &String, -) -> Result { - let mut command = if let Some(target) = arguments.target { - let mut command = create_sync_ssh_command(target, arguments.modifiers)?; - - // force ssh to use our pesudo terminal - command.arg("-tt"); - - command - } else { - let mut command = portable_pty::CommandBuilder::new("sh"); - - command.arg("-c"); - - command - }; - - if let Some(escalation_command) = &arguments.privilege_escalation_command { - command.arg(format!("{escalation_command} sh -c '{command_string}'")); - } else { - command.arg(command_string); - } - - Ok(command) -} - -impl CompletionStatus { - const fn new() -> Self { - CompletionStatus { - completed: Mutex::new(false), - success: Mutex::new(None), - condvar: Condvar::new(), - } - } - - fn mark_completed(&self, was_successful: bool) { - let mut completed = self.completed.lock().unwrap(); - let mut success = self.success.lock().unwrap(); - - *completed = true; - *success = Some(was_successful); - - self.condvar.notify_all(); - } - - fn wait(&self) -> Option { - let mut completed = self.completed.lock().unwrap(); - - while !*completed { - completed = self.condvar.wait(completed).unwrap(); - } - - *self.success.lock().unwrap() - } -} - -impl WireCommandChip for InteractiveChildChip { - type ExitStatus = (portable_pty::ExitStatus, String); - - #[instrument(skip_all)] - async fn wait_till_success(mut self) -> Result { - drop(self.write_stdin_pipe_w); - - let exit_status = tokio::task::spawn_blocking(move || self.child.wait()) - .await - .map_err(CommandError::JoinError)? - .map_err(CommandError::WaitForStatus)?; - - debug!("exit_status: {exit_status:?}"); - - self.stdout_handle - .join() - .map_err(|_| CommandError::ThreadPanic)??; - let success = self.completion_status.wait(); - let _ = posix_write(&self.cancel_stdin_pipe_w, THREAD_QUIT_SIGNAL); - - if let Some(true) = success { - let logs = self - .stdout_collection - .lock() - .unwrap() - .iter() - .rev() - .map(|x| x.trim()) - .join("\n"); - - return Ok((exit_status, logs)); - } - - debug!("child did not succeed"); - - let logs = self - .stderr_collection - .lock() - .unwrap() - .iter() - .rev() - .join("\n"); - - Err(CommandError::CommandFailed { - command_ran: self.original_command, - logs, - code: format!("code {}", exit_status.exit_code()), - reason: match success { - Some(_) => "marked-unsuccessful", - None => "child-crashed-before-succeeding", - }, - }) - } - - async fn write_stdin(&mut self, data: Vec) -> Result<(), HiveLibError> { - trace!("Writing {} bytes to stdin", data.len()); - - posix_write(&self.write_stdin_pipe_w, &data) - .map_err(|x| HiveLibError::CommandError(CommandError::PosixPipe(x)))?; - - Ok(()) - } -} - -impl StdinTermiosAttrGuard { - fn new() -> Result { - let stdin = std::io::stdin(); - let stdin_fd = stdin.as_fd(); - - let mut termios = tcgetattr(stdin_fd).map_err(CommandError::TermAttrs)?; - let original_termios = termios.clone(); - - termios.local_flags &= !(LocalFlags::ECHO | LocalFlags::ICANON); - tcsetattr(stdin_fd, SetArg::TCSANOW, &termios).map_err(CommandError::TermAttrs)?; - - Ok(StdinTermiosAttrGuard(original_termios)) - } -} - -impl Drop for StdinTermiosAttrGuard { - fn drop(&mut self) { - let stdin = std::io::stdin(); - let stdin_fd = stdin.as_fd(); - - let _ = tcsetattr(stdin_fd, SetArg::TCSANOW, &self.0); - } -} - -fn create_sync_ssh_command( - target: &Target, - modifiers: SubCommandModifiers, -) -> Result { - let mut command = portable_pty::CommandBuilder::new("ssh"); - command.args(target.create_ssh_args(modifiers, false, false)?); - command.arg(target.get_preferred_host()?.to_string()); - Ok(command) -} - -#[instrument(skip_all, name = "log", parent = arguments.span)] -fn dynamic_watch_sudo_stdout(arguments: WatchStdoutArguments) -> Result<(), CommandError> { - let WatchStdoutArguments { - began_tx, - mut reader, - succeed_needle, - failed_needle, - start_needle, - output_mode, - stdout_collection, - stderr_collection, - completion_status, - log_stdout, - .. - } = arguments; - - let aho_corasick = AhoCorasick::builder() - .ascii_case_insensitive(false) - .match_kind(aho_corasick::MatchKind::LeftmostFirst) - .build([ - start_needle.as_ref(), - succeed_needle.as_ref(), - failed_needle.as_ref(), - ]) - .unwrap(); - - let mut buffer = [0u8; 1024]; - let mut stderr = std::io::stderr(); - let mut began = false; - let mut log_buffer = LogBuffer::new(); - let mut raw_mode_buffer = Vec::new(); - let mut belled = false; - - 'outer: loop { - match reader.read(&mut buffer) { - Ok(0) => break 'outer, - Ok(n) => { - if !began { - let findings = handle_rawmode_data( - &mut stderr, - &buffer, - n, - &mut raw_mode_buffer, - &aho_corasick, - &completion_status, - &began_tx, - )?; - - match findings { - SearchFindings::Terminate => break 'outer, - SearchFindings::Started => { - began = true; - continue; - } - SearchFindings::None => {} - } - - if belled { - continue; - } - - stderr - .write(b"\x07") - .map_err(CommandError::WritingClientStderr)?; - stderr.flush().map_err(CommandError::WritingClientStderr)?; - - belled = true; - - continue; - } - - log_buffer.process_slice(&buffer[..n]); - - while let Some(mut line) = log_buffer.next_line() { - let findings = - search_string(&aho_corasick, &line, &completion_status, &began_tx); - - match findings { - SearchFindings::Terminate => break 'outer, - SearchFindings::Started => { - began = true; - continue; - } - SearchFindings::None => {} - } - - handle_normal_data( - &stderr_collection, - &stdout_collection, - &mut line, - log_stdout, - output_mode, - ); - } - } - Err(e) => { - eprintln!("Error reading from PTY: {e}"); - break; - } - } - } - - let _ = began_tx.send(()); - - // failsafe if there were errors or the reader stopped - if !*completion_status.completed.lock().unwrap() { - completion_status.mark_completed(false); - } - - debug!("stdout: goodbye"); - - Ok(()) -} - -fn handle_normal_data( - stderr_collection: &Arc>>, - stdout_collection: &Arc>>, - line: &mut [u8], - log_stdout: bool, - output_mode: ChildOutputMode, -) { - if line.starts_with(b"#") { - let stripped = &mut line[1..]; - - if log_stdout { - output_mode.trace_slice(stripped); - } - - let mut queue = stdout_collection.lock().unwrap(); - queue.push_front(String::from_utf8_lossy(stripped).to_string()); - return; - } - - let log = output_mode.trace_slice(line); - - if let Some(error_msg) = log { - let mut queue = stderr_collection.lock().unwrap(); - - // add at most 20 message to the front, drop the rest. - queue.push_front(error_msg); - queue.truncate(20); - } -} - -fn handle_rawmode_data( - stderr: &mut W, - buffer: &[u8], - n: usize, - raw_mode_buffer: &mut Vec, - aho_corasick: &AhoCorasick, - completion_status: &CompletionStatus, - began_tx: &Sender<()>, -) -> Result { - raw_mode_buffer.extend_from_slice(&buffer[..n]); - - let findings = search_string(aho_corasick, raw_mode_buffer, completion_status, began_tx); - - if !matches!(findings, SearchFindings::None) { - return Ok(findings); - } - - stderr - .write_all(&buffer[..n]) - .map_err(CommandError::WritingClientStderr)?; - - stderr.flush().map_err(CommandError::WritingClientStderr)?; - - Ok(findings) -} - -/// returns true if the command is considered stopped -fn search_string( - aho_corasick: &AhoCorasick, - haystack: &[u8], - completion_status: &CompletionStatus, - began_tx: &Sender<()>, -) -> SearchFindings { - let searched = aho_corasick - .find_iter(haystack) - .map(|x| x.pattern()) - .collect::>(); - - let started = if searched.contains(&STARTED_PATTERN) { - debug!("start needle was found, switching mode..."); - let _ = began_tx.send(()); - true - } else { - false - }; - - let succeeded = if searched.contains(&SUCCEEDED_PATTERN) { - debug!("succeed needle was found, marking child as succeeding."); - completion_status.mark_completed(true); - true - } else { - false - }; - - let failed = if searched.contains(&FAILED_PATTERN) { - debug!("failed needle was found, elevated child did not succeed."); - completion_status.mark_completed(false); - true - } else { - false - }; - - if succeeded || failed { - return SearchFindings::Terminate; - } - - if started { - return SearchFindings::Started; - } - - SearchFindings::None -} - -/// Exits on any data written to `cancel_pipe_r` -#[instrument(skip_all, level = "trace", parent = span)] -fn watch_stdin_from_user( - cancel_pipe_r: &OwnedFd, - mut master_writer: MasterWriter, - write_pipe_r: &OwnedFd, - span: Span, -) -> Result<(), CommandError> { - const WRITER_POSITION: usize = 0; - const SIGNAL_POSITION: usize = 1; - const USER_POSITION: usize = 2; - - let mut buffer = [0u8; 1024]; - let stdin = std::io::stdin(); - let mut cancel_pipe_buf = [0u8; 1]; - - let user_stdin_fd = std::os::fd::AsFd::as_fd(&stdin); - let cancel_pipe_r_fd = cancel_pipe_r.as_fd(); - - let mut all_fds = vec![ - PollFd::new(write_pipe_r.as_fd(), PollFlags::POLLIN), - PollFd::new(cancel_pipe_r.as_fd(), PollFlags::POLLIN), - PollFd::new(user_stdin_fd, PollFlags::POLLIN), - ]; - - loop { - match poll(&mut all_fds, PollTimeout::NONE) { - Ok(0) => {} // timeout, impossible - Ok(_) => { - // The user stdin pipe can be removed - if all_fds.get(USER_POSITION).is_some() - && let Some(events) = all_fds[USER_POSITION].revents() - && events.contains(PollFlags::POLLIN) - { - trace!("Got stdin from user..."); - let n = - posix_read(user_stdin_fd, &mut buffer).map_err(CommandError::PosixPipe)?; - master_writer - .write_all(&buffer[..n]) - .map_err(CommandError::WritingMasterStdout)?; - master_writer - .flush() - .map_err(CommandError::WritingMasterStdout)?; - } - - if let Some(events) = all_fds[WRITER_POSITION].revents() - && events.contains(PollFlags::POLLIN) - { - trace!("Got stdin from writer..."); - let n = - posix_read(write_pipe_r, &mut buffer).map_err(CommandError::PosixPipe)?; - master_writer - .write_all(&buffer[..n]) - .map_err(CommandError::WritingMasterStdout)?; - master_writer - .flush() - .map_err(CommandError::WritingMasterStdout)?; - } - - if let Some(events) = all_fds[SIGNAL_POSITION].revents() - && events.contains(PollFlags::POLLIN) - { - let n = posix_read(cancel_pipe_r_fd, &mut cancel_pipe_buf) - .map_err(CommandError::PosixPipe)?; - let message = &cancel_pipe_buf[..n]; - - trace!("Got byte from signal pipe: {message:?}"); - - if message == THREAD_QUIT_SIGNAL { - return Ok(()); - } - - if message == THREAD_BEGAN_SIGNAL { - all_fds.remove(USER_POSITION); - } - } - } - Err(e) => { - error!("Poll error: {e}"); - break; - } - } - } - - debug!("stdin_thread: goodbye"); - Ok(()) -} - -#[cfg(test)] -mod tests { - use super::*; - use std::{assert_matches::assert_matches, sync::mpsc::TryRecvError}; - - #[test] - fn test_rawmode_data() { - let aho_corasick = AhoCorasick::builder() - .ascii_case_insensitive(false) - .match_kind(aho_corasick::MatchKind::LeftmostFirst) - .build(["START_NEEDLE", "SUCCEEDED_NEEDLE", "FAILED_NEEDLE"]) - .unwrap(); - let mut stderr = vec![]; - let (began_tx, began_rx) = mpsc::channel::<()>(); - let completion_status = CompletionStatus::new(); - - // each "Bla" is 4 bytes. - let buffer = "bla bla bla START_NEEDLE bla bla bla".as_bytes(); - let mut raw_mode_buffer = vec![]; - - // handle 1 "bla" - assert_matches!( - handle_rawmode_data( - &mut stderr, - buffer, - 4, - &mut raw_mode_buffer, - &aho_corasick, - &completion_status, - &began_tx - ), - Ok(SearchFindings::None) - ); - assert_eq!(raw_mode_buffer, b"bla "); - assert_matches!(began_rx.try_recv(), Err(TryRecvError::Empty)); - assert!(!*completion_status.completed.lock().unwrap()); - - let buffer = &buffer[4..]; - - // handle 2 "bla"'s and half a "START_NEEDLE" - let n = 4 + 4 + 6; - assert_matches!( - handle_rawmode_data( - &mut stderr, - buffer, - n, - &mut raw_mode_buffer, - &aho_corasick, - &completion_status, - &began_tx - ), - Ok(SearchFindings::None) - ); - assert_matches!(began_rx.try_recv(), Err(TryRecvError::Empty)); - assert_eq!(raw_mode_buffer, b"bla bla bla START_"); - assert!(!*completion_status.completed.lock().unwrap()); - - let buffer = &buffer[n..]; - - // handle rest of the data - let n = buffer.len(); - assert_matches!( - handle_rawmode_data( - &mut stderr, - buffer, - n, - &mut raw_mode_buffer, - &aho_corasick, - &completion_status, - &began_tx - ), - Ok(SearchFindings::Started) - ); - assert_matches!(began_rx.try_recv(), Ok(())); - assert_eq!(raw_mode_buffer, b"bla bla bla START_NEEDLE bla bla bla"); - assert!(!*completion_status.completed.lock().unwrap()); - - // test failed needle - let buffer = "bla FAILED_NEEDLE bla".as_bytes(); - let mut raw_mode_buffer = vec![]; - - let n = buffer.len(); - assert_matches!( - handle_rawmode_data( - &mut stderr, - buffer, - n, - &mut raw_mode_buffer, - &aho_corasick, - &completion_status, - &began_tx - ), - Ok(SearchFindings::Terminate) - ); - assert_matches!(*completion_status.success.lock().unwrap(), Some(false)); - - // test succeed needle - let buffer = "bla SUCCEEDED_NEEDLE bla".as_bytes(); - let mut raw_mode_buffer = vec![]; - let completion_status = CompletionStatus::new(); - - let n = buffer.len(); - assert_matches!( - handle_rawmode_data( - &mut stderr, - buffer, - n, - &mut raw_mode_buffer, - &aho_corasick, - &completion_status, - &began_tx - ), - Ok(SearchFindings::Terminate) - ); - assert_matches!(*completion_status.success.lock().unwrap(), Some(true)); - } -} diff --git a/wire/lib/src/commands/mod.rs b/wire/lib/src/commands/mod.rs index 08e87a8..3103280 100644 --- a/wire/lib/src/commands/mod.rs +++ b/wire/lib/src/commands/mod.rs @@ -1,6 +1,7 @@ // SPDX-License-Identifier: AGPL-3.0-or-later // Copyright 2024-2025 wire Contributors +use crate::commands::pty::{InteractiveChildChip, interactive_command_with_env}; use std::{collections::HashMap, str::from_utf8, sync::LazyLock}; use aho_corasick::AhoCorasick; @@ -12,18 +13,14 @@ use tracing::{debug, error, info, trace, warn}; use crate::{ SubCommandModifiers, - commands::{ - interactive::{InteractiveChildChip, interactive_command_with_env}, - noninteractive::{NonInteractiveChildChip, non_interactive_command_with_env}, - }, + commands::noninteractive::{NonInteractiveChildChip, non_interactive_command_with_env}, errors::{CommandError, HiveLibError}, hive::node::{Node, Target}, }; pub(crate) mod common; -pub(crate) mod interactive; -pub(crate) mod interactive_logbuffer; pub(crate) mod noninteractive; +pub(crate) mod pty; #[derive(Copy, Clone, Debug)] pub(crate) enum ChildOutputMode { @@ -101,13 +98,13 @@ impl<'a, S: AsRef> CommandArguments<'a, S> { } } -pub(crate) fn run_command>( +pub(crate) async fn run_command>( arguments: &CommandArguments<'_, S>, ) -> Result, HiveLibError> { - run_command_with_env(arguments, HashMap::new()) + run_command_with_env(arguments, HashMap::new()).await } -pub(crate) fn run_command_with_env>( +pub(crate) async fn run_command_with_env>( arguments: &CommandArguments<'_, S>, envs: HashMap, ) -> Result, HiveLibError> { @@ -118,7 +115,9 @@ pub(crate) fn run_command_with_env>( )?)); } - Ok(Either::Left(interactive_command_with_env(arguments, envs)?)) + Ok(Either::Left( + interactive_command_with_env(arguments, envs).await?, + )) } pub(crate) trait WireCommandChip { diff --git a/wire/lib/src/commands/pty/input.rs b/wire/lib/src/commands/pty/input.rs new file mode 100644 index 0000000..7c96ba2 --- /dev/null +++ b/wire/lib/src/commands/pty/input.rs @@ -0,0 +1,102 @@ +// SPDX-License-Identifier: AGPL-3.0-or-later +// Copyright 2024-2025 wire Contributors + +use std::os::fd::{AsFd, OwnedFd}; + +use nix::{ + poll::{PollFd, PollFlags, PollTimeout, poll}, + unistd::read, +}; +use tracing::{Span, debug, error, instrument, trace}; + +use crate::{ + commands::pty::{MasterWriter, THREAD_BEGAN_SIGNAL, THREAD_QUIT_SIGNAL}, + errors::CommandError, +}; + +/// Exits on any data written to `cancel_pipe_r` +/// A pipe is used to cancel the function. +#[instrument(skip_all, level = "trace", parent = span)] +pub(super) fn watch_stdin_from_user( + cancel_pipe_r: &OwnedFd, + mut master_writer: MasterWriter, + write_pipe_r: &OwnedFd, + span: Span, +) -> Result<(), CommandError> { + const WRITER_POSITION: usize = 0; + const SIGNAL_POSITION: usize = 1; + const USER_POSITION: usize = 2; + + let mut buffer = [0u8; 1024]; + let stdin = std::io::stdin(); + let mut cancel_pipe_buf = [0u8; 1]; + + let user_stdin_fd = stdin.as_fd(); + let cancel_pipe_r_fd = cancel_pipe_r.as_fd(); + + let mut all_fds = vec![ + PollFd::new(write_pipe_r.as_fd(), PollFlags::POLLIN), + PollFd::new(cancel_pipe_r.as_fd(), PollFlags::POLLIN), + PollFd::new(user_stdin_fd, PollFlags::POLLIN), + ]; + + loop { + match poll(&mut all_fds, PollTimeout::NONE) { + Ok(0) => {} // timeout, impossible + Ok(_) => { + // The user stdin pipe can be removed + if all_fds.get(USER_POSITION).is_some() + && let Some(events) = all_fds[USER_POSITION].revents() + && events.contains(PollFlags::POLLIN) + { + trace!("Got stdin from user..."); + let n = read(user_stdin_fd, &mut buffer).map_err(CommandError::PosixPipe)?; + master_writer + .write_all(&buffer[..n]) + .map_err(CommandError::WritingMasterStdout)?; + master_writer + .flush() + .map_err(CommandError::WritingMasterStdout)?; + } + + if let Some(events) = all_fds[WRITER_POSITION].revents() + && events.contains(PollFlags::POLLIN) + { + trace!("Got stdin from writer..."); + let n = read(write_pipe_r, &mut buffer).map_err(CommandError::PosixPipe)?; + master_writer + .write_all(&buffer[..n]) + .map_err(CommandError::WritingMasterStdout)?; + master_writer + .flush() + .map_err(CommandError::WritingMasterStdout)?; + } + + if let Some(events) = all_fds[SIGNAL_POSITION].revents() + && events.contains(PollFlags::POLLIN) + { + let n = read(cancel_pipe_r_fd, &mut cancel_pipe_buf) + .map_err(CommandError::PosixPipe)?; + let message = &cancel_pipe_buf[..n]; + + trace!("Got byte from signal pipe: {message:?}"); + + if message == THREAD_QUIT_SIGNAL { + return Ok(()); + } + + if message == THREAD_BEGAN_SIGNAL { + all_fds.remove(USER_POSITION); + } + } + } + Err(e) => { + error!("Poll error: {e}"); + break; + } + } + } + + debug!("stdin_thread: goodbye"); + Ok(()) +} diff --git a/wire/lib/src/commands/interactive_logbuffer.rs b/wire/lib/src/commands/pty/logbuffer.rs similarity index 100% rename from wire/lib/src/commands/interactive_logbuffer.rs rename to wire/lib/src/commands/pty/logbuffer.rs diff --git a/wire/lib/src/commands/pty/mod.rs b/wire/lib/src/commands/pty/mod.rs new file mode 100644 index 0000000..f904fd5 --- /dev/null +++ b/wire/lib/src/commands/pty/mod.rs @@ -0,0 +1,560 @@ +// SPDX-License-Identifier: AGPL-3.0-or-later +// Copyright 2024-2025 wire Contributors + +use crate::commands::pty::output::{WatchStdoutArguments, handle_pty_stdout}; +use aho_corasick::PatternID; +use itertools::Itertools; +use nix::sys::termios::{LocalFlags, SetArg, Termios, tcgetattr, tcsetattr}; +use nix::unistd::pipe; +use nix::unistd::write as posix_write; +use portable_pty::{CommandBuilder, NativePtySystem, PtyPair, PtySize}; +use rand::distr::Alphabetic; +use std::collections::VecDeque; +use std::sync::{LazyLock, Mutex}; +use std::{ + io::{Read, Write}, + os::fd::{AsFd, OwnedFd}, + sync::Arc, +}; +use tokio::sync::{oneshot, watch}; +use tracing::instrument; +use tracing::{Span, debug, trace}; + +use crate::commands::CommandArguments; +use crate::commands::pty::input::watch_stdin_from_user; +use crate::errors::CommandError; +use crate::{STDIN_CLOBBER_LOCK, SubCommandModifiers}; +use crate::{ + commands::{ChildOutputMode, WireCommandChip}, + errors::HiveLibError, + hive::node::Target, +}; + +mod input; +mod logbuffer; +mod output; + +type MasterWriter = Box; +type MasterReader = Box; + +/// the underlying command began +const THREAD_BEGAN_SIGNAL: &[u8; 1] = b"b"; +const THREAD_QUIT_SIGNAL: &[u8; 1] = b"q"; + +type Child = Box; + +pub(crate) struct InteractiveChildChip { + child: Child, + + cancel_stdin_pipe_w: OwnedFd, + write_stdin_pipe_w: OwnedFd, + + stderr_collection: Arc>>, + stdout_collection: Arc>>, + + original_command: String, + + status_receiver: watch::Receiver, + stdout_handle: tokio::task::JoinHandle>, +} + +/// sets and reverts terminal options (the terminal user interaction is performed) +/// reverts data when dropped +struct StdinTermiosAttrGuard(Termios); + +#[derive(Debug)] +enum Status { + Running, + Done { success: bool }, +} + +#[derive(Debug)] +enum SearchFindings { + None, + Started, + Terminate, +} + +static STARTED_PATTERN: LazyLock = LazyLock::new(|| PatternID::must(0)); +static SUCCEEDED_PATTERN: LazyLock = LazyLock::new(|| PatternID::must(1)); +static FAILED_PATTERN: LazyLock = LazyLock::new(|| PatternID::must(2)); + +/// substitutes STDOUT with #$line. stdout is far less common than stderr. +const IO_SUBS: &str = "1> >(while IFS= read -r line; do echo \"#$line\"; done)"; + +fn create_ending_segment>( + arguments: &CommandArguments<'_, S>, + needles: &Needles, +) -> String { + let Needles { + succeed, + fail, + start, + } = needles; + + format!( + "echo -e '{succeed}' || echo '{failed}'", + succeed = if matches!(arguments.output_mode, ChildOutputMode::Interactive) { + format!( + "{start}\\n{succeed}", + start = String::from_utf8_lossy(start), + succeed = String::from_utf8_lossy(succeed) + ) + } else { + String::from_utf8_lossy(succeed).to_string() + }, + failed = String::from_utf8_lossy(fail) + ) +} + +fn create_starting_segment>( + arguments: &CommandArguments<'_, S>, + start_needle: &Arc>, +) -> String { + if matches!(arguments.output_mode, ChildOutputMode::Interactive) { + String::new() + } else { + format!( + "echo '{start}' && ", + start = String::from_utf8_lossy(start_needle) + ) + } +} + +#[instrument(skip_all, name = "run-int", fields(elevated = %arguments.is_elevated()))] +pub(crate) async fn interactive_command_with_env>( + arguments: &CommandArguments<'_, S>, + envs: std::collections::HashMap, +) -> Result { + print_authenticate_warning(arguments)?; + + let needles = create_needles(); + let pty_system = NativePtySystem::default(); + let pty_pair = portable_pty::PtySystem::openpty(&pty_system, PtySize::default()).unwrap(); + setup_master(&pty_pair)?; + + let command_string = &format!( + "{starting}{command} {flags} {IO_SUBS} && {ending}", + command = arguments.command_string.as_ref(), + flags = match arguments.output_mode { + ChildOutputMode::Nix => "--log-format internal-json", + ChildOutputMode::Generic | ChildOutputMode::Interactive => "", + }, + starting = create_starting_segment(arguments, &needles.start), + ending = create_ending_segment(arguments, &needles) + ); + + debug!("{command_string}"); + + let mut command = build_command(arguments, command_string)?; + + // give command all env vars + for (key, value) in envs { + command.env(key, value); + } + + let clobber_guard = STDIN_CLOBBER_LOCK.acquire().await.unwrap(); + let _guard = StdinTermiosAttrGuard::new().map_err(HiveLibError::CommandError)?; + let child = pty_pair + .slave + .spawn_command(command) + .map_err(|x| HiveLibError::CommandError(CommandError::PortablePty(x)))?; + + // Release any handles owned by the slave: we don't need it now + // that we've spawned the child. + drop(pty_pair.slave); + + let reader = pty_pair + .master + .try_clone_reader() + .map_err(|x| HiveLibError::CommandError(CommandError::PortablePty(x)))?; + let master_writer = pty_pair + .master + .take_writer() + .map_err(|x| HiveLibError::CommandError(CommandError::PortablePty(x)))?; + + let stderr_collection = Arc::new(Mutex::new(VecDeque::::with_capacity(10))); + let stdout_collection = Arc::new(Mutex::new(VecDeque::::with_capacity(10))); + let (began_tx, began_rx) = oneshot::channel::<()>(); + let (status_sender, status_receiver) = watch::channel(Status::Running); + + let stdout_handle = { + let arguments = WatchStdoutArguments { + began_tx, + reader, + needles, + output_mode: arguments.output_mode, + stderr_collection: stderr_collection.clone(), + stdout_collection: stdout_collection.clone(), + span: Span::current(), + log_stdout: arguments.log_stdout, + status_sender, + }; + + tokio::task::spawn_blocking(move || handle_pty_stdout(arguments)) + }; + + let (write_stdin_pipe_r, write_stdin_pipe_w) = + pipe().map_err(|x| HiveLibError::CommandError(CommandError::PosixPipe(x)))?; + let (cancel_stdin_pipe_r, cancel_stdin_pipe_w) = + pipe().map_err(|x| HiveLibError::CommandError(CommandError::PosixPipe(x)))?; + + tokio::task::spawn_blocking(move || { + watch_stdin_from_user( + &cancel_stdin_pipe_r, + master_writer, + &write_stdin_pipe_r, + Span::current(), + ) + }); + + debug!("Setup threads"); + + let () = began_rx + .await + .map_err(|x| HiveLibError::CommandError(CommandError::OneshotRecvError(x)))?; + + drop(clobber_guard); + + if arguments.keep_stdin_open { + trace!("Sending THREAD_BEGAN_SIGNAL"); + + posix_write(&cancel_stdin_pipe_w, THREAD_BEGAN_SIGNAL) + .map_err(|x| HiveLibError::CommandError(CommandError::PosixPipe(x)))?; + } else { + trace!("Sending THREAD_QUIT_SIGNAL"); + + posix_write(&cancel_stdin_pipe_w, THREAD_QUIT_SIGNAL) + .map_err(|x| HiveLibError::CommandError(CommandError::PosixPipe(x)))?; + } + + Ok(InteractiveChildChip { + child, + cancel_stdin_pipe_w, + write_stdin_pipe_w, + stderr_collection, + stdout_collection, + original_command: arguments.command_string.as_ref().to_string(), + status_receiver, + stdout_handle, + }) +} + +fn print_authenticate_warning>( + arguments: &CommandArguments, +) -> Result<(), HiveLibError> { + if !arguments.is_elevated() { + return Ok(()); + } + + eprintln!( + "{} | Authenticate for \"sudo {}\":", + arguments + .target + .map_or(Ok("localhost (!)".to_string()), |target| Ok(format!( + "{}@{}:{}", + target.user, + target.get_preferred_host()?, + target.port + )))?, + arguments.command_string.as_ref() + ); + + Ok(()) +} + +struct Needles { + succeed: Arc>, + fail: Arc>, + start: Arc>, +} + +fn create_needles() -> Needles { + let tmp_prefix = rand::distr::SampleString::sample_string(&Alphabetic, &mut rand::rng(), 5); + + Needles { + succeed: Arc::new(format!("{tmp_prefix}_W_Q").as_bytes().to_vec()), + fail: Arc::new(format!("{tmp_prefix}_W_F").as_bytes().to_vec()), + start: Arc::new(format!("{tmp_prefix}_W_S").as_bytes().to_vec()), + } +} + +fn setup_master(pty_pair: &PtyPair) -> Result<(), HiveLibError> { + if let Some(fd) = pty_pair.master.as_raw_fd() { + // convert raw fd to a BorrowedFd + // safe as `fd` is dropped well before `pty_pair.master` + let fd = unsafe { std::os::unix::io::BorrowedFd::borrow_raw(fd) }; + let mut termios = + tcgetattr(fd).map_err(|x| HiveLibError::CommandError(CommandError::TermAttrs(x)))?; + + termios.local_flags &= !LocalFlags::ECHO; + // Key agent does not work well without canonical mode + termios.local_flags &= !LocalFlags::ICANON; + // Actually quit + termios.local_flags &= !LocalFlags::ISIG; + + tcsetattr(fd, SetArg::TCSANOW, &termios) + .map_err(|x| HiveLibError::CommandError(CommandError::TermAttrs(x)))?; + } + + Ok(()) +} + +fn build_command>( + arguments: &CommandArguments<'_, S>, + command_string: &String, +) -> Result { + let mut command = if let Some(target) = arguments.target { + let mut command = create_int_ssh_command(target, arguments.modifiers)?; + + // force ssh to use our pesudo terminal + command.arg("-tt"); + + command + } else { + let mut command = portable_pty::CommandBuilder::new("sh"); + + command.arg("-c"); + + command + }; + + if arguments.is_elevated() { + command.arg(format!("sudo -u root -- sh -c '{command_string}'")); + } else { + command.arg(command_string); + } + + Ok(command) +} + +impl WireCommandChip for InteractiveChildChip { + type ExitStatus = (portable_pty::ExitStatus, String); + + #[instrument(skip_all)] + async fn wait_till_success(mut self) -> Result { + drop(self.write_stdin_pipe_w); + + let exit_status = tokio::task::spawn_blocking(move || self.child.wait()) + .await + .map_err(CommandError::JoinError)? + .map_err(CommandError::WaitForStatus)?; + + debug!("exit_status: {exit_status:?}"); + + self.stdout_handle + .await + .map_err(|_| CommandError::ThreadPanic)??; + + let status = self + .status_receiver + .wait_for(|value| matches!(value, Status::Done { .. })) + .await + .unwrap(); + + let _ = posix_write(&self.cancel_stdin_pipe_w, THREAD_QUIT_SIGNAL); + + if let Status::Done { success: true } = *status { + let logs = self + .stdout_collection + .lock() + .unwrap() + .iter() + .rev() + .map(|x| x.trim()) + .join("\n"); + + return Ok((exit_status, logs)); + } + + debug!("child did not succeed"); + + let logs = self + .stderr_collection + .lock() + .unwrap() + .iter() + .rev() + .join("\n"); + + Err(CommandError::CommandFailed { + command_ran: self.original_command, + logs, + code: format!("code {}", exit_status.exit_code()), + reason: match *status { + Status::Done { .. } => "marked-unsuccessful", + Status::Running => "child-crashed-before-succeeding", + }, + }) + } + + async fn write_stdin(&mut self, data: Vec) -> Result<(), HiveLibError> { + trace!("Writing {} bytes to stdin", data.len()); + + posix_write(&self.write_stdin_pipe_w, &data) + .map_err(|x| HiveLibError::CommandError(CommandError::PosixPipe(x)))?; + + Ok(()) + } +} + +impl StdinTermiosAttrGuard { + fn new() -> Result { + let stdin = std::io::stdin(); + let stdin_fd = stdin.as_fd(); + + let mut termios = tcgetattr(stdin_fd).map_err(CommandError::TermAttrs)?; + let original_termios = termios.clone(); + + termios.local_flags &= !(LocalFlags::ECHO | LocalFlags::ICANON); + tcsetattr(stdin_fd, SetArg::TCSANOW, &termios).map_err(CommandError::TermAttrs)?; + + Ok(StdinTermiosAttrGuard(original_termios)) + } +} + +impl Drop for StdinTermiosAttrGuard { + fn drop(&mut self) { + let stdin = std::io::stdin(); + let stdin_fd = stdin.as_fd(); + + let _ = tcsetattr(stdin_fd, SetArg::TCSANOW, &self.0); + } +} + +fn create_int_ssh_command( + target: &Target, + modifiers: SubCommandModifiers, +) -> Result { + let mut command = portable_pty::CommandBuilder::new("ssh"); + command.args(target.create_ssh_args(modifiers, false, false)?); + command.arg(target.get_preferred_host()?.to_string()); + Ok(command) +} + +#[cfg(test)] +mod tests { + use aho_corasick::AhoCorasick; + use tokio::sync::oneshot::error::TryRecvError; + + use crate::commands::pty::output::handle_rawmode_data; + + use super::*; + use std::assert_matches::assert_matches; + + #[test] + fn test_rawmode_data() { + let aho_corasick = AhoCorasick::builder() + .ascii_case_insensitive(false) + .match_kind(aho_corasick::MatchKind::LeftmostFirst) + .build(["START_NEEDLE", "SUCCEEDED_NEEDLE", "FAILED_NEEDLE"]) + .unwrap(); + let mut stderr = vec![]; + let (began_tx, mut began_rx) = oneshot::channel::<()>(); + let mut began_tx = Some(began_tx); + let (status_sender, _) = watch::channel(Status::Running); + + // each "Bla" is 4 bytes. + let buffer = "bla bla bla START_NEEDLE bla bla bla".as_bytes(); + let mut raw_mode_buffer = vec![]; + + // handle 1 "bla" + assert_matches!( + handle_rawmode_data( + &mut stderr, + buffer, + 4, + &mut raw_mode_buffer, + &aho_corasick, + &status_sender, + &mut began_tx + ), + Ok(SearchFindings::None) + ); + assert_matches!(began_rx.try_recv(), Err(TryRecvError::Empty)); + assert!(began_tx.is_some()); + assert_eq!(raw_mode_buffer, b"bla "); + assert_matches!(*status_sender.borrow(), Status::Running); + + let buffer = &buffer[4..]; + + // handle 2 "bla"'s and half a "START_NEEDLE" + let n = 4 + 4 + 6; + assert_matches!( + handle_rawmode_data( + &mut stderr, + buffer, + n, + &mut raw_mode_buffer, + &aho_corasick, + &status_sender, + &mut began_tx + ), + Ok(SearchFindings::None) + ); + assert_matches!(began_rx.try_recv(), Err(TryRecvError::Empty)); + assert!(began_tx.is_some()); + assert_matches!(*status_sender.borrow(), Status::Running); + assert_eq!(raw_mode_buffer, b"bla bla bla START_"); + + let buffer = &buffer[n..]; + + // handle rest of the data + let n = buffer.len(); + assert_matches!( + handle_rawmode_data( + &mut stderr, + buffer, + n, + &mut raw_mode_buffer, + &aho_corasick, + &status_sender, + &mut began_tx + ), + Ok(SearchFindings::Started) + ); + assert_matches!(began_rx.try_recv(), Ok(())); + assert_matches!(began_tx, None); + assert_eq!(raw_mode_buffer, b"bla bla bla START_NEEDLE bla bla bla"); + assert_matches!(*status_sender.borrow(), Status::Running); + + // test failed needle + let buffer = "bla FAILED_NEEDLE bla".as_bytes(); + let mut raw_mode_buffer = vec![]; + + let n = buffer.len(); + assert_matches!( + handle_rawmode_data( + &mut stderr, + buffer, + n, + &mut raw_mode_buffer, + &aho_corasick, + &status_sender, + &mut began_tx + ), + Ok(SearchFindings::Terminate) + ); + assert_matches!(*status_sender.borrow(), Status::Done { success: false }); + + // test succeed needle + let buffer = "bla SUCCEEDED_NEEDLE bla".as_bytes(); + let mut raw_mode_buffer = vec![]; + let (status_sender, _) = watch::channel(Status::Running); + + let n = buffer.len(); + assert_matches!( + handle_rawmode_data( + &mut stderr, + buffer, + n, + &mut raw_mode_buffer, + &aho_corasick, + &status_sender, + &mut began_tx + ), + Ok(SearchFindings::Terminate) + ); + assert_matches!(*status_sender.borrow(), Status::Done { success: true }); + } +} diff --git a/wire/lib/src/commands/pty/output.rs b/wire/lib/src/commands/pty/output.rs new file mode 100644 index 0000000..6b010ea --- /dev/null +++ b/wire/lib/src/commands/pty/output.rs @@ -0,0 +1,259 @@ +// SPDX-License-Identifier: AGPL-3.0-or-later +// Copyright 2024-2025 wire Contributors + +use crate::{ + commands::{ + ChildOutputMode, + pty::{ + FAILED_PATTERN, Needles, STARTED_PATTERN, SUCCEEDED_PATTERN, SearchFindings, Status, + logbuffer::LogBuffer, + }, + }, + errors::CommandError, +}; +use aho_corasick::AhoCorasick; +use std::{ + collections::VecDeque, + io::Write, + sync::{Arc, Mutex}, +}; +use tokio::sync::{oneshot, watch}; +use tracing::{Span, debug, instrument}; + +pub(super) struct WatchStdoutArguments { + pub began_tx: oneshot::Sender<()>, + pub reader: super::MasterReader, + pub needles: Needles, + pub output_mode: ChildOutputMode, + pub stderr_collection: Arc>>, + pub stdout_collection: Arc>>, + pub status_sender: watch::Sender, + pub span: Span, + pub log_stdout: bool, +} + +/// Handles data from the PTY, and logs or prompts the user depending on the state +/// of the command. +/// +/// Emits a message on the `began_tx` when the command is considered started. +/// +/// Records stderr and stdout when it is considered notable (all stdout, last few stderr messages) +#[instrument(skip_all, name = "log", parent = arguments.span)] +pub(super) fn handle_pty_stdout(arguments: WatchStdoutArguments) -> Result<(), CommandError> { + let WatchStdoutArguments { + began_tx, + mut reader, + needles, + output_mode, + stdout_collection, + stderr_collection, + status_sender, + log_stdout, + .. + } = arguments; + + let aho_corasick = AhoCorasick::builder() + .ascii_case_insensitive(false) + .match_kind(aho_corasick::MatchKind::LeftmostFirst) + .build([ + needles.start.as_ref(), + needles.succeed.as_ref(), + needles.fail.as_ref(), + ]) + .unwrap(); + + let mut buffer = [0u8; 1024]; + let mut stderr = std::io::stderr(); + let mut began = false; + let mut log_buffer = LogBuffer::new(); + let mut raw_mode_buffer = Vec::new(); + let mut belled = false; + let mut began_tx = Some(began_tx); + + 'outer: loop { + match reader.read(&mut buffer) { + Ok(0) => break 'outer, + Ok(n) => { + if !began { + let findings = handle_rawmode_data( + &mut stderr, + &buffer, + n, + &mut raw_mode_buffer, + &aho_corasick, + &status_sender, + &mut began_tx, + )?; + + match findings { + SearchFindings::Terminate => break 'outer, + SearchFindings::Started => { + began = true; + continue; + } + SearchFindings::None => {} + } + + if belled { + continue; + } + + stderr + .write(b"\x07") + .map_err(CommandError::WritingClientStderr)?; + stderr.flush().map_err(CommandError::WritingClientStderr)?; + + belled = true; + + continue; + } + + log_buffer.process_slice(&buffer[..n]); + + while let Some(mut line) = log_buffer.next_line() { + let findings = + search_string(&aho_corasick, &line, &status_sender, &mut began_tx); + + match findings { + SearchFindings::Terminate => break 'outer, + SearchFindings::Started => { + began = true; + continue; + } + SearchFindings::None => {} + } + + handle_normal_data( + &stderr_collection, + &stdout_collection, + &mut line, + log_stdout, + output_mode, + ); + } + } + Err(e) => { + eprintln!("Error reading from PTY: {e}"); + break; + } + } + } + + began_tx.map(|began_tx| began_tx.send(())); + + // failsafe if there were errors or the reader stopped + if matches!(*status_sender.borrow(), Status::Running) { + status_sender.send_replace(Status::Done { success: false }); + } + + debug!("stdout: goodbye"); + + Ok(()) +} + +/// handles raw data, prints to stderr when a prompt is detected +pub(super) fn handle_rawmode_data( + stderr: &mut W, + buffer: &[u8], + n: usize, + raw_mode_buffer: &mut Vec, + aho_corasick: &AhoCorasick, + status_sender: &watch::Sender, + began_tx: &mut Option>, +) -> Result { + raw_mode_buffer.extend_from_slice(&buffer[..n]); + + let findings = search_string(aho_corasick, raw_mode_buffer, status_sender, began_tx); + + if !matches!(findings, SearchFindings::None) { + return Ok(findings); + } + + stderr + .write_all(&buffer[..n]) + .map_err(CommandError::WritingClientStderr)?; + + stderr.flush().map_err(CommandError::WritingClientStderr)?; + + Ok(findings) +} + +/// handles data when the command is considered "started", logs and records errors as appropriate +fn handle_normal_data( + stderr_collection: &Arc>>, + stdout_collection: &Arc>>, + line: &mut [u8], + log_stdout: bool, + output_mode: ChildOutputMode, +) { + if line.starts_with(b"#") { + let stripped = &mut line[1..]; + + if log_stdout { + output_mode.trace_slice(stripped); + } + + let mut queue = stdout_collection.lock().unwrap(); + queue.push_front(String::from_utf8_lossy(stripped).to_string()); + return; + } + + let log = output_mode.trace_slice(line); + + if let Some(error_msg) = log { + let mut queue = stderr_collection.lock().unwrap(); + + // add at most 20 message to the front, drop the rest. + queue.push_front(error_msg); + queue.truncate(20); + } +} + +/// returns true if the command is considered stopped +fn search_string( + aho_corasick: &AhoCorasick, + haystack: &[u8], + status_sender: &watch::Sender, + began_tx: &mut Option>, +) -> SearchFindings { + let searched = aho_corasick + .find_iter(haystack) + .map(|x| x.pattern()) + .collect::>(); + + let started = if searched.contains(&STARTED_PATTERN) { + debug!("start needle was found, switching mode..."); + if let Some(began_tx) = began_tx.take() { + let _ = began_tx.send(()); + } + true + } else { + false + }; + + let succeeded = if searched.contains(&SUCCEEDED_PATTERN) { + debug!("succeed needle was found, marking child as succeeding."); + status_sender.send_replace(Status::Done { success: true }); + true + } else { + false + }; + + let failed = if searched.contains(&FAILED_PATTERN) { + debug!("failed needle was found, elevated child did not succeed."); + status_sender.send_replace(Status::Done { success: false }); + true + } else { + false + }; + + if succeeded || failed { + return SearchFindings::Terminate; + } + + if started { + return SearchFindings::Started; + } + + SearchFindings::None +} diff --git a/wire/lib/src/errors.rs b/wire/lib/src/errors.rs index cdc16cb..f4ffb1a 100644 --- a/wire/lib/src/errors.rs +++ b/wire/lib/src/errors.rs @@ -276,6 +276,13 @@ pub enum CommandError { )] #[error("$XDG_RUNTIME_DIR could not be used.")] RuntimeDirectoryMissing(#[source] std::env::VarError), + + #[diagnostic( + code(wire::command::OneshotRecvError), + url("{DOCS_URL}#{}", self.code().unwrap()) + )] + #[error("Error waiting for begin message")] + OneshotRecvError(#[source] tokio::sync::oneshot::error::RecvError), } #[derive(Debug, Diagnostic, Error)] diff --git a/wire/lib/src/hive/node.rs b/wire/lib/src/hive/node.rs index 116253a..9f93ed6 100644 --- a/wire/lib/src/hive/node.rs +++ b/wire/lib/src/hive/node.rs @@ -218,7 +218,8 @@ impl Node { &CommandArguments::new(command_string, modifiers) .log_stdout() .mode(crate::commands::ChildOutputMode::Interactive), - )?; + ) + .await?; output.wait_till_success().await.map_err(|source| { HiveLibError::NetworkError(NetworkError::HostUnreachable { diff --git a/wire/lib/src/hive/steps/activate.rs b/wire/lib/src/hive/steps/activate.rs index 298c656..45b4190 100644 --- a/wire/lib/src/hive/steps/activate.rs +++ b/wire/lib/src/hive/steps/activate.rs @@ -58,7 +58,8 @@ async fn set_profile( Some(&ctx.node.target) }) .elevated(ctx.node), - )?; + ) + .await?; let _ = child .wait_till_success() @@ -113,7 +114,8 @@ impl ExecuteStep for SwitchToConfiguration { }) .elevated(ctx.node) .log_stdout(), - )?; + ) + .await?; let result = child.wait_till_success().await; @@ -136,7 +138,8 @@ impl ExecuteStep for SwitchToConfiguration { .log_stdout() .on_target(Some(&ctx.node.target)) .elevated(ctx.node), - )?; + ) + .await?; // consume result, impossible to know if the machine failed to reboot or we // simply disconnected diff --git a/wire/lib/src/hive/steps/build.rs b/wire/lib/src/hive/steps/build.rs index 64629af..25416bf 100644 --- a/wire/lib/src/hive/steps/build.rs +++ b/wire/lib/src/hive/steps/build.rs @@ -47,7 +47,8 @@ impl ExecuteStep for Build { .mode(crate::commands::ChildOutputMode::Nix) .log_stdout(), std::collections::HashMap::new(), - )? + ) + .await? .wait_till_success() .await .map_err(|source| HiveLibError::NixBuildError { diff --git a/wire/lib/src/hive/steps/keys.rs b/wire/lib/src/hive/steps/keys.rs index 7fa7263..eb7f619 100644 --- a/wire/lib/src/hive/steps/keys.rs +++ b/wire/lib/src/hive/steps/keys.rs @@ -261,7 +261,8 @@ impl ExecuteStep for Keys { .elevated(ctx.node) .keep_stdin_open() .log_stdout(), - )?; + ) + .await?; let mut writer = SimpleLengthDelimWriter::new(async |data| child.write_stdin(data).await); diff --git a/wire/lib/src/lib.rs b/wire/lib/src/lib.rs index 298e6d8..9e2e6e0 100644 --- a/wire/lib/src/lib.rs +++ b/wire/lib/src/lib.rs @@ -4,10 +4,9 @@ #![feature(assert_matches)] #![feature(iter_intersperse)] -use std::{ - io::IsTerminal, - sync::{Arc, LazyLock, Mutex}, -}; +use std::{io::IsTerminal, sync::LazyLock}; + +use tokio::sync::Semaphore; use crate::{errors::HiveLibError, hive::node::Name}; @@ -54,5 +53,4 @@ pub enum EvalGoal<'a> { GetTopLevel(&'a Name), } -pub static STDIN_CLOBBER_LOCK: LazyLock>> = - LazyLock::new(|| Arc::new(Mutex::new(()))); +pub static STDIN_CLOBBER_LOCK: LazyLock = LazyLock::new(|| Semaphore::new(1));