diff --git a/.tegami/add-opt-in-flow-control.md b/.tegami/add-opt-in-flow-control.md new file mode 100644 index 0000000..9564976 --- /dev/null +++ b/.tegami/add-opt-in-flow-control.md @@ -0,0 +1,11 @@ +--- +packages: + et: + type: patch +--- + +## Add opt-in terminal flow control + +Clients can now select lossless backpressure or oldest-output discard when +terminal output outruns the network, keeping Ctrl-C and prompt responses +bounded without changing the default session behavior. diff --git a/Cargo.lock b/Cargo.lock index e463cc8..6e34cfd 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -215,6 +215,21 @@ dependencies = [ "libc", ] +[[package]] +name = "crossbeam-channel" +version = "0.5.16" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d85363c37faeca707aef026efa9f3b34d077bce547e48f770770625c6013679e" +dependencies = [ + "crossbeam-utils", +] + +[[package]] +name = "crossbeam-utils" +version = "0.8.22" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "61803da095bee82a81bb1a452ecc25d3b2f1416d1897eb86430c6159ef717c17" + [[package]] name = "crossterm" version = "0.29.0" @@ -367,6 +382,7 @@ dependencies = [ "prost", "rustix", "signal-hook", + "socket2", "sysinfo", "wait-timeout", "windows-spawn", @@ -407,6 +423,7 @@ dependencies = [ name = "et-net" version = "0.0.19" dependencies = [ + "crossbeam-channel", "et-core", "privdrop", "prost", diff --git a/crates/et-bin/Cargo.toml b/crates/et-bin/Cargo.toml index a35c501..8a041e7 100644 --- a/crates/et-bin/Cargo.toml +++ b/crates/et-bin/Cargo.toml @@ -41,6 +41,9 @@ ctrlc = "3.5" # Safe CreateProcessW wrapper with an explicit inherited-handle allowlist. windows-spawn = "0.1.0" +[dev-dependencies] +socket2 = "0.6.1" + [target.'cfg(target_os = "linux")'.dev-dependencies] nix = { version = "0.31.3", default-features = false, features = ["fs", "poll"] } rustix = { version = "1.1.4", features = ["process"] } diff --git a/crates/et-bin/src/client.rs b/crates/et-bin/src/client.rs index f652f1c..e91e968 100644 --- a/crates/et-bin/src/client.rs +++ b/crates/et-bin/src/client.rs @@ -119,9 +119,12 @@ fn run_client( .into(), ); } - let mut locale_environment = - bounded_locale_environment(ssh_locale_environment(), &reserved_environment) - .map_err(crate::forward_config::ForwardConfigError::EnvironmentPacketTooLarge)?; + let mut locale_environment = bounded_locale_environment( + ssh_locale_environment(), + &reserved_environment, + forward_config.initial_payload.flowcontrol, + ) + .map_err(crate::forward_config::ForwardConfigError::EnvironmentPacketTooLarge)?; if args.jumphost.is_some() { bound_jumphost_locale_environment( &forward_config.initial_payload, @@ -305,6 +308,7 @@ fn run_client( command: args.command.as_deref(), no_exit: args.no_exit, keepalive: args.keepalive, + flow_control: args.flow_control, terminal_enabled: !args.no_terminal, lines: crate::client_terminal::RemoteLines::from(remote_mode.terminal_shell), connection_name: &request.host_alias, diff --git a/crates/et-bin/src/client_environment.rs b/crates/et-bin/src/client_environment.rs index 24dbd5d..81fbaa0 100644 --- a/crates/et-bin/src/client_environment.rs +++ b/crates/et-bin/src/client_environment.rs @@ -1,8 +1,8 @@ use std::collections::BTreeMap; use et_cli::tunnel::MAX_UNIX_SOCKET_PATH; -use et_core::packet::{Packet, HEADER_LEN}; -use et_core::proto::{InitialPayload, TerminalPacketType}; +use et_core::packet::Packet; +use et_core::proto::{InitialPayload, TermInit, TerminalPacketType}; use et_net::local_packet::MAX_LOCAL_PACKET_LEN; use prost::Message; @@ -83,8 +83,18 @@ pub(crate) fn locale_environment_capacity(reserved: usize) -> usize { pub(crate) fn bounded_locale_environment( candidates: impl IntoIterator, reserved: &BTreeMap, + flowcontrol: Option, ) -> Result, usize> { - let mut packet_len = HEADER_LEN; + let mut packet_len = Packet::new( + TerminalPacketType::TerminalInit as u8, + TermInit { + environmentnames: Vec::new(), + environmentvalues: Vec::new(), + flowcontrol, + } + .encode_to_vec(), + ) + .wire_len(); for (name, value_len) in reserved { packet_len = packet_len .saturating_add(encoded_string_field_len(name.len())) @@ -163,7 +173,8 @@ mod tests { }; use et_core::packet::Packet; use et_core::proto::{ - InitialPayload, PortForwardSourceRequest, SocketEndpoint, TerminalPacketType, + FlowControlMode, InitialPayload, PortForwardSourceRequest, SocketEndpoint, TermInit, + TerminalPacketType, }; use et_net::local_packet::MAX_LOCAL_PACKET_LEN; use prost::Message; @@ -189,6 +200,7 @@ mod tests { jumphost: Some(false), reversetunnels: vec![request; 128], environmentvariables: Default::default(), + flowcontrol: None, }; let mut locale = vec![ ("LC_ALL".to_owned(), "C".to_owned()), @@ -224,6 +236,71 @@ mod tests { .wire_len() } + #[test] + fn terminal_environment_budget_includes_opted_in_flow_control_bytes() { + for mode in [FlowControlMode::Backpressure, FlowControlMode::Discard] { + // Given: one locale value whose flow-control-free TermInit lands + // exactly on the local framing limit. + let reserved = std::collections::BTreeMap::new(); + let name = "LANG".to_owned(); + let value_len = (MAX_LOCAL_PACKET_LEN - 32..MAX_LOCAL_PACKET_LEN) + .find(|value_len| { + let init = TermInit { + environmentnames: vec![name.clone()], + environmentvalues: vec!["x".repeat(*value_len)], + flowcontrol: None, + }; + Packet::new(TerminalPacketType::TerminalInit as u8, init.encode_to_vec()) + .wire_len() + == MAX_LOCAL_PACKET_LEN + }) + .unwrap(); + + // When: locale selection budgets the opt-in enum field. + let selected = super::bounded_locale_environment( + [(name, "x".repeat(value_len))], + &reserved, + Some(mode as i32), + ) + .unwrap(); + + // Then: the server's actual TermInit remains within the local cap. + let init = TermInit { + environmentnames: selected.iter().map(|(name, _)| name.clone()).collect(), + environmentvalues: selected.into_iter().map(|(_, value)| value).collect(), + flowcontrol: Some(mode as i32), + }; + assert!( + Packet::new(TerminalPacketType::TerminalInit as u8, init.encode_to_vec(),) + .wire_len() + <= MAX_LOCAL_PACKET_LEN + ); + } + } + + #[test] + fn absent_flow_control_keeps_the_existing_exact_boundary() { + let reserved = std::collections::BTreeMap::new(); + let name = "LANG".to_owned(); + let value_len = (MAX_LOCAL_PACKET_LEN - 32..MAX_LOCAL_PACKET_LEN) + .find(|value_len| { + let init = TermInit { + environmentnames: vec![name.clone()], + environmentvalues: vec!["x".repeat(*value_len)], + flowcontrol: None, + }; + Packet::new(TerminalPacketType::TerminalInit as u8, init.encode_to_vec()).wire_len() + == MAX_LOCAL_PACKET_LEN + }) + .unwrap(); + + let selected = + super::bounded_locale_environment([(name, "x".repeat(value_len))], &reserved, None) + .unwrap(); + + assert_eq!(selected[0].1.len(), value_len); + } + #[test] fn locale_capacity_reserves_terminal_environment_entries() { assert_eq!(locale_environment_capacity(0), 128); diff --git a/crates/et-bin/src/client_output.rs b/crates/et-bin/src/client_output.rs new file mode 100644 index 0000000..4dbd641 --- /dev/null +++ b/crates/et-bin/src/client_output.rs @@ -0,0 +1,607 @@ +//! Bounded, nonblocking local console-output worker for opt-in flow control. + +use std::collections::VecDeque; +#[cfg(unix)] +use std::fs::File; +#[cfg(unix)] +use std::io::Read; +use std::io::{self, Write}; +#[cfg(unix)] +use std::os::fd::AsFd; +#[cfg(windows)] +use std::process::{ChildStderr, ChildStdin, Command, Stdio}; +use std::sync::{Arc, Condvar, Mutex}; +use std::thread; + +use et_cli::client::FlowControlMode; +#[cfg(unix)] +use et_net::local::LocalStream; + +const OUTPUT_BYTES: usize = 64 * 1024; +const OUTPUT_PACKETS: usize = 4096; + +struct OutputEntry { + bytes: Vec, + terminal_modes: crate::client_terminal::TerminalModeState, +} + +struct State { + queue: VecDeque, + bytes: usize, + stopping: bool, + error: Option, + cursor_reports: usize, + worker_done: bool, +} + +struct Shared { + state: Mutex, + wake: Condvar, +} + +#[derive(Clone, Copy)] +pub(crate) enum ConsoleCompletion { + RemoteSessionEnded, + #[cfg(any(unix, test))] + LocalInputClosed, +} + +pub(crate) struct ConsoleOutput { + mode: FlowControlMode, + shared: Option>, + #[cfg(unix)] + capacity_wake: LocalStream, + #[cfg(unix)] + status_wake: LocalStream, + #[cfg(unix)] + _idle_signals: Option<(LocalStream, LocalStream)>, + cancel: Option>, + graceful_finish: Option io::Result<()> + Send>>, + worker: Option>, +} + +impl ConsoleOutput { + pub(crate) fn stdout(mode: FlowControlMode) -> io::Result { + if mode == FlowControlMode::None { + return Self::new_with_lifecycle( + mode, + Box::new(io::stdout()), + Box::new(|| {}), + Box::new(|| Ok(())), + ); + } + #[cfg(unix)] + { + let file = File::from(rustix::io::dup(io::stdout().lock().as_fd())?); + let (cancel_reader, mut cancel_writer) = et_net::local::wake_pair()?; + let cancel = Box::new(move || { + let _ = cancel_writer.write_all(&[1]); + }); + Self::new_with_lifecycle( + mode, + Box::new(CancellableStdout { + file, + cancel: cancel_reader, + }), + cancel, + Box::new(|| Ok(())), + ) + } + #[cfg(windows)] + { + let mut child = Command::new(std::env::current_exe()?) + .arg("__et-console-writer") + .stdin(Stdio::piped()) + .stdout(Stdio::inherit()) + .stderr(Stdio::piped()) + .spawn()?; + let input = child + .stdin + .take() + .ok_or_else(|| io::Error::other("console helper stdin unavailable"))?; + let ack = child + .stderr + .take() + .ok_or_else(|| io::Error::other("console helper acknowledgement unavailable"))?; + let child = Arc::new(Mutex::new(child)); + let cancel_child = Arc::clone(&child); + let graceful_child = Arc::clone(&child); + Self::new_with_lifecycle( + mode, + Box::new(WindowsHelperWriter { input, ack }), + Box::new(move || cancel_windows_helper(&cancel_child)), + Box::new(move || wait_windows_helper(&graceful_child)), + ) + } + } + + #[cfg(test)] + pub(crate) fn new(mode: FlowControlMode, writer: Box) -> io::Result { + Self::new_with_lifecycle(mode, writer, Box::new(|| {}), Box::new(|| Ok(()))) + } + + #[cfg(test)] + pub(crate) fn new_with_cancel( + mode: FlowControlMode, + writer: Box, + cancel: Box, + ) -> io::Result { + Self::new_with_lifecycle(mode, writer, cancel, Box::new(|| Ok(()))) + } + + pub(crate) fn new_with_lifecycle( + mode: FlowControlMode, + mut writer: Box, + cancel: Box, + graceful_finish: Box io::Result<()> + Send>, + ) -> io::Result { + #[cfg(unix)] + let (capacity_wake, mut capacity_signal) = { + let (wake, signal) = et_net::local::wake_pair()?; + wake.set_nonblocking(true)?; + signal.set_nonblocking(true)?; + (wake, signal) + }; + #[cfg(unix)] + let (status_wake, mut status_signal) = { + let (wake, signal) = et_net::local::wake_pair()?; + wake.set_nonblocking(true)?; + signal.set_nonblocking(true)?; + (wake, signal) + }; + match mode { + FlowControlMode::None => { + return Ok(Self { + mode, + shared: None, + #[cfg(unix)] + capacity_wake, + #[cfg(unix)] + status_wake, + #[cfg(unix)] + _idle_signals: Some((capacity_signal, status_signal)), + cancel: None, + graceful_finish: None, + worker: None, + }); + } + FlowControlMode::Backpressure | FlowControlMode::Discard => {} + } + let shared = Arc::new(Shared { + state: Mutex::new(State { + queue: VecDeque::new(), + bytes: 0, + stopping: false, + error: None, + cursor_reports: 0, + worker_done: false, + }), + wake: Condvar::new(), + }); + let worker_shared = Arc::clone(&shared); + let worker = thread::Builder::new() + .name("et-console-output".to_owned()) + .spawn(move || { + run_writer( + &worker_shared, + &mut writer, + #[cfg(unix)] + &mut capacity_signal, + #[cfg(unix)] + &mut status_signal, + ); + })?; + Ok(Self { + mode, + shared: Some(shared), + #[cfg(unix)] + capacity_wake, + #[cfg(unix)] + status_wake, + #[cfg(unix)] + _idle_signals: None, + cancel: Some(cancel), + graceful_finish: Some(graceful_finish), + worker: Some(worker), + }) + } + + /// Attempt to admit one complete terminal-output packet without waiting. + /// + /// `Ok(false)` leaves ownership with the caller, which must retry the same + /// packet before reading another server packet. + pub(crate) fn try_write( + &self, + bytes: &[u8], + terminal_modes: &crate::client_terminal::TerminalModeState, + ) -> io::Result { + let Some(shared) = &self.shared else { + io::stdout() + .lock() + .write_all(bytes) + .and_then(|()| io::stdout().lock().flush())?; + terminal_modes.observe(bytes); + return Ok(true); + }; + if self.mode == FlowControlMode::Backpressure && bytes.len() > OUTPUT_BYTES { + return Err(io::Error::new( + io::ErrorKind::InvalidInput, + "terminal output packet exceeds console queue capacity", + )); + } + let retained = if bytes.len() > OUTPUT_BYTES { + &bytes[bytes.len() - OUTPUT_BYTES..] + } else { + bytes + }; + let mut state = shared + .state + .lock() + .map_err(|_| io::Error::other("console output worker unavailable"))?; + if let Some(error) = state.error.take() { + return Err(error); + } + if state.stopping { + return Err(io::Error::new( + io::ErrorKind::BrokenPipe, + "console output stopped", + )); + } + match self.mode { + FlowControlMode::None => unreachable!("none has no shared output queue"), + FlowControlMode::Backpressure + if state.bytes.saturating_add(retained.len()) > OUTPUT_BYTES + || state.queue.len() >= OUTPUT_PACKETS => + { + return Ok(false); + } + FlowControlMode::Backpressure => {} + FlowControlMode::Discard => { + while state.bytes.saturating_add(retained.len()) > OUTPUT_BYTES + || state.queue.len() >= OUTPUT_PACKETS + { + let Some(removed) = state.queue.pop_front() else { + break; + }; + state.bytes -= removed.bytes.len(); + } + } + } + state.bytes += retained.len(); + state.queue.push_back(OutputEntry { + bytes: retained.to_vec(), + terminal_modes: terminal_modes.clone(), + }); + drop(state); + shared.wake.notify_one(); + Ok(true) + } + + pub(crate) fn is_async(&self) -> bool { + self.shared.is_some() + } + + pub(crate) fn take_cursor_reports(&self) -> io::Result { + let Some(shared) = &self.shared else { + return Ok(0); + }; + let mut state = shared + .state + .lock() + .map_err(|_| io::Error::other("console output worker unavailable"))?; + Ok(std::mem::take(&mut state.cursor_reports)) + } + + #[cfg(test)] + pub(crate) fn wait_worker_done(&self) { + let Some(shared) = &self.shared else { + return; + }; + let state = shared.state.lock().unwrap(); + drop( + shared + .wake + .wait_while(state, |state| !state.worker_done) + .unwrap(), + ); + } + + pub(crate) fn complete(mut self, completion: ConsoleCompletion) -> io::Result<()> { + match completion { + ConsoleCompletion::RemoteSessionEnded => self.finish_gracefully(), + #[cfg(any(unix, test))] + ConsoleCompletion::LocalInputClosed => Ok(()), + } + } + + pub(crate) fn finish_gracefully(&mut self) -> io::Result<()> { + let Some(shared) = &self.shared else { + return Ok(()); + }; + { + let mut state = shared + .state + .lock() + .map_err(|_| io::Error::other("console output worker unavailable"))?; + state.stopping = true; + shared.wake.notify_all(); + } + if let Some(worker) = self.worker.take() { + worker + .join() + .map_err(|_| io::Error::other("console output worker panicked"))?; + } + self.check_error()?; + if let Some(finish) = self.graceful_finish.take() { + finish()?; + } + self.cancel = None; + self.shared = None; + Ok(()) + } + + pub(crate) fn check_error(&self) -> io::Result<()> { + let Some(shared) = &self.shared else { + return Ok(()); + }; + let mut state = shared + .state + .lock() + .map_err(|_| io::Error::other("console output worker unavailable"))?; + match state.error.take() { + Some(error) => Err(error), + None => Ok(()), + } + } + + #[cfg(unix)] + pub(crate) fn wake(&self) -> &LocalStream { + &self.capacity_wake + } + + #[cfg(unix)] + pub(crate) fn status_wake(&self) -> &LocalStream { + &self.status_wake + } + + #[cfg(unix)] + pub(crate) fn drain_wake(&mut self) -> io::Result<()> { + drain_stream(&mut self.capacity_wake) + } + + #[cfg(unix)] + pub(crate) fn drain_status_wake(&mut self) -> io::Result<()> { + drain_stream(&mut self.status_wake) + } +} + +impl Drop for ConsoleOutput { + fn drop(&mut self) { + if let Some(shared) = &self.shared { + if let Ok(mut state) = shared.state.lock() { + state.stopping = true; + shared.wake.notify_all(); + } + } + self.graceful_finish = None; + if let Some(cancel) = self.cancel.take() { + cancel(); + } + if let Some(worker) = self.worker.take() { + let _ = worker.join(); + } + } +} + +struct WorkerDone<'a>(&'a Shared); + +impl Drop for WorkerDone<'_> { + fn drop(&mut self) { + if let Ok(mut state) = self.0.state.lock() { + state.worker_done = true; + self.0.wake.notify_all(); + } + } +} + +fn run_writer( + shared: &Shared, + writer: &mut dyn Write, + #[cfg(unix)] capacity_signal: &mut LocalStream, + #[cfg(unix)] status_signal: &mut LocalStream, +) { + let _done = WorkerDone(shared); + loop { + let bytes = { + let Ok(state) = shared.state.lock() else { + return; + }; + let Ok(mut state) = shared + .wake + .wait_while(state, |state| state.queue.is_empty() && !state.stopping) + else { + return; + }; + let Some(bytes) = state.queue.pop_front() else { + return; + }; + state.bytes -= bytes.bytes.len(); + bytes + }; + #[cfg(unix)] + signal_capacity(capacity_signal); + if let Err(error) = writer.write_all(&bytes.bytes).and_then(|()| writer.flush()) { + if let Ok(mut state) = shared.state.lock() { + state.error = Some(error); + state.stopping = true; + state.worker_done = true; + shared.wake.notify_all(); + } + #[cfg(unix)] + signal_capacity(status_signal); + return; + } + bytes.terminal_modes.observe(&bytes.bytes); + if crate::client_terminal::contains_cursor_report_request(&bytes.bytes) { + if let Ok(mut state) = shared.state.lock() { + state.cursor_reports += 1; + } + #[cfg(unix)] + signal_capacity(capacity_signal); + } + } +} + +#[cfg(windows)] +struct WindowsHelperWriter { + input: ChildStdin, + ack: ChildStderr, +} + +#[cfg(windows)] +impl Write for WindowsHelperWriter { + fn write(&mut self, bytes: &[u8]) -> io::Result { + let length = u32::try_from(bytes.len()) + .map_err(|_| io::Error::new(io::ErrorKind::InvalidInput, "console packet too large"))?; + self.input.write_all(&length.to_le_bytes())?; + self.input.write_all(bytes)?; + self.input.flush()?; + let mut ack = [0u8; 1]; + std::io::Read::read_exact(&mut self.ack, &mut ack)?; + if ack != [1] { + return Err(io::Error::new( + io::ErrorKind::InvalidData, + "console helper acknowledgement is invalid", + )); + } + Ok(bytes.len()) + } + + fn flush(&mut self) -> io::Result<()> { + Ok(()) + } +} + +#[cfg(windows)] +fn wait_windows_helper(child: &Mutex) -> io::Result<()> { + let mut child = child + .lock() + .map_err(|_| io::Error::other("console helper unavailable"))?; + let status = child.wait()?; + if status.success() { + Ok(()) + } else { + Err(io::Error::other(format!( + "console helper exited with {status}" + ))) + } +} + +#[cfg(windows)] +fn cancel_windows_helper(child: &Mutex) { + if let Ok(mut child) = child.lock() { + let _ = child.kill(); + let _ = child.wait(); + } +} + +#[cfg(windows)] +pub(crate) fn run_windows_helper() -> i32 { + let mut input = io::stdin().lock(); + let mut output = io::stdout().lock(); + let mut acknowledgements = io::stderr().lock(); + loop { + let mut length = [0u8; 4]; + match std::io::Read::read_exact(&mut input, &mut length) { + Ok(()) => {} + Err(error) if error.kind() == io::ErrorKind::UnexpectedEof => return 0, + Err(_) => return 1, + } + let mut bytes = vec![0u8; u32::from_le_bytes(length) as usize]; + if std::io::Read::read_exact(&mut input, &mut bytes).is_err() + || output + .write_all(&bytes) + .and_then(|()| output.flush()) + .is_err() + || acknowledgements + .write_all(&[1]) + .and_then(|()| acknowledgements.flush()) + .is_err() + { + return 1; + } + } +} + +#[cfg(unix)] +fn drain_stream(stream: &mut LocalStream) -> io::Result<()> { + let mut bytes = [0u8; 64]; + loop { + match stream.read(&mut bytes) { + Ok(0) => return Ok(()), + Ok(_) => {} + Err(error) if error.kind() == io::ErrorKind::WouldBlock => return Ok(()), + Err(error) if error.kind() == io::ErrorKind::Interrupted => {} + Err(error) => return Err(error), + } + } +} + +#[cfg(unix)] +struct CancellableStdout { + file: File, + cancel: LocalStream, +} + +#[cfg(unix)] +impl Write for CancellableStdout { + fn write(&mut self, bytes: &[u8]) -> io::Result { + use rustix::event::{poll, PollFd, PollFlags}; + let mut descriptors = [ + PollFd::new(&self.file, PollFlags::OUT), + PollFd::new(&self.cancel, PollFlags::IN | PollFlags::HUP), + ]; + loop { + match poll(&mut descriptors, None) { + Ok(_) + if descriptors[1] + .revents() + .intersects(PollFlags::IN | PollFlags::HUP) => + { + return Err(io::Error::new( + io::ErrorKind::BrokenPipe, + "console output cancelled", + )); + } + Ok(_) if descriptors[0].revents().contains(PollFlags::OUT) => { + return rustix::io::write(&self.file, &bytes[..bytes.len().min(4096)]) + .map_err(io::Error::from); + } + Ok(_) => {} + Err(error) if error == rustix::io::Errno::INTR => {} + Err(error) => return Err(io::Error::from(error)), + } + } + } + + fn flush(&mut self) -> io::Result<()> { + Ok(()) + } +} + +#[cfg(unix)] +fn signal_capacity(signal: &mut LocalStream) { + match signal.write(&[1]) { + Ok(_) => {} + Err(error) + if matches!( + error.kind(), + io::ErrorKind::WouldBlock | io::ErrorKind::Interrupted + ) => {} + Err(_) => {} + } +} + +#[cfg(test)] +#[path = "client_output_tests.rs"] +mod tests; diff --git a/crates/et-bin/src/client_output_tests.rs b/crates/et-bin/src/client_output_tests.rs new file mode 100644 index 0000000..9138ca1 --- /dev/null +++ b/crates/et-bin/src/client_output_tests.rs @@ -0,0 +1,230 @@ +use super::*; +use crate::client_terminal::TerminalModeState; +use std::sync::mpsc; + +struct GatedWriter { + entered: mpsc::SyncSender, + release: mpsc::Receiver<()>, +} + +impl Write for GatedWriter { + fn write(&mut self, bytes: &[u8]) -> io::Result { + self.entered + .send(bytes.len()) + .map_err(|_| io::Error::other("test observer closed"))?; + self.release + .recv() + .map_err(|_| io::Error::other("test release closed"))?; + Ok(bytes.len()) + } + + fn flush(&mut self) -> io::Result<()> { + Ok(()) + } +} + +#[test] +fn full_backpressure_queue_does_not_block_control_progress() { + let (entered_tx, entered_rx) = mpsc::sync_channel(0); + let (release_tx, release_rx) = mpsc::sync_channel(0); + let output = ConsoleOutput::new( + FlowControlMode::Backpressure, + Box::new(GatedWriter { + entered: entered_tx, + release: release_rx, + }), + ) + .unwrap(); + let modes = TerminalModeState::default(); + assert!(output.try_write(&vec![1; OUTPUT_BYTES], &modes).unwrap()); + assert_eq!(entered_rx.recv().unwrap(), OUTPUT_BYTES); + assert!(output.try_write(&vec![2; OUTPUT_BYTES], &modes).unwrap()); + + assert!(!output.try_write(&[3], &modes).unwrap()); + let (control_tx, control_rx) = mpsc::sync_channel(1); + control_tx.send("ctrl-c").unwrap(); + assert_eq!(control_rx.recv().unwrap(), "ctrl-c"); + release_tx.send(()).unwrap(); + assert_eq!(entered_rx.recv().unwrap(), OUTPUT_BYTES); + release_tx.send(()).unwrap(); + drop(output); +} + +#[test] +fn discard_eviction_does_not_change_confirmed_terminal_mode() { + let (entered_tx, entered_rx) = mpsc::sync_channel(0); + let (release_tx, release_rx) = mpsc::sync_channel(0); + let output = ConsoleOutput::new( + FlowControlMode::Discard, + Box::new(GatedWriter { + entered: entered_tx, + release: release_rx, + }), + ) + .unwrap(); + let modes = TerminalModeState::default(); + assert!(output.try_write(b"visible", &modes).unwrap()); + assert_eq!(entered_rx.recv().unwrap(), 7); + assert!(output.try_write(b"\x1b[?1049h", &modes).unwrap()); + assert!(output.try_write(&[b'n'; OUTPUT_BYTES], &modes).unwrap()); + + assert!(!modes.alternate_screen()); + release_tx.send(()).unwrap(); + assert_eq!(entered_rx.recv().unwrap(), OUTPUT_BYTES); + release_tx.send(()).unwrap(); + drop(output); + assert!(!modes.alternate_screen()); +} + +#[test] +fn evicted_alternate_leave_preserves_confirmed_enter_state() { + let (entered_tx, entered_rx) = mpsc::sync_channel(0); + let (release_tx, release_rx) = mpsc::sync_channel(0); + let output = ConsoleOutput::new( + FlowControlMode::Discard, + Box::new(GatedWriter { + entered: entered_tx, + release: release_rx, + }), + ) + .unwrap(); + let modes = TerminalModeState::default(); + assert!(output.try_write(b"\x1b[?1049h", &modes).unwrap()); + assert_eq!(entered_rx.recv().unwrap(), 8); + assert!(output.try_write(b"blocking", &modes).unwrap()); + release_tx.send(()).unwrap(); + assert_eq!(entered_rx.recv().unwrap(), 8); + assert!(modes.alternate_screen()); + assert!(output.try_write(b"\x1b[?1049l", &modes).unwrap()); + assert!(output.try_write(&[b'n'; OUTPUT_BYTES], &modes).unwrap()); + + release_tx.send(()).unwrap(); + assert_eq!(entered_rx.recv().unwrap(), OUTPUT_BYTES); + release_tx.send(()).unwrap(); + drop(output); + assert!(modes.alternate_screen()); +} + +#[test] +fn clean_remote_session_end_drains_admitted_output_before_returning() { + let (entered_tx, entered_rx) = mpsc::sync_channel(0); + let (release_tx, release_rx) = mpsc::sync_channel(0); + let output = ConsoleOutput::new( + FlowControlMode::Backpressure, + Box::new(GatedWriter { + entered: entered_tx, + release: release_rx, + }), + ) + .unwrap(); + let modes = TerminalModeState::default(); + assert!(output.try_write(b"first", &modes).unwrap()); + assert_eq!(entered_rx.recv().unwrap(), 5); + assert!(output.try_write(b"second", &modes).unwrap()); + let (done_tx, done_rx) = mpsc::sync_channel(0); + std::thread::spawn(move || { + done_tx + .send(output.complete(ConsoleCompletion::RemoteSessionEnded)) + .unwrap(); + }); + + assert!(done_rx.try_recv().is_err()); + release_tx.send(()).unwrap(); + assert_eq!(entered_rx.recv().unwrap(), 6); + assert!(done_rx.try_recv().is_err()); + release_tx.send(()).unwrap(); + assert!(done_rx.recv().unwrap().is_ok()); +} + +#[test] +fn graceful_finish_is_idempotent() { + let mut output = + ConsoleOutput::new(FlowControlMode::Backpressure, Box::new(Vec::::new())).unwrap(); + assert!(output.finish_gracefully().is_ok()); + assert!(output.finish_gracefully().is_ok()); +} + +#[test] +fn graceful_finish_surfaces_last_write_error() { + struct Broken; + impl Write for Broken { + fn write(&mut self, _bytes: &[u8]) -> io::Result { + Err(io::ErrorKind::BrokenPipe.into()) + } + fn flush(&mut self) -> io::Result<()> { + Ok(()) + } + } + let mut output = ConsoleOutput::new(FlowControlMode::Backpressure, Box::new(Broken)).unwrap(); + let modes = TerminalModeState::default(); + assert!(output.try_write(b"last", &modes).unwrap()); + + assert_eq!( + output.finish_gracefully().unwrap_err().kind(), + io::ErrorKind::BrokenPipe + ); +} + +#[test] +fn local_input_close_cancels_blocked_output_before_join() { + enum Gate { + Cancel, + } + struct CancelWriter { + entered: mpsc::SyncSender<()>, + gate: mpsc::Receiver, + } + impl Write for CancelWriter { + fn write(&mut self, _bytes: &[u8]) -> io::Result { + self.entered.send(()).unwrap(); + match self.gate.recv().unwrap() { + Gate::Cancel => Err(io::Error::new(io::ErrorKind::BrokenPipe, "cancelled")), + } + } + fn flush(&mut self) -> io::Result<()> { + Ok(()) + } + } + let (entered_tx, entered_rx) = mpsc::sync_channel(0); + let (gate_tx, gate_rx) = mpsc::channel(); + let output = ConsoleOutput::new_with_cancel( + FlowControlMode::Backpressure, + Box::new(CancelWriter { + entered: entered_tx, + gate: gate_rx, + }), + Box::new(move || gate_tx.send(Gate::Cancel).unwrap()), + ) + .unwrap(); + let modes = TerminalModeState::default(); + assert!(output.try_write(b"blocked", &modes).unwrap()); + entered_rx.recv().unwrap(); + + let (done_tx, done_rx) = mpsc::sync_channel(0); + std::thread::spawn(move || { + let result = output.complete(ConsoleCompletion::LocalInputClosed); + done_tx.send(result).unwrap(); + }); + assert!(done_rx.recv().unwrap().is_ok()); +} + +#[test] +fn last_packet_broken_pipe_is_reported_without_another_packet() { + struct Broken; + impl Write for Broken { + fn write(&mut self, _bytes: &[u8]) -> io::Result { + Err(io::ErrorKind::BrokenPipe.into()) + } + fn flush(&mut self) -> io::Result<()> { + Ok(()) + } + } + let output = ConsoleOutput::new(FlowControlMode::Discard, Box::new(Broken)).unwrap(); + let modes = TerminalModeState::default(); + assert!(output.try_write(b"last", &modes).unwrap()); + output.wait_worker_done(); + assert_eq!( + output.check_error().unwrap_err().kind(), + io::ErrorKind::BrokenPipe + ); +} diff --git a/crates/et-bin/src/client_terminal.rs b/crates/et-bin/src/client_terminal.rs index 0073660..1f3e453 100644 --- a/crates/et-bin/src/client_terminal.rs +++ b/crates/et-bin/src/client_terminal.rs @@ -6,7 +6,7 @@ use std::thread; use crossterm::terminal::{disable_raw_mode, enable_raw_mode, size}; use et_core::proto::{TerminalBuffer, TerminalInfo, TerminalPacketType}; -use et_net::connection::{ConnError, Connection}; +use et_net::connection::{ConnError, Connection, WritePacketError}; use et_net::forward::Forwarder; use prost::Message; #[cfg(unix)] @@ -44,6 +44,7 @@ pub struct TerminalOptions<'a> { pub command: Option<&'a str>, pub no_exit: bool, pub keepalive: u32, + pub flow_control: et_cli::client::FlowControlMode, pub terminal_enabled: bool, pub lines: RemoteLines, pub connection_name: &'a str, @@ -52,7 +53,7 @@ pub struct TerminalOptions<'a> { pub fn run( mut connection: Connection, options: TerminalOptions<'_>, - forwarder: Forwarder, + mut forwarder: Forwarder, mut reconnect: F, ) -> Result<(), ClientError> where @@ -62,6 +63,7 @@ where command, no_exit, keepalive, + flow_control, terminal_enabled, lines, connection_name, @@ -78,30 +80,33 @@ where let mut terminal_modes = TerminalModeState::default(); // The network can disappear immediately after the initial handshake (a // laptop waking up is particularly prone to this). These writes are - // replay-buffered by `Connection`, so once recovery succeeds they must - // not be sent again; doing so would duplicate a `--command`. Recover the - // transport and let the buffered packet be replayed instead. + // replay-buffered by `Connection`. Retry plaintext only when admission + // failed before replay ownership; after admission, recover and let the + // buffered packet replay instead of duplicating a `--command`. if terminal_enabled { - let initial_size = send_size(&mut connection); - if !recover_initial_transport( - &mut connection, - &mut reconnect, - terminal_enabled, - initial_size, - )? { - return raw_mode.finish(Ok(()), close_message, terminal_modes.alternate_screen); + if let Some(initial_size) = terminal_size_payload()? { + if matches!( + write_terminal_size_recovering(&mut connection, &initial_size, &mut reconnect)?, + OwnedWriteOutcome::SessionEnded + ) { + return raw_mode.finish(Ok(()), close_message, terminal_modes.alternate_screen()); + } } } if terminal_enabled { if let Some(command) = command { - let initial_command = send_command(&mut connection, command, no_exit, lines); - if !recover_initial_transport( - &mut connection, - &mut reconnect, - terminal_enabled, - initial_command, - )? { - return raw_mode.finish(Ok(()), close_message, terminal_modes.alternate_screen); + let initial_command = command_payload(command, no_exit, lines)?; + if matches!( + write_owned_recovering( + &mut connection, + TerminalPacketType::TerminalBuffer as u8, + &initial_command, + &mut reconnect, + terminal_enabled, + )?, + OwnedWriteOutcome::SessionEnded + ) { + return raw_mode.finish(Ok(()), close_message, terminal_modes.alternate_screen()); } } } @@ -146,11 +151,12 @@ where crate::client_terminal_loop::PumpOptions { read_stdin, keepalive_seconds: keepalive, + flow_control, terminal_enabled, auto_cursor_report, terminal_modes: &mut terminal_modes, }, - &forwarder, + &mut forwarder, reconnect, ); signal_handle.close(); @@ -163,14 +169,15 @@ where crate::client_terminal_loop::PumpOptions { read_stdin, keepalive_seconds: keepalive, + flow_control, terminal_enabled, auto_cursor_report, terminal_modes: &mut terminal_modes, }, - &forwarder, + &mut forwarder, reconnect, ); - raw_mode.finish(result, close_message, terminal_modes.alternate_screen) + raw_mode.finish(result, close_message, terminal_modes.alternate_screen()) } /// Device Status Report request (`ESC [ 6 n`). @@ -184,10 +191,65 @@ pub(crate) const CURSOR_REPORT_REPLY: &[u8] = b"\x1b[1;1R"; /// emits that request on startup and waits for the answer, which an interactive /// terminal emulator provides. Non-interactive sessions have no emulator to /// answer, so the caller replies on their behalf (see `auto_cursor_report`). -pub(crate) fn display_packet( +pub(crate) enum DisplayOutcome { + Displayed { cursor_report: bool }, + Pending(et_core::packet::Packet), +} + +pub(crate) struct RetainedCompletion { + terminal: Option, + forwarding: Option, +} + +pub(crate) fn classify_forward_completion( + current_outbound: Option, + abandoned: bool, +) -> Result<(), ClientError> { + if current_outbound.is_some() || abandoned { + return Err(terminal_text( + "remote session ended with undeliverable outbound forwarding data", + )); + } + Ok(()) +} + +impl RetainedCompletion { + pub(crate) fn new( + terminal: Option, + forwarding: Option, + ) -> Self { + Self { + terminal, + forwarding, + } + } + + pub(crate) fn advance( + &mut self, + mut admit_terminal: T, + mut admit_forwarding: F, + ) -> Result + where + T: FnMut(et_core::packet::Packet) -> Result, ClientError>, + F: FnMut(et_core::packet::Packet) -> Result, ClientError>, + { + if let Some(packet) = self.forwarding.take() { + self.forwarding = admit_forwarding(packet)?; + } + if let Some(packet) = self.terminal.take() { + self.terminal = admit_terminal(packet)?; + } + Ok(self.forwarding.is_none() && self.terminal.is_none()) + } +} + +pub(crate) fn display_packet_with( packet: et_core::packet::Packet, - terminal_modes: &mut TerminalModeState, -) -> Result { + mut output: F, +) -> Result +where + F: FnMut(&[u8]) -> Result, +{ match packet.header() { value if value == TerminalPacketType::TerminalBuffer as u8 => { let message = TerminalBuffer::decode(packet.payload()) @@ -195,31 +257,31 @@ pub(crate) fn display_packet( let bytes = message .buffer .ok_or_else(|| terminal_text("terminal output is missing bytes"))?; - terminal_modes.observe(&bytes); - io::stdout() - .lock() - .write_all(&bytes) - .and_then(|()| io::stdout().lock().flush()) - .map_err(|error| terminal_io("writing terminal output", error))?; - Ok(contains_cursor_report_request(&bytes)) + if !output(&bytes)? { + return Ok(DisplayOutcome::Pending(packet)); + } + Ok(DisplayOutcome::Displayed { + cursor_report: contains_cursor_report_request(&bytes), + }) } - value if value == TerminalPacketType::KeepAlive as u8 => Ok(false), + value if value == TerminalPacketType::KeepAlive as u8 => Ok(DisplayOutcome::Displayed { + cursor_report: false, + }), _ => Err(terminal_text("server sent an unsupported terminal packet")), } } -fn contains_cursor_report_request(bytes: &[u8]) -> bool { +pub(crate) fn contains_cursor_report_request(bytes: &[u8]) -> bool { bytes .windows(CURSOR_REPORT_REQUEST.len()) .any(|window| window == CURSOR_REPORT_REQUEST) } -fn send_command( - connection: &mut Connection, +fn command_payload( command: &str, no_exit: bool, lines: RemoteLines, -) -> Result<(), ClientError> { +) -> Result, ClientError> { if command.contains('\0') || command.len() > 64 * 1024 { return Err(terminal_text("remote command is invalid or too large")); } @@ -236,26 +298,143 @@ fn send_command( let mut bytes = Vec::with_capacity(command.len() + suffix.len()); bytes.extend_from_slice(command.as_bytes()); bytes.extend_from_slice(suffix.as_bytes()); - send_buffer_checked(connection, &bytes) + Ok(encoded_buffer(&bytes)) } -pub(crate) fn send_buffer(connection: &mut Connection, bytes: &[u8]) -> Result<(), ConnError> { - let message = TerminalBuffer { +pub(crate) fn encoded_buffer(bytes: &[u8]) -> Vec { + TerminalBuffer { buffer: Some(bytes.to_vec()), - }; + } + .encode_to_vec() +} + +#[cfg(test)] +pub(crate) fn send_buffer(connection: &mut Connection, bytes: &[u8]) -> Result<(), ConnError> { connection.write_packet( TerminalPacketType::TerminalBuffer as u8, - &message.encode_to_vec(), + &encoded_buffer(bytes), ) } -fn send_buffer_checked(connection: &mut Connection, bytes: &[u8]) -> Result<(), ClientError> { - send_buffer(connection, bytes).map_err(ClientError::Transport) +pub(crate) enum OwnedWriteOutcome { + Written, + Recovered, + SessionEnded, } -pub(crate) fn send_size(connection: &mut Connection) -> Result<(), ClientError> { +#[derive(Clone, Copy)] +pub(crate) enum OwnedWritePolicy { + ExactPlaintext, + ReplaceableTerminalSize, +} + +pub(crate) fn write_owned_recovering( + connection: &mut Connection, + header: u8, + payload: &[u8], + reconnect: &mut F, + send_terminal_size: bool, +) -> Result +where + F: FnMut(&mut Connection) -> Result, +{ + write_owned_recovering_with( + connection, + header, + payload, + reconnect, + send_terminal_size, + |connection, header, payload| connection.write_packet_owned(header, payload), + ) +} + +fn write_owned_recovering_with( + connection: &mut Connection, + header: u8, + payload: &[u8], + reconnect: &mut F, + send_terminal_size: bool, + write: W, +) -> Result +where + F: FnMut(&mut Connection) -> Result, + W: FnMut(&mut Connection, u8, &[u8]) -> Result<(), WritePacketError>, +{ + write_owned_with_policy( + connection, + header, + payload, + OwnedWritePolicy::ExactPlaintext, + write, + |connection, _policy| recover_transport(connection, reconnect, send_terminal_size), + ) +} + +fn write_owned_with_policy( + connection: &mut Connection, + header: u8, + payload: &[u8], + policy: OwnedWritePolicy, + mut write: W, + mut recover: R, +) -> Result +where + W: FnMut(&mut Connection, u8, &[u8]) -> Result<(), WritePacketError>, + R: FnMut(&mut Connection, OwnedWritePolicy) -> Result, +{ + let mut recovered = false; + loop { + match write(connection, header, payload) { + Ok(()) => { + return Ok(if recovered { + OwnedWriteOutcome::Recovered + } else { + OwnedWriteOutcome::Written + }); + } + Err(WritePacketError::BeforeReplay(error)) if connection_ended(&error) => { + connection.disconnect(); + if !recover(connection, policy)? { + return Ok(OwnedWriteOutcome::SessionEnded); + } + recovered = true; + if matches!(policy, OwnedWritePolicy::ReplaceableTerminalSize) { + return Ok(OwnedWriteOutcome::Recovered); + } + } + Err(WritePacketError::ReplayOwned(error)) if connection_ended(&error) => { + return Ok(if recover(connection, policy)? { + OwnedWriteOutcome::Recovered + } else { + OwnedWriteOutcome::SessionEnded + }); + } + Err(error) => return Err(terminal_error(error.into_inner())), + } + } +} + +pub(crate) fn write_terminal_size_recovering( + connection: &mut Connection, + payload: &[u8], + reconnect: &mut F, +) -> Result +where + F: FnMut(&mut Connection) -> Result, +{ + write_owned_with_policy( + connection, + TerminalPacketType::TerminalInfo as u8, + payload, + OwnedWritePolicy::ReplaceableTerminalSize, + |connection, header, payload| connection.write_packet_owned(header, payload), + |connection, _policy| recover_transport(connection, reconnect, true), + ) +} + +pub(crate) fn terminal_size_payload() -> Result>, ClientError> { if !io::stdout().is_terminal() { - return Ok(()); + return Ok(None); } let (columns, rows) = size().map_err(|error| terminal_io("reading terminal size", error))?; let message = TerminalInfo { @@ -265,12 +444,7 @@ pub(crate) fn send_size(connection: &mut Connection) -> Result<(), ClientError> width: Some(0), height: Some(0), }; - connection - .write_packet( - TerminalPacketType::TerminalInfo as u8, - &message.encode_to_vec(), - ) - .map_err(ClientError::Transport) + Ok(Some(message.encode_to_vec())) } /// Finish an initial client write without losing a session to a race between @@ -279,6 +453,7 @@ pub(crate) fn send_size(connection: &mut Connection) -> Result<(), ClientError> /// `Connection` has already retained a packet whose socket write failed. /// Recovery replays that exact packet, so this deliberately does not retry /// `result` after reconnecting (retrying a command could execute it twice). +#[cfg(test)] fn recover_initial_transport( connection: &mut Connection, reconnect: &mut F, @@ -313,15 +488,30 @@ pub(crate) fn recover_transport( where F: FnMut(&mut Connection) -> Result, { + let mut replay_owns_size = false; loop { match reconnect(connection)? { ReconnectOutcome::SessionEnded => return Ok(false), - ReconnectOutcome::Recovered if !send_terminal_size => return Ok(true), - ReconnectOutcome::Recovered => match send_size(connection) { - Ok(()) => return Ok(true), - Err(ClientError::Transport(error)) if connection_ended(&error) => {} - Err(error) => return Err(error), - }, + ReconnectOutcome::Recovered if !send_terminal_size || replay_owns_size => { + return Ok(true); + } + ReconnectOutcome::Recovered => { + let Some(payload) = terminal_size_payload()? else { + return Ok(true); + }; + match connection + .write_packet_owned(TerminalPacketType::TerminalInfo as u8, &payload) + { + Ok(()) => return Ok(true), + Err(WritePacketError::BeforeReplay(error)) if connection_ended(&error) => { + connection.disconnect(); + } + Err(WritePacketError::ReplayOwned(error)) if connection_ended(&error) => { + replay_owns_size = true; + } + Err(error) => return Err(terminal_error(error.into_inner())), + } + } } } } @@ -350,29 +540,39 @@ const GRACEFUL_TERMINAL_MODE_RESET: &[u8] = b"\x1b[<64u\x1b[=0;1u\ \x1b[>4;0m\x1b[?2004l\x1b[?1004l\x1b[?1000l\x1b[?1002l\x1b[?1003l\x1b[?1006l\x1b[?25h"; #[derive(Default)] -pub(crate) struct TerminalModeState { +struct TerminalModeInner { alternate_prefix_len: usize, alternate_screen: bool, } +#[derive(Clone, Default)] +pub(crate) struct TerminalModeState(std::sync::Arc>); + impl TerminalModeState { - fn observe(&mut self, bytes: &[u8]) { + pub(crate) fn observe(&self, bytes: &[u8]) { const PREFIX: &[u8] = b"\x1b[?1049"; + let Ok(mut state) = self.0.lock() else { + return; + }; for &byte in bytes { - if self.alternate_prefix_len == PREFIX.len() { + if state.alternate_prefix_len == PREFIX.len() { match byte { - b'h' => self.alternate_screen = true, - b'l' => self.alternate_screen = false, + b'h' => state.alternate_screen = true, + b'l' => state.alternate_screen = false, _ => {} } - self.alternate_prefix_len = usize::from(byte == PREFIX[0]); - } else if byte == PREFIX[self.alternate_prefix_len] { - self.alternate_prefix_len += 1; + state.alternate_prefix_len = usize::from(byte == PREFIX[0]); + } else if byte == PREFIX[state.alternate_prefix_len] { + state.alternate_prefix_len += 1; } else { - self.alternate_prefix_len = usize::from(byte == PREFIX[0]); + state.alternate_prefix_len = usize::from(byte == PREFIX[0]); } } } + + pub(crate) fn alternate_screen(&self) -> bool { + self.0.lock().is_ok_and(|state| state.alternate_screen) + } } #[derive(Clone, Copy, Debug, PartialEq, Eq)] diff --git a/crates/et-bin/src/client_terminal_loop.rs b/crates/et-bin/src/client_terminal_loop.rs index 58962bf..979e37f 100644 --- a/crates/et-bin/src/client_terminal_loop.rs +++ b/crates/et-bin/src/client_terminal_loop.rs @@ -16,11 +16,16 @@ use rustix::event::{poll, PollFd, PollFlags}; #[cfg(unix)] use rustix::time::Timespec; +#[cfg(unix)] +use crate::client_output::ConsoleCompletion; +#[cfg(unix)] +use crate::client_terminal::DisplayOutcome; use crate::client_terminal::TerminalModeState; #[cfg(unix)] use crate::client_terminal::{ - connection_ended, display_packet, recover_transport, send_buffer, send_size, terminal_error, - terminal_io, terminal_text, + classify_forward_completion, connection_ended, encoded_buffer, recover_transport, + terminal_error, terminal_io, terminal_size_payload, terminal_text, write_owned_recovering, + write_terminal_size_recovering, OwnedWriteOutcome, RetainedCompletion, }; #[cfg(unix)] use crate::error::ClientError; @@ -75,6 +80,7 @@ impl PumpProbe { pub(crate) struct PumpOptions<'a> { pub(crate) read_stdin: bool, pub(crate) keepalive_seconds: u32, + pub(crate) flow_control: et_cli::client::FlowControlMode, pub(crate) terminal_enabled: bool, pub(crate) auto_cursor_report: bool, pub(crate) terminal_modes: &'a mut TerminalModeState, @@ -85,7 +91,7 @@ pub fn pump( connection: &mut Connection, wake: &mut UnixStream, options: PumpOptions<'_>, - forwarder: &Forwarder, + forwarder: &mut Forwarder, mut reconnect: F, ) -> Result<(), ClientError> where @@ -94,11 +100,14 @@ where let PumpOptions { read_stdin, keepalive_seconds, + flow_control, terminal_enabled, auto_cursor_report, terminal_modes, } = options; let stdin = io::stdin(); + let mut console_output = crate::client_output::ConsoleOutput::stdout(flow_control) + .map_err(|error| terminal_io("starting console output worker", error))?; let interval = Duration::from_secs(u64::from(keepalive_seconds.max(1))); let silence = interval.saturating_mul(MISSED_KEEPALIVES); let mut last_received = Instant::now(); @@ -108,6 +117,10 @@ where // further session packets are read (ordering) and the network fd is not // watched for readability (a readable socket would busy-loop the poll). let mut pending_forward: Option = None; + // A server terminal packet that could not enter the bounded local output + // queue. Keep exact ownership and stop reading the ordered server stream, + // while stdin, outbound forwarding, keepalive, and recovery stay live. + let mut pending_output: Option = None; let forward_wake = forwarder .wake() .map_err(|error| terminal_text(error.to_string()))?; @@ -129,17 +142,52 @@ where .try_receive(packet) .map_err(|error| terminal_text(error.to_string()))?; } - let network_flags = if pending_forward.is_none() { + if let Some(packet) = pending_output.take() { + match route_server_packet(packet, terminal_enabled, terminal_modes, &console_output)? { + DisplayOutcome::Displayed { cursor_report } + if cursor_report && auto_cursor_report && !console_output.is_async() => + { + if matches!( + write_cursor_report( + connection, + &mut reconnect, + &mut stream, + terminal_enabled, + )?, + OwnedWriteOutcome::SessionEnded + ) { + return finish_remote_completion( + console_output, + pending_output, + pending_forward, + terminal_enabled, + terminal_modes, + forwarder, + None, + ); + } + } + DisplayOutcome::Displayed { .. } => {} + DisplayOutcome::Pending(packet) => pending_output = Some(packet), + } + } + let network_flags = if pending_forward.is_none() && pending_output.is_none() { PollFlags::IN | PollFlags::HUP | PollFlags::ERR } else { PollFlags::HUP | PollFlags::ERR }; - let deadline = next_keepalive.min(last_received + silence); - let (network, resize, forwarding, input) = { + let deadline = if pending_output.is_some() { + next_keepalive + } else { + next_keepalive.min(last_received + silence) + }; + let (network, resize, forwarding, output_ready, output_status, input) = { let mut descriptors = vec![ PollFd::new(&stream, network_flags), PollFd::new(&*wake, PollFlags::IN | PollFlags::HUP), PollFd::new(forward_wake, PollFlags::IN | PollFlags::HUP), + PollFd::new(console_output.wake(), PollFlags::IN | PollFlags::HUP), + PollFd::new(console_output.status_wake(), PollFlags::IN | PollFlags::HUP), ]; if read_stdin { descriptors.push(PollFd::new( @@ -179,8 +227,10 @@ where descriptors[0].revents(), descriptors[1].revents(), descriptors[2].revents(), + descriptors[3].revents(), + descriptors[4].revents(), descriptors - .get(3) + .get(5) .map(PollFd::revents) .unwrap_or(PollFlags::empty()), ) @@ -189,21 +239,70 @@ where probe.progressed()?; } let mut reconnect_needed = network.intersects(PollFlags::HUP | PollFlags::ERR); + if output_ready.intersects(PollFlags::IN | PollFlags::HUP) { + console_output + .drain_wake() + .map_err(|error| terminal_io("draining console output wakeup", error))?; + console_output + .check_error() + .map_err(|error| terminal_io("writing terminal output", error))?; + if auto_cursor_report { + for _ in 0..console_output + .take_cursor_reports() + .map_err(|error| terminal_io("reading console confirmations", error))? + { + if matches!( + write_cursor_report( + connection, + &mut reconnect, + &mut stream, + terminal_enabled, + )?, + OwnedWriteOutcome::SessionEnded + ) { + return finish_remote_completion( + console_output, + pending_output, + pending_forward, + terminal_enabled, + terminal_modes, + forwarder, + None, + ); + } + } + } + } + if output_status.intersects(PollFlags::IN | PollFlags::HUP) { + console_output + .drain_status_wake() + .map_err(|error| terminal_io("draining console status wakeup", error))?; + console_output + .check_error() + .map_err(|error| terminal_io("writing terminal output", error))?; + } if resize.intersects(PollFlags::IN | PollFlags::HUP) { drain(wake)?; - match if terminal_enabled { - send_size(connection) - } else { - Ok(()) - } { - Ok(()) => {} - Err(ClientError::Transport(error)) if connection_ended(&error) => { - reconnect_needed = true; + if terminal_enabled { + if let Some(payload) = terminal_size_payload()? { + if matches!( + write_terminal_size(connection, &payload, &mut reconnect, &mut stream)?, + OwnedWriteOutcome::SessionEnded + ) { + return finish_remote_completion( + console_output, + pending_output, + pending_forward, + terminal_enabled, + terminal_modes, + forwarder, + None, + ); + } } - Err(error) => return Err(error), } } - while pending_forward.is_none() { + while pending_forward.is_none() && pending_output.is_none() { match connection.try_read_packet() { Ok(Some(packet)) => { last_received = Instant::now(); @@ -216,11 +315,41 @@ where pending_forward = forwarder .try_receive(packet) .map_err(|error| terminal_text(error.to_string()))?; - } else if route_server_packet(packet, terminal_enabled, terminal_modes)? - && auto_cursor_report - { - let _ = - send_buffer(connection, crate::client_terminal::CURSOR_REPORT_REPLY); + } else { + match route_server_packet( + packet, + terminal_enabled, + terminal_modes, + &console_output, + )? { + DisplayOutcome::Displayed { cursor_report } + if cursor_report + && auto_cursor_report + && !console_output.is_async() => + { + if matches!( + write_cursor_report( + connection, + &mut reconnect, + &mut stream, + terminal_enabled, + )?, + OwnedWriteOutcome::SessionEnded + ) { + return finish_remote_completion( + console_output, + pending_output, + pending_forward, + terminal_enabled, + terminal_modes, + forwarder, + None, + ); + } + } + DisplayOutcome::Displayed { .. } => {} + DisplayOutcome::Pending(packet) => pending_output = Some(packet), + } } } Ok(None) => break, @@ -236,23 +365,48 @@ where .try_outbound() .map_err(|error| terminal_text(error.to_string()))? { - match connection.write_packet(packet.header(), packet.payload()) { - Ok(()) => {} - Err(error) if connection_ended(&error) => { - reconnect_needed = true; - break; + match write_owned( + connection, + packet.header(), + packet.payload(), + &mut reconnect, + &mut stream, + terminal_enabled, + )? { + OwnedWriteOutcome::Written => {} + OwnedWriteOutcome::Recovered => { + last_received = Instant::now(); + next_keepalive = last_received + interval; + } + OwnedWriteOutcome::SessionEnded => { + return finish_remote_completion( + console_output, + pending_output, + pending_forward, + terminal_enabled, + terminal_modes, + forwarder, + Some(packet), + ); } - Err(error) => return Err(terminal_error(error)), } } } let now = Instant::now(); - if now.saturating_duration_since(last_received) >= silence { + if pending_output.is_none() && now.saturating_duration_since(last_received) >= silence { reconnect_needed = true; } if reconnect_needed { if !recover(connection, &mut reconnect, &mut stream, terminal_enabled)? { - return Ok(()); + return finish_remote_completion( + console_output, + pending_output, + pending_forward, + terminal_enabled, + terminal_modes, + forwarder, + None, + ); } last_received = Instant::now(); next_keepalive = last_received + interval; @@ -265,34 +419,67 @@ where .read(&mut bytes) .map_err(|error| terminal_io("reading terminal input", error))?; if count == 0 { - return Ok(()); + return console_output + .complete(ConsoleCompletion::LocalInputClosed) + .map_err(|error| terminal_io("stopping terminal output", error)); } - match send_buffer(connection, &bytes[..count]) { - Ok(()) => {} - Err(error) if connection_ended(&error) => { - if !recover(connection, &mut reconnect, &mut stream, terminal_enabled)? { - return Ok(()); - } + let payload = encoded_buffer(&bytes[..count]); + match write_owned( + connection, + TerminalPacketType::TerminalBuffer as u8, + &payload, + &mut reconnect, + &mut stream, + terminal_enabled, + )? { + OwnedWriteOutcome::Written => {} + OwnedWriteOutcome::Recovered => { last_received = Instant::now(); next_keepalive = last_received + interval; } - Err(error) => return Err(terminal_error(error)), + OwnedWriteOutcome::SessionEnded => { + return finish_remote_completion( + console_output, + pending_output, + pending_forward, + terminal_enabled, + terminal_modes, + forwarder, + None, + ); + } } } if input.intersects(PollFlags::HUP | PollFlags::ERR) { - return Ok(()); + return console_output + .complete(ConsoleCompletion::LocalInputClosed) + .map_err(|error| terminal_io("stopping terminal output", error)); } let now = Instant::now(); if now >= next_keepalive { // The payload acknowledges everything read so far, so the server // can trim its replay backup; legacy servers ignore it. let ack = connection.keepalive_ack(); - if connection - .write_packet(TerminalPacketType::KeepAlive as u8, &ack) - .is_err() - && !recover(connection, &mut reconnect, &mut stream, terminal_enabled)? - { - return Ok(()); + if matches!( + write_owned( + connection, + TerminalPacketType::KeepAlive as u8, + &ack, + &mut reconnect, + &mut stream, + terminal_enabled, + )?, + OwnedWriteOutcome::SessionEnded + ) { + return finish_remote_completion( + console_output, + pending_output, + pending_forward, + terminal_enabled, + terminal_modes, + forwarder, + None, + ); } next_keepalive = Instant::now() + interval; } @@ -305,18 +492,165 @@ fn route_server_packet( packet: et_core::packet::Packet, terminal_enabled: bool, terminal_modes: &mut TerminalModeState, -) -> Result { + output: &crate::client_output::ConsoleOutput, +) -> Result { if terminal_enabled || packet.header() == TerminalPacketType::KeepAlive as u8 { - return display_packet(packet, terminal_modes); + return crate::client_terminal::display_packet_with(packet, |bytes| { + output + .try_write(bytes, terminal_modes) + .map_err(|error| terminal_io("writing terminal output", error)) + }); } if packet.header() == TerminalPacketType::TerminalBuffer as u8 { - return Ok(false); + return Ok(DisplayOutcome::Displayed { + cursor_report: false, + }); } Err(terminal_text( "server sent an unsupported no-terminal packet", )) } +#[cfg(unix)] +fn finish_remote_completion( + mut output: crate::client_output::ConsoleOutput, + pending_output: Option, + pending_forward: Option, + terminal_enabled: bool, + terminal_modes: &mut TerminalModeState, + forwarder: &mut Forwarder, + current_outbound: Option, +) -> Result<(), ClientError> { + let mut retained = RetainedCompletion::new(pending_output, pending_forward); + loop { + output + .check_error() + .map_err(|error| terminal_io("writing retained terminal output", error))?; + if retained.advance( + |packet| match route_server_packet(packet, terminal_enabled, terminal_modes, &output)? { + DisplayOutcome::Displayed { .. } => Ok(None), + DisplayOutcome::Pending(packet) => Ok(Some(packet)), + }, + |packet| { + forwarder + .try_receive(packet) + .map_err(|error| terminal_text(error.to_string())) + }, + )? { + let abandoned = forwarder + .shutdown_hard() + .map_err(|error| terminal_text(error.to_string()))?; + classify_forward_completion(current_outbound, abandoned)?; + return output + .complete(ConsoleCompletion::RemoteSessionEnded) + .map_err(|error| terminal_io("draining terminal output", error)); + } + if let Some(packet) = forwarder + .try_outbound() + .map_err(|error| terminal_text(error.to_string()))? + { + let abandoned = forwarder + .shutdown_hard() + .map_err(|error| terminal_text(error.to_string()))?; + classify_forward_completion(Some(packet), abandoned)?; + } + + let (output_ready, output_status) = { + let mut descriptors = [ + PollFd::new( + forwarder + .wake() + .map_err(|error| terminal_text(error.to_string()))?, + PollFlags::IN | PollFlags::HUP, + ), + PollFd::new(output.wake(), PollFlags::IN | PollFlags::HUP), + PollFd::new(output.status_wake(), PollFlags::IN | PollFlags::HUP), + ]; + let timeout = Timespec::try_from(Duration::from_millis(100)) + .map_err(|_| terminal_text("completion wait exceeds poll range"))?; + match poll(&mut descriptors, Some(&timeout)) { + Ok(_) => {} + Err(error) if error == rustix::io::Errno::INTR => continue, + Err(error) => { + return Err(terminal_io( + "waiting to drain retained session packets", + io::Error::from(error), + )); + } + } + (descriptors[1].revents(), descriptors[2].revents()) + }; + if output_ready.intersects(PollFlags::IN | PollFlags::HUP) { + output + .drain_wake() + .map_err(|error| terminal_io("draining console output wakeup", error))?; + } + if output_status.intersects(PollFlags::IN | PollFlags::HUP) { + output + .drain_status_wake() + .map_err(|error| terminal_io("draining console status wakeup", error))?; + } + } +} + +#[cfg(unix)] +fn write_cursor_report( + connection: &mut Connection, + reconnect: &mut F, + stream: &mut std::net::TcpStream, + send_terminal_size: bool, +) -> Result +where + F: FnMut(&mut Connection) -> Result, +{ + let payload = encoded_buffer(crate::client_terminal::CURSOR_REPORT_REPLY); + write_owned( + connection, + TerminalPacketType::TerminalBuffer as u8, + &payload, + reconnect, + stream, + send_terminal_size, + ) +} + +#[cfg(unix)] +fn write_terminal_size( + connection: &mut Connection, + payload: &[u8], + reconnect: &mut F, + stream: &mut std::net::TcpStream, +) -> Result +where + F: FnMut(&mut Connection) -> Result, +{ + let outcome = write_terminal_size_recovering(connection, payload, reconnect)?; + if matches!(outcome, OwnedWriteOutcome::Recovered) { + *stream = connection.try_clone_stream().map_err(terminal_error)?; + } + Ok(outcome) +} + +#[cfg(unix)] +fn write_owned( + connection: &mut Connection, + header: u8, + payload: &[u8], + reconnect: &mut F, + stream: &mut std::net::TcpStream, + send_terminal_size: bool, +) -> Result +where + F: FnMut(&mut Connection) -> Result, +{ + let outcome = + write_owned_recovering(connection, header, payload, reconnect, send_terminal_size)?; + if matches!(outcome, OwnedWriteOutcome::Recovered) { + *stream = connection.try_clone_stream().map_err(terminal_error)?; + } + Ok(outcome) +} + #[cfg(unix)] fn recover( connection: &mut Connection, @@ -347,3 +681,7 @@ fn drain(wake: &mut UnixStream) -> Result<(), ClientError> { } } } + +#[cfg(test)] +#[path = "client_terminal_loop_tests.rs"] +mod tests; diff --git a/crates/et-bin/src/client_terminal_loop_tests.rs b/crates/et-bin/src/client_terminal_loop_tests.rs new file mode 100644 index 0000000..150c1fc --- /dev/null +++ b/crates/et-bin/src/client_terminal_loop_tests.rs @@ -0,0 +1,83 @@ +#![cfg(unix)] + +use super::*; +use std::io::Write; +use std::net::{Ipv4Addr, TcpListener, TcpStream}; +use std::sync::mpsc; +use std::thread; + +use crate::client_terminal::send_buffer; +use et_core::proto::TerminalBuffer; +use prost::Message; + +struct GatedConsole { + entered: mpsc::SyncSender, + release: mpsc::Receiver<()>, +} + +impl Write for GatedConsole { + fn write(&mut self, bytes: &[u8]) -> io::Result { + self.entered + .send(bytes.len()) + .map_err(|_| io::Error::other("console observer closed"))?; + self.release + .recv() + .map_err(|_| io::Error::other("console release closed"))?; + Ok(bytes.len()) + } + + fn flush(&mut self) -> io::Result<()> { + Ok(()) + } +} + +#[test] +fn pending_console_output_does_not_block_terminal_input() { + // Given: a deliberately blocked console and a completely full queue. + let (entered_tx, entered_rx) = mpsc::sync_channel(0); + let (release_tx, release_rx) = mpsc::sync_channel(0); + let output = crate::client_output::ConsoleOutput::new( + et_cli::client::FlowControlMode::Backpressure, + Box::new(GatedConsole { + entered: entered_tx, + release: release_rx, + }), + ) + .unwrap(); + let modes = TerminalModeState::default(); + assert!(output.try_write(&vec![1; 64 * 1024], &modes).unwrap()); + assert_eq!(entered_rx.recv().unwrap(), 64 * 1024); + assert!(output.try_write(&vec![2; 64 * 1024], &modes).unwrap()); + let packet = et_core::packet::Packet::new( + TerminalPacketType::TerminalBuffer as u8, + TerminalBuffer { + buffer: Some(b"pending".to_vec()), + } + .encode_to_vec(), + ); + let mut modes = TerminalModeState::default(); + assert!(matches!( + route_server_packet(packet, true, &mut modes, &output).unwrap(), + DisplayOutcome::Pending(_) + )); + + // When: Ctrl-C input is sent while output remains blocked. + let listener = TcpListener::bind((Ipv4Addr::LOCALHOST, 0)).unwrap(); + let address = listener.local_addr().unwrap(); + let connector = thread::spawn(move || TcpStream::connect(address).unwrap()); + let (server_stream, _) = listener.accept().unwrap(); + let client_stream = connector.join().unwrap(); + let key = [31u8; 32]; + let mut sender = Connection::new_client(client_stream, &key); + let mut receiver = Connection::new_server(server_stream, &key); + send_buffer(&mut sender, b"\x03").unwrap(); + + // Then: input arrives before either console write is released. + let received = receiver.read_packet().unwrap(); + let input = TerminalBuffer::decode(received.payload()).unwrap(); + assert_eq!(input.buffer.as_deref(), Some(b"\x03".as_slice())); + release_tx.send(()).unwrap(); + assert_eq!(entered_rx.recv().unwrap(), 64 * 1024); + release_tx.send(()).unwrap(); + drop(output); +} diff --git a/crates/et-bin/src/client_terminal_tests.rs b/crates/et-bin/src/client_terminal_tests.rs index c2f2810..b9e4c88 100644 --- a/crates/et-bin/src/client_terminal_tests.rs +++ b/crates/et-bin/src/client_terminal_tests.rs @@ -1,18 +1,138 @@ use std::net::{Ipv4Addr, TcpListener, TcpStream}; use et_core::crypto::KEY_LEN; -use et_core::proto::{TerminalBuffer, TerminalPacketType}; -use et_net::connection::{ConnError, Connection}; +use et_core::packet::Packet; +use et_core::proto::{TerminalBuffer, TerminalInfo, TerminalPacketType}; +use et_net::connection::{ConnError, Connection, WritePacketError}; use prost::Message; use super::{ - recover_initial_transport, recover_transport, send_command, TerminalModeState, TerminalReset, - GRACEFUL_TERMINAL_MODE_RESET, TERMINAL_MODE_RESET, + classify_forward_completion, command_payload, recover_initial_transport, recover_transport, + write_owned_recovering_with, write_owned_with_policy, OwnedWriteOutcome, OwnedWritePolicy, + RetainedCompletion, TerminalModeState, TerminalReset, GRACEFUL_TERMINAL_MODE_RESET, + TERMINAL_MODE_RESET, }; use crate::client_terminal::{connection_ended, RemoteLines}; use crate::error::ClientError; use crate::initial_connect::ReconnectOutcome; +#[test] +fn outbound_forwarding_prevents_successful_remote_completion() { + let current = Packet::new( + TerminalPacketType::PortForwardData as u8, + b"current".as_slice(), + ); + assert!(classify_forward_completion(Some(current), false).is_err()); + assert!(classify_forward_completion(None, true).is_err()); + assert!(classify_forward_completion(None, false).is_ok()); +} + +#[test] +fn retained_completion_waits_for_terminal_and_forward_capacity() { + let terminal = Packet::new( + TerminalPacketType::TerminalBuffer as u8, + b"terminal".as_slice(), + ); + let forwarding = Packet::new(91, b"forwarding".as_slice()); + let mut completion = RetainedCompletion::new(Some(terminal), Some(forwarding)); + let (terminal_attempt_tx, terminal_attempt_rx) = std::sync::mpsc::channel(); + let (forward_attempt_tx, forward_attempt_rx) = std::sync::mpsc::channel(); + let (terminal_release_tx, terminal_release_rx) = std::sync::mpsc::channel(); + let (forward_release_tx, forward_release_rx) = std::sync::mpsc::channel(); + + assert!(!completion + .advance( + |packet| { + terminal_attempt_tx.send(packet.header()).unwrap(); + Ok(match terminal_release_rx.try_recv() { + Ok(()) => None, + Err(std::sync::mpsc::TryRecvError::Empty) => Some(packet), + Err(error) => panic!("terminal release channel: {error}"), + }) + }, + |packet| { + forward_attempt_tx.send(packet.header()).unwrap(); + Ok(match forward_release_rx.try_recv() { + Ok(()) => None, + Err(std::sync::mpsc::TryRecvError::Empty) => Some(packet), + Err(error) => panic!("forward release channel: {error}"), + }) + }, + ) + .unwrap()); + assert_eq!( + terminal_attempt_rx.recv().unwrap(), + TerminalPacketType::TerminalBuffer as u8 + ); + assert_eq!(forward_attempt_rx.recv().unwrap(), 91); + + terminal_release_tx.send(()).unwrap(); + forward_release_tx.send(()).unwrap(); + assert!(completion + .advance( + |packet| { + terminal_attempt_tx.send(packet.header()).unwrap(); + terminal_release_rx.recv().unwrap(); + Ok(None) + }, + |packet| { + forward_attempt_tx.send(packet.header()).unwrap(); + forward_release_rx.recv().unwrap(); + Ok(None) + }, + ) + .unwrap()); + assert_eq!( + terminal_attempt_rx.recv().unwrap(), + TerminalPacketType::TerminalBuffer as u8 + ); + assert_eq!(forward_attempt_rx.recv().unwrap(), 91); +} + +#[test] +fn replaceable_terminal_size_never_retries_stale_payload_after_recovery() { + let (stream, _peer) = tcp_pair(); + let mut connection = Connection::new_client(stream, &[7u8; KEY_LEN]); + let old = TerminalInfo { + row: Some(24), + column: Some(80), + ..Default::default() + } + .encode_to_vec(); + let new = TerminalInfo { + row: Some(50), + column: Some(160), + ..Default::default() + } + .encode_to_vec(); + let mut size_provider = [old.clone(), new.clone()].into_iter(); + let initial = size_provider.next().unwrap(); + let sent = std::cell::RefCell::new(Vec::new()); + let outcome = write_owned_with_policy( + &mut connection, + TerminalPacketType::TerminalInfo as u8, + &initial, + OwnedWritePolicy::ReplaceableTerminalSize, + |_, _, payload| { + sent.borrow_mut().push(("initial", payload.to_vec())); + Err(WritePacketError::BeforeReplay(ConnError::Io( + std::io::ErrorKind::ConnectionReset.into(), + ))) + }, + |_, policy| { + assert!(matches!(policy, OwnedWritePolicy::ReplaceableTerminalSize)); + sent.borrow_mut() + .push(("recovery", size_provider.next().unwrap())); + Ok(true) + }, + ) + .unwrap(); + + assert!(matches!(outcome, OwnedWriteOutcome::Recovered)); + assert_eq!(sent.into_inner(), vec![("initial", old), ("recovery", new)]); + assert!(size_provider.next().is_none()); +} + #[test] fn command_exit_suffix_matches_no_exit_flag_and_remote_shell() { for (lines, no_exit, expected) in [ @@ -22,18 +142,9 @@ fn command_exit_suffix_matches_no_exit_flag_and_remote_shell() { (RemoteLines::Cmd, false, b"printf ok & exit\r\n".as_slice()), (RemoteLines::Cmd, true, b"printf ok\r\n".as_slice()), ] { - let (client_stream, server_stream) = tcp_pair(); - let key = [7u8; KEY_LEN]; - let worker = std::thread::spawn(move || { - let mut server = Connection::new_server(server_stream, &key); - server.read_packet().unwrap() - }); - let mut client = Connection::new_client(client_stream, &key); - send_command(&mut client, "printf ok", no_exit, lines).unwrap(); - let packet = worker.join().unwrap(); - assert_eq!(packet.header(), TerminalPacketType::TerminalBuffer as u8); + let payload = command_payload("printf ok", no_exit, lines).unwrap(); assert_eq!( - TerminalBuffer::decode(packet.payload()) + TerminalBuffer::decode(payload.as_slice()) .unwrap() .buffer .as_deref(), @@ -42,6 +153,72 @@ fn command_exit_suffix_matches_no_exit_flag_and_remote_shell() { } } +#[test] +fn before_replay_client_write_retries_plaintext_once_after_recovery() { + let (stream, _peer) = tcp_pair(); + let mut connection = Connection::new_client(stream, &[7u8; KEY_LEN]); + let payload = command_payload("echo once", false, RemoteLines::Posix).unwrap(); + let mut writes = 0; + let mut recoveries = 0; + let outcome = write_owned_recovering_with( + &mut connection, + TerminalPacketType::TerminalBuffer as u8, + &payload, + &mut |_| { + recoveries += 1; + Ok(ReconnectOutcome::Recovered) + }, + false, + |_, _, actual| { + writes += 1; + assert_eq!(actual, payload); + if writes == 1 { + Err(WritePacketError::BeforeReplay(ConnError::Io( + std::io::ErrorKind::ConnectionReset.into(), + ))) + } else { + Ok(()) + } + }, + ) + .unwrap(); + + assert!(matches!(outcome, OwnedWriteOutcome::Recovered)); + assert_eq!((writes, recoveries), (2, 1)); +} + +#[test] +fn replay_owned_client_write_recovers_without_plaintext_retry() { + let (stream, _peer) = tcp_pair(); + let mut connection = Connection::new_client(stream, &[7u8; KEY_LEN]); + let payload = TerminalBuffer { + buffer: Some(b"input-once".to_vec()), + } + .encode_to_vec(); + let mut writes = 0; + let mut recoveries = 0; + let outcome = write_owned_recovering_with( + &mut connection, + TerminalPacketType::TerminalBuffer as u8, + &payload, + &mut |_| { + recoveries += 1; + Ok(ReconnectOutcome::Recovered) + }, + false, + |_, _, _| { + writes += 1; + Err(WritePacketError::ReplayOwned(ConnError::Io( + std::io::ErrorKind::ConnectionReset.into(), + ))) + }, + ) + .unwrap(); + + assert!(matches!(outcome, OwnedWriteOutcome::Recovered)); + assert_eq!((writes, recoveries), (1, 1)); +} + #[test] fn every_socket_error_reconnects_including_sleep_timeout() { // macOS reports a stale post-sleep TCP flow as ETIMEDOUT (os error 60). @@ -144,15 +321,15 @@ fn observed_alternate_screen_selects_graceful_or_abrupt_reset() { #[test] fn alternate_screen_tracking_handles_split_enter_and_leave_sequences() { - let mut modes = TerminalModeState::default(); + let modes = TerminalModeState::default(); modes.observe(b"before\x1b[?10"); - assert!(!modes.alternate_screen); + assert!(!modes.alternate_screen()); modes.observe(b"49hinside"); - assert!(modes.alternate_screen); + assert!(modes.alternate_screen()); modes.observe(b"\x1b[?104"); - assert!(modes.alternate_screen); + assert!(modes.alternate_screen()); modes.observe(b"9lafter"); - assert!(!modes.alternate_screen); + assert!(!modes.alternate_screen()); } fn tcp_pair() -> (TcpStream, TcpStream) { diff --git a/crates/et-bin/src/client_terminal_windows.rs b/crates/et-bin/src/client_terminal_windows.rs index 8a59159..854ab35 100644 --- a/crates/et-bin/src/client_terminal_windows.rs +++ b/crates/et-bin/src/client_terminal_windows.rs @@ -18,9 +18,12 @@ use et_core::proto::TerminalPacketType; use et_net::connection::Connection; use et_net::forward::{is_forward_packet, Forwarder}; +use crate::client_output::ConsoleCompletion; use crate::client_terminal::{ - connection_ended, display_packet, recover_transport, send_buffer, send_size, terminal_error, - terminal_text, TerminalModeState, + classify_forward_completion, connection_ended, encoded_buffer, recover_transport, + terminal_error, terminal_io, terminal_size_payload, terminal_text, write_owned_recovering, + write_terminal_size_recovering, DisplayOutcome, OwnedWriteOutcome, RetainedCompletion, + TerminalModeState, }; use crate::error::ClientError; use crate::initial_connect::ReconnectOutcome; @@ -32,7 +35,7 @@ const POLL_INTERVAL: Duration = Duration::from_millis(10); pub fn pump( connection: &mut Connection, options: crate::client_terminal_loop::PumpOptions<'_>, - forwarder: &Forwarder, + forwarder: &mut Forwarder, mut reconnect: F, ) -> Result<(), ClientError> where @@ -41,10 +44,13 @@ where let crate::client_terminal_loop::PumpOptions { read_stdin, keepalive_seconds, + flow_control, terminal_enabled, auto_cursor_report, terminal_modes, } = options; + let console_output = crate::client_output::ConsoleOutput::stdout(flow_control) + .map_err(|error| terminal_io("starting console output worker", error))?; let interval = Duration::from_secs(u64::from(keepalive_seconds.max(1))); let silence = interval.saturating_mul(MISSED_KEEPALIVES); let mut last_received = Instant::now(); @@ -55,7 +61,32 @@ where // A forwarding packet the worker had no room for. While it is held, no // further session packets are read so forwarding data stays ordered. let mut pending_forward: Option = None; + let mut pending_output: Option = None; loop { + console_output + .check_error() + .map_err(|error| terminal_io("writing terminal output", error))?; + if auto_cursor_report { + for _ in 0..console_output + .take_cursor_reports() + .map_err(|error| terminal_io("reading console confirmations", error))? + { + if matches!( + write_cursor_report(connection, &mut reconnect, terminal_enabled)?, + OwnedWriteOutcome::SessionEnded + ) { + return finish_remote_completion( + console_output, + pending_output, + pending_forward, + terminal_enabled, + terminal_modes, + forwarder, + None, + ); + } + } + } let mut reconnect_needed = false; // Retry the held packet first: draining the forwarder's outbound // queue below is what frees worker capacity, so this makes progress @@ -65,6 +96,30 @@ where .try_receive(packet) .map_err(|error| terminal_text(error.to_string()))?; } + if let Some(packet) = pending_output.take() { + match route_server_packet(packet, terminal_enabled, terminal_modes, &console_output)? { + DisplayOutcome::Displayed { cursor_report } + if cursor_report && auto_cursor_report && !console_output.is_async() => + { + if matches!( + write_cursor_report(connection, &mut reconnect, terminal_enabled)?, + OwnedWriteOutcome::SessionEnded + ) { + return finish_remote_completion( + console_output, + pending_output, + pending_forward, + terminal_enabled, + terminal_modes, + forwarder, + None, + ); + } + } + DisplayOutcome::Displayed { .. } => {} + DisplayOutcome::Pending(packet) => pending_output = Some(packet), + } + } // 1. Console input and resize notifications. if read_stdin { @@ -79,19 +134,52 @@ where if bytes.is_empty() { continue; } - match send_buffer(connection, &bytes) { - Ok(()) => {} - Err(error) if connection_ended(&error) => reconnect_needed = true, - Err(error) => return Err(terminal_error(error)), + let payload = encoded_buffer(&bytes); + match write_owned_recovering( + connection, + TerminalPacketType::TerminalBuffer as u8, + &payload, + &mut reconnect, + terminal_enabled, + )? { + OwnedWriteOutcome::Written => {} + OwnedWriteOutcome::Recovered => reconnect_needed = false, + OwnedWriteOutcome::SessionEnded => { + return finish_remote_completion( + console_output, + pending_output, + pending_forward, + terminal_enabled, + terminal_modes, + forwarder, + None, + ); + } } } - Event::Resize(_, _) if terminal_enabled => match send_size(connection) { - Ok(()) => {} - Err(ClientError::Transport(error)) if connection_ended(&error) => { - reconnect_needed = true; + Event::Resize(_, _) if terminal_enabled => { + if let Some(payload) = terminal_size_payload()? { + match write_terminal_size_recovering( + connection, + &payload, + &mut reconnect, + )? { + OwnedWriteOutcome::Written => {} + OwnedWriteOutcome::Recovered => reconnect_needed = false, + OwnedWriteOutcome::SessionEnded => { + return finish_remote_completion( + console_output, + pending_output, + pending_forward, + terminal_enabled, + terminal_modes, + forwarder, + None, + ); + } + } } - Err(error) => return Err(error), - }, + } // Upstream forwards neither mouse nor focus records. Event::Mouse(mouse) if matches!(mouse.kind, MouseEventKind::Moved) => {} _ => {} @@ -100,7 +188,7 @@ where } // 2. Server packets. - while pending_forward.is_none() { + while pending_forward.is_none() && pending_output.is_none() { match connection.try_read_packet() { Ok(Some(packet)) => { last_received = Instant::now(); @@ -113,11 +201,40 @@ where pending_forward = forwarder .try_receive(packet) .map_err(|error| terminal_text(error.to_string()))?; - } else if route_server_packet(packet, terminal_enabled, terminal_modes)? - && auto_cursor_report - { - let _ = - send_buffer(connection, crate::client_terminal::CURSOR_REPORT_REPLY); + } else { + match route_server_packet( + packet, + terminal_enabled, + terminal_modes, + &console_output, + )? { + DisplayOutcome::Displayed { cursor_report } + if cursor_report + && auto_cursor_report + && !console_output.is_async() => + { + if matches!( + write_cursor_report( + connection, + &mut reconnect, + terminal_enabled, + )?, + OwnedWriteOutcome::SessionEnded + ) { + return finish_remote_completion( + console_output, + pending_output, + pending_forward, + terminal_enabled, + terminal_modes, + forwarder, + None, + ); + } + } + DisplayOutcome::Displayed { .. } => {} + DisplayOutcome::Pending(packet) => pending_output = Some(packet), + } } } Ok(None) => break, @@ -134,23 +251,44 @@ where .try_outbound() .map_err(|error| terminal_text(error.to_string()))? { - match connection.write_packet(packet.header(), packet.payload()) { - Ok(()) => {} - Err(error) if connection_ended(&error) => { - reconnect_needed = true; - break; + match write_owned_recovering( + connection, + packet.header(), + packet.payload(), + &mut reconnect, + terminal_enabled, + )? { + OwnedWriteOutcome::Written => {} + OwnedWriteOutcome::Recovered => reconnect_needed = false, + OwnedWriteOutcome::SessionEnded => { + return finish_remote_completion( + console_output, + pending_output, + pending_forward, + terminal_enabled, + terminal_modes, + forwarder, + Some(packet), + ); } - Err(error) => return Err(terminal_error(error)), } } let now = Instant::now(); - if now.saturating_duration_since(last_received) >= silence { + if pending_output.is_none() && now.saturating_duration_since(last_received) >= silence { reconnect_needed = true; } if reconnect_needed { if !recover_transport(connection, &mut reconnect, terminal_enabled)? { - return Ok(()); + return finish_remote_completion( + console_output, + pending_output, + pending_forward, + terminal_enabled, + terminal_modes, + forwarder, + None, + ); } last_received = Instant::now(); next_keepalive = last_received + interval; @@ -160,12 +298,25 @@ where // The payload acknowledges everything read so far, so the server // can trim its replay backup; legacy servers ignore it. let ack = connection.keepalive_ack(); - if connection - .write_packet(TerminalPacketType::KeepAlive as u8, &ack) - .is_err() - && !recover_transport(connection, &mut reconnect, terminal_enabled)? - { - return Ok(()); + if matches!( + write_owned_recovering( + connection, + TerminalPacketType::KeepAlive as u8, + &ack, + &mut reconnect, + terminal_enabled, + )?, + OwnedWriteOutcome::SessionEnded + ) { + return finish_remote_completion( + console_output, + pending_output, + pending_forward, + terminal_enabled, + terminal_modes, + forwarder, + None, + ); } next_keepalive = Instant::now() + interval; } @@ -180,17 +331,88 @@ where } } +fn finish_remote_completion( + output: crate::client_output::ConsoleOutput, + pending_output: Option, + pending_forward: Option, + terminal_enabled: bool, + terminal_modes: &mut TerminalModeState, + forwarder: &mut Forwarder, + current_outbound: Option, +) -> Result<(), ClientError> { + let mut retained = RetainedCompletion::new(pending_output, pending_forward); + loop { + output + .check_error() + .map_err(|error| terminal_io("writing retained terminal output", error))?; + if retained.advance( + |packet| match route_server_packet(packet, terminal_enabled, terminal_modes, &output)? { + DisplayOutcome::Displayed { .. } => Ok(None), + DisplayOutcome::Pending(packet) => Ok(Some(packet)), + }, + |packet| { + forwarder + .try_receive(packet) + .map_err(|error| terminal_text(error.to_string())) + }, + )? { + let abandoned = forwarder + .shutdown_hard() + .map_err(|error| terminal_text(error.to_string()))?; + classify_forward_completion(current_outbound, abandoned)?; + return output + .complete(ConsoleCompletion::RemoteSessionEnded) + .map_err(|error| terminal_io("draining terminal output", error)); + } + if let Some(packet) = forwarder + .try_outbound() + .map_err(|error| terminal_text(error.to_string()))? + { + let abandoned = forwarder + .shutdown_hard() + .map_err(|error| terminal_text(error.to_string()))?; + classify_forward_completion(Some(packet), abandoned)?; + } + std::thread::sleep(POLL_INTERVAL); + } +} + +fn write_cursor_report( + connection: &mut Connection, + reconnect: &mut F, + send_terminal_size: bool, +) -> Result +where + F: FnMut(&mut Connection) -> Result, +{ + let payload = encoded_buffer(crate::client_terminal::CURSOR_REPORT_REPLY); + write_owned_recovering( + connection, + TerminalPacketType::TerminalBuffer as u8, + &payload, + reconnect, + send_terminal_size, + ) +} + /// Returns `true` when a cursor position report must be sent back. fn route_server_packet( packet: et_core::packet::Packet, terminal_enabled: bool, terminal_modes: &mut TerminalModeState, -) -> Result { + output: &crate::client_output::ConsoleOutput, +) -> Result { if terminal_enabled || packet.header() == TerminalPacketType::KeepAlive as u8 { - return display_packet(packet, terminal_modes); + return crate::client_terminal::display_packet_with(packet, |bytes| { + output + .try_write(bytes, terminal_modes) + .map_err(|error| terminal_io("writing terminal output", error)) + }); } if packet.header() == TerminalPacketType::TerminalBuffer as u8 { - return Ok(false); + return Ok(DisplayOutcome::Displayed { + cursor_report: false, + }); } Err(terminal_text( "server sent an unsupported no-terminal packet", diff --git a/crates/et-bin/src/forward_config.rs b/crates/et-bin/src/forward_config.rs index 2480fb3..b4a327d 100644 --- a/crates/et-bin/src/forward_config.rs +++ b/crates/et-bin/src/forward_config.rs @@ -148,6 +148,7 @@ pub fn build( jumphost: Some(false), reversetunnels: reverse_tunnels, environmentvariables: std::collections::HashMap::new(), + flowcontrol: args.flow_control.protocol_value(), }, }) } diff --git a/crates/et-bin/src/main.rs b/crates/et-bin/src/main.rs index 266b8b3..dd29c4e 100644 --- a/crates/et-bin/src/main.rs +++ b/crates/et-bin/src/main.rs @@ -16,6 +16,10 @@ use std::ffi::OsString; fn main() { + #[cfg(windows)] + if std::env::args_os().nth(1).as_deref() == Some(std::ffi::OsStr::new("__et-console-writer")) { + std::process::exit(crate::client_output::run_windows_helper()); + } #[cfg(unix)] if let Some(code) = et_net::user_socket_ops::maybe_run_helper() { std::process::exit(code); @@ -83,6 +87,7 @@ fn role(name: &str, args: &[OsString]) -> Result { mod bootstrap; mod client; mod client_environment; +mod client_output; mod client_terminal; mod client_terminal_loop; #[cfg(windows)] diff --git a/crates/et-bin/src/terminal_jump.rs b/crates/et-bin/src/terminal_jump.rs index 0703421..30a3620 100644 --- a/crates/et-bin/src/terminal_jump.rs +++ b/crates/et-bin/src/terminal_jump.rs @@ -7,20 +7,20 @@ //! the same id/passkey, and then relays packets verbatim in both directions. use et_net::local::LocalStream; -use std::io::{self, Read}; +use std::io::{self, Read, Write}; use std::net::ToSocketAddrs; use std::time::Duration; use et_core::keys::passkey_to_key; use et_core::packet::Packet; use et_core::proto::{ - ConnectResponse, ConnectStatus, EtPacketType, InitialPayload, InitialResponse, + ConnectResponse, ConnectStatus, EtPacketType, FlowControlMode, InitialPayload, InitialResponse, TerminalPacketType, }; use et_net::connection::Connection; use et_net::framing_io::{read_proto_limited, write_proto}; use et_net::handshake::{client_request, MAX_HANDSHAKE_PROTO_LEN}; -use et_net::local_packet::{write_local_packet, LocalPacketDecoder}; +use et_net::local_packet::{encode_local_packet, write_local_packet, LocalPacketDecoder}; use prost::Message; #[cfg(unix)] use rustix::event::{poll, PollFd, PollFlags}; @@ -32,6 +32,22 @@ const READ_BUFFER: usize = 16 * 1024; /// Upstream retries the destination connection three times before failing. const CONNECT_ATTEMPTS: usize = 3; +#[cfg(test)] +fn run( + router: LocalStream, + input: &CredentialInput, + destination_host: &str, + destination_port: u16, +) -> Result { + run_with_startup( + router, + input, + destination_host, + destination_port, + |_| Ok(()), + ) +} + pub fn run_with_startup( mut router: LocalStream, input: &CredentialInput, @@ -46,6 +62,20 @@ where if !payload.jumphost.unwrap_or(false) { return Err("Jumphost should be set by the initial client".to_owned()); } + let flow_control = payload + .flowcontrol + .and_then(|value| FlowControlMode::try_from(value).ok()) + .unwrap_or(FlowControlMode::None); + let bounded_output = match flow_control { + FlowControlMode::None => false, + FlowControlMode::Backpressure | FlowControlMode::Discard => true, + }; + if bounded_output { + // Destination output enters the jumphost router through this terminal- + // side sender. Keep pressure in the server's bounded application lanes. + et_net::local::minimize_terminal_output_buffering(&router) + .map_err(|error| format!("could not bound jumphost output buffering: {error}"))?; + } // The destination runs a real terminal, so the relayed payload must not // ask it to start another jumphost. let mut payload = payload; @@ -124,6 +154,20 @@ fn try_connect_once( port: u16, payload: &InitialPayload, ) -> Result { + try_connect_once_observed(id, key, host, port, payload, |_| Ok(())) +} + +fn try_connect_once_observed( + id: &str, + key: &[u8; 32], + host: &str, + port: u16, + payload: &InitialPayload, + observe_before_payload: F, +) -> Result +where + F: FnOnce(&Connection) -> Result<(), String>, +{ let addresses = (host, port) .to_socket_addrs() .map_err(|error| format!("could not resolve {host}: {error}"))?; @@ -169,6 +213,17 @@ fn try_connect_once( None => return Err("destination sent an unknown connect status".to_owned()), } let mut connection = Connection::new_client(stream, key); + match payload + .flowcontrol + .and_then(|value| FlowControlMode::try_from(value).ok()) + .unwrap_or(FlowControlMode::None) + { + FlowControlMode::None => {} + FlowControlMode::Backpressure | FlowControlMode::Discard => connection + .minimize_output_buffering() + .map_err(|error| format!("could not bound destination output buffering: {error}"))?, + } + observe_before_payload(&connection)?; connection .write_packet_live(EtPacketType::InitialPayload as u8, &payload.encode_to_vec()) .map_err(|error| format!("could not send INITIAL_PAYLOAD: {error}"))?; @@ -197,7 +252,15 @@ fn is_acknowledgement(packet: &Packet) -> bool { } /// Relay packets verbatim between the local router and the destination. -fn relay(mut router: LocalStream, destination: &mut Connection) -> Result { +fn relay(router: LocalStream, destination: &mut Connection) -> Result { + relay_with_output_observer(router, destination, || {}) +} + +fn relay_with_output_observer( + mut router: LocalStream, + destination: &mut Connection, + mut output_pending: impl FnMut(), +) -> Result { router .set_nonblocking(true) .map_err(|error| format!("could not configure the router socket: {error}"))?; @@ -205,8 +268,10 @@ fn relay(mut router: LocalStream, destination: &mut Connection) -> Result, usize)> = None; + #[cfg(windows)] loop { - let mut progress = false; + let mut progress = write_pending_local(&mut router, &mut pending_output)?; if let Some(packet) = read_router_packet(&mut router, &mut decoder)? { progress = true; decoder = LocalPacketDecoder::new(); @@ -218,13 +283,20 @@ fn relay(mut router: LocalStream, destination: &mut Connection) -> Result { progress = true; let packet = router_packet(destination, packet); - if write_local_packet(&mut router, &packet).is_err() { - return Ok(0); + pending_output = Some(( + encode_local_packet(&packet).map_err(|error| { + format!("could not frame destination output: {error}") + })?, + 0, + )); + let _ = write_pending_local(&mut router, &mut pending_output)?; + if pending_output.is_some() { + output_pending(); } } Ok(None) => break, @@ -236,31 +308,66 @@ fn relay(mut router: LocalStream, destination: &mut Connection) -> Result, usize)> = None; + #[cfg(unix)] + let mut resume_destination_drain = false; + #[cfg(unix)] + let mut destination_closed = false; + #[cfg(unix)] + let mut destination_drained = false; + #[cfg(unix)] loop { - let client = destination - .try_clone_stream() - .map_err(|error| format!("could not poll the destination: {error}"))?; - let mut descriptors = [ - PollFd::new(&router, PollFlags::IN | PollFlags::HUP | PollFlags::ERR), - PollFd::new(&client, PollFlags::IN | PollFlags::HUP | PollFlags::ERR), - ]; - // poll() is never restarted by SA_RESTART; retry on EINTR so a stray - // signal cannot kill the jump-host bridge. - loop { - match poll(&mut descriptors, None) { - Ok(_) => break, - Err(error) if error == rustix::io::Errno::INTR => {} - Err(error) => return Err(format!("poll failed: {error}")), - } + if destination_drained && pending_output.is_none() { + return Ok(0); } - let router_events = descriptors[0].revents(); - let client_events = descriptors[1].revents(); - drop(client); + let router_flags = router_poll_flags(destination_closed, pending_output.is_some()); + let (router_events, client_events) = if destination_closed { + let mut descriptors = [PollFd::new(&router, router_flags)]; + loop { + match poll(&mut descriptors, None) { + Ok(_) => break, + Err(error) if error == rustix::io::Errno::INTR => {} + Err(error) => return Err(format!("poll failed: {error}")), + } + } + (descriptors[0].revents(), PollFlags::empty()) + } else { + let client = destination + .try_clone_stream() + .map_err(|error| format!("could not poll the destination: {error}"))?; + let client_flags = if pending_output.is_none() { + PollFlags::IN | PollFlags::HUP | PollFlags::ERR + } else { + PollFlags::HUP | PollFlags::ERR + }; + let mut descriptors = [ + PollFd::new(&router, router_flags), + PollFd::new(&client, client_flags), + ]; + // poll() is never restarted by SA_RESTART; retry on EINTR so a + // stray signal cannot kill the jump-host bridge. + loop { + match poll(&mut descriptors, None) { + Ok(_) => break, + Err(error) if error == rustix::io::Errno::INTR => {} + Err(error) => return Err(format!("poll failed: {error}")), + } + } + (descriptors[0].revents(), descriptors[1].revents()) + }; + if client_events.intersects(PollFlags::HUP | PollFlags::ERR) { + destination_closed = true; + } if router_events.intersects(PollFlags::HUP | PollFlags::ERR) { return Ok(0); } - if router_events.contains(PollFlags::IN) { + if router_events.contains(PollFlags::OUT) + && write_pending_local(&mut router, &mut pending_output).is_err() + { + return Ok(0); + } + if !destination_closed && router_events.contains(PollFlags::IN) { if let Some(packet) = read_router_packet(&mut router, &mut decoder)? { decoder = LocalPacketDecoder::new(); let packet = destination_packet(destination, packet); @@ -272,22 +379,83 @@ fn relay(mut router: LocalStream, destination: &mut Connection) -> Result { + resume_destination_drain = true; let packet = router_packet(destination, packet); - if write_local_packet(&mut router, &packet).is_err() { + pending_output = Some(( + encode_local_packet(&packet).map_err(|error| { + format!("could not frame destination output: {error}") + })?, + 0, + )); + if write_pending_local(&mut router, &mut pending_output).is_err() { return Ok(0); } + if pending_output.is_some() { + output_pending(); + } + } + Ok(None) => { + resume_destination_drain = false; + destination_drained = destination_closed; + break; + } + Err(_) => { + destination_closed = true; + destination_drained = true; + break; } - Ok(None) => break, - Err(_) => return Ok(0), } } } - if client_events.intersects(PollFlags::HUP | PollFlags::ERR) { - return Ok(0); + } +} + +#[cfg(unix)] +fn router_poll_flags(destination_closed: bool, output_pending: bool) -> PollFlags { + let input = if destination_closed { + PollFlags::empty() + } else { + PollFlags::IN + }; + input + | PollFlags::HUP + | PollFlags::ERR + | if output_pending { + PollFlags::OUT + } else { + PollFlags::empty() + } +} + +fn write_pending_local( + router: &mut LocalStream, + pending: &mut Option<(Vec, usize)>, +) -> Result { + let Some((frame, offset)) = pending.as_mut() else { + return Ok(false); + }; + loop { + match router.write(&frame[*offset..]) { + Ok(0) => return Err("jumphost router stopped accepting output".to_owned()), + Ok(count) => { + *offset += count; + if *offset == frame.len() { + *pending = None; + return Ok(true); + } + } + Err(error) if error.kind() == io::ErrorKind::Interrupted => {} + Err(error) if error.kind() == io::ErrorKind::WouldBlock => return Ok(false), + Err(error) => return Err(format!("could not write destination output: {error}")), } } } @@ -347,25 +515,5 @@ fn read_router_packet( } #[cfg(test)] -mod tests { - use super::downstream_requires_ack; - use et_core::proto::InitialResponse; - use prost::Message; - - #[test] - fn malformed_downstream_initial_response_fails_closed() { - assert_eq!( - downstream_requires_ack(&[0xff]), - Err("destination sent a malformed INITIAL_RESPONSE".to_owned()) - ); - assert_eq!( - downstream_requires_ack( - &InitialResponse { - error: Some("ordinary fatal".to_owned()), - } - .encode_to_vec() - ), - Ok(true) - ); - } -} +#[path = "terminal_jump_tests.rs"] +mod tests; diff --git a/crates/et-bin/src/terminal_jump_tests.rs b/crates/et-bin/src/terminal_jump_tests.rs new file mode 100644 index 0000000..2a0940a --- /dev/null +++ b/crates/et-bin/src/terminal_jump_tests.rs @@ -0,0 +1,419 @@ +use super::*; +use std::net::{Ipv4Addr, TcpListener}; +use std::sync::mpsc; +use std::thread; + +use socket2::SockRef; + +const TEST_TIMEOUT: Duration = Duration::from_secs(3); +const PTY_CONTRACT: Duration = Duration::from_secs(5); +const ID: &str = "abcdefghijklmnop"; +const KEY_TEXT: &str = "ABCDEFGHIJKLMNOPQRSTUVWXYZabcdef"; + +#[test] +fn malformed_downstream_initial_response_fails_closed() { + assert_eq!( + downstream_requires_ack(&[0xff]), + Err("destination sent a malformed INITIAL_RESPONSE".to_owned()) + ); + assert_eq!( + downstream_requires_ack( + &InitialResponse { + error: Some("ordinary fatal".to_owned()), + } + .encode_to_vec() + ), + Ok(true) + ); +} + +#[test] +fn jumphost_clamps_destination_before_typed_initial_payload() { + // Given: a destination requiring proof of the clamp before it accepts + // INITIAL_PAYLOAD and starts terminal output. + let listener = TcpListener::bind((Ipv4Addr::LOCALHOST, 0)).unwrap(); + let port = listener.local_addr().unwrap().port(); + let (clamped_tx, clamped_rx) = mpsc::sync_channel(1); + let key = passkey_to_key(KEY_TEXT).unwrap(); + let destination = thread::spawn(move || { + let (mut stream, _) = listener.accept().unwrap(); + stream.set_read_timeout(Some(TEST_TIMEOUT)).unwrap(); + stream.set_write_timeout(Some(TEST_TIMEOUT)).unwrap(); + let _: et_core::proto::ConnectRequest = + read_proto_limited(&mut stream, MAX_HANDSHAKE_PROTO_LEN).unwrap(); + write_proto( + &mut stream, + &ConnectResponse { + status: Some(ConnectStatus::NewClient as i32), + error: None, + }, + ) + .unwrap(); + let observed = clamped_rx.recv_timeout(TEST_TIMEOUT).unwrap(); + assert!( + observed <= et_net::connection::FLOW_CONTROL_SOCKET_BUFFER_BYTES * 2, + "jumphost destination send buffer remained {observed} bytes" + ); + let mut connection = Connection::new_server(stream, &key); + let packet = connection.read_packet().unwrap(); + assert_eq!(packet.header(), EtPacketType::InitialPayload as u8); + let payload = InitialPayload::decode(packet.payload()).unwrap(); + assert_eq!( + payload.flowcontrol, + Some(FlowControlMode::Backpressure as i32) + ); + assert_eq!(payload.jumphost, Some(false)); + connection + .write_packet( + EtPacketType::InitialResponse as u8, + &InitialResponse { error: None }.encode_to_vec(), + ) + .unwrap(); + }); + let payload = InitialPayload { + jumphost: Some(false), + flowcontrol: Some(FlowControlMode::Backpressure as i32), + ..Default::default() + }; + + // When: the jumphost establishes its destination side. + let connection = + try_connect_once_observed(ID, &key, "127.0.0.1", port, &payload, |connection| { + let stream = connection + .try_clone_stream() + .map_err(|error| error.to_string())?; + let size = SockRef::from(&stream) + .send_buffer_size() + .map_err(|error| error.to_string())?; + clamped_tx.send(size).map_err(|error| error.to_string()) + }) + .unwrap(); + + // Then: initialization completed after bounded pressure and mode retention. + drop(connection); + destination.join().unwrap(); +} + +#[test] +fn slow_jumphost_output_does_not_block_ctrl_c_toward_destination() { + // Given: a real encrypted destination hop and a bounded local router link + // whose client side deliberately does not consume destination output. + let listener = TcpListener::bind((Ipv4Addr::LOCALHOST, 0)).unwrap(); + let address = listener.local_addr().unwrap(); + let connector = thread::spawn(move || std::net::TcpStream::connect(address).unwrap()); + let (server_stream, _) = listener.accept().unwrap(); + let client_stream = connector.join().unwrap(); + let key = [29u8; 32]; + let mut destination_client = Connection::new_client(client_stream, &key); + let mut destination_server = Connection::new_server(server_stream, &key); + let (relay_router, mut router_peer) = et_net::local::wake_pair().unwrap(); + saturate_router_output(&relay_router, &router_peer); + let (output_sent_tx, output_sent_rx) = mpsc::sync_channel(0); + let (pending_tx, pending_rx) = mpsc::sync_channel(0); + let (input_tx, input_rx) = mpsc::sync_channel(0); + let destination = thread::spawn(move || { + destination_server + .write_packet(71, &vec![b'p'; 60 * 1024]) + .unwrap(); + output_sent_tx.send(()).unwrap(); + let input = destination_server.read_packet().unwrap(); + input_tx + .send((input.header(), input.payload().to_vec())) + .unwrap(); + destination_server.shutdown().unwrap(); + }); + let relay = thread::spawn(move || { + relay_with_output_observer(relay_router, &mut destination_client, || { + pending_tx.send(()).unwrap(); + }) + }); + output_sent_rx.recv_timeout(PTY_CONTRACT).unwrap(); + pending_rx.recv_timeout(PTY_CONTRACT).unwrap(); + + // When: Ctrl-C arrives while ownership of the blocked prompt/output frame + // remains in the destination-to-router direction. + write_local_packet( + &mut router_peer, + &Packet::new(TerminalPacketType::TerminalBuffer as u8, vec![3]), + ) + .unwrap(); + + // Then: input reaches the destination within the unchanged five-second PTY contract. + assert_eq!( + input_rx.recv_timeout(PTY_CONTRACT).unwrap(), + (TerminalPacketType::TerminalBuffer as u8, vec![3]) + ); + destination.join().unwrap(); + drop(router_peer); + assert_eq!(relay.join().unwrap().unwrap(), 0); +} + +#[cfg(unix)] +#[test] +fn jumphost_drains_coalesced_destination_packets_without_new_socket_readiness() { + // Given: several small encrypted frames are already queued before the + // relay reads once, allowing one recv() to move all of them into the + // connection's userspace BackedReader. + let listener = TcpListener::bind((Ipv4Addr::LOCALHOST, 0)).unwrap(); + let address = listener.local_addr().unwrap(); + let connector = thread::spawn(move || std::net::TcpStream::connect(address).unwrap()); + let (server_stream, _) = listener.accept().unwrap(); + let client_stream = connector.join().unwrap(); + let key = [31u8; 32]; + let mut destination_client = Connection::new_client(client_stream, &key); + let mut destination_server = Connection::new_server(server_stream, &key); + for header in 91..96 { + destination_server.write_packet(header, &[header]).unwrap(); + } + let (relay_router, mut router_peer) = et_net::local::wake_pair().unwrap(); + router_peer.set_read_timeout(Some(TEST_TIMEOUT)).unwrap(); + let relay = thread::spawn(move || relay(relay_router, &mut destination_client)); + + // When: the router consumes every frame without sending input that could + // create a fresh poll event. + let mut relayed = Vec::new(); + for _ in 0..5 { + match et_net::local_packet::read_local_packet(&mut router_peer) { + Ok(packet) => relayed.push((packet.header(), packet.payload().to_vec())), + Err(_) => break, + } + } + drop(router_peer); + assert_eq!(relay.join().unwrap().unwrap(), 0); + + // Then: userspace-buffered frames were drained even after kernel POLLIN + // readiness disappeared. + assert_eq!( + relayed, + (91..96) + .map(|header| (header, vec![header])) + .collect::>() + ); +} + +#[cfg(unix)] +#[test] +fn jumphost_resumes_coalesced_destination_packets_after_router_backpressure() { + // Given: five destination packets coalesce in userspace while the router + // output queue is already full. + let listener = TcpListener::bind((Ipv4Addr::LOCALHOST, 0)).unwrap(); + let address = listener.local_addr().unwrap(); + let connector = thread::spawn(move || std::net::TcpStream::connect(address).unwrap()); + let (server_stream, _) = listener.accept().unwrap(); + let client_stream = connector.join().unwrap(); + let key = [37u8; 32]; + let mut destination_client = Connection::new_client(client_stream, &key); + let mut destination_server = Connection::new_server(server_stream, &key); + for header in 101..106 { + destination_server.write_packet(header, &[header]).unwrap(); + } + let (relay_router, mut router_peer) = et_net::local::wake_pair().unwrap(); + let filler_bytes = saturate_router_output(&relay_router, &router_peer); + router_peer.set_read_timeout(Some(TEST_TIMEOUT)).unwrap(); + let (pending_tx, pending_rx) = mpsc::sync_channel(0); + let relay = thread::spawn(move || { + relay_with_output_observer(relay_router, &mut destination_client, || { + pending_tx.send(()).unwrap(); + }) + }); + pending_rx.recv_timeout(TEST_TIMEOUT).unwrap(); + + // When: the router resumes reading after the first destination frame was + // interrupted by real socket backpressure. + router_peer + .read_exact(&mut vec![0u8; filler_bytes]) + .unwrap(); + let mut relayed = Vec::new(); + for _ in 0..5 { + match et_net::local_packet::read_local_packet(&mut router_peer) { + Ok(packet) => relayed.push((packet.header(), packet.payload().to_vec())), + Err(_) => break, + } + } + drop(router_peer); + assert_eq!(relay.join().unwrap().unwrap(), 0); + + // Then: clearing pending output resumes the userspace drain without a new + // destination POLLIN event. + assert_eq!( + relayed, + (101..106) + .map(|header| (header, vec![header])) + .collect::>() + ); +} + +#[cfg(unix)] +#[test] +fn closed_destination_poll_mask_ignores_readable_router_input() { + // Given: destination closure is latched while output remains pending. + // When: the router poll subscription is constructed. + let flags = router_poll_flags(true, true); + + // Then: readable input cannot wake a loop that deliberately ignores it, + // while output progress and local closure remain observable. + assert!(!flags.contains(PollFlags::IN)); + assert!(flags.contains(PollFlags::OUT)); + assert!(flags.contains(PollFlags::HUP)); + assert!(flags.contains(PollFlags::ERR)); +} + +#[cfg(unix)] +#[test] +fn jumphost_drains_buffered_destination_packets_after_hup() { + // Given: destination frames are coalesced before the peer closes, while + // the first local frame is held by real router backpressure. + let listener = TcpListener::bind((Ipv4Addr::LOCALHOST, 0)).unwrap(); + let address = listener.local_addr().unwrap(); + let connector = thread::spawn(move || std::net::TcpStream::connect(address).unwrap()); + let (server_stream, _) = listener.accept().unwrap(); + let client_stream = connector.join().unwrap(); + let key = [41u8; 32]; + let mut destination_client = Connection::new_client(client_stream, &key); + let mut destination_server = Connection::new_server(server_stream, &key); + for header in 111..116 { + destination_server.write_packet(header, &[header]).unwrap(); + } + let reset = destination_server.try_clone_stream().unwrap(); + SockRef::from(&reset) + .set_linger(Some(Duration::ZERO)) + .unwrap(); + drop(reset); + drop(destination_server); + let (relay_router, mut router_peer) = et_net::local::wake_pair().unwrap(); + let filler_bytes = saturate_router_output(&relay_router, &router_peer); + router_peer.set_read_timeout(Some(TEST_TIMEOUT)).unwrap(); + let (pending_tx, pending_rx) = mpsc::sync_channel(0); + let relay = thread::spawn(move || { + relay_with_output_observer(relay_router, &mut destination_client, || { + pending_tx.send(()).unwrap(); + }) + }); + pending_rx.recv_timeout(TEST_TIMEOUT).unwrap(); + + // When: the router resumes after destination HUP has already been observed. + router_peer + .read_exact(&mut vec![0u8; filler_bytes]) + .unwrap(); + let mut relayed = Vec::new(); + for _ in 0..5 { + let packet = et_net::local_packet::read_local_packet(&mut router_peer).unwrap(); + relayed.push((packet.header(), packet.payload().to_vec())); + } + + // Then: the pending frame and every packet already buffered in the + // destination BackedReader leave in order before the relay exits. + assert_eq!(relay.join().unwrap().unwrap(), 0); + assert_eq!( + relayed, + (111..116) + .map(|header| (header, vec![header])) + .collect::>() + ); +} + +#[test] +fn jumphost_run_bounds_router_sender_before_destination_output() { + // Given: the real jump relay is paused after the destination receives its + // typed initialization but before that destination may produce output. + let (run_router, mut router_peer) = et_net::local::wake_pair().unwrap(); + let router_observer = run_router.try_clone().unwrap(); + let payload = InitialPayload { + jumphost: Some(true), + flowcontrol: Some(FlowControlMode::Discard as i32), + ..Default::default() + }; + write_local_packet( + &mut router_peer, + &Packet::new( + TerminalPacketType::JumphostInit as u8, + payload.encode_to_vec(), + ), + ) + .unwrap(); + let listener = TcpListener::bind((Ipv4Addr::LOCALHOST, 0)).unwrap(); + let port = listener.local_addr().unwrap().port(); + let (initialized_tx, initialized_rx) = mpsc::sync_channel(1); + let (output_release_tx, output_release_rx) = mpsc::sync_channel(0); + let key = passkey_to_key(KEY_TEXT).unwrap(); + let destination = thread::spawn(move || { + let (mut stream, _) = listener.accept().unwrap(); + stream.set_read_timeout(Some(TEST_TIMEOUT)).unwrap(); + stream.set_write_timeout(Some(TEST_TIMEOUT)).unwrap(); + let _: et_core::proto::ConnectRequest = + read_proto_limited(&mut stream, MAX_HANDSHAKE_PROTO_LEN).unwrap(); + write_proto( + &mut stream, + &ConnectResponse { + status: Some(ConnectStatus::NewClient as i32), + error: None, + }, + ) + .unwrap(); + let mut connection = Connection::new_server(stream, &key); + let packet = connection.read_packet().unwrap(); + let received = InitialPayload::decode(packet.payload()).unwrap(); + assert_eq!(received.flowcontrol, Some(FlowControlMode::Discard as i32)); + assert_eq!(received.jumphost, Some(false)); + initialized_tx.send(()).unwrap(); + output_release_rx.recv_timeout(TEST_TIMEOUT).unwrap(); + connection + .write_packet( + EtPacketType::InitialResponse as u8, + &InitialResponse { error: None }.encode_to_vec(), + ) + .unwrap(); + }); + let input = crate::terminal_credentials::CredentialInput { + id: ID.to_owned(), + passkey: KEY_TEXT.to_owned(), + term: "xterm-256color".to_owned(), + }; + let relay = thread::spawn(move || run(run_router, &input, "127.0.0.1", port)); + + // When: destination startup reaches the exact pre-output barrier. + initialized_rx.recv_timeout(TEST_TIMEOUT).unwrap(); + + // Then: the jumphost router sender is already bounded and mode survived both hops. + let send_buffer = SockRef::from(&router_observer).send_buffer_size().unwrap(); + assert!( + send_buffer <= et_net::local::FLOW_CONTROL_SEND_BUFFER_BYTES * 2, + "jumphost router send buffer remained {send_buffer} bytes" + ); + output_release_tx.send(()).unwrap(); + destination.join().unwrap(); + let response = et_net::local_packet::read_local_packet(&mut router_peer).unwrap(); + assert_eq!(response.header(), TerminalPacketType::JumphostInit as u8); + assert_eq!( + InitialResponse::decode(response.payload()).unwrap().error, + None + ); + drop(router_peer); + assert_eq!(relay.join().unwrap().unwrap(), 0); +} + +fn saturate_router_output( + router: &et_net::local::LocalStream, + peer: &et_net::local::LocalStream, +) -> usize { + et_net::local::minimize_terminal_output_buffering(router).unwrap(); + SockRef::from(router) + .set_send_buffer_size(2 * 1024) + .unwrap(); + SockRef::from(peer).set_recv_buffer_size(2 * 1024).unwrap(); + let mut saturator = router.try_clone().unwrap(); + saturator.set_nonblocking(true).unwrap(); + let filler = [0u8; 16 * 1024]; + let mut saturated_bytes = 0usize; + loop { + match saturator.write(&filler) { + Ok(0) => panic!("router output closed before reaching backpressure"), + Ok(written) => saturated_bytes += written, + Err(error) if error.kind() == io::ErrorKind::Interrupted => {} + Err(error) if error.kind() == io::ErrorKind::WouldBlock => break, + Err(error) => panic!("could not saturate router output: {error}"), + } + } + assert!(saturated_bytes > 0); + saturated_bytes +} diff --git a/crates/et-bin/src/terminal_protocol.rs b/crates/et-bin/src/terminal_protocol.rs index 3a60e04..32e11f3 100644 --- a/crates/et-bin/src/terminal_protocol.rs +++ b/crates/et-bin/src/terminal_protocol.rs @@ -2,7 +2,7 @@ use et_net::local::LocalStream; use std::io::{self, Read, Write}; use et_core::packet::Packet; -use et_core::proto::{TermInit, TerminalBuffer, TerminalInfo, TerminalPacketType}; +use et_core::proto::{FlowControlMode, TermInit, TerminalBuffer, TerminalInfo, TerminalPacketType}; use et_net::local_packet::{read_local_packet, LocalPacketDecoder}; use portable_pty::{MasterPty, PtySize}; use prost::Message; @@ -11,7 +11,12 @@ pub(crate) const MAX_ENVIRONMENT: usize = 128; pub(crate) const MAX_ENV_VALUE: usize = 4096; const READ_BUFFER: usize = 16 * 1024; -pub fn read_initial_environment(router: &mut LocalStream) -> Result, String> { +pub struct TerminalInitialization { + pub environment: Vec<(String, String)>, + pub flow_control: FlowControlMode, +} + +pub fn read_initialization(router: &mut LocalStream) -> Result { let packet = read_local_packet(router) .map_err(|error| format!("could not read terminal initialization: {error}"))?; if packet.is_encrypted() || packet.header() != TerminalPacketType::TerminalInit as u8 { @@ -24,7 +29,12 @@ pub fn read_initial_environment(router: &mut LocalStream) -> Result Result>()?; + Ok(TerminalInitialization { + environment, + flow_control, + }) } pub fn read_ready_packet( @@ -182,18 +196,22 @@ mod tests { TermInit { environmentnames: vec!["A".to_owned()], environmentvalues: Vec::new(), + flowcontrol: None, }, TermInit { environmentnames: vec!["BAD-NAME".to_owned()], environmentvalues: vec!["value".to_owned()], + flowcontrol: None, }, TermInit { environmentnames: vec!["VALID".to_owned()], environmentvalues: vec!["bad\0value".to_owned()], + flowcontrol: None, }, TermInit { environmentnames: vec!["VALID".to_owned()], environmentvalues: vec!["x".repeat(MAX_ENV_VALUE + 1)], + flowcontrol: None, }, ] { let packet = Packet::new(TerminalPacketType::TerminalInit as u8, init.encode_to_vec()); @@ -204,6 +222,24 @@ mod tests { fn read_environment_packet(packet: Packet) -> Result, String> { let (mut reader, mut writer) = et_net::local::wake_pair().unwrap(); write_local_packet(&mut writer, &packet).unwrap(); - read_initial_environment(&mut reader) + read_initialization(&mut reader).map(|initialization| initialization.environment) + } + + #[test] + fn initialization_retains_the_typed_flow_control_mode() { + let (mut terminal, mut server) = et_net::local::wake_pair().unwrap(); + let init = TermInit { + environmentnames: Vec::new(), + environmentvalues: Vec::new(), + flowcontrol: Some(FlowControlMode::Discard as i32), + }; + write_local_packet( + &mut server, + &Packet::new(TerminalPacketType::TerminalInit as u8, init.encode_to_vec()), + ) + .unwrap(); + + let initialization = read_initialization(&mut terminal).unwrap(); + assert_eq!(initialization.flow_control, FlowControlMode::Discard); } } diff --git a/crates/et-bin/src/terminal_pty.rs b/crates/et-bin/src/terminal_pty.rs index f4dda7e..c1d3f15 100644 --- a/crates/et-bin/src/terminal_pty.rs +++ b/crates/et-bin/src/terminal_pty.rs @@ -6,7 +6,7 @@ use std::thread; use std::time::{Duration, Instant}; use et_core::packet::Packet; -use et_core::proto::{TerminalBuffer, TerminalPacketType}; +use et_core::proto::{FlowControlMode, TerminalBuffer, TerminalPacketType}; use et_net::local::LocalStream; use et_net::local_packet::{write_local_packet_until_cancelled, LocalPacketDecoder}; #[cfg(unix)] @@ -23,7 +23,7 @@ use sysinfo::{Pid as SystemPid, ProcessesToUpdate, Signal as SystemSignal, Syste const MAX_OUTPUT_CHUNK: usize = 16 * 1024; const FINAL_OUTPUT_DRAIN_TIMEOUT: Duration = Duration::from_secs(2); -use crate::terminal_protocol::{handle_packet, read_initial_environment, read_ready_packet}; +use crate::terminal_protocol::{handle_packet, read_initialization, read_ready_packet}; enum WorkerEvent { Output(Result<(), String>), @@ -54,7 +54,13 @@ fn run_with_command( where F: FnOnce(&mut LocalStream) -> Result<(), String>, { - let environment = read_initial_environment(&mut router)?; + let initialization = read_initialization(&mut router)?; + if initialization.flow_control != FlowControlMode::None { + // Keep terminal output in the server's bounded application queue, + // rather than a large opaque local-socket queue (upstream PR #730). + et_net::local::minimize_terminal_output_buffering(&router) + .map_err(|error| format!("could not bound terminal output buffering: {error}"))?; + } let pair = native_pty_system() .openpty(PtySize { rows: 24, @@ -64,7 +70,7 @@ where }) .map_err(|error| format!("could not open PTY: {error}"))?; command.env("TERM", term); - for (name, value) in environment { + for (name, value) in initialization.environment { command.env(name, value); } // Complete every fallible descriptor allocation before creating the shell. diff --git a/crates/et-bin/tests/client_bootstrap.rs b/crates/et-bin/tests/client_bootstrap.rs index 5b28299..54d7321 100644 --- a/crates/et-bin/tests/client_bootstrap.rs +++ b/crates/et-bin/tests/client_bootstrap.rs @@ -742,6 +742,7 @@ fn posix_client_bounds_locale_to_local_terminal_packet() { let term_init = TermInit { environmentnames: environment.keys().cloned().collect(), environmentvalues: environment.values().cloned().collect(), + flowcontrol: None, }; let packet = Packet::new( TerminalPacketType::TerminalInit as u8, diff --git a/crates/et-bin/tests/flow_control_tty_qa.rs b/crates/et-bin/tests/flow_control_tty_qa.rs new file mode 100644 index 0000000..111627d --- /dev/null +++ b/crates/et-bin/tests/flow_control_tty_qa.rs @@ -0,0 +1,177 @@ +#![cfg(unix)] +#![forbid(unsafe_code)] + +mod flow_control_tty_support; + +use std::fs; +use std::io::{Read, Write}; +use std::sync::mpsc; +use std::thread; +use std::time::{Duration, Instant}; + +use flow_control_tty_support::{ + receive_until, Stack, ThrottleProxy, MAX_PROMPT_LATENCY, SATURATION_BYTES, + THROTTLE_BYTES_PER_SECOND, +}; +use portable_pty::{native_pty_system, CommandBuilder, PtySize}; + +#[test] +fn flow_control_keeps_ctrl_c_and_prompt_responsive_on_a_slow_link() { + let evidence = std::env::var_os("ET_FLOW_QA_EVIDENCE_DIR").map(std::path::PathBuf::from); + if let Some(directory) = &evidence { + fs::create_dir_all(directory).unwrap(); + } + + for mode in ["none", "backpressure", "discard"] { + let stack = Stack::start(); + let bytes_per_second = THROTTLE_BYTES_PER_SECOND; + let proxy = ThrottleProxy::start(stack.port, bytes_per_second, SATURATION_BYTES); + let pair = native_pty_system() + .openpty(PtySize { + rows: 24, + cols: 80, + pixel_width: 800, + pixel_height: 480, + }) + .unwrap(); + let mut client = CommandBuilder::new(env!("CARGO_BIN_EXE_et")); + client.args([ + "--flow-control", + mode, + "--terminal-path", + stack.terminal.to_str().unwrap(), + "--serverfifo", + stack.router.to_str().unwrap(), + "-p", + &proxy.port.to_string(), + "127.0.0.1", + ]); + client.env( + "PATH", + format!( + "{}:{}", + stack.directory.display(), + std::env::var("PATH").unwrap() + ), + ); + client.env("TERM", "xterm-256color"); + let mut child = pair.slave.spawn_command(client).unwrap(); + drop(pair.slave); + + let mut writer = pair.master.take_writer().unwrap(); + let mut reader = pair.master.try_clone_reader().unwrap(); + // Keep the test harness from adding its own half-megabyte output + // queue on top of the ET pipeline being measured. + let (sender, receiver) = mpsc::sync_channel(32); + let reader_thread = thread::spawn(move || { + let mut chunk = [0u8; 8192]; + loop { + match reader.read(&mut chunk) { + Ok(0) | Err(_) => return, + Ok(count) if sender.send(chunk[..count].to_vec()).is_err() => return, + Ok(_) => {} + } + } + }); + + writeln!( + writer, + "interrupted=; \ + trap 'interrupted=1; printf \"\\nFLOW-INTERRUPTED\\n\"' INT; \ + printf 'FLOW-%s\\n' START; \ + IFS= read -r flow_release; \ + i=0; while [ -z \"$interrupted\" ] && [ \"$i\" -lt 65536 ]; do \ + printf '0123456789abcdef0123456789abcdef0123456789abcdef0123456789abcdef'; \ + i=$((i + 1)); done; trap - INT; printf 'FLOW-%s\\n' READY" + ) + .unwrap(); + let startup_timeout = Duration::from_secs(10); + let output = match receive_until(&receiver, Vec::new(), b"FLOW-START\r\n", startup_timeout) + { + Ok(output) => output, + Err(error) => { + child.kill().unwrap(); + drop(writer); + let _ = child.wait(); + reader_thread.join().unwrap(); + panic!("{mode}: waiting for FLOW-START: {error}"); + } + }; + writer.write_all(b"FLOW-RELEASE\n").unwrap(); + writer.flush().unwrap(); + proxy + .wait_saturated(Duration::from_secs(40)) + .unwrap_or_else(|error| panic!("{mode}: exact saturation event: {error}")); + let mut output = output; + while let Ok(chunk) = receiver.try_recv() { + output.extend(chunk); + } + let interrupted = Instant::now(); + let prompt_timeout = MAX_PROMPT_LATENCY; + let deadline = interrupted + prompt_timeout; + writer.write_all(b"\x03").unwrap(); + writer.flush().unwrap(); + let interrupt = receive_until( + &receiver, + output.clone(), + b"FLOW-INTERRUPTED\r\n", + prompt_timeout, + ); + let prompt = interrupt.and_then(|output| { + let remaining = deadline + .checked_duration_since(Instant::now()) + .ok_or(mpsc::RecvTimeoutError::Timeout)?; + let output = receive_until(&receiver, output, b"FLOW-READY\r\n", remaining)?; + writer + .write_all(b"printf 'FLOW-%s\\n' PROMPT\n") + .map_err(|_| mpsc::RecvTimeoutError::Disconnected)?; + writer + .flush() + .map_err(|_| mpsc::RecvTimeoutError::Disconnected)?; + let remaining = deadline + .checked_duration_since(Instant::now()) + .ok_or(mpsc::RecvTimeoutError::Timeout)?; + receive_until(&receiver, output, b"FLOW-PROMPT\r\n", remaining) + }); + let latency = interrupted.elapsed(); + let latency_failed = prompt.is_err(); + if mode == "none" { + // None is the unbounded baseline: latency may fail, but a fast + // host/kernel is also allowed to drain enough for it to pass. + if let Ok(prompt) = prompt { + output = prompt; + } + } else { + output = prompt.unwrap_or_else(|error| { + panic!("{mode}: waiting for Ctrl-C prompt within {prompt_timeout:?}: {error}") + }); + assert!( + latency <= MAX_PROMPT_LATENCY, + "{mode} Ctrl-C-to-prompt latency {latency:?} exceeded {MAX_PROMPT_LATENCY:?}" + ); + } + child.kill().unwrap(); + drop(writer); + let _ = child.wait(); + while let Ok(chunk) = receiver.recv() { + output.extend(chunk); + } + reader_thread.join().unwrap(); + proxy.finish().unwrap(); + + if let Some(directory) = &evidence { + fs::write(directory.join(format!("{mode}.ansi")), &output).unwrap(); + fs::write( + directory.join(format!("{mode}.json")), + format!( + "{{\"mode\":\"{mode}\",\"rate_bytes_per_second\":{bytes_per_second},\ + \"saturation_bytes\":{SATURATION_BYTES},\"ctrl_c_prompt_millis\":{},\ + \"latency_failed\":{},\"scenario_pass\":true}}\n", + latency.as_millis(), + latency_failed + ), + ) + .unwrap(); + } + } +} diff --git a/crates/et-bin/tests/flow_control_tty_support.rs b/crates/et-bin/tests/flow_control_tty_support.rs new file mode 100644 index 0000000..60d7226 --- /dev/null +++ b/crates/et-bin/tests/flow_control_tty_support.rs @@ -0,0 +1,247 @@ +#![cfg(unix)] +#![allow(dead_code)] + +use socket2::{Domain, Protocol, SockAddr, Socket, Type}; +use std::fs; +use std::io::{self, BufRead, BufReader, Read, Write}; +use std::net::{Ipv4Addr, Shutdown, SocketAddrV4, TcpListener, TcpStream}; +use std::os::unix::fs::{symlink, PermissionsExt}; +use std::process::{Command, Stdio}; +use std::sync::atomic::{AtomicU64, Ordering}; +use std::sync::mpsc; +use std::thread; +use std::time::{Duration, Instant}; + +use nix::sys::signal::{kill, Signal}; +use nix::unistd::Pid; +use wait_timeout::ChildExt; + +const TIMEOUT: Duration = Duration::from_secs(10); +pub const MAX_PROMPT_LATENCY: Duration = Duration::from_secs(5); +pub const THROTTLE_BYTES_PER_SECOND: usize = 100 * 1024; +pub const SATURATION_BYTES: usize = 128 * 1024; +const PROXY_RECEIVE_BUFFER_BYTES: usize = 64 * 1024; +static STACK_SEQUENCE: AtomicU64 = AtomicU64::new(0); + +pub fn receive_until( + receiver: &mpsc::Receiver>, + mut output: Vec, + marker: &[u8], + timeout: Duration, +) -> Result, mpsc::RecvTimeoutError> { + let deadline = Instant::now() + timeout; + while !output.windows(marker.len()).any(|window| window == marker) { + let Some(remaining) = deadline.checked_duration_since(Instant::now()) else { + return Err(mpsc::RecvTimeoutError::Timeout); + }; + output.extend(receiver.recv_timeout(remaining)?); + } + Ok(output) +} + +pub fn receive_bytes( + receiver: &mpsc::Receiver>, + mut output: Vec, + additional: usize, + timeout: Duration, +) -> Result, mpsc::RecvTimeoutError> { + let target = output.len() + additional; + let deadline = Instant::now() + timeout; + while output.len() < target { + let Some(remaining) = deadline.checked_duration_since(Instant::now()) else { + return Err(mpsc::RecvTimeoutError::Timeout); + }; + output.extend(receiver.recv_timeout(remaining)?); + } + Ok(output) +} + +pub struct ThrottleProxy { + pub port: u16, + saturated: mpsc::Receiver<()>, + stop: mpsc::Receiver, + worker: thread::JoinHandle>, +} + +impl ThrottleProxy { + pub fn start(server_port: u16, bytes_per_second: usize, saturation_bytes: usize) -> Self { + let listener = TcpListener::bind((Ipv4Addr::LOCALHOST, 0)).unwrap(); + let port = listener.local_addr().unwrap().port(); + let (saturated_tx, saturated) = mpsc::sync_channel(1); + let (stop_tx, stop) = mpsc::sync_channel(1); + let worker = thread::spawn(move || { + let (mut client, _) = listener.accept()?; + let socket = Socket::new(Domain::IPV4, Type::STREAM, Some(Protocol::TCP))?; + socket.set_recv_buffer_size(PROXY_RECEIVE_BUFFER_BYTES)?; + socket.connect(&SockAddr::from(SocketAddrV4::new( + Ipv4Addr::LOCALHOST, + server_port, + )))?; + let mut server = TcpStream::from(socket); + stop_tx + .send(server.try_clone()?) + .map_err(|_| io::Error::other("proxy stop receiver closed"))?; + let mut client_read = client.try_clone()?; + let mut server_write = server.try_clone()?; + let upstream = thread::spawn(move || io::copy(&mut client_read, &mut server_write)); + + let downstream = (|| { + let started = Instant::now(); + let mut transferred = 0usize; + let mut saturation_signalled = false; + let mut chunk = [0u8; 8192]; + loop { + let count = server.read(&mut chunk)?; + if count == 0 { + return Ok(()); + } + client.write_all(&chunk[..count])?; + transferred += count; + if !saturation_signalled && transferred >= saturation_bytes { + saturated_tx + .send(()) + .map_err(|_| io::Error::other("saturation receiver closed"))?; + saturation_signalled = true; + } + let expected = + Duration::from_secs_f64(transferred as f64 / bytes_per_second as f64); + if let Some(remaining) = expected.checked_sub(started.elapsed()) { + thread::sleep(remaining); + } + } + })(); + let _ = client.shutdown(Shutdown::Both); + let _ = server.shutdown(Shutdown::Both); + let upstream = upstream + .join() + .map_err(|_| io::Error::other("proxy upload thread panicked"))?; + normalize_proxy_close(downstream)?; + normalize_proxy_close(upstream.map(|_| ())) + }); + Self { + port, + saturated, + stop, + worker, + } + } + + pub fn wait_saturated(&self, timeout: Duration) -> Result<(), mpsc::RecvTimeoutError> { + self.saturated.recv_timeout(timeout) + } + + pub fn finish(self) -> io::Result<()> { + let stream = self + .stop + .recv_timeout(TIMEOUT) + .map_err(|_| io::Error::other("proxy did not accept client"))?; + let _ = stream.shutdown(Shutdown::Both); + self.worker + .join() + .map_err(|_| io::Error::other("proxy worker panicked"))? + } +} + +fn normalize_proxy_close(result: io::Result<()>) -> io::Result<()> { + match result { + Ok(()) => Ok(()), + Err(error) + if matches!( + error.kind(), + io::ErrorKind::BrokenPipe + | io::ErrorKind::ConnectionAborted + | io::ErrorKind::ConnectionReset + | io::ErrorKind::NotConnected + ) => + { + Ok(()) + } + Err(error) => Err(error), + } +} + +pub struct Stack { + pub directory: std::path::PathBuf, + pub router: std::path::PathBuf, + pub terminal: std::path::PathBuf, + pub port: u16, + server: std::process::Child, +} + +impl Stack { + pub fn start() -> Self { + let sequence = STACK_SEQUENCE.fetch_add(1, Ordering::Relaxed); + let directory = std::env::temp_dir().join(format!( + "et-rs-flow-control-qa-{}-{sequence}", + std::process::id() + )); + let _ = fs::remove_dir_all(&directory); + fs::create_dir(&directory).unwrap(); + fs::set_permissions(&directory, fs::Permissions::from_mode(0o700)).unwrap(); + let router = directory.join("router.sock"); + let reserved = TcpListener::bind((Ipv4Addr::LOCALHOST, 0)).unwrap(); + let port = reserved.local_addr().unwrap().port(); + drop(reserved); + let config = directory.join("et.cfg"); + fs::write( + &config, + format!( + "[Networking]\nport={port}\nbind_ip=127.0.0.1\n[Debug]\nserverfifo={}\n", + router.display() + ), + ) + .unwrap(); + let mut server = Command::new(env!("CARGO_BIN_EXE_et")) + .args(["server", "--cfgfile"]) + .arg(&config) + .stdout(Stdio::piped()) + .stderr(Stdio::piped()) + .spawn() + .unwrap(); + wait_ready(&mut server, port, &router); + let ssh = directory.join("ssh"); + fs::write( + &ssh, + "#!/bin/sh\nif [ \"$1\" = \"-G\" ]; then\n\ + printf 'hostname 127.0.0.1\\nuser tester\\n'; exit 0; fi\n\ + for last do :; done\nexec /bin/sh -c \"$last\"\n", + ) + .unwrap(); + fs::set_permissions(&ssh, fs::Permissions::from_mode(0o755)).unwrap(); + let terminal = directory.join("etterminal"); + symlink(env!("CARGO_BIN_EXE_et"), &terminal).unwrap(); + Self { + directory, + router, + terminal, + port, + server, + } + } +} + +impl Drop for Stack { + fn drop(&mut self) { + let pid = Pid::from_raw(i32::try_from(self.server.id()).unwrap()); + let _ = kill(pid, Signal::SIGTERM); + let _ = self.server.wait_timeout(TIMEOUT); + let _ = fs::remove_dir_all(&self.directory); + } +} + +fn wait_ready(server: &mut std::process::Child, port: u16, router: &std::path::Path) { + let stdout = server.stdout.take().unwrap(); + let (sender, receiver) = mpsc::sync_channel(1); + thread::spawn(move || { + let mut line = String::new(); + let result = BufReader::new(stdout).read_line(&mut line).map(|_| line); + let _ = sender.send(result); + }); + assert_eq!( + receiver.recv_timeout(TIMEOUT).unwrap().unwrap(), + format!( + "ETSERVER_READY tcp=127.0.0.1:{port} router={}\n", + router.display() + ) + ); +} diff --git a/crates/et-bin/tests/reconnect_outage_end_to_end.rs b/crates/et-bin/tests/reconnect_outage_end_to_end.rs index 28d5810..e956244 100644 --- a/crates/et-bin/tests/reconnect_outage_end_to_end.rs +++ b/crates/et-bin/tests/reconnect_outage_end_to_end.rs @@ -22,8 +22,6 @@ use portable_pty::{native_pty_system, CommandBuilder, PtySize}; use reconnect_stack::Stack; const TIMEOUT: Duration = Duration::from_secs(10); -/// Long enough for several reconnect attempts (retry delay is one second). -const OUTAGE: Duration = Duration::from_secs(4); #[test] fn client_retries_reconnect_through_network_outage() { @@ -85,7 +83,7 @@ fn client_retries_reconnect_through_network_outage() { // Drop the link and refuse every reconnect attempt for a while, like a // laptop that wakes from sleep before its network is back. proxy.outage(); - thread::sleep(OUTAGE); + proxy.wait_for_refused(2); proxy.restore(); // The client announces the retry loop on the first failed attempt. @@ -150,6 +148,7 @@ struct OutageProxy { port: u16, outage: mpsc::SyncSender<()>, restore: mpsc::SyncSender<()>, + refused: mpsc::Receiver, worker: Option>>, } @@ -159,6 +158,7 @@ impl OutageProxy { let port = listener.local_addr().unwrap().port(); let (outage_tx, outage_rx) = mpsc::sync_channel(1); let (restore_tx, restore_rx) = mpsc::sync_channel::<()>(1); + let (refused_tx, refused_rx) = mpsc::sync_channel(2); let worker = thread::spawn(move || { let (first, _) = listener.accept()?; let backend = TcpStream::connect((Ipv4Addr::LOCALHOST, backend_port))?; @@ -178,6 +178,7 @@ impl OutageProxy { } let _ = attempt.shutdown(Shutdown::Both); refused += 1; + let _ = refused_tx.try_send(refused); }; let backend = TcpStream::connect((Ipv4Addr::LOCALHOST, backend_port))?; let relays = relay(&recovered, &backend)?; @@ -194,6 +195,7 @@ impl OutageProxy { port, outage: outage_tx, restore: restore_tx, + refused: refused_rx, worker: Some(worker), } } @@ -206,6 +208,15 @@ impl OutageProxy { self.restore.send(()).unwrap(); } + fn wait_for_refused(&self, minimum: usize) { + loop { + let refused = self.refused.recv_timeout(TIMEOUT).unwrap(); + if refused >= minimum { + return; + } + } + } + fn join(mut self) -> usize { self.worker.take().unwrap().join().unwrap().unwrap() } diff --git a/crates/et-bin/tests/terminal_adversarial.rs b/crates/et-bin/tests/terminal_adversarial.rs index bb4767d..cfde775 100644 --- a/crates/et-bin/tests/terminal_adversarial.rs +++ b/crates/et-bin/tests/terminal_adversarial.rs @@ -58,6 +58,7 @@ fn malformed_initialization_and_shell_spawn_failure_are_typed() { &TermInit { environmentnames: Vec::new(), environmentvalues: Vec::new(), + flowcontrol: None, }, ); assert!(!failed_child @@ -94,6 +95,7 @@ fn shell_exit_reaps_background_process_group() { &TermInit { environmentnames: Vec::new(), environmentvalues: Vec::new(), + flowcontrol: None, }, ); let mut output = String::new(); diff --git a/crates/et-bin/tests/terminal_runtime.rs b/crates/et-bin/tests/terminal_runtime.rs index 648b101..63770e7 100644 --- a/crates/et-bin/tests/terminal_runtime.rs +++ b/crates/et-bin/tests/terminal_runtime.rs @@ -62,6 +62,7 @@ fn bootstrap_parent_reports_marker_and_leaves_registered_session_running() { &TermInit { environmentnames: Vec::new(), environmentvalues: Vec::new(), + flowcontrol: None, }, ); expect_startup(&mut router); @@ -175,6 +176,7 @@ fn new_terminal_uses_legacy_sequence_with_old_router() { &TermInit { environmentnames: Vec::new(), environmentvalues: Vec::new(), + flowcontrol: None, }, ); send( @@ -213,6 +215,7 @@ fn real_terminal_registers_runs_shell_and_resizes_pty() { &TermInit { environmentnames: vec!["G004_VALUE".to_owned()], environmentvalues: vec!["literal-value".to_owned()], + flowcontrol: None, }, ); expect_startup(&mut router); @@ -270,6 +273,7 @@ fn pty_output_backpressure_longer_than_two_seconds_preserves_session_and_order() &TermInit { environmentnames: Vec::new(), environmentvalues: Vec::new(), + flowcontrol: None, }, ); expect_startup(&mut router); @@ -375,6 +379,7 @@ fn router_disconnect_terminates_the_shell() { &TermInit { environmentnames: Vec::new(), environmentvalues: Vec::new(), + flowcontrol: None, }, ); expect_startup(&mut router); @@ -400,6 +405,7 @@ fn real_terminal_starts_login_shell_and_loads_profile_color() { &TermInit { environmentnames: Vec::new(), environmentvalues: Vec::new(), + flowcontrol: None, }, ); expect_startup(&mut router); @@ -444,6 +450,7 @@ fn real_terminal_login_shell_preserves_term_without_colorterm() { &TermInit { environmentnames: Vec::new(), environmentvalues: Vec::new(), + flowcontrol: None, }, ); expect_startup(&mut router); @@ -508,6 +515,7 @@ fn real_terminal_emits_motd_before_login_shell_output() { &TermInit { environmentnames: Vec::new(), environmentvalues: Vec::new(), + flowcontrol: None, }, ); expect_startup(&mut router); @@ -572,6 +580,7 @@ fn real_terminal_suppresses_motd_when_home_has_hushlogin() { &TermInit { environmentnames: Vec::new(), environmentvalues: Vec::new(), + flowcontrol: None, }, ); expect_startup(&mut router); diff --git a/crates/et-bin/tests/tunnels_end_to_end.rs b/crates/et-bin/tests/tunnels_end_to_end.rs index d7d1ea4..a043ccb 100644 --- a/crates/et-bin/tests/tunnels_end_to_end.rs +++ b/crates/et-bin/tests/tunnels_end_to_end.rs @@ -16,8 +16,6 @@ use reconnect_stack::{mkfifo, shell_quote, Stack}; use tunnel_support::SingleCutProxy; use wait_timeout::ChildExt; -// Process creation can legitimately exceed ten seconds under the load this -// suite exercises; every wait remains bounded by the exact FIFO/process event. const TIMEOUT: Duration = Duration::from_secs(30); #[test] @@ -820,10 +818,36 @@ fn spawn_client( } fn spawn_tcp_echo_once(listener: TcpListener, expected: &'static [u8]) -> thread::JoinHandle<()> { + let address = listener.local_addr().unwrap(); + let cancelled = std::sync::Arc::new(std::sync::atomic::AtomicBool::new(false)); + let worker_cancelled = std::sync::Arc::clone(&cancelled); + let (completed_tx, completed_rx) = mpsc::sync_channel(1); + let worker = thread::spawn(move || { + // Accept multiple times so readiness probes do not exhaust a one-shot echo. + loop { + let (mut stream, _) = listener.accept().unwrap(); + stream.set_read_timeout(Some(TIMEOUT)).unwrap(); + let mut payload = vec![0u8; expected.len()]; + if stream.read_exact(&mut payload).is_err() { + if worker_cancelled.load(std::sync::atomic::Ordering::Acquire) { + return; + } + continue; + } + assert_eq!(payload, expected); + stream.write_all(&payload).unwrap(); + completed_tx.send(()).unwrap(); + return; + } + }); thread::spawn(move || { - let (mut stream, _) = listener.accept().unwrap(); - stream.set_read_timeout(Some(TIMEOUT)).unwrap(); - echo_once(&mut stream, expected); + if completed_rx.recv_timeout(TIMEOUT + TIMEOUT).is_err() { + cancelled.store(true, std::sync::atomic::Ordering::Release); + drop(TcpStream::connect(address)); + worker.join().unwrap(); + panic!("TCP echo did not receive its payload before the deadline"); + } + worker.join().unwrap(); }) } diff --git a/crates/et-cli/src/client.rs b/crates/et-cli/src/client.rs index ee6bd54..4f1116a 100644 --- a/crates/et-cli/src/client.rs +++ b/crates/et-cli/src/client.rs @@ -91,6 +91,18 @@ pub struct ClientArgs { #[arg(short = 'l', long = "logdir")] pub logdir: Option, + #[arg( + long = "flow-control", + value_enum, + default_value_t = FlowControlMode::None, + help = "Bound terminal output when it outruns the network", + long_help = "Choose how terminal output behaves when it outruns the network.\n\n\ + none preserves the existing replay behavior. backpressure pauses the remote \ + producer at a bounded queue without losing output. discard drops the oldest \ + terminal output while preserving control traffic so the display stays current." + )] + pub flow_control: FlowControlMode, + #[arg(long = "logtostdout")] pub logtostdout: bool, @@ -135,6 +147,29 @@ pub enum RemoteShellKind { Powershell, } +/// Terminal-output behavior when the producer outruns the network. +#[derive(clap::ValueEnum, Clone, Copy, Debug, Default, PartialEq, Eq)] +pub enum FlowControlMode { + /// Preserve the existing unbounded replay behavior. + #[default] + None, + /// Pause the terminal producer while the bounded output queue is full. + Backpressure, + /// Drop the oldest queued terminal output to keep the display current. + Discard, +} + +impl FlowControlMode { + /// Optional protobuf value; `none` stays absent for legacy wire parity. + pub const fn protocol_value(self) -> Option { + match self { + Self::None => None, + Self::Backpressure => Some(et_core::proto::FlowControlMode::Backpressure as i32), + Self::Discard => Some(et_core::proto::FlowControlMode::Discard as i32), + } + } +} + impl ClientArgs { /// Effective remote shell grammar. pub fn effective_remote_shell(&self) -> RemoteShellKind { @@ -211,6 +246,23 @@ mod tests { assert_eq!(a.keepalive, MAX_KEEPALIVE); } + #[test] + fn flow_control_defaults_to_none() { + let a = ClientArgs::try_parse_from(["et", "host"]).unwrap(); + assert_eq!(a.flow_control, FlowControlMode::None); + } + + #[test] + fn flow_control_parses_opt_in_modes() { + let backpressure = + ClientArgs::try_parse_from(["et", "host", "--flow-control", "backpressure"]).unwrap(); + assert_eq!(backpressure.flow_control, FlowControlMode::Backpressure); + + let discard = + ClientArgs::try_parse_from(["et", "host", "--flow-control", "discard"]).unwrap(); + assert_eq!(discard.flow_control, FlowControlMode::Discard); + } + #[test] fn upstream_long_flag_spellings_parse() { let a = ClientArgs::try_parse_from([ diff --git a/crates/et-core/proto/ETerminal.proto b/crates/et-core/proto/ETerminal.proto index 6b83186..28584b8 100644 --- a/crates/et-core/proto/ETerminal.proto +++ b/crates/et-core/proto/ETerminal.proto @@ -16,6 +16,12 @@ enum TerminalPacketType { JUMPHOST_INIT = 10; } +enum FlowControlMode { + NONE = 0; + BACKPRESSURE = 1; + DISCARD = 2; +} + message TerminalBuffer { optional bytes buffer = 1; } @@ -61,6 +67,7 @@ message InitialPayload { optional bool jumphost = 1 [default = false]; repeated PortForwardSourceRequest reversetunnels = 2; map environmentvariables = 3; + optional FlowControlMode flowcontrol = 4 [default = NONE]; } message InitialResponse { @@ -75,6 +82,7 @@ message ConfigParams { message TermInit { repeated string environmentnames = 1; repeated string environmentvalues = 2; + optional FlowControlMode flowcontrol = 3 [default = NONE]; } message TerminalUserInfo { @@ -83,4 +91,4 @@ message TerminalUserInfo { optional int64 uid = 3; optional int64 gid = 4; optional int64 fd = 5; -} \ No newline at end of file +} diff --git a/crates/et-core/src/flow_control.rs b/crates/et-core/src/flow_control.rs new file mode 100644 index 0000000..ce7f7ee --- /dev/null +++ b/crates/et-core/src/flow_control.rs @@ -0,0 +1,201 @@ +//! Bounded, fair output lanes used by opt-in flow-control sessions. + +use std::collections::VecDeque; + +use prost::Message; + +use crate::packet::Packet; +use crate::proto::{TerminalBuffer, TerminalPacketType}; + +const FRAME_BYTES: usize = std::mem::size_of::(); +pub(crate) const MAX_PACKETS_PER_LANE: usize = 4096; + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub enum FlowControlMode { + Backpressure, + Discard, +} + +#[derive(Debug, PartialEq, Eq)] +pub enum QueuePushError { + Full(Packet), + Oversized(Packet), +} + +pub struct OutputQueue { + mode: FlowControlMode, + limit: usize, + terminal_bytes: usize, + control_bytes: usize, + terminal_packets: usize, + control_packets: usize, + terminal: VecDeque, + control: VecDeque, + prefer_terminal: bool, +} + +impl OutputQueue { + pub fn new(mode: FlowControlMode, limit: usize) -> Self { + Self { + mode, + limit, + terminal_bytes: 0, + control_bytes: 0, + terminal_packets: 0, + control_packets: 0, + terminal: VecDeque::new(), + control: VecDeque::new(), + prefer_terminal: true, + } + } + + pub fn push(&mut self, packet: Packet) -> Result<(), QueuePushError> { + if is_terminal_output(&packet) { + self.push_terminal(packet) + } else { + self.push_control(packet) + } + } + + fn push_terminal(&mut self, mut packet: Packet) -> Result<(), QueuePushError> { + if packet_cost(&packet) > self.limit { + if self.mode == FlowControlMode::Backpressure { + return Err(QueuePushError::Oversized(packet)); + } + packet = truncate_terminal(packet, self.limit).map_err(QueuePushError::Oversized)?; + } + let wanted = packet_cost(&packet); + if self.mode == FlowControlMode::Backpressure + && (self.terminal_bytes.saturating_add(wanted) > self.limit + || self.terminal_packets >= MAX_PACKETS_PER_LANE) + { + return Err(QueuePushError::Full(packet)); + } + while self.terminal_bytes.saturating_add(wanted) > self.limit + || self.terminal_packets >= MAX_PACKETS_PER_LANE + { + let Some(removed) = self.terminal.pop_front() else { + return Err(QueuePushError::Full(packet)); + }; + self.terminal_bytes -= packet_cost(&removed); + self.terminal_packets -= 1; + } + self.terminal_bytes += wanted; + self.terminal_packets += 1; + self.terminal.push_back(packet); + Ok(()) + } + + fn push_control(&mut self, packet: Packet) -> Result<(), QueuePushError> { + let wanted = packet_cost(&packet); + if wanted > self.limit { + return Err(QueuePushError::Oversized(packet)); + } + if self.control_bytes.saturating_add(wanted) > self.limit + || self.control_packets >= MAX_PACKETS_PER_LANE + { + return Err(QueuePushError::Full(packet)); + } + self.control_bytes += wanted; + self.control_packets += 1; + self.control.push_back(packet); + Ok(()) + } + + pub fn take(&mut self) -> Option { + let packet = if self.prefer_terminal { + self.terminal + .pop_front() + .or_else(|| self.control.pop_front()) + } else { + self.control + .pop_front() + .or_else(|| self.terminal.pop_front()) + }?; + self.prefer_terminal = !self.prefer_terminal; + Some(packet) + } + + pub fn complete(&mut self, packet: &Packet) { + if is_terminal_output(packet) { + self.terminal_bytes -= packet_cost(packet); + self.terminal_packets -= 1; + } else { + self.control_bytes -= packet_cost(packet); + self.control_packets -= 1; + } + } + + pub fn restore_front(&mut self, packet: Packet) { + if is_terminal_output(&packet) { + self.terminal.push_front(packet); + } else { + self.control.push_front(packet); + } + } + + pub fn pop(&mut self) -> Option { + let packet = self.take()?; + self.complete(&packet); + Some(packet) + } + + pub fn can_accept_terminal(&self, payload_bytes: usize) -> bool { + let wanted = payload_bytes.saturating_add(crate::packet::HEADER_LEN + FRAME_BYTES); + match self.mode { + FlowControlMode::Backpressure => { + self.terminal_packets < MAX_PACKETS_PER_LANE + && self.terminal_bytes.saturating_add(wanted) <= self.limit + } + FlowControlMode::Discard => wanted <= self.limit, + } + } + + pub fn is_empty(&self) -> bool { + self.terminal.is_empty() && self.control.is_empty() + } + + pub fn bytes(&self) -> usize { + self.terminal_bytes + self.control_bytes + } +} + +fn packet_cost(packet: &Packet) -> usize { + packet.wire_len().saturating_add(FRAME_BYTES) +} + +fn truncate_terminal(packet: Packet, limit: usize) -> Result { + let Ok(message) = TerminalBuffer::decode(packet.payload()) else { + return Err(packet); + }; + let Some(bytes) = message.buffer else { + return Err(packet); + }; + let mut low = 0usize; + let mut high = bytes.len(); + while low < high { + let middle = low + (high - low).div_ceil(2); + let encoded = TerminalBuffer { + buffer: Some(bytes[bytes.len() - middle..].to_vec()), + } + .encode_to_vec(); + if encoded + .len() + .saturating_add(crate::packet::HEADER_LEN + FRAME_BYTES) + <= limit + { + low = middle; + } else { + high = middle - 1; + } + } + let encoded = TerminalBuffer { + buffer: Some(bytes[bytes.len() - low..].to_vec()), + } + .encode_to_vec(); + Ok(Packet::new(packet.header(), encoded)) +} + +pub fn is_terminal_output(packet: &Packet) -> bool { + packet.header() == TerminalPacketType::TerminalBuffer as u8 +} diff --git a/crates/et-core/src/flow_control_tests.rs b/crates/et-core/src/flow_control_tests.rs new file mode 100644 index 0000000..f1b47b2 --- /dev/null +++ b/crates/et-core/src/flow_control_tests.rs @@ -0,0 +1,181 @@ +use prost::Message; + +use crate::flow_control::{FlowControlMode, OutputQueue, QueuePushError, MAX_PACKETS_PER_LANE}; +use crate::packet::Packet; +use crate::proto::{TerminalBuffer, TerminalPacketType}; + +const LIMIT: usize = 64; + +fn terminal(bytes: &[u8]) -> Packet { + Packet::new( + TerminalPacketType::TerminalBuffer as u8, + TerminalBuffer { + buffer: Some(bytes.to_vec()), + } + .encode_to_vec(), + ) +} + +fn control(payload: &[u8]) -> Packet { + Packet::new(TerminalPacketType::KeepAlive as u8, payload) +} + +fn forwarding(payload: &[u8]) -> Packet { + Packet::new(TerminalPacketType::PortForwardData as u8, payload) +} + +#[test] +fn backpressure_is_lossless_and_refuses_capacity_overflow() { + let mut queue = OutputQueue::new(FlowControlMode::Backpressure, LIMIT); + let first = terminal(&[1; 40]); + let second = terminal(&[2; 40]); + + assert!(queue.push(first.clone()).is_ok()); + assert_eq!( + queue.push(second.clone()), + Err(QueuePushError::Full(second)) + ); + assert_eq!(queue.pop(), Some(first)); + assert_eq!(queue.pop(), None); +} + +#[test] +fn discard_drops_oldest_terminal_output_and_keeps_newest() { + let mut queue = OutputQueue::new(FlowControlMode::Discard, LIMIT); + let oldest = terminal(&[1; 40]); + let newest = terminal(&[2; 40]); + + assert!(queue.push(oldest).is_ok()); + assert!(queue.push(newest.clone()).is_ok()); + assert_eq!(queue.pop(), Some(newest)); + assert_eq!(queue.pop(), None); +} + +#[test] +fn control_has_reserved_capacity_and_does_not_evict_terminal_output() { + let mut queue = OutputQueue::new(FlowControlMode::Discard, LIMIT); + let output = terminal(&[1; 40]); + let keepalive = control(&[9; 40]); + + assert!(queue.push(output.clone()).is_ok()); + assert!(queue.push(keepalive.clone()).is_ok()); + assert_eq!(queue.pop(), Some(output)); + assert_eq!(queue.pop(), Some(keepalive)); +} + +#[test] +fn oversized_terminal_packet_truncates_decoded_bytes_and_reencodes() { + let mut queue = OutputQueue::new(FlowControlMode::Discard, LIMIT); + let oversized = terminal(&(0u8..100).collect::>()); + + assert!(queue.push(oversized).is_ok()); + let packet = queue.pop().unwrap(); + let decoded = TerminalBuffer::decode(packet.payload()).unwrap(); + let retained = decoded.buffer.unwrap(); + assert_eq!(retained, (44u8..100).collect::>()); + assert!(queue.bytes() <= LIMIT); +} + +#[test] +fn oversized_control_is_permanent_not_temporary_full() { + let packet = control(&[0; LIMIT]); + let mut queue = OutputQueue::new(FlowControlMode::Backpressure, LIMIT); + + assert_eq!( + queue.push(packet.clone()), + Err(QueuePushError::Oversized(packet)) + ); +} + +#[test] +fn framing_bytes_are_included_in_capacity() { + let packet = control(&[0; LIMIT - crate::packet::HEADER_LEN]); + let mut queue = OutputQueue::new(FlowControlMode::Backpressure, LIMIT); + + assert_eq!( + queue.push(packet.clone()), + Err(QueuePushError::Oversized(packet)) + ); +} + +#[test] +fn header_only_packets_are_packet_aware_and_never_look_empty() { + let mut queue = OutputQueue::new(FlowControlMode::Backpressure, LIMIT); + let packet = control(&[]); + + assert!(queue.push(packet.clone()).is_ok()); + assert!(!queue.is_empty()); + assert_eq!(queue.pop(), Some(packet)); + assert!(queue.is_empty()); +} + +#[test] +fn full_control_lane_does_not_prevent_discard_terminal_progress() { + let mut queue = OutputQueue::new(FlowControlMode::Discard, 64); + let held = control(&[1; 58]); + queue.push(held.clone()).unwrap(); + let rejected = control(&[]); + assert_eq!( + queue.push(rejected.clone()), + Err(QueuePushError::Full(rejected)) + ); + let output = terminal(&[2; 40]); + queue.push(output.clone()).unwrap(); + + assert_eq!(queue.pop(), Some(output)); + assert_eq!(queue.pop(), Some(held)); +} + +#[test] +fn shell_output_and_keepalive_progress_during_sustained_forwarding() { + let mut queue = OutputQueue::new(FlowControlMode::Backpressure, 1024); + let first = terminal(b"shell-one"); + let second = terminal(b"shell-two"); + queue.push(forwarding(&[0])).unwrap(); + queue.push(forwarding(&[1])).unwrap(); + queue.push(control(b"keepalive")).unwrap(); + queue.push(forwarding(&[2])).unwrap(); + queue.push(first.clone()).unwrap(); + queue.push(second.clone()).unwrap(); + + assert_eq!(queue.pop(), Some(first)); + assert_eq!(queue.pop(), Some(forwarding(&[0]))); + assert_eq!(queue.pop(), Some(second)); + assert_eq!(queue.pop(), Some(forwarding(&[1]))); + assert_eq!(queue.pop(), Some(control(b"keepalive"))); +} + +#[test] +fn in_flight_packet_retains_exact_packet_count_capacity() { + let limit = (MAX_PACKETS_PER_LANE + 1) * 8; + let mut queue = OutputQueue::new(FlowControlMode::Backpressure, limit); + let packet = control(&[]); + for _ in 0..MAX_PACKETS_PER_LANE { + queue.push(packet.clone()).unwrap(); + } + let in_flight = queue.take().unwrap(); + + assert_eq!( + queue.push(packet.clone()), + Err(QueuePushError::Full(packet)) + ); + queue.complete(&in_flight); + assert!(queue.push(control(&[])).is_ok()); +} + +#[test] +fn failed_send_restoration_retains_its_capacity_reservation() { + let mut queue = OutputQueue::new(FlowControlMode::Backpressure, LIMIT); + let first = terminal(&[1; 40]); + let second = terminal(&[2; 40]); + queue.push(first.clone()).unwrap(); + let in_flight = queue.take().unwrap(); + + assert_eq!( + queue.push(second.clone()), + Err(QueuePushError::Full(second)) + ); + queue.restore_front(in_flight); + assert!(queue.bytes() <= LIMIT); + assert_eq!(queue.pop(), Some(first)); +} diff --git a/crates/et-core/src/lib.rs b/crates/et-core/src/lib.rs index cb85c73..aaf386a 100644 --- a/crates/et-core/src/lib.rs +++ b/crates/et-core/src/lib.rs @@ -9,6 +9,7 @@ pub mod backed_reader; pub mod backed_writer; pub mod crypto; +pub mod flow_control; pub mod framing; pub mod keepalive; pub mod keys; @@ -16,3 +17,6 @@ pub mod packet; pub mod proto; pub const PROTOCOL_VERSION: i32 = 6; + +#[cfg(test)] +mod flow_control_tests; diff --git a/crates/et-core/tests/proto_golden.rs b/crates/et-core/tests/proto_golden.rs index bf7b44f..4189cc8 100644 --- a/crates/et-core/tests/proto_golden.rs +++ b/crates/et-core/tests/proto_golden.rs @@ -108,3 +108,29 @@ fn socket_endpoint() { }; assert_eq!(enc(&se), *f.get("proto_socketendpoint_unix").unwrap()); } + +#[test] +fn flow_control_defaults_are_absent_and_wire_compatible() { + let initial = et::InitialPayload::default(); + let term = et::TermInit::default(); + + assert_eq!(initial.flowcontrol, None); + assert_eq!(term.flowcontrol, None); + assert!(enc(&initial).is_empty()); + assert!(enc(&term).is_empty()); +} + +#[test] +fn flow_control_opt_in_modes_use_additive_proto_fields() { + let initial = et::InitialPayload { + flowcontrol: Some(et::FlowControlMode::Backpressure as i32), + ..Default::default() + }; + let term = et::TermInit { + flowcontrol: Some(et::FlowControlMode::Discard as i32), + ..Default::default() + }; + + assert_eq!(enc(&initial), [0x20, 0x01]); + assert_eq!(enc(&term), [0x18, 0x02]); +} diff --git a/crates/et-net/Cargo.toml b/crates/et-net/Cargo.toml index 6f6af91..1cc7d0f 100644 --- a/crates/et-net/Cargo.toml +++ b/crates/et-net/Cargo.toml @@ -12,10 +12,11 @@ repository.workspace = true workspace = true [dependencies] +crossbeam-channel = "0.5" et-core = { path = "../et-core" } prost = "0.13" rustix = { version = "1.1.4", features = ["event"] } -socket2 = "0.6.1" +socket2 = { version = "0.6.1", features = ["all"] } wait-timeout = "0.2" [target.'cfg(unix)'.dependencies] diff --git a/crates/et-net/src/connection.rs b/crates/et-net/src/connection.rs index f3af7c4..14919de 100644 --- a/crates/et-net/src/connection.rs +++ b/crates/et-net/src/connection.rs @@ -8,6 +8,7 @@ use et_core::crypto::{ CryptoHandler, EncryptError, DIR_CLIENT_TO_SERVER, DIR_SERVER_TO_CLIENT, KEY_LEN, }; use et_core::packet::Packet; +use socket2::SockRef; #[path = "connection_recovery.rs"] mod recovery; @@ -22,6 +23,8 @@ pub use recovery::{DEFAULT_RECOVERY_TIMEOUT, MAX_RECOVERY_PROTO_LEN}; /// [`write_all_until`] with this deadline so the transport soft-disconnects /// and recovery can proceed. pub const DEFAULT_LIVE_WRITE_TIMEOUT: Duration = Duration::from_secs(2); +pub const FLOW_CONTROL_LIVE_WRITE_TIMEOUT: Duration = Duration::from_secs(5); +pub const FLOW_CONTROL_SOCKET_BUFFER_BYTES: usize = 32 * 1024; #[derive(Debug)] pub enum ConnError { @@ -30,6 +33,7 @@ pub enum ConnError { Recover(RecoverError), Encrypt(EncryptError), Backpressure, + PacketTooLarge, SequenceOutOfRange(i64), InvalidRecoverySequence(Option), } @@ -38,6 +42,45 @@ pub struct Connection { stream: TcpStream, writer: BackedWriter, reader: BackedReader, + live_write_timeout: Duration, +} + +pub struct PreparedWrite { + live: Option<(TcpStream, Vec, Duration)>, +} + +#[derive(Debug)] +pub enum WritePacketError { + BeforeReplay(ConnError), + ReplayOwned(ConnError), +} + +impl WritePacketError { + pub fn into_inner(self) -> ConnError { + match self { + Self::BeforeReplay(error) | Self::ReplayOwned(error) => error, + } + } +} + +impl PreparedWrite { + pub fn send(self) -> Result<(), WritePacketError> { + let Some((mut stream, frame, timeout)) = self.live else { + return Ok(()); + }; + write_live_frame(&mut stream, &frame, timeout) + .map_err(ConnError::Io) + .map_err(WritePacketError::ReplayOwned) + } + + fn send_until(self, deadline: Instant) -> Result<(), WritePacketError> { + let Some((mut stream, frame, _timeout)) = self.live else { + return Ok(()); + }; + write_live_frame_until(&mut stream, &frame, deadline) + .map_err(ConnError::Io) + .map_err(WritePacketError::ReplayOwned) + } } impl Connection { @@ -54,43 +97,88 @@ impl Connection { stream, writer: BackedWriter::new(CryptoHandler::new(key, encrypt), true), reader: BackedReader::new(CryptoHandler::new(key, decrypt), true), + live_write_timeout: DEFAULT_LIVE_WRITE_TIMEOUT, } } pub fn write_packet(&mut self, header: u8, payload: &[u8]) -> Result<(), ConnError> { - self.write_packet_with_deadline(header, payload, None) + self.write_packet_owned(header, payload) + .map_err(WritePacketError::into_inner) } - fn write_packet_with_deadline( + pub fn write_packet_owned( + &mut self, + header: u8, + payload: &[u8], + ) -> Result<(), WritePacketError> { + self.write_packet_owned_until(header, payload, None) + } + + fn write_packet_owned_until( &mut self, header: u8, payload: &[u8], deadline: Option, - ) -> Result<(), ConnError> { + ) -> Result<(), WritePacketError> { + let prepared = self.prepare_write_packet(header, payload)?; + let result = match deadline { + Some(deadline) => prepared.send_until(deadline), + None => prepared.send(), + }; + match result { + Ok(()) => Ok(()), + Err(error) => { + self.disconnect(); + Err(error) + } + } + } + + pub fn prepare_write_packet( + &mut self, + header: u8, + payload: &[u8], + ) -> Result { + self.prepare_write_packet_with(header, payload, TcpStream::try_clone) + } + + pub fn prepare_write_packet_with( + &mut self, + header: u8, + payload: &[u8], + clone_stream: F, + ) -> Result + where + F: FnOnce(&TcpStream) -> io::Result, + { // Probe first so a half-closed peer (laptop sleep, Wi-Fi drop) moves // the writer into the disconnected catch-up buffer before we try to // push bytes onto a dead socket. - if let Err(_error) = self.refresh_connectivity() { + if self.refresh_connectivity().is_err() { self.disconnect(); } - match self.writer.write_packet(header, payload)? { - WriterOutcome::Send(frame) => { - let deadline = deadline.unwrap_or_else(|| { - Instant::now() - .checked_add(DEFAULT_LIVE_WRITE_TIMEOUT) - .unwrap_or_else(Instant::now) - }); - if let Err(_error) = self.write_live_frame_until(&frame, deadline) { - // The encrypted packet is already in the replay backup - // (BackedWriter pushes before returning Send). Mark the - // transport disconnected so further writes buffer for a - // returning client instead of tearing the session down. - self.disconnect(); - } - Ok(()) - } - WriterOutcome::BufferedOnly => Ok(()), - WriterOutcome::Skipped => Err(ConnError::Backpressure), + // Cloning is the only fallible transport preparation. Do it before + // encryption advances the nonce/sequence and inserts replay history. + let live_stream = self + .writer + .connected() + .then(|| { + clone_stream(&self.stream) + .map_err(ConnError::Io) + .map_err(WritePacketError::BeforeReplay) + }) + .transpose()?; + match self + .writer + .write_packet(header, payload) + .map_err(ConnError::Encrypt) + .map_err(WritePacketError::BeforeReplay)? + { + WriterOutcome::Send(frame) => Ok(PreparedWrite { + live: live_stream.map(|stream| (stream, frame, self.live_write_timeout)), + }), + WriterOutcome::BufferedOnly => Ok(PreparedWrite { live: None }), + WriterOutcome::Skipped => Err(WritePacketError::BeforeReplay(ConnError::Backpressure)), } } @@ -103,7 +191,8 @@ impl Connection { if !self.connected() { return Err(io::Error::from(io::ErrorKind::NotConnected).into()); } - self.write_packet(header, payload)?; + self.write_packet_owned(header, payload) + .map_err(WritePacketError::into_inner)?; if self.connected() { Ok(()) } else { @@ -125,7 +214,8 @@ impl Connection { if !self.connected() { return Err(io::Error::from(io::ErrorKind::NotConnected).into()); } - self.write_packet_with_deadline(header, payload, Some(deadline))?; + self.write_packet_owned_until(header, payload, Some(deadline)) + .map_err(WritePacketError::into_inner)?; if self.connected() { Ok(()) } else { @@ -133,32 +223,6 @@ impl Connection { } } - /// Write a framed packet to a still-connected peer with a bounded timeout. - /// - /// Uses a write loop (not bare `write_all`) so each `write` is capped by - /// the remaining deadline. On any incomplete write we error out and the - /// caller soft-disconnects: the old TCP path is abandoned (and shut down - /// by the soft-disconnect path) so a partial frame left on the wire cannot - /// desync a later recovery, which always uses a new stream. - /// - /// Restores a cleared write timeout afterwards so recovery / handshake - /// code that sets its own deadlines is not left with a stale value. - fn write_live_frame_until(&mut self, frame: &[u8], deadline: Instant) -> io::Result<()> { - let result = write_all_until(&mut self.stream, frame, deadline); - // Best-effort restore: a failed clear must not hide a write error. - let clear = self.stream.set_write_timeout(None); - match (result, clear) { - (Ok(()), Ok(())) => Ok(()), - (Err(error), _) => { - // Force the peer off the half-written frame so it reconnects - // rather than blocking on the rest of a truncated record. - let _ = self.stream.shutdown(Shutdown::Both); - Err(error) - } - (Ok(()), Err(error)) => Err(error), - } - } - pub fn read_packet(&mut self) -> Result { loop { match self.reader.pop() { @@ -292,6 +356,21 @@ impl Connection { Ok(()) } + /// Keep opt-in flow-control backlog in the application queue instead of + /// allowing the kernel send queue to autotune to multiple megabytes. + pub fn minimize_output_buffering(&mut self) -> Result<(), ConnError> { + self.live_write_timeout = FLOW_CONTROL_LIVE_WRITE_TIMEOUT; + let socket = SockRef::from(&self.stream); + socket + .set_send_buffer_size(FLOW_CONTROL_SOCKET_BUFFER_BYTES) + .map_err(ConnError::Io)?; + #[cfg(target_os = "linux")] + socket + .set_tcp_notsent_lowat(FLOW_CONTROL_SOCKET_BUFFER_BYTES as u32) + .map_err(ConnError::Io)?; + Ok(()) + } + pub fn writer_sequence(&self) -> i64 { self.writer.sequence() } @@ -339,6 +418,37 @@ impl Connection { } } +fn write_live_frame(stream: &mut TcpStream, frame: &[u8], timeout: Duration) -> io::Result<()> { + let deadline = Instant::now() + .checked_add(timeout) + .ok_or_else(|| io::Error::new(io::ErrorKind::TimedOut, "live write deadline"))?; + write_live_frame_until(stream, frame, deadline) +} + +/// Write a framed packet to a still-connected peer before an absolute deadline. +/// +/// On an incomplete write, shut down the abandoned transport so a partial +/// frame cannot desynchronize a later recovery on a replacement stream. +fn write_live_frame_until( + stream: &mut TcpStream, + frame: &[u8], + deadline: Instant, +) -> io::Result<()> { + let result = write_all_until(stream, frame, deadline); + // Best-effort restore: a failed clear must not hide a write error. + let clear = stream.set_write_timeout(None); + match (result, clear) { + (Ok(()), Ok(())) => Ok(()), + (Err(error), _) => { + // Force the peer off the half-written frame so it reconnects + // rather than blocking on the rest of a truncated record. + let _ = stream.shutdown(Shutdown::Both); + Err(error) + } + (Ok(()), Err(error)) => Err(error), + } +} + /// Write the full buffer before `deadline`, refreshing the socket write /// timeout on each attempt so a blackholed peer cannot pin the caller. fn write_all_until(stream: &mut TcpStream, mut buffer: &[u8], deadline: Instant) -> io::Result<()> { @@ -359,3 +469,40 @@ fn write_all_until(stream: &mut TcpStream, mut buffer: &[u8], deadline: Instant) } Ok(()) } + +#[cfg(test)] +mod tests { + use super::*; + use std::net::{Ipv4Addr, TcpListener}; + use std::thread; + + #[test] + fn clone_failure_does_not_advance_writer_nonce_or_sequence() { + // Given: an authenticated connection whose transport clone fails. + let listener = TcpListener::bind((Ipv4Addr::LOCALHOST, 0)).unwrap(); + let address = listener.local_addr().unwrap(); + let connector = thread::spawn(move || TcpStream::connect(address).unwrap()); + let (server_stream, _) = listener.accept().unwrap(); + let client_stream = connector.join().unwrap(); + let key = [19u8; KEY_LEN]; + let mut sender = Connection::new_client(client_stream, &key); + let mut receiver = Connection::new_server(server_stream, &key); + + // When: preparation fails, then the same logical packet is retried. + let failure = sender.prepare_write_packet_with(7, b"once", |_| { + Err(io::Error::other("injected clone failure")) + }); + assert!(matches!( + failure, + Err(WritePacketError::BeforeReplay(ConnError::Io(_))) + )); + assert_eq!(sender.writer_sequence(), 0); + sender.write_packet(7, b"once").unwrap(); + + // Then: the peer decrypts sequence zero exactly once. + let packet = receiver.read_packet().unwrap(); + assert_eq!((packet.header(), packet.payload()), (7, b"once".as_slice())); + assert_eq!(sender.writer_sequence(), 1); + assert_eq!(receiver.reader_sequence(), 1); + } +} diff --git a/crates/et-net/src/connection_error.rs b/crates/et-net/src/connection_error.rs index 8ac5d82..e7122e6 100644 --- a/crates/et-net/src/connection_error.rs +++ b/crates/et-net/src/connection_error.rs @@ -38,6 +38,7 @@ impl std::fmt::Display for ConnError { Self::Recover(error) => write!(f, "recover: {error}"), Self::Encrypt(error) => write!(f, "encrypt: {error}"), Self::Backpressure => write!(f, "disconnected write buffer is full"), + Self::PacketTooLarge => write!(f, "packet exceeds the bounded output lane"), Self::SequenceOutOfRange(sequence) => { write!(f, "sequence number {sequence} exceeds the wire format") } diff --git a/crates/et-net/src/connection_nonblocking.rs b/crates/et-net/src/connection_nonblocking.rs index 4c5de78..14b1d08 100644 --- a/crates/et-net/src/connection_nonblocking.rs +++ b/crates/et-net/src/connection_nonblocking.rs @@ -1,4 +1,6 @@ -use std::io::{self, Read}; +use std::io; +#[cfg(windows)] +use std::io::Read; use std::net::TcpStream; use et_core::backed_reader::{BackedReader, ReadItem}; @@ -14,10 +16,18 @@ pub(crate) fn try_read( ReadItem::Packet(packet) => return Ok(Some(packet)), ReadItem::NeedMore => {} } - stream.set_nonblocking(true)?; let mut buffer = [0u8; 8192]; - let read = stream.read(&mut buffer); - stream.set_nonblocking(false)?; + #[cfg(unix)] + let read = rustix::net::recv(stream, &mut buffer, rustix::net::RecvFlags::DONTWAIT) + .map(|(count, _)| count) + .map_err(io::Error::from); + #[cfg(windows)] + let read = { + stream.set_nonblocking(true)?; + let read = stream.read(&mut buffer); + stream.set_nonblocking(false)?; + read + }; match read { Ok(0) => Err(ConnError::Io(io::ErrorKind::UnexpectedEof.into())), Ok(count) => { diff --git a/crates/et-net/src/connection_recovery.rs b/crates/et-net/src/connection_recovery.rs index 0fad792..12c2287 100644 --- a/crates/et-net/src/connection_recovery.rs +++ b/crates/et-net/src/connection_recovery.rs @@ -42,6 +42,7 @@ impl Connection { stream: new_stream, writer: self.writer.clone(), reader: self.reader.clone(), + live_write_timeout: self.live_write_timeout, }; candidate.disconnect(); candidate diff --git a/crates/et-net/src/forward.rs b/crates/et-net/src/forward.rs index 1e8034e..b13f3ef 100644 --- a/crates/et-net/src/forward.rs +++ b/crates/et-net/src/forward.rs @@ -7,6 +7,8 @@ use std::os::unix::net::UnixStream; use std::sync::atomic::{AtomicBool, Ordering}; use std::sync::{mpsc, Arc, Condvar, Mutex, OnceLock}; use std::thread::JoinHandle; + +use crossbeam_channel as channel; use std::time::{Duration, Instant}; use et_core::packet::Packet; @@ -14,7 +16,9 @@ use et_core::proto::{PortForwardSourceRequest, TerminalPacketType}; use crate::forward_endpoint::{Endpoint, ResolvedEndpoint}; use crate::forward_io::BoundSource; -use crate::forward_worker::{command_channel, run, Command, CommandSender, TryCommandError}; +use crate::forward_worker::{ + command_channel, run, Command, CommandSender, TryCommandError, WorkerChannels, +}; const CHANNEL_CAPACITY: usize = 256; /// Maximum reverse listeners owned by one terminal session, after DNS fanout @@ -302,7 +306,8 @@ pub struct SkippedForward { pub struct Forwarder { commands: CommandSender, - outbound: mpsc::Receiver, + outbound: channel::Receiver, + cancel: Option>, /// Readiness channel for outbound packets. Unix callers poll it exactly /// like upstream's `select()`; Windows callers drain [`Forwarder::try_outbound`] /// on the client loop's 10ms cadence instead, because a socket pair created @@ -311,6 +316,7 @@ pub struct Forwarder { wake: UnixStream, shutdown: Arc, worker: Option>, + abandoned: Arc, } impl Forwarder { @@ -404,7 +410,10 @@ fn start_forwarder_hook( ensure_setup_deadline(deadline)?; let session_user = owner; let (commands_tx, commands_rx) = command_channel(CHANNEL_CAPACITY); - let (outbound_tx, outbound_rx) = mpsc::sync_channel(CHANNEL_CAPACITY); + let (outbound_tx, outbound_rx) = channel::bounded(CHANNEL_CAPACITY); + let (cancel_tx, cancel_rx) = channel::bounded(1); + let abandoned = Arc::new(AtomicBool::new(false)); + let worker_abandoned = abandoned.clone(); #[cfg(unix)] let (wake, wake_writer) = { let (reader, writer) = UnixStream::pair()?; @@ -426,9 +435,13 @@ fn start_forwarder_hook( worker_start(); run( sources, - commands_rx, - worker_commands, - outbound_tx, + WorkerChannels { + receiver: commands_rx, + sender: worker_commands, + outbound: outbound_tx, + cancel: cancel_rx, + abandoned: worker_abandoned, + }, #[cfg(unix)] wake_writer, (listener_stop_reader, session_user, worker_shutdown), @@ -442,10 +455,12 @@ fn start_forwarder_hook( let forwarder = Forwarder { commands: commands_tx, outbound: outbound_rx, + cancel: Some(cancel_tx), #[cfg(unix)] wake, shutdown, worker: Some(worker), + abandoned, }; before_publish(); ensure_setup_deadline(deadline)?; @@ -495,8 +510,8 @@ impl Forwarder { drain_wake(&self.wake)?; match self.outbound.try_recv() { Ok(result) => result.map(Some), - Err(mpsc::TryRecvError::Empty) => Ok(None), - Err(mpsc::TryRecvError::Disconnected) => Err(ForwardError::Unavailable), + Err(channel::TryRecvError::Empty) => Ok(None), + Err(channel::TryRecvError::Disconnected) => Err(ForwardError::Unavailable), } } @@ -510,6 +525,27 @@ impl Forwarder { self.stop() } + /// Cancel independently of bounded command/output capacity and join the + /// worker. Returns true when queued commands or outbound packets could not + /// be completed and were explicitly abandoned. + pub fn shutdown_hard(&mut self) -> Result { + self.shutdown.store(true, Ordering::Release); + self.cancel.take(); + let mut abandoned = !self.commands.is_empty() || !self.outbound.is_empty(); + self.commands.shutdown(); + if let Some(worker) = self.worker.take() { + worker.join().map_err(|_| ForwardError::Unavailable)?; + } + while self.outbound.try_recv().is_ok() { + abandoned = true; + } + if !self.commands.is_empty() { + abandoned = true; + } + abandoned |= self.abandoned.load(Ordering::Acquire); + Ok(abandoned) + } + fn stop(&mut self) -> Result<(), ForwardError> { if let Some(worker) = self.worker.take() { self.shutdown.store(true, Ordering::Release); @@ -522,7 +558,9 @@ impl Forwarder { impl Drop for Forwarder { fn drop(&mut self) { - let _ = self.stop(); + // Drop is an abort path and must not block behind bounded forwarding + // queues. Callers that require graceful completion use `shutdown`. + let _ = self.shutdown_hard(); } } diff --git a/crates/et-net/src/forward_endpoint.rs b/crates/et-net/src/forward_endpoint.rs index 79f8f33..7159a13 100644 --- a/crates/et-net/src/forward_endpoint.rs +++ b/crates/et-net/src/forward_endpoint.rs @@ -613,6 +613,20 @@ impl ForwardStream { Self::Unix(stream) => stream.shutdown(Shutdown::Read), }; } + + #[cfg(windows)] + pub(crate) fn set_read_timeout(&self, timeout: Option) -> io::Result<()> { + match self { + Self::Tcp(stream) => stream.set_read_timeout(timeout), + } + } + + #[cfg(windows)] + pub(crate) fn set_write_timeout(&self, timeout: Option) -> io::Result<()> { + match self { + Self::Tcp(stream) => stream.set_write_timeout(timeout), + } + } } impl Read for ForwardStream { diff --git a/crates/et-net/src/forward_io.rs b/crates/et-net/src/forward_io.rs index adf43cf..7d7a290 100644 --- a/crates/et-net/src/forward_io.rs +++ b/crates/et-net/src/forward_io.rs @@ -1,10 +1,10 @@ use std::io::{self, Read, Write}; -#[cfg(windows)] -use std::sync::atomic::AtomicBool; -use std::sync::atomic::{AtomicI32, Ordering}; -use std::sync::{mpsc, Arc}; +use std::sync::atomic::{AtomicBool, AtomicI32, AtomicUsize, Ordering}; +use std::sync::Arc; use std::thread::{self, JoinHandle}; +use crossbeam_channel as channel; + #[cfg(unix)] use rustix::event::{poll, PollFd, PollFlags}; @@ -14,6 +14,8 @@ use et_core::proto::SocketEndpoint; use super::forward_worker::{Command, CommandSender, Role}; const READ_CHUNK: usize = 16 * 1024; +#[cfg(windows)] +const IO_CANCEL_INTERVAL: std::time::Duration = std::time::Duration::from_millis(100); pub(crate) struct BoundSource { pub(crate) listener: ForwardListener, @@ -21,8 +23,11 @@ pub(crate) struct BoundSource { } pub(crate) struct ActiveIo { - pub(crate) writer: mpsc::SyncSender, + pub(crate) writer: channel::Sender, pub(crate) control: ForwardStream, + pub(crate) cancel: channel::Receiver<()>, + pub(crate) pending_bytes: Arc, + pub(crate) abandoned: Arc, } pub(crate) enum WriteCommand { @@ -50,6 +55,7 @@ const ACCEPT_INTERVAL: std::time::Duration = std::time::Duration::from_millis(10 pub(crate) fn spawn_listener( source: BoundSource, commands: CommandSender, + cancel: channel::Receiver<()>, stop: ListenerStop, next_client_fd: Arc, ) -> JoinHandle<()> { @@ -94,7 +100,10 @@ pub(crate) fn spawn_listener( Ok(stream) => { accepted_any = true; let client_fd = next_client_fd.fetch_add(1, Ordering::Relaxed); - if client_fd <= 0 + if client_fd <= 0 { + return; + } + if cancellation_requested(&cancel) || commands .send(Command::Accepted { client_fd, @@ -126,15 +135,18 @@ pub(crate) fn spawn_connector( socket_id: i32, destination: Endpoint, commands: CommandSender, + cancel: channel::Receiver<()>, session_user: Option<(u32, u32)>, ) -> JoinHandle<()> { thread::spawn(move || { let result = destination.connect_with_user(session_user); - let _ = commands.send(Command::Connected { - client_fd, - socket_id, - result, - }); + if !cancellation_requested(&cancel) { + let _ = commands.send(Command::Connected { + client_fd, + socket_id, + result, + }); + } }) } @@ -143,81 +155,269 @@ pub(crate) fn spawn_io( socket_id: i32, stream: ForwardStream, commands: CommandSender, + cancel: channel::Receiver<()>, + abandoned: Arc, ) -> io::Result<(ActiveIo, [JoinHandle<()>; 2])> { + #[cfg(windows)] + { + // Winsock does not reliably interrupt an in-flight synchronous I/O + // call when another handle for the same socket is shut down. Finite + // deadlines make both sibling threads observe hard cancellation. + stream.set_read_timeout(Some(IO_CANCEL_INTERVAL))?; + stream.set_write_timeout(Some(IO_CANCEL_INTERVAL))?; + } let mut reader = stream.try_clone()?; let control = stream.try_clone()?; - let (writer_tx, writer_rx) = mpsc::sync_channel(64); + let (writer_tx, writer_rx) = channel::bounded(64); let reader_commands = commands.clone(); + let reader_cancel = cancel.clone(); let reader_handle = thread::spawn(move || { let mut buffer = [0u8; READ_CHUNK]; loop { + #[cfg(windows)] + if cancellation_requested(&reader_cancel) { + return; + } match reader.read(&mut buffer) { Ok(0) => { - let _ = reader_commands.send(Command::Closed { role, socket_id }); + if !cancellation_requested(&reader_cancel) { + let _ = reader_commands.send(Command::Closed { role, socket_id }); + } return; } Ok(count) => { - if reader_commands - .send(Command::Read { - role, - socket_id, - buffer: buffer[..count].to_vec(), - }) - .is_err() + if cancellation_requested(&reader_cancel) + || reader_commands + .send(Command::Read { + role, + socket_id, + buffer: buffer[..count].to_vec(), + }) + .is_err() { return; } } Err(error) if error.kind() == io::ErrorKind::Interrupted => {} + #[cfg(windows)] + Err(error) + if matches!( + error.kind(), + io::ErrorKind::WouldBlock | io::ErrorKind::TimedOut + ) => {} Err(error) => { - let _ = reader_commands.send(Command::IoFailed { - role, - socket_id, - error, - }); + if !cancellation_requested(&reader_cancel) { + let _ = reader_commands.send(Command::IoFailed { + role, + socket_id, + error, + }); + } return; } } } }); let writer_commands = commands; + let writer_cancel = cancel.clone(); + let pending_bytes = Arc::new(AtomicUsize::new(0)); + let writer_pending_bytes = pending_bytes.clone(); + let writer_abandoned = abandoned.clone(); let writer_handle = thread::spawn(move || { let mut writer = stream; - while let Ok(command) = writer_rx.recv() { + loop { + let command = channel::select! { + recv(writer_rx) -> command => match command { + Ok(command) => command, + Err(_) => break, + }, + recv(writer_cancel) -> _ => break, + }; match command { WriteCommand::Data(buffer) => { - if let Err(error) = writer.write_all(&buffer) { - let _ = writer_commands.send(Command::IoFailed { - role, - socket_id, - error, - }); - return; + #[cfg(windows)] + let result = write_all_cancellable(&mut writer, &buffer, &writer_cancel); + #[cfg(not(windows))] + let result = writer.write_all(&buffer).map(|()| true); + let delivered = match result { + Ok(delivered) => delivered, + Err(error) => { + if !cancellation_requested(&writer_cancel) { + let _ = writer_commands.send(Command::IoFailed { + role, + socket_id, + error, + }); + } + break; + } + }; + if !delivered { + break; } + writer_pending_bytes.fetch_sub(buffer.len(), Ordering::AcqRel); } WriteCommand::Stop => { // Perform the final shutdown here so every Data command // queued before Stop is flushed to the socket first. writer.shutdown(); - return; + break; } } } + if writer_pending_bytes.load(Ordering::Acquire) != 0 { + writer_abandoned.store(true, Ordering::Release); + } }); Ok(( ActiveIo { writer: writer_tx, control, + cancel, + pending_bytes, + abandoned, }, [reader_handle, writer_handle], )) } +#[cfg(windows)] +fn write_all_cancellable( + writer: &mut ForwardStream, + mut remaining: &[u8], + cancel: &channel::Receiver<()>, +) -> io::Result { + while !remaining.is_empty() { + if cancellation_requested(cancel) { + return Ok(false); + } + match writer.write(remaining) { + Ok(0) => return Err(io::ErrorKind::WriteZero.into()), + Ok(written) => remaining = &remaining[written..], + Err(error) if error.kind() == io::ErrorKind::Interrupted => {} + Err(error) + if matches!( + error.kind(), + io::ErrorKind::WouldBlock | io::ErrorKind::TimedOut + ) => {} + Err(error) => return Err(error), + } + } + Ok(true) +} + +fn cancellation_requested(cancel: &channel::Receiver<()>) -> bool { + !matches!(cancel.try_recv(), Err(channel::TryRecvError::Empty)) +} + pub(crate) fn stop_io(io: ActiveIo) { - // Queue Stop before touching the socket: the writer thread drains any - // pending Data commands in FIFO order and then closes the socket. Only - // shut down the read half here to wake the reader thread; a full - // shutdown would discard writes that are still queued. - let _ = io.writer.send(WriteCommand::Stop); - io.control.shutdown_read(); + // Keep the control socket alive while Stop waits for queue capacity so + // hard cancellation can still abort an in-flight write and release it. + let stop_admitted = channel::select! { + send(io.writer, WriteCommand::Stop) -> result => result.is_ok(), + recv(io.cancel) -> _ => false, + }; + if stop_admitted { + io.control.shutdown_read(); + } else { + if io.pending_bytes.load(Ordering::Acquire) != 0 { + io.abandoned.store(true, Ordering::Release); + } + io.control.shutdown(); + } +} + +pub(crate) fn abort_io(io: ActiveIo) { + // Hard cancellation abandons queued output and closes both socket halves + // before joining, waking a writer blocked inside write_all. + io.control.shutdown(); +} + +#[cfg(all(test, unix))] +mod tests { + use std::io::{self, Write}; + use std::os::unix::net::UnixStream; + use std::sync::atomic::AtomicBool; + use std::sync::{mpsc, Arc}; + use std::time::Duration; + + use crossbeam_channel as channel; + + use super::{spawn_io, stop_io, ForwardStream, Role, WriteCommand}; + use crate::forward_worker::command_channel; + + const EVENT_TIMEOUT: Duration = Duration::from_secs(3); + + #[test] + fn stop_io_cancellation_bypasses_a_full_writer_queue() { + // Given: the writer owns one command in a blocked socket write and its + // bounded queue is full. Admission of the final command proves this + // state without scheduler timing assumptions. + let (stream, peer) = UnixStream::pair().unwrap(); + socket2::SockRef::from(&stream) + .set_send_buffer_size(2 * 1024) + .unwrap(); + socket2::SockRef::from(&peer) + .set_recv_buffer_size(2 * 1024) + .unwrap(); + let mut saturator = stream.try_clone().unwrap(); + saturator.set_nonblocking(true).unwrap(); + loop { + match saturator.write(&[0u8; 16 * 1024]) { + Ok(0) => panic!("forward socket closed before reaching backpressure"), + Ok(_) => {} + Err(error) if error.kind() == io::ErrorKind::Interrupted => {} + Err(error) if error.kind() == io::ErrorKind::WouldBlock => break, + Err(error) => panic!("could not saturate forward socket: {error}"), + } + } + saturator.set_nonblocking(false).unwrap(); + let control = stream.try_clone().unwrap(); + let (commands, _command_receiver) = command_channel(1); + let (cancel, cancel_receiver) = channel::bounded(1); + let abandoned = Arc::new(AtomicBool::new(false)); + let (active, handles) = spawn_io( + Role::Destination, + 1, + ForwardStream::Unix(stream), + commands, + cancel_receiver, + abandoned, + ) + .unwrap(); + for _ in 0..64 { + active.writer.send(WriteCommand::Data(vec![1])).unwrap(); + } + let writer = active.writer.clone(); + let (admitted_tx, admitted_rx) = mpsc::sync_channel(0); + let admission = std::thread::spawn(move || { + writer.send(WriteCommand::Data(vec![1])).unwrap(); + admitted_tx.send(()).unwrap(); + }); + admitted_rx.recv_timeout(EVENT_TIMEOUT).unwrap(); + admission.join().unwrap(); + + // When: hard cancellation races graceful Stop admission. + let (done_tx, done_rx) = mpsc::sync_channel(0); + let stopping = std::thread::spawn(move || { + stop_io(active); + done_tx.send(()).unwrap(); + }); + drop(cancel); + let completed_before_abort = done_rx.recv_timeout(EVENT_TIMEOUT).is_ok(); + + // Then: cancellation itself must complete removal; the retained clone + // is used only to guarantee cleanup after observing a regression. + let _ = control.shutdown(std::net::Shutdown::Both); + let _ = done_rx.recv_timeout(EVENT_TIMEOUT); + stopping.join().unwrap(); + for handle in handles { + handle.join().unwrap(); + } + drop(peer); + assert!( + completed_before_abort, + "stop_io remained blocked behind the full writer queue after cancellation" + ); + } } diff --git a/crates/et-net/src/forward_worker.rs b/crates/et-net/src/forward_worker.rs index e1187fa..af4fa6f 100644 --- a/crates/et-net/src/forward_worker.rs +++ b/crates/et-net/src/forward_worker.rs @@ -5,14 +5,18 @@ use std::io::{self}; #[cfg(unix)] use std::os::unix::net::UnixStream; use std::sync::atomic::{AtomicBool, AtomicI32, Ordering}; -use std::sync::{mpsc, Arc, Condvar, Mutex}; +#[cfg(all(test, unix))] +use std::sync::mpsc; +use std::sync::{Arc, Condvar, Mutex}; use std::thread::JoinHandle; +use crossbeam_channel as channel; + use crate::forward::{ForwardError, Outbound}; use crate::forward_endpoint::ForwardStream; use crate::forward_io::{ - spawn_connector, spawn_io, spawn_listener, stop_io, ActiveIo, BoundSource, ListenerStop, - WriteCommand, + abort_io, spawn_connector, spawn_io, spawn_listener, stop_io, ActiveIo, BoundSource, + ListenerStop, WriteCommand, }; use et_core::packet::Packet; use et_core::proto::SocketEndpoint; @@ -152,7 +156,14 @@ impl CommandSender { self.queue.changed.notify_all(); } - #[cfg(test)] + pub(crate) fn is_empty(&self) -> bool { + self.queue + .state + .lock() + .map_or(true, |state| state.commands.is_empty()) + } + + #[cfg(all(test, unix))] pub(crate) fn wait_shutdown_timeout(&self, timeout: std::time::Duration) -> bool { let deadline = std::time::Instant::now() + timeout; let mut state = self.queue.state.lock().unwrap(); @@ -188,7 +199,7 @@ impl CommandReceiver { } } - #[cfg(test)] + #[cfg(all(test, unix))] fn recv_timeout( &self, timeout: std::time::Duration, @@ -223,29 +234,46 @@ impl Drop for CommandReceiver { } } +pub(crate) struct WorkerChannels { + pub(crate) receiver: CommandReceiver, + pub(crate) sender: CommandSender, + pub(crate) outbound: channel::Sender, + pub(crate) cancel: channel::Receiver<()>, + pub(crate) abandoned: Arc, +} + pub(crate) fn run( sources: Vec, - commands: CommandReceiver, - command_sender: CommandSender, - outbound: mpsc::SyncSender, + channels: WorkerChannels, #[cfg(unix)] mut outbound_wake: UnixStream, control: (ListenerStop, Option<(u32, u32)>, Arc), ) { + let WorkerChannels { + receiver: commands, + sender: command_sender, + outbound, + cancel, + abandoned, + } = channels; let (listener_stop, session_user, shutdown) = control; #[cfg(unix)] let result = Worker::new( command_sender, outbound.clone(), outbound_wake.try_clone().ok(), - shutdown.clone(), + cancel.clone(), + abandoned, ) .and_then(|mut worker| worker.run(sources, commands, listener_stop, session_user)); #[cfg(windows)] - let result = Worker::new(command_sender, outbound.clone(), shutdown.clone()) + let result = Worker::new(command_sender, outbound.clone(), cancel.clone(), abandoned) .and_then(|mut worker| worker.run(sources, commands, listener_stop, session_user)); if let Err(error) = result { if !shutdown.load(Ordering::Acquire) { - let _ = outbound.try_send(Err(error)); + channel::select! { + send(outbound, Err(error)) -> _ => {} + recv(cancel) -> _ => {} + } } #[cfg(unix)] let _ = outbound_wake.write(&[1]); @@ -254,7 +282,9 @@ pub(crate) fn run( struct Worker { commands: CommandSender, - outbound: mpsc::SyncSender, + outbound: channel::Sender, + cancel: channel::Receiver<()>, + abandoned: Arc, #[cfg(unix)] outbound_wake: UnixStream, pending: HashMap, @@ -264,15 +294,15 @@ struct Worker { threads: Vec>, next_socket_id: i32, session_user: Option<(u32, u32)>, - shutdown: Arc, } impl Worker { fn new( commands: CommandSender, - outbound: mpsc::SyncSender, + outbound: channel::Sender, #[cfg(unix)] outbound_wake: Option, - shutdown: Arc, + cancel: channel::Receiver<()>, + abandoned: Arc, ) -> Result { #[cfg(unix)] let outbound_wake = { @@ -283,6 +313,8 @@ impl Worker { Ok(Self { commands, outbound, + cancel, + abandoned, #[cfg(unix)] outbound_wake, pending: HashMap::new(), @@ -292,7 +324,6 @@ impl Worker { threads: Vec::new(), next_socket_id: 1, session_user: None, - shutdown, }) } @@ -313,14 +344,14 @@ impl Worker { self.threads.push(spawn_listener( source, self.commands.clone(), + self.cancel.clone(), stop, next_client_fd.clone(), )); } let result = loop { - let command = match commands.recv() { - Some(command) => command, - None => break Ok(()), + let Some(command) = commands.recv() else { + break Ok(()); }; let step = match command { Command::Packet(packet) => self.handle_packet(packet), @@ -364,8 +395,13 @@ impl Worker { for (_, stream) in self.pending.drain() { stream.shutdown(); } + let hard_cancelled = !matches!(self.cancel.try_recv(), Err(channel::TryRecvError::Empty)); for (_, io) in self.sources.drain().chain(self.destinations.drain()) { - stop_io(io); + if hard_cancelled { + abort_io(io); + } else { + stop_io(io); + } } for thread in self.threads.drain(..) { let _ = thread.join(); @@ -380,9 +416,10 @@ mod state; mod tests { use std::os::unix::net::UnixStream; use std::sync::atomic::AtomicBool; - use std::sync::{mpsc, Arc}; + use std::sync::Arc; use std::time::Duration; + use crossbeam_channel as channel; use et_core::packet::Packet; use et_core::proto::{ PortForwardDestinationRequest, PortForwardDestinationResponse, SocketEndpoint, @@ -397,21 +434,25 @@ mod tests { fn worker() -> ( Worker, CommandReceiver, - mpsc::Receiver, + channel::Receiver, + channel::Sender<()>, ) { let (commands, command_receiver) = command_channel(MAX_ACTIVE_SOCKETS + 1); - let (outbound, outbound_receiver) = mpsc::sync_channel(MAX_ACTIVE_SOCKETS + 1); + let (outbound, outbound_receiver) = channel::bounded(MAX_ACTIVE_SOCKETS + 1); + let (cancel, cancel_receiver) = channel::bounded(1); let (_wake_reader, wake_writer) = UnixStream::pair().unwrap(); ( Worker::new( commands, outbound, Some(wake_writer), + cancel_receiver, Arc::new(AtomicBool::new(false)), ) .unwrap(), command_receiver, outbound_receiver, + cancel, ) } @@ -457,7 +498,7 @@ mod tests { fn in_flight_connectors_count_toward_socket_limit_and_failure_releases_slot() { let (directory, _cleanup) = test_directory("connector-failure-test"); let missing_socket = directory.join("missing.sock"); - let (mut worker, commands, outbound) = worker(); + let (mut worker, commands, outbound, _cancel) = worker(); for fd in 1..=MAX_ACTIVE_SOCKETS as i32 { worker @@ -505,7 +546,7 @@ mod tests { let (directory, _cleanup) = test_directory("connector-success-test"); let path = directory.join("destination.sock"); let listener = std::os::unix::net::UnixListener::bind(&path).unwrap(); - let (mut worker, commands, _outbound) = worker(); + let (mut worker, commands, _outbound, _cancel) = worker(); let request = Packet::new( TerminalPacketType::PortForwardDestinationRequest as u8, PortForwardDestinationRequest { diff --git a/crates/et-net/src/forward_worker_state.rs b/crates/et-net/src/forward_worker_state.rs index 31e3100..7e417e5 100644 --- a/crates/et-net/src/forward_worker_state.rs +++ b/crates/et-net/src/forward_worker_state.rs @@ -2,6 +2,7 @@ use std::io::Write; use std::io::{self}; +use crossbeam_channel as channel; use et_core::packet::Packet; use et_core::proto::{ PortForwardData, PortForwardDestinationRequest, PortForwardDestinationResponse, SocketEndpoint, @@ -111,6 +112,7 @@ impl Worker { socket_id, destination, self.commands.clone(), + self.cancel.clone(), self.session_user, )); Ok(()) @@ -168,10 +170,22 @@ impl Worker { let Some(active) = self.map_ref(role).get(&socket_id) else { return Ok(()); }; + let byte_count = buffer.len(); active - .writer - .send(WriteCommand::Data(buffer)) - .map_err(|_| ForwardError::Unavailable) + .pending_bytes + .fetch_add(byte_count, std::sync::atomic::Ordering::AcqRel); + let admitted = channel::select! { + send(active.writer, WriteCommand::Data(buffer)) -> result => result.is_ok(), + recv(self.cancel) -> _ => false, + }; + if admitted { + Ok(()) + } else { + active + .pending_bytes + .fetch_sub(byte_count, std::sync::atomic::Ordering::AcqRel); + Err(ForwardError::Unavailable) + } } pub(super) fn send_data( @@ -211,8 +225,15 @@ impl Worker { socket_id: i32, stream: ForwardStream, ) -> Result<(), ForwardError> { - let (active, handles) = - spawn_io(role, socket_id, stream, self.commands.clone()).map_err(ForwardError::Io)?; + let (active, handles) = spawn_io( + role, + socket_id, + stream, + self.commands.clone(), + self.cancel.clone(), + self.abandoned.clone(), + ) + .map_err(ForwardError::Io)?; self.map(role).insert(socket_id, active); self.threads.extend(handles); Ok(()) @@ -252,21 +273,12 @@ impl Worker { } fn emit(&mut self, header: u8, message: M) -> Result<(), ForwardError> { - let mut outbound = Ok(Packet::new(header, message.encode_to_vec())); - loop { - if self.shutdown.load(std::sync::atomic::Ordering::Acquire) { - return Err(ForwardError::Unavailable); - } - match self.outbound.try_send(outbound) { - Ok(()) => break, - Err(std::sync::mpsc::TrySendError::Full(value)) => { - outbound = value; - std::thread::yield_now(); - } - Err(std::sync::mpsc::TrySendError::Disconnected(_)) => { - return Err(ForwardError::Unavailable); - } + let packet = Ok(Packet::new(header, message.encode_to_vec())); + channel::select! { + send(self.outbound, packet) -> result => { + result.map_err(|_| ForwardError::Unavailable)?; } + recv(self.cancel) -> _ => return Err(ForwardError::Unavailable), } // Unix consumers poll the wake socket; Windows consumers drain // `try_outbound` on the client loop's 10ms cadence. diff --git a/crates/et-net/src/local.rs b/crates/et-net/src/local.rs index b2b564f..0a8597e 100644 --- a/crates/et-net/src/local.rs +++ b/crates/et-net/src/local.rs @@ -29,12 +29,26 @@ use std::path::PathBuf; const REGISTRATION_ACK_CAPABILITY: &str = "et-registration-ack-v1"; +use socket2::SockRef; + +/// Terminal-side kernel queue bound for opted-in flow-control sessions. +pub const FLOW_CONTROL_SEND_BUFFER_BYTES: usize = 64 * 1024; + /// Stream type used for local server/terminal IPC. #[cfg(unix)] pub type LocalStream = std::os::unix::net::UnixStream; #[cfg(windows)] pub type LocalStream = std::net::TcpStream; +/// Bound terminal-to-server buffering on the sending endpoint. +/// +/// On Unix this configures the `etterminal` Unix socket, matching upstream +/// PR #730. On Windows `LocalStream` is loopback TCP, where `SO_SNDBUF` is the +/// corresponding bound on the same terminal-side hop. +pub fn minimize_terminal_output_buffering(stream: &LocalStream) -> io::Result<()> { + SockRef::from(stream).set_send_buffer_size(FLOW_CONTROL_SEND_BUFFER_BYTES) +} + /// Length of the hex-encoded Windows registration token. #[cfg(windows)] pub const TOKEN_LEN: usize = 64; diff --git a/crates/et-net/src/local_packet.rs b/crates/et-net/src/local_packet.rs index 49f5a0e..73ad8e2 100644 --- a/crates/et-net/src/local_packet.rs +++ b/crates/et-net/src/local_packet.rs @@ -65,6 +65,22 @@ pub fn read_local_packet(reader: &mut R) -> Result io::Result> { + let serialized = packet.serialize(); + if serialized.len() > MAX_LOCAL_PACKET_LEN { + return Err(io::Error::new( + io::ErrorKind::InvalidInput, + "local packet exceeds 64 KiB", + )); + } + let length = i64::try_from(serialized.len()) + .map_err(|_| io::Error::new(io::ErrorKind::InvalidInput, "local packet too large"))?; + let mut frame = Vec::with_capacity(PREFIX_LEN + serialized.len()); + frame.extend_from_slice(&length.to_ne_bytes()); + frame.extend_from_slice(&serialized); + Ok(frame) +} + pub fn status_packet(header: u8, result: Result<(), &str>) -> Packet { let mut payload = Vec::new(); match result { @@ -100,9 +116,6 @@ pub fn write_local_packet(writer: &mut W, packet: &Packet) -> io::Resu write_local_packet_with(writer, packet, || false) } -/// Write one frame while allowing an owner to cancel blocked backpressure. -/// Once cancellation begins the caller must close the channel because a -/// partially written frame is intentionally abandoned during teardown. pub fn write_local_packet_cancelled( writer: &mut W, packet: &Packet, @@ -114,8 +127,6 @@ pub fn write_local_packet_cancelled( }) } -/// Write one complete local frame under ordinary backpressure, stopping only -/// when teardown explicitly requests cancellation. pub fn write_local_packet_until_cancelled( writer: &mut W, packet: &Packet, @@ -129,17 +140,8 @@ fn write_local_packet_with( packet: &Packet, cancelled: impl Fn() -> bool, ) -> io::Result<()> { - let serialized = packet.serialize(); - if serialized.len() > MAX_LOCAL_PACKET_LEN { - return Err(io::Error::new( - io::ErrorKind::InvalidInput, - "local packet exceeds 64 KiB", - )); - } - let length = i64::try_from(serialized.len()) - .map_err(|_| io::Error::new(io::ErrorKind::InvalidInput, "local packet too large"))?; - write_all_blocking(writer, &length.to_ne_bytes(), &cancelled)?; - write_all_blocking(writer, &serialized, &cancelled)?; + let frame = encode_local_packet(packet)?; + write_all_blocking(writer, &frame, &cancelled)?; flush_blocking(writer, &cancelled) } diff --git a/crates/et-net/tests/connection_resilience.rs b/crates/et-net/tests/connection_resilience.rs index 6e29f9a..207d7f9 100644 --- a/crates/et-net/tests/connection_resilience.rs +++ b/crates/et-net/tests/connection_resilience.rs @@ -5,7 +5,9 @@ use std::net::{Shutdown, TcpListener, TcpStream}; use std::thread; use std::time::{Duration, Instant}; -use et_net::connection::{Connection, DEFAULT_LIVE_WRITE_TIMEOUT, MAX_RECOVERY_PROTO_LEN}; +use et_net::connection::{ + ConnError, Connection, DEFAULT_LIVE_WRITE_TIMEOUT, MAX_RECOVERY_PROTO_LEN, +}; #[test] fn eof_invalidates_connection_so_future_output_is_buffered() { @@ -111,10 +113,26 @@ fn blackholed_peer_write_soft_disconnects_within_live_timeout() { let mut disconnected = false; // Fill the socket send buffer until the bounded write times out. for _ in 0..4_096 { - connection.write_packet(7, &payload).unwrap(); - if !connection.connected() { - disconnected = true; - break; + match connection.write_packet(7, &payload) { + Ok(()) if connection.connected() => {} + Ok(()) => { + disconnected = true; + break; + } + Err(ConnError::Io(error)) + if matches!( + error.kind(), + std::io::ErrorKind::TimedOut | std::io::ErrorKind::WouldBlock + ) => + { + assert!( + !connection.connected(), + "timed-out live write did not soft-disconnect" + ); + disconnected = true; + break; + } + Err(error) => panic!("blackhole write failed unexpectedly: {error}"), } assert!( started.elapsed() < DEFAULT_LIVE_WRITE_TIMEOUT + Duration::from_secs(3), diff --git a/crates/et-net/tests/forward.rs b/crates/et-net/tests/forward.rs index cfa0a63..14f808a 100644 --- a/crates/et-net/tests/forward.rs +++ b/crates/et-net/tests/forward.rs @@ -21,6 +21,8 @@ use et_net::forward::{ForwardOrigin, ForwardSource}; use prost::Message; const TIMEOUT: Duration = Duration::from_secs(3); +const REFUSED_DESTINATION_TIMEOUT: Duration = Duration::from_secs(7); +const HARD_SHUTDOWN_TIMEOUT: Duration = Duration::from_secs(10); #[test] fn two_forwarders_relay_a_real_tcp_round_trip() { @@ -74,7 +76,11 @@ fn refused_destination_closes_the_accepted_source() { .receive(source.wait_outbound(TIMEOUT).unwrap()) .unwrap(); source - .receive(destination.wait_outbound(TIMEOUT).unwrap()) + .receive( + destination + .wait_outbound(REFUSED_DESTINATION_TIMEOUT) + .unwrap(), + ) .unwrap(); let mut byte = [0u8; 1]; assert_eq!(application.read(&mut byte).unwrap(), 0); @@ -88,6 +94,175 @@ fn refused_destination_closes_the_accepted_source() { /// its next packet — wedging the session permanently. `try_receive` must /// report a full worker instead of blocking, and draining outbound packets /// (the session loop's next step) must make the held packet deliverable. +#[test] +fn hard_shutdown_cancels_worker_blocked_on_full_command_and_outbound_queues() { + let mut forwarder = Forwarder::start(Vec::new()).unwrap(); + let request = |fd: i32| { + Packet::new( + TerminalPacketType::PortForwardDestinationRequest as u8, + PortForwardDestinationRequest { + destination: Some(SocketEndpoint { + name: None, + port: Some(0), + }), + fd: Some(fd), + } + .encode_to_vec(), + ) + }; + + let mut held = None; + for fd in 1..=4096 { + if let Some(packet) = forwarder.try_receive(request(fd)).unwrap() { + held = Some(packet); + break; + } + } + assert!(held.is_some(), "bounded forwarding queues never filled"); + + let (done_tx, done_rx) = std::sync::mpsc::sync_channel(0); + let worker = thread::spawn(move || done_tx.send(forwarder.shutdown_hard()).unwrap()); + assert!(done_rx + .recv_timeout(HARD_SHUTDOWN_TIMEOUT) + .unwrap() + .unwrap()); + worker.join().unwrap(); +} + +#[test] +fn hard_shutdown_cancels_active_destination_write_after_socket_backpressure() { + // Given: a real forwarding destination accepts but never drains its socket. + let listener = TcpListener::bind((Ipv4Addr::LOCALHOST, 0)).unwrap(); + let destination_port = listener.local_addr().unwrap().port(); + let (accepted_tx, accepted_rx) = std::sync::mpsc::sync_channel(0); + let (release_tx, release_rx) = std::sync::mpsc::sync_channel(0); + let peer = thread::spawn(move || { + let (stream, _) = listener.accept().unwrap(); + socket2::SockRef::from(&stream) + .set_recv_buffer_size(4096) + .unwrap(); + accepted_tx.send(()).unwrap(); + release_rx.recv().unwrap(); + }); + let source_port = reserve_port(); + let mut source = Forwarder::start(vec![request(source_port, destination_port)]).unwrap(); + let mut destination = Forwarder::start(Vec::new()).unwrap(); + let mut application = TcpStream::connect((Ipv4Addr::LOCALHOST, source_port)).unwrap(); + destination + .receive(source.wait_outbound(TIMEOUT).unwrap()) + .unwrap(); + source + .receive(destination.wait_outbound(TIMEOUT).unwrap()) + .unwrap(); + accepted_rx.recv_timeout(TIMEOUT).unwrap(); + let application_writer = thread::spawn(move || { + let payload = vec![7u8; 64 * 1024 * 1024]; + let _ = application.write_all(&payload); + }); + + // When: destination writer ownership is proven blocked by its full bounded + // write-command queue, hard cancellation must close the socket before join. + let held = loop { + let packet = source.wait_outbound(TIMEOUT).unwrap(); + if let Some(held) = destination.try_receive(packet).unwrap() { + break held; + } + }; + drop(held); + let (done_tx, done_rx) = std::sync::mpsc::sync_channel(0); + let shutdown = thread::spawn(move || done_tx.send(destination.shutdown_hard()).unwrap()); + + // Then: completion is the exact shutdown result, not elapsed-time inference. + assert!(done_rx + .recv_timeout(HARD_SHUTDOWN_TIMEOUT) + .unwrap() + .unwrap()); + shutdown.join().unwrap(); + source.shutdown_hard().unwrap(); + application_writer.join().unwrap(); + release_tx.send(()).unwrap(); + peer.join().unwrap(); +} + +#[test] +fn hard_shutdown_reports_admitted_socket_bytes_abandoned() { + // Given: a destination writer is blocked by a peer that never drains, but + // every forwarding command has crossed the worker boundary. The trailing + // response is a FIFO barrier proving the outer command queue is empty. + let listener = TcpListener::bind((Ipv4Addr::LOCALHOST, 0)).unwrap(); + let destination_port = listener.local_addr().unwrap().port(); + let (accepted_tx, accepted_rx) = std::sync::mpsc::sync_channel(0); + let (release_tx, release_rx) = std::sync::mpsc::sync_channel(0); + let peer = thread::spawn(move || { + let (stream, _) = listener.accept().unwrap(); + socket2::SockRef::from(&stream) + .set_recv_buffer_size(4096) + .unwrap(); + accepted_tx.send(()).unwrap(); + release_rx.recv().unwrap(); + }); + let source_port = reserve_port(); + let mut source = Forwarder::start(vec![request(source_port, destination_port)]).unwrap(); + let mut destination = Forwarder::start(Vec::new()).unwrap(); + let application = TcpStream::connect((Ipv4Addr::LOCALHOST, source_port)).unwrap(); + destination + .receive(source.wait_outbound(TIMEOUT).unwrap()) + .unwrap(); + let response_packet = destination.wait_outbound(TIMEOUT).unwrap(); + let response = + et_core::proto::PortForwardDestinationResponse::decode(response_packet.payload()).unwrap(); + let socket_id = response.socketid.unwrap(); + source.receive(response_packet).unwrap(); + accepted_rx.recv_timeout(TIMEOUT).unwrap(); + let data_packet = || { + Packet::new( + TerminalPacketType::PortForwardData as u8, + et_core::proto::PortForwardData { + sourcetodestination: Some(true), + socketid: Some(socket_id), + buffer: Some(vec![9u8; 64 * 1024]), + error: None, + closed: None, + } + .encode_to_vec(), + ) + }; + for _ in 0..65 { + destination.receive(data_packet()).unwrap(); + } + destination + .receive(Packet::new( + TerminalPacketType::PortForwardDestinationRequest as u8, + PortForwardDestinationRequest { + destination: Some(SocketEndpoint { + name: None, + port: Some(0), + }), + fd: Some(777), + } + .encode_to_vec(), + )) + .unwrap(); + let barrier = destination.wait_outbound(TIMEOUT).unwrap(); + let barrier = + et_core::proto::PortForwardDestinationResponse::decode(barrier.payload()).unwrap(); + assert_eq!(barrier.clientfd, Some(777)); + assert!(barrier.error.is_some()); + + // When: hard shutdown aborts the admitted writer queue and in-flight write. + let abandoned = destination.shutdown_hard().unwrap(); + + // Then: payload ownership loss is reported even though outer lanes drained. + release_tx.send(()).unwrap(); + peer.join().unwrap(); + drop(application); + source.shutdown_hard().unwrap(); + assert!( + abandoned, + "hard shutdown discarded admitted socket bytes without reporting abandonment" + ); +} + #[test] fn try_receive_reports_backpressure_and_outbound_drain_recovers() { let forwarder = Forwarder::start(Vec::new()).unwrap(); diff --git a/crates/et-net/tests/local_packet.rs b/crates/et-net/tests/local_packet.rs index 8b9794a..fac8dc7 100644 --- a/crates/et-net/tests/local_packet.rs +++ b/crates/et-net/tests/local_packet.rs @@ -1,5 +1,7 @@ #![forbid(unsafe_code)] +#[cfg(target_os = "linux")] +use std::io::Write; use std::io::{Cursor, ErrorKind, Read}; use et_core::packet::Packet; @@ -115,3 +117,34 @@ fn writer_survives_wouldblock_backpressure_on_a_nonblocking_stream() { assert_eq!(packet.payload(), vec![index as u8; PAYLOAD]); } } + +#[cfg(target_os = "linux")] +#[test] +fn opted_in_terminal_sender_hits_backpressure_near_the_configured_bound() { + let (_server, mut terminal) = et_net::local::wake_pair().unwrap(); + et_net::local::minimize_terminal_output_buffering(&terminal).unwrap(); + terminal.set_nonblocking(true).unwrap(); + + let chunk = [0u8; 16 * 1024]; + let mut queued = 0usize; + loop { + match terminal.write(&chunk) { + Ok(0) => panic!("local stream stopped accepting output before backpressure"), + Ok(count) => queued += count, + Err(error) if error.kind() == ErrorKind::WouldBlock => break, + Err(error) => panic!("unexpected terminal output error: {error}"), + } + } + + // Linux reports SO_SNDBUF at twice the requested value for bookkeeping. + // One write may straddle the threshold, so allow one chunk of headroom. + let configured = et_net::local::FLOW_CONTROL_SEND_BUFFER_BYTES; + assert!( + queued >= configured, + "backpressure arrived too early: {queued}" + ); + assert!( + queued <= configured * 2 + chunk.len(), + "sender queued {queued} bytes past the configured {configured}-byte bound" + ); +} diff --git a/crates/et-server/src/runtime.rs b/crates/et-server/src/runtime.rs index e940c10..999f373 100644 --- a/crates/et-server/src/runtime.rs +++ b/crates/et-server/src/runtime.rs @@ -21,6 +21,10 @@ use crate::runtime_state::{ }; use crate::session_table::SessionTable; +#[cfg(all(test, unix))] +#[path = "runtime_recovery_test.rs"] +mod recovery_tests; + pub struct Runtime { core: Arc, router: Option, @@ -294,7 +298,8 @@ mod tests { let (assignment_tx, assignment_rx) = mpsc::sync_channel(1); let (release_tx, release_rx) = mpsc::sync_channel(1); let (scan_tx, scan_rx) = mpsc::sync_channel(1); - crate::runtime_handler::install_raw_assignment_hook(ID, assignment_tx, release_rx); + let identity = runtime.core.registry.get(ID).unwrap().unwrap().identity(); + crate::runtime_handler::install_raw_assignment_hook(identity, assignment_tx, release_rx); crate::runtime_lifecycle::install_raw_scan_hook(ID, scan_tx); let mut client = connect_request(address); diff --git a/crates/et-server/src/runtime_handler.rs b/crates/et-server/src/runtime_handler.rs index 318f5b9..126848d 100644 --- a/crates/et-server/src/runtime_handler.rs +++ b/crates/et-server/src/runtime_handler.rs @@ -101,9 +101,9 @@ pub(crate) fn handle( return; } }; - #[cfg(test)] - run_before_raw_assignment_hook(&id); let registration_identity = registration.identity(); + #[cfg(test)] + run_before_raw_assignment_hook(®istration_identity); if guard.assign(registration_identity.clone()).is_err() { crate::diag::info(format!( "drop {peer} id={id}: could not track raw socket for registration" @@ -180,12 +180,20 @@ pub(crate) fn handle( return; } }; + if guard.own_session().is_err() { + crate::diag::info(format!( + "id={id}: drop recover from {peer}: could not protect returning socket" + )); + return; + } if send_status(&mut stream, ConnectStatus::ReturningClient).is_err() { crate::diag::info(format!( "id={id}: failed to send ReturningClient status to {peer}" )); return; } + #[cfg(test)] + run_after_returning_status_hook(&id); match permit.complete(stream) { Ok(()) => { crate::diag::info(format!("id={id}: session recover accepted from {peer}")) @@ -374,6 +382,7 @@ fn handle_new( let term_init = TermInit { environmentnames: environment.keys().cloned().collect(), environmentvalues: environment.values().cloned().collect(), + flowcontrol: payload.flowcontrol, }; let init_packet = et_core::packet::Packet::new( TerminalPacketType::TerminalInit as u8, @@ -417,7 +426,7 @@ fn handle_new( )); return; } - let active = match ActiveSession::new(connection, &terminal) { + let active = match ActiveSession::new(connection, &terminal, payload.flowcontrol) { Ok(active) => active, Err(error) => { crate::diag::info(format!( @@ -427,6 +436,7 @@ fn handle_new( } }; let active = Arc::new(active); + active.start_flow_writer(); if start.activate(active.clone()).is_err() { crate::diag::info(format!("id={id}: could not activate session for {peer}")); return; @@ -567,7 +577,7 @@ fn run_jumphost( { return; } - let active = match ActiveSession::new(connection, &terminal) { + let active = match ActiveSession::new(connection, &terminal, payload.flowcontrol) { Ok(active) => active, Err(error) => { crate::diag::info(format!( @@ -577,6 +587,7 @@ fn run_jumphost( } }; let active = Arc::new(active); + active.start_flow_writer(); if start.activate(active.clone()).is_err() { crate::diag::info(format!("id={id}: jumphost could not activate session")); return; @@ -669,7 +680,7 @@ fn valid_id(id: &str) -> bool { #[cfg(test)] struct RawAssignmentHook { - id: String, + identity: crate::registry::RegistrationIdentity, reached: std::sync::mpsc::SyncSender<()>, release: std::sync::mpsc::Receiver<()>, } @@ -683,21 +694,67 @@ fn raw_assignment_hook() -> &'static std::sync::Mutex> #[cfg(test)] pub(crate) fn install_raw_assignment_hook( - id: &str, + identity: crate::registry::RegistrationIdentity, reached: std::sync::mpsc::SyncSender<()>, release: std::sync::mpsc::Receiver<()>, ) { *raw_assignment_hook().lock().unwrap() = Some(RawAssignmentHook { - id: id.to_owned(), + identity, reached, release, }); } #[cfg(test)] -fn run_before_raw_assignment_hook(id: &str) { +fn run_before_raw_assignment_hook(identity: &crate::registry::RegistrationIdentity) { let hook = { let mut installed = raw_assignment_hook().lock().unwrap(); + if installed + .as_ref() + .is_some_and(|hook| hook.identity.same_generation(identity)) + { + installed.take() + } else { + None + } + }; + if let Some(hook) = hook { + hook.reached.send(()).unwrap(); + hook.release.recv().unwrap(); + } +} + +#[cfg(test)] +struct ReturningStatusHook { + id: String, + reached: std::sync::mpsc::SyncSender<()>, + release: std::sync::mpsc::Receiver<()>, +} + +#[cfg(test)] +fn returning_status_hook() -> &'static std::sync::Mutex> { + static HOOK: std::sync::OnceLock>> = + std::sync::OnceLock::new(); + HOOK.get_or_init(|| std::sync::Mutex::new(None)) +} + +#[cfg(test)] +pub(crate) fn install_returning_status_hook( + id: &str, + reached: std::sync::mpsc::SyncSender<()>, + release: std::sync::mpsc::Receiver<()>, +) { + *returning_status_hook().lock().unwrap() = Some(ReturningStatusHook { + id: id.to_owned(), + reached, + release, + }); +} + +#[cfg(test)] +fn run_after_returning_status_hook(id: &str) { + let hook = { + let mut installed = returning_status_hook().lock().unwrap(); if installed.as_ref().is_some_and(|hook| hook.id == id) { installed.take() } else { diff --git a/crates/et-server/src/runtime_lifecycle.rs b/crates/et-server/src/runtime_lifecycle.rs index 0a224e2..ae3bdaa 100644 --- a/crates/et-server/src/runtime_lifecycle.rs +++ b/crates/et-server/src/runtime_lifecycle.rs @@ -83,37 +83,25 @@ pub(crate) fn run( } #[cfg(test)] -struct RawScanHook { - id: String, - complete: std::sync::mpsc::SyncSender<()>, -} - -#[cfg(test)] -fn raw_scan_hook() -> &'static std::sync::Mutex> { - static HOOK: std::sync::OnceLock>> = - std::sync::OnceLock::new(); - HOOK.get_or_init(|| std::sync::Mutex::new(None)) +fn raw_scan_hooks( +) -> &'static std::sync::Mutex>> { + static HOOKS: std::sync::OnceLock< + std::sync::Mutex>>, + > = std::sync::OnceLock::new(); + HOOKS.get_or_init(|| std::sync::Mutex::new(std::collections::HashMap::new())) } #[cfg(test)] pub(crate) fn install_raw_scan_hook(id: &str, complete: std::sync::mpsc::SyncSender<()>) { - *raw_scan_hook().lock().unwrap() = Some(RawScanHook { - id: id.to_owned(), - complete, - }); + raw_scan_hooks() + .lock() + .unwrap() + .insert(id.to_owned(), complete); } #[cfg(test)] fn notify_raw_scan_complete(id: &str) { - let complete = { - let mut installed = raw_scan_hook().lock().unwrap(); - if installed.as_ref().is_some_and(|hook| hook.id == id) { - installed.take().map(|hook| hook.complete) - } else { - None - } - }; - if let Some(complete) = complete { + if let Some(complete) = raw_scan_hooks().lock().unwrap().remove(id) { let _ = complete.send(()); } } diff --git a/crates/et-server/src/runtime_recovery_test.rs b/crates/et-server/src/runtime_recovery_test.rs new file mode 100644 index 0000000..4091b0b --- /dev/null +++ b/crates/et-server/src/runtime_recovery_test.rs @@ -0,0 +1,218 @@ +use std::collections::HashMap; +use std::fs; +use std::io::{self, Read}; +use std::net::{IpAddr, Ipv4Addr, SocketAddr, TcpStream}; +use std::os::unix::fs::PermissionsExt; +use std::path::{Path, PathBuf}; +use std::sync::atomic::{AtomicU64, Ordering}; +use std::sync::mpsc; +use std::time::Duration; + +use et_core::keys::passkey_to_key; +use et_core::packet::Packet; +use et_core::proto::{ + ConnectResponse, ConnectStatus, FlowControlMode, InitialPayload, InitialResponse, + TerminalBuffer, TerminalPacketType, TerminalUserInfo, +}; +use et_net::connection::Connection; +use et_net::framing_io::{read_proto_limited, write_proto}; +use et_net::handshake::client_request; +use et_net::local_packet::{read_local_packet, write_local_packet}; +use prost::Message; + +use super::Runtime; +use crate::path::select_router_path_for; + +const ID: &str = "flowpause0000001"; +const KEY: &str = "ABCDEFGHIJKLMNOPQRSTUVWXYZabcdef"; +const TIMEOUT: Duration = Duration::from_secs(3); +static NEXT_DIRECTORY: AtomicU64 = AtomicU64::new(0); + +struct TestDirectory(PathBuf); + +impl TestDirectory { + fn new() -> Self { + let serial = NEXT_DIRECTORY.fetch_add(1, Ordering::Relaxed); + let path = std::env::temp_dir().join(format!( + "et-rs-recovery-pause-test-{}-{serial}", + std::process::id() + )); + fs::create_dir(&path).unwrap(); + fs::set_permissions(&path, fs::Permissions::from_mode(0o700)).unwrap(); + Self(path) + } + + fn socket(&self) -> PathBuf { + self.0.join("router.sock") + } +} + +impl Drop for TestDirectory { + fn drop(&mut self) { + let _ = fs::remove_dir_all(&self.0); + } +} + +#[test] +fn terminal_hup_after_returning_status_delivers_final_output_only_to_recovered_connection() { + // Given: an active flow-controlled session and a recovery handler blocked + // immediately after exposing ReturningClient but before permit completion. + let directory = TestDirectory::new(); + let router_path = select_router_path_for( + rustix::process::getuid().as_raw(), + Some(&directory.socket()), + None, + None, + ) + .unwrap(); + let mut runtime = Runtime::start(IpAddr::V4(Ipv4Addr::LOCALHOST), 0, router_path).unwrap(); + let handle = runtime.handle(); + let address = runtime.tcp_addresses()[0]; + let mut terminal = register(&directory.socket(), &handle); + let (stream, response) = handshake(address); + assert_eq!(response.status, Some(ConnectStatus::NewClient as i32)); + let mut client = Connection::new_client(stream, &passkey_to_key(KEY).unwrap()); + let payload = InitialPayload { + jumphost: Some(false), + reversetunnels: Vec::new(), + environmentvariables: HashMap::new(), + flowcontrol: Some(FlowControlMode::Backpressure as i32), + }; + client.write_packet(253, &payload.encode_to_vec()).unwrap(); + let initial = client.read_packet().unwrap(); + assert_eq!( + InitialResponse::decode(initial.payload()).unwrap().error, + None + ); + assert_eq!( + read_local_packet(&mut terminal).unwrap().header(), + TerminalPacketType::TerminalInit as u8 + ); + let mut old_stream = client.try_clone_stream().unwrap(); + old_stream.set_read_timeout(Some(TIMEOUT)).unwrap(); + let (status_tx, status_rx) = mpsc::sync_channel(1); + let (release_tx, release_rx) = mpsc::sync_channel(1); + crate::runtime_handler::install_returning_status_hook(ID, status_tx, release_rx); + let (returning, response) = handshake(address); + assert_eq!(response.status, Some(ConnectStatus::ReturningClient as i32)); + status_rx + .recv_timeout(TIMEOUT) + .expect("recovery handler did not reach the post-status barrier"); + let (client_tx, client_rx) = mpsc::sync_channel(1); + let recovering = std::thread::spawn(move || { + client.recover(returning).unwrap(); + client + .write_packet(TerminalPacketType::KeepAlive as u8, &[]) + .unwrap(); + client_tx.send(client).unwrap(); + }); + + // When: the terminal queues its final output while permit completion + // remains blocked at that exact barrier, then the real terminal socket + // reaches HUP and lifecycle scans the returning raw socket. + let (queued_tx, queued_rx) = mpsc::sync_channel(1); + runtime + .core + .sessions + .active(ID) + .unwrap() + .unwrap() + .install_flow_enqueue_hook(queued_tx); + let final_output = TerminalBuffer { + buffer: Some(b"final-after-returning-status".to_vec()), + }; + write_local_packet( + &mut terminal, + &Packet::new( + TerminalPacketType::TerminalBuffer as u8, + final_output.encode_to_vec(), + ), + ) + .unwrap(); + queued_rx + .recv_timeout(TIMEOUT) + .expect("terminal output was not admitted to the flow queue"); + let (scan_tx, scan_rx) = mpsc::sync_channel(1); + crate::runtime_lifecycle::install_raw_scan_hook(ID, scan_tx); + drop(terminal); + scan_rx + .recv_timeout(TIMEOUT) + .expect("terminal lifecycle did not complete its raw-socket scan"); + + // Then: releasing recovery installs the protected candidate, delivers the + // packet exactly once there, and retires the old stream without bytes. + release_tx.send(()).unwrap(); + let mut client = client_rx + .recv_timeout(TIMEOUT) + .expect("client recovery did not complete after terminal HUP"); + recovering.join().unwrap(); + let mut byte = [0u8; 1]; + let old_read = old_stream.read(&mut byte); + assert!( + match &old_read { + Ok(0) => true, + Err(error) + if matches!( + error.kind(), + io::ErrorKind::ConnectionReset | io::ErrorKind::ConnectionAborted + ) => + { + true + } + Ok(_) | Err(_) => false, + }, + "final output reached the old connection after terminal HUP: {old_read:?}" + ); + assert_eq!( + client.read_packet().unwrap().header(), + TerminalPacketType::KeepAlive as u8 + ); + let delivered = client.read_packet().unwrap(); + assert_eq!(delivered.header(), TerminalPacketType::TerminalBuffer as u8); + assert_eq!( + TerminalBuffer::decode(delivered.payload()).unwrap(), + final_output + ); + if let Ok(packet) = client.read_packet() { + assert_eq!( + packet.header(), + TerminalPacketType::KeepAlive as u8, + "final terminal output was duplicated before EOF" + ); + assert!( + client.read_packet().is_err(), + "terminal EOF must follow the recovery keepalive" + ); + } + runtime.shutdown().unwrap(); +} + +fn register( + path: &Path, + handle: &crate::runtime_handle::RuntimeHandle, +) -> et_net::local::LocalStream { + let mut stream = et_net::local::connect(path).unwrap(); + let packet = Packet::new( + TerminalPacketType::TerminalUserInfo as u8, + TerminalUserInfo { + id: Some(ID.to_owned()), + passkey: Some(KEY.to_owned()), + uid: Some(i64::from(rustix::process::getuid().as_raw())), + gid: Some(i64::from(rustix::process::getgid().as_raw())), + fd: None, + } + .encode_to_vec(), + ); + write_local_packet(&mut stream, &packet).unwrap(); + handle.wait_registered(ID, TIMEOUT).unwrap(); + stream +} + +fn handshake(address: SocketAddr) -> (TcpStream, ConnectResponse) { + let mut stream = TcpStream::connect(address).unwrap(); + stream.set_read_timeout(Some(TIMEOUT)).unwrap(); + stream.set_write_timeout(Some(TIMEOUT)).unwrap(); + write_proto(&mut stream, &client_request(ID)).unwrap(); + let response = read_proto_limited(&mut stream, 64 * 1024).unwrap(); + (stream, response) +} diff --git a/crates/et-server/src/session.rs b/crates/et-server/src/session.rs index 27bd007..883fffe 100644 --- a/crates/et-server/src/session.rs +++ b/crates/et-server/src/session.rs @@ -1,15 +1,15 @@ use et_net::local::LocalStream; -use std::io::{self, Write}; +use std::io; use std::net::{Shutdown, TcpStream}; use std::sync::atomic::{AtomicBool, AtomicU64, Ordering}; -use std::sync::{Arc, Condvar, Mutex, MutexGuard, TryLockError}; -use std::time::{Duration, Instant}; +use std::sync::{Arc, Condvar, Mutex}; +use std::time::Duration; use et_core::backed_writer::{ MAX_BACKUP_PACKETS, MAX_DISCONNECT_PACKETS, MAX_RECOVERY_BACKUP_BYTES, }; -use et_core::proto::TerminalPacketType; -use et_net::connection::DEFAULT_RECOVERY_TIMEOUT; +use et_core::flow_control::FlowControlMode as QueueMode; +use et_core::packet::Packet; use et_net::connection::{ConnError, Connection}; /// How long `recover` may wait for the session mutexes. @@ -20,7 +20,19 @@ use et_net::connection::{ConnError, Connection}; /// deadline, piled-up returning clients fail fast and retry instead of /// blocking the accept path for minutes. Recovery network I/O itself runs /// *without* the connection mutex (see [`ActiveSession::recover_body`]). -const RECOVERY_LOCK_TIMEOUT: Duration = DEFAULT_RECOVERY_TIMEOUT; +const RECOVERY_LOCK_TIMEOUT: Duration = et_net::connection::DEFAULT_RECOVERY_TIMEOUT; +const FLOW_CONTROL_BUFFER_BYTES: usize = 64 * 1024; + +#[path = "session_flow.rs"] +mod session_flow; +#[cfg(test)] +#[path = "session_flow_test.rs"] +mod session_flow_tests; +#[path = "session_io.rs"] +mod session_io; +#[path = "session_recovery.rs"] +mod session_recovery; +use session_flow::FlowControl; pub(crate) struct ActiveSession { connection: Mutex, @@ -38,6 +50,8 @@ pub(crate) struct ActiveSession { /// to queue on the connection mutex for minutes after a blackhole write. recovering: AtomicBool, connection_generation: AtomicU64, + flow_control: Option>, + flow_writer: Mutex>>, bridge_generation: Mutex, bridge_changed: Condvar, } @@ -47,6 +61,20 @@ pub(crate) enum SessionConnection { Active(Arc), } +#[derive(Debug)] +pub enum SessionWriteError { + BeforeReplay(SessionError), + ReplayOwned(SessionError), +} + +impl SessionWriteError { + fn into_inner(self) -> SessionError { + match self { + Self::BeforeReplay(error) | Self::ReplayOwned(error) => error, + } + } +} + #[derive(Debug)] pub enum SessionError { Connection(ConnError), @@ -83,9 +111,16 @@ impl std::error::Error for SessionError { impl ActiveSession { pub(crate) fn new( - connection: Connection, + mut connection: Connection, terminal: &LocalStream, + flow_control: Option, ) -> Result { + let queue_mode = queue_mode(flow_control); + if queue_mode.is_some() { + connection + .minimize_output_buffering() + .map_err(SessionError::Connection)?; + } let control = connection .try_clone_stream() .map_err(SessionError::Connection)?; @@ -102,38 +137,108 @@ impl ActiveSession { shutdown: AtomicBool::new(false), recovering: AtomicBool::new(false), connection_generation: AtomicU64::new(0), + flow_control: queue_mode.map(|mode| Arc::new(FlowControl::new(mode))), + flow_writer: Mutex::new(None), bridge_generation: Mutex::new(0), bridge_changed: Condvar::new(), }) } + pub(crate) fn start_flow_writer(self: &Arc) { + let Some(state) = self.flow_control.clone() else { + return; + }; + let session = Arc::downgrade(self); + let handle = std::thread::spawn(move || session_flow::run_writer(session, state)); + if let Ok(mut writer) = self.flow_writer.lock() { + *writer = Some(handle); + } + } + + #[cfg(test)] + pub(crate) fn install_flow_enqueue_hook(&self, reached: std::sync::mpsc::SyncSender<()>) { + self.flow_control + .as_ref() + .expect("test session must enable flow control") + .install_enqueue_hook(reached); + } + pub(crate) fn send_packet(&self, header: u8, payload: &[u8]) -> Result<(), SessionError> { + match self.send_packet_owned(header, payload) { + Ok(()) => Ok(()), + Err(SessionWriteError::BeforeReplay(SessionError::Connection(ConnError::Io(_)))) => { + self.connection + .lock() + .map_err(|_| SessionError::Unavailable)? + .disconnect(); + self.send_packet_owned(header, payload) + .map_err(SessionWriteError::into_inner) + } + Err(error) => Err(error.into_inner()), + } + } + + pub(crate) fn send_packet_owned( + &self, + header: u8, + payload: &[u8], + ) -> Result<(), SessionWriteError> { + self.send_packet_owned_with(header, payload, |connection, header, payload| { + connection.write_packet_owned(header, payload) + }) + } + + fn send_packet_owned_with( + &self, + header: u8, + payload: &[u8], + mut write: W, + ) -> Result<(), SessionWriteError> + where + W: FnMut(&mut Connection, u8, &[u8]) -> Result<(), et_net::connection::WritePacketError>, + { + if let Some(state) = &self.flow_control { + return state + .enqueue(Packet::new(header, payload)) + .map_err(SessionWriteError::BeforeReplay); + } // While a recover holds the single-flight permit, queue terminal // output instead of contending on the connection mutex for the // recovery network RTT. Flushed after the new stream is installed. - if self.queue_if_recovering(header, payload)? { + if self + .queue_if_recovering(header, payload) + .map_err(SessionWriteError::BeforeReplay)? + { return Ok(()); } let mut connection = self .connection .lock() - .map_err(|_| SessionError::Unavailable)?; + .map_err(|_| SessionWriteError::BeforeReplay(SessionError::Unavailable))?; // Recover may have started after the fast path check. Drop the // connection lock before taking `recover_hold` (flush takes hold // then connection — reverse order deadlocks). if self.recovering.load(Ordering::Acquire) { drop(connection); - if self.queue_if_recovering(header, payload)? { + if self + .queue_if_recovering(header, payload) + .map_err(SessionWriteError::BeforeReplay)? + { return Ok(()); } connection = self .connection .lock() - .map_err(|_| SessionError::Unavailable)?; + .map_err(|_| SessionWriteError::BeforeReplay(SessionError::Unavailable))?; } - connection - .write_packet(header, payload) - .map_err(SessionError::Connection) + write(&mut connection, header, payload).map_err(|error| match error { + et_net::connection::WritePacketError::BeforeReplay(error) => { + SessionWriteError::BeforeReplay(SessionError::Connection(error)) + } + et_net::connection::WritePacketError::ReplayOwned(error) => { + SessionWriteError::ReplayOwned(SessionError::Connection(error)) + } + }) } fn queue_if_recovering(&self, header: u8, payload: &[u8]) -> Result { @@ -170,114 +275,9 @@ impl ActiveSession { Ok(true) } - /// Acquire the single-flight recover permit without speaking on the wire. - /// - /// Callers must send `ReturningClient` only after this succeeds, so a - /// concurrent recover does not commit the peer to sequence exchange and - /// then fail with `RecoverBusy`. The permit releases the flag on drop - /// (including panic unwind). - pub(crate) fn try_begin_recover(&self) -> Result, SessionError> { - if self - .recovering - .compare_exchange(false, true, Ordering::AcqRel, Ordering::Acquire) - .is_err() - { - return Err(SessionError::RecoverBusy); - } - Ok(RecoverPermit { session: self }) - } - - /// Prepare → network handshake off-lock → install → flush hold. - /// - /// The connection mutex is held only for soft-disconnect/snapshot and for - /// installing the new stream, not for sequence exchange or peer auth. - fn recover_body(&self, stream: TcpStream) -> Result<(), SessionError> { - // Phase 1: soft-disconnect and snapshot under a short lock. - let mut candidate = { - let connection = lock_timeout(&self.connection, RECOVERY_LOCK_TIMEOUT)?; - // Snapshot onto the new stream without closing or disconnecting - // the live victim socket (ET #784 / ANT-2026-VAMER5RC). Terminal - // output during the off-lock handshake is queued in - // `recover_hold` because `recovering` is set. A failed recover - // must leave the existing session intact. - connection.prepare_recovery_candidate(stream) - }; - - // Phase 2: recovery network I/O without the session connection lock. - candidate - .run_recovery_handshake(DEFAULT_RECOVERY_TIMEOUT) - .map_err(SessionError::Connection)?; - // Any packet that decrypts with the session key authenticates the - // returning client; it is requeued and handled by the session loop. - candidate - .authenticate_peer(DEFAULT_RECOVERY_TIMEOUT) - .map_err(SessionError::Connection)?; - let ack = candidate.keepalive_ack(); - candidate - .write_packet_live(TerminalPacketType::KeepAlive as u8, &ack) - .map_err(SessionError::Connection)?; - let new_control = candidate - .try_clone_stream() - .map_err(SessionError::Connection)?; - - // Phase 3: install under a short lock. - { - let mut control = lock_timeout(&self.control, RECOVERY_LOCK_TIMEOUT)?; - let mut connection = lock_timeout(&self.connection, RECOVERY_LOCK_TIMEOUT)?; - let old_control = std::mem::replace(&mut *control, new_control); - let _ = old_control.shutdown(Shutdown::Both); - *connection = candidate; - self.connection_generation.fetch_add(1, Ordering::Release); - } - - // Phase 4: drain terminal output queued while the handshake ran. - // Still under `recovering` so concurrent send_packet keeps queuing - // until the permit drops; Drop flushes once more after clearing. - self.flush_recover_hold() - } - - fn flush_recover_hold(&self) -> Result<(), SessionError> { - loop { - let batch = { - let mut hold = self - .recover_hold - .lock() - .map_err(|_| SessionError::Unavailable)?; - if hold.is_empty() { - return Ok(()); - } - std::mem::take(&mut *hold) - }; - let batch_bytes = batch.iter().fold(0u64, |total, (_, payload)| { - total.saturating_add(payload.len() as u64) - }); - let mut connection = lock_timeout(&self.connection, RECOVERY_LOCK_TIMEOUT)?; - let mut remaining = batch.into_iter(); - while let Some((header, payload)) = remaining.next() { - if let Err(error) = connection.write_packet(header, &payload) { - // Release the connection lock before taking `recover_hold` - // (send_packet may hold hold → connection; reverse deadlocks). - drop(connection); - // Put the failed packet and unwritten tail back ahead of - // anything concurrent senders queued after we took `batch`. - let mut hold = self - .recover_hold - .lock() - .map_err(|_| SessionError::Unavailable)?; - let concurrent = std::mem::take(&mut *hold); - hold.push((header, payload)); - hold.extend(remaining); - hold.extend(concurrent); - return Err(SessionError::Connection(error)); - } - } - self.recover_hold_bytes - .fetch_sub(batch_bytes, Ordering::AcqRel); - } - } - pub(crate) fn finish_terminal(&self) -> Result<(), SessionError> { self.shutdown.store(true, Ordering::Release); + let flow_result = self.join_flow_writer(true); let _ = self.signal(); let terminal = self .terminal_control @@ -286,15 +286,31 @@ impl ActiveSession { let _ = terminal.shutdown(Shutdown::Both); drop(terminal); let control = self.control.lock().map_err(|_| SessionError::Unavailable)?; - match control.shutdown(Shutdown::Write) { - Ok(()) => Ok(()), - Err(error) if error.kind() == io::ErrorKind::NotConnected => Ok(()), - Err(error) => Err(SessionError::Io(error)), + let control_result = if flow_result.is_err() { + control.shutdown(Shutdown::Both) + } else { + control.shutdown(Shutdown::Write) + }; + match control_result { + Ok(()) => {} + Err(error) if error.kind() == io::ErrorKind::NotConnected => {} + Err(error) => return Err(SessionError::Io(error)), + } + drop(control); + if let Err(error) = flow_result { + let mut connection = self + .connection + .lock() + .map_err(|_| SessionError::Unavailable)?; + let _ = connection.shutdown(); + return Err(error); } + Ok(()) } pub(crate) fn shutdown(&self) -> Result<(), SessionError> { self.shutdown.store(true, Ordering::Release); + self.join_flow_writer(false)?; let _ = self.signal(); let terminal = self .terminal_control @@ -316,206 +332,54 @@ impl ActiveSession { connection.shutdown().map_err(SessionError::Connection) } - pub(crate) fn take_wake_reader(&self) -> Result { - self.wake_reader - .lock() - .map_err(|_| SessionError::Unavailable)? - .take() - .ok_or(SessionError::Unavailable) - } - - #[cfg_attr(windows, allow(dead_code))] - pub(crate) fn try_clone_stream(&self) -> Result<(TcpStream, u64), SessionError> { - let connection = self - .connection - .lock() - .map_err(|_| SessionError::Unavailable)?; - let generation = self.connection_generation.load(Ordering::Acquire); - let stream = connection - .try_clone_stream() - .map_err(SessionError::Connection)?; - Ok((stream, generation)) - } - - pub(crate) fn try_read_packet(&self) -> Result, SessionError> { - self.connection - .lock() - .map_err(|_| SessionError::Unavailable)? - .try_read_packet() - .map_err(SessionError::Connection) - } - - pub(crate) fn note_bridge_generation(&self, generation: u64) -> Result<(), SessionError> { - let mut observed = self - .bridge_generation - .lock() - .map_err(|_| SessionError::Unavailable)?; - if generation > *observed { - *observed = generation; - self.bridge_changed.notify_all(); - } - Ok(()) - } - - pub(crate) fn wait_for_bridge_generation( - &self, - expected: u64, - timeout: Duration, - ) -> Result<(), SessionError> { - let deadline = Instant::now() - .checked_add(timeout) - .ok_or(SessionError::Unavailable)?; - let mut observed = self - .bridge_generation - .lock() - .map_err(|_| SessionError::Unavailable)?; - while *observed < expected { - let remaining = deadline - .checked_duration_since(Instant::now()) - .ok_or(SessionError::RecoverBusy)?; - let (next, result) = self - .bridge_changed - .wait_timeout(observed, remaining) - .map_err(|_| SessionError::Unavailable)?; - observed = next; - if result.timed_out() && *observed < expected { - return Err(SessionError::RecoverBusy); + fn join_flow_writer(&self, graceful: bool) -> Result<(), SessionError> { + if let Some(state) = &self.flow_control { + if graceful { + state.stop_gracefully(); + } else { + state.stop_hard(); } } - Ok(()) - } - - pub(crate) fn connection_state(&self) -> Result<(bool, u64), SessionError> { - let connection = self - .connection - .lock() - .map_err(|_| SessionError::Unavailable)?; - Ok(( - connection.connected(), - self.connection_generation.load(Ordering::Acquire), - )) - } - - /// Soft-drop the encrypted client transport without killing the terminal. - /// - /// Used when the client TCP path dies (sleep, Wi-Fi, NAT) so terminal - /// output keeps buffering and a returning client can recover the same - /// session. Does not set the session shutdown flag or close the terminal. - pub(crate) fn mark_client_disconnected( - &self, - expected_generation: u64, - ) -> Result { - let mut connection = self - .connection - .lock() - .map_err(|_| SessionError::Unavailable)?; - if self.connection_generation.load(Ordering::Acquire) != expected_generation { - return Ok(false); - } - connection.disconnect(); - Ok(true) - } - - /// Apply a client delivery acknowledgement to the replay backup. - pub(crate) fn acknowledge_delivery(&self, sequence: i64) -> Result<(), SessionError> { - self.connection + let handle = self + .flow_writer .lock() .map_err(|_| SessionError::Unavailable)? - .acknowledge_delivery(sequence); - Ok(()) - } - - /// Keep-alive payload acknowledging everything read from the client. - pub(crate) fn keepalive_ack( - &self, - ) -> Result<[u8; et_core::keepalive::ACK_PAYLOAD_LEN], SessionError> { - Ok(self - .connection - .lock() - .map_err(|_| SessionError::Unavailable)? - .keepalive_ack()) - } - - pub(crate) fn can_buffer_write(&self, bytes: i64) -> Result { - let hold = self - .recover_hold - .lock() - .map_err(|_| SessionError::Unavailable)?; - if hold.len() >= MAX_BACKUP_PACKETS + MAX_DISCONNECT_PACKETS { - return Ok(false); + .take(); + if handle.is_some_and(|handle| handle.join().is_err()) { + return Err(SessionError::Unavailable); } - let held = - i64::try_from(self.recover_hold_bytes.load(Ordering::Acquire)).unwrap_or(i64::MAX); - let requested = held.checked_add(bytes).unwrap_or(i64::MAX); - if requested > MAX_RECOVERY_BACKUP_BYTES { - return Ok(false); + if graceful + && self + .flow_control + .as_ref() + .is_some_and(|state| state.unrecoverable()) + { + return Err(SessionError::Connection(ConnError::Io(io::Error::new( + io::ErrorKind::BrokenPipe, + "terminal ended before retained flow output could be delivered", + )))); } - Ok(self - .connection - .lock() - .map_err(|_| SessionError::Unavailable)? - .can_buffer_write(requested)) - } - - pub(crate) fn is_shutting_down(&self) -> bool { - self.shutdown.load(Ordering::Acquire) - } - - fn signal(&self) -> Result<(), SessionError> { - self.wake_writer - .lock() - .map_err(|_| SessionError::Unavailable)? - .write_all(&[1]) - .map_err(SessionError::Io) + Ok(()) } -} -/// Single-flight recover permit. Dropping it (normally or on panic) always -/// clears [`ActiveSession::recovering`] and wakes the terminal bridge. -pub(crate) struct RecoverPermit<'a> { - session: &'a ActiveSession, -} - -impl RecoverPermit<'_> { - /// Run the recovery handshake and install the new stream. - pub(crate) fn complete(self, stream: TcpStream) -> Result<(), SessionError> { - // `self` drops after this returns (or panics), clearing `recovering` - // and flushing any straggler hold packets. - self.session.recover_body(stream) + fn stop_flow_writer(&self) { + if let Some(state) = &self.flow_control { + state.stop_hard(); + } } } -impl Drop for RecoverPermit<'_> { +impl Drop for ActiveSession { fn drop(&mut self) { - // Flush while still marked recovering so send_packet keeps queuing - // rather than racing into a half-installed connection. - let _ = self.session.flush_recover_hold(); - self.session.recovering.store(false, Ordering::Release); - // Catch anything that observed `recovering` and queued after the first - // flush but before the flag cleared (re-check is under the hold lock). - let _ = self.session.flush_recover_hold(); - // Wake the bridge even on failure so it re-checks connection state. - let _ = self.session.signal(); + self.stop_flow_writer(); } } -/// Acquire a [`Mutex`] with a deadline so recover cannot park forever behind a -/// bridge thread blocked in a live write. -fn lock_timeout(mutex: &Mutex, timeout: Duration) -> Result, SessionError> { - let deadline = Instant::now() - .checked_add(timeout) - .ok_or(SessionError::RecoverBusy)?; - loop { - match mutex.try_lock() { - Ok(guard) => return Ok(guard), - Err(TryLockError::Poisoned(_)) => return Err(SessionError::Unavailable), - Err(TryLockError::WouldBlock) => { - if Instant::now() >= deadline { - return Err(SessionError::RecoverBusy); - } - std::thread::sleep(Duration::from_millis(5)); - } - } +fn queue_mode(value: Option) -> Option { + match value.and_then(|value| et_core::proto::FlowControlMode::try_from(value).ok()) { + None | Some(et_core::proto::FlowControlMode::None) => None, + Some(et_core::proto::FlowControlMode::Backpressure) => Some(QueueMode::Backpressure), + Some(et_core::proto::FlowControlMode::Discard) => Some(QueueMode::Discard), } } @@ -545,7 +409,7 @@ mod tests { let _peer = peer.join().unwrap(); let (terminal, _terminal_peer) = et_net::local::wake_pair().unwrap(); let session = - ActiveSession::new(Connection::new_server(stream, &[7; 32]), &terminal).unwrap(); + ActiveSession::new(Connection::new_server(stream, &[7; 32]), &terminal, None).unwrap(); let permit = session.try_begin_recover().unwrap(); for index in 0..4u8 { diff --git a/crates/et-server/src/session_flow.rs b/crates/et-server/src/session_flow.rs new file mode 100644 index 0000000..b329fdb --- /dev/null +++ b/crates/et-server/src/session_flow.rs @@ -0,0 +1,290 @@ +use std::sync::{Arc, Condvar, Mutex, Weak}; + +use et_core::flow_control::{FlowControlMode, OutputQueue, QueuePushError}; +use et_core::packet::Packet; +use et_net::connection::ConnError; + +use super::{ActiveSession, SessionError, FLOW_CONTROL_BUFFER_BYTES}; + +#[cfg(test)] +#[path = "session_flow_hook.rs"] +mod test_hook; +#[path = "session_flow_write.rs"] +pub(super) mod writer; +#[cfg(test)] +pub(super) use writer::FlowWriteResult; + +#[derive(Clone, Copy, PartialEq, Eq)] +enum StopMode { + Running, + Graceful, + Hard, +} + +struct WriterState { + queue: OutputQueue, + connected: bool, + paused: bool, + in_flight: bool, + reader_waiting: bool, + stop: StopMode, + unrecoverable: bool, +} + +pub(super) struct FlowControl { + state: Mutex, + wake: Condvar, + #[cfg(test)] + enqueue_hook: Mutex>>, +} + +impl FlowControl { + pub(super) fn new(mode: FlowControlMode) -> Self { + Self { + state: Mutex::new(WriterState { + queue: OutputQueue::new(mode, FLOW_CONTROL_BUFFER_BYTES), + connected: true, + paused: false, + in_flight: false, + reader_waiting: false, + stop: StopMode::Running, + unrecoverable: false, + }), + wake: Condvar::new(), + #[cfg(test)] + enqueue_hook: Mutex::new(None), + } + } + + pub(super) fn enqueue(&self, packet: Packet) -> Result<(), SessionError> { + let mut state = self.state.lock().map_err(|_| SessionError::Unavailable)?; + if state.stop != StopMode::Running { + return Err(SessionError::Unavailable); + } + match state.queue.push(packet) { + Ok(()) => { + #[cfg(test)] + self.run_enqueue_hook(); + drop(state); + self.wake.notify_one(); + Ok(()) + } + Err(QueuePushError::Full(_packet)) => { + Err(SessionError::Connection(ConnError::Backpressure)) + } + Err(QueuePushError::Oversized(_packet)) => { + Err(SessionError::Connection(ConnError::PacketTooLarge)) + } + } + } + + pub(super) fn can_accept_terminal(&self, bytes: usize) -> Result { + let state = self.state.lock().map_err(|_| SessionError::Unavailable)?; + Ok(state.stop == StopMode::Running && state.queue.can_accept_terminal(bytes)) + } + + pub(super) fn pause(&self) -> Result<(), SessionError> { + let mut state = self.state.lock().map_err(|_| SessionError::Unavailable)?; + state.paused = true; + let state = match self.wake.wait_while(state, |state| state.in_flight) { + Ok(state) => state, + Err(error) => { + error.into_inner().paused = false; + self.wake.notify_all(); + return Err(SessionError::Unavailable); + } + }; + drop(state); + Ok(()) + } + + pub(super) fn resume(&self, connected: bool) { + if let Ok(mut state) = self.state.lock() { + state.connected = connected; + state.paused = false; + if !connected && state.stop == StopMode::Graceful { + state.unrecoverable = true; + state.stop = StopMode::Hard; + } + self.wake.notify_all(); + } + } + + pub(super) fn disconnected(&self) { + if let Ok(mut state) = self.state.lock() { + state.connected = false; + self.wake.notify_all(); + } + } + + pub(super) fn set_reader_waiting(&self, waiting: bool) { + if let Ok(mut state) = self.state.lock() { + state.reader_waiting = waiting; + self.wake.notify_all(); + } + } + + fn wait_for_reader(&self) -> bool { + let Ok(state) = self.state.lock() else { + return false; + }; + let Ok(state) = self.wake.wait_while(state, |state| { + state.reader_waiting && state.stop != StopMode::Hard + }) else { + return false; + }; + state.stop != StopMode::Hard + } + + pub(super) fn stop_gracefully(&self) { + if let Ok(mut state) = self.state.lock() { + // A recovery pause owns the connection snapshot until its permit + // installs the candidate (or safely abandons it). Terminal EOF + // must not bypass that pause and drain queued output onto the old + // stream; RecoverPermit::drop resumes the writer atomically. + if state.stop == StopMode::Running { + if state.connected || state.paused { + state.stop = StopMode::Graceful; + } else { + state.unrecoverable = true; + state.stop = StopMode::Hard; + } + self.wake.notify_all(); + } + } + } + + pub(super) fn unrecoverable(&self) -> bool { + self.state.lock().map_or(true, |state| state.unrecoverable) + } + + pub(super) fn stop_hard(&self) { + if let Ok(mut state) = self.state.lock() { + state.stop = StopMode::Hard; + self.wake.notify_all(); + } + } + + #[cfg(test)] + pub(super) fn wait_in_flight(&self) { + let state = self.state.lock().unwrap(); + drop( + self.wake + .wait_while(state, |state| !state.in_flight) + .unwrap(), + ); + } + + #[cfg(test)] + pub(super) fn wait_for_stop(&self, graceful: bool) { + let state = self.state.lock().unwrap(); + drop( + self.wake + .wait_while(state, |state| { + state.stop + != if graceful { + StopMode::Graceful + } else { + StopMode::Hard + } + }) + .unwrap(), + ); + } + + fn is_hard_stopped(&self) -> bool { + self.state + .lock() + .map_or(true, |state| state.stop == StopMode::Hard) + } + + pub(super) fn next_packet(&self) -> Option { + let state = self.state.lock().ok()?; + let mut state = self + .wake + .wait_while(state, |state| match state.stop { + StopMode::Hard => false, + StopMode::Graceful => { + state.paused + || (!state.connected && !state.queue.is_empty()) + || (state.queue.is_empty() && state.in_flight) + } + StopMode::Running => state.paused || !state.connected || state.queue.is_empty(), + }) + .ok()?; + match state.stop { + StopMode::Hard => None, + StopMode::Graceful if state.queue.is_empty() => None, + StopMode::Running | StopMode::Graceful => { + let packet = state.queue.take()?; + state.in_flight = true; + self.wake.notify_all(); + Some(packet) + } + } + } + + pub(super) fn complete( + &self, + packet: Packet, + result: &writer::FlowWriteResult, + connected: bool, + ) -> bool { + let Ok(mut state) = self.state.lock() else { + return false; + }; + state.in_flight = false; + state.connected = connected; + match result { + writer::FlowWriteResult::Delivered => { + state.queue.complete(&packet); + if !connected && state.stop == StopMode::Graceful { + state.unrecoverable = true; + state.stop = StopMode::Hard; + } + } + writer::FlowWriteResult::BeforeReplay(_error) => { + state.queue.restore_front(packet); + state.connected = false; + if state.stop == StopMode::Graceful { + state.unrecoverable = true; + state.stop = StopMode::Hard; + } + } + writer::FlowWriteResult::ReplayOwned(_error) => { + state.queue.complete(&packet); + state.connected = false; + if state.stop == StopMode::Graceful { + state.unrecoverable = true; + state.stop = StopMode::Hard; + } + } + writer::FlowWriteResult::Fatal(error) => { + state.queue.complete(&packet); + crate::diag::info(format!("flow-control writer stopped: {error}")); + state.stop = StopMode::Hard; + } + } + self.wake.notify_all(); + state.stop != StopMode::Hard + } +} + +pub(super) fn run_writer(session: Weak, flow: Arc) { + while let Some(packet) = flow.next_packet() { + if !flow.wait_for_reader() { + return; + } + let Some(session) = session.upgrade() else { + return; + }; + // Popping this packet may have reopened bounded queue capacity. + // Wake the bridge so it polls terminal output again instead of + // sleeping indefinitely with terminal readability disabled. + let _ = session.signal(); + let (result, connected) = writer::write_packet(&session, &flow, &packet); + if !flow.complete(packet, &result, connected) { + return; + } + } +} diff --git a/crates/et-server/src/session_flow_hook.rs b/crates/et-server/src/session_flow_hook.rs new file mode 100644 index 0000000..8156d6f --- /dev/null +++ b/crates/et-server/src/session_flow_hook.rs @@ -0,0 +1,13 @@ +use super::FlowControl; + +impl FlowControl { + pub(crate) fn install_enqueue_hook(&self, reached: std::sync::mpsc::SyncSender<()>) { + *self.enqueue_hook.lock().unwrap() = Some(reached); + } + + pub(super) fn run_enqueue_hook(&self) { + if let Some(reached) = self.enqueue_hook.lock().unwrap().take() { + reached.send(()).unwrap(); + } + } +} diff --git a/crates/et-server/src/session_flow_test.rs b/crates/et-server/src/session_flow_test.rs new file mode 100644 index 0000000..00f8b9b --- /dev/null +++ b/crates/et-server/src/session_flow_test.rs @@ -0,0 +1,632 @@ +#![cfg(test)] + +use std::net::{Ipv4Addr, TcpListener, TcpStream}; +use std::sync::{mpsc, Arc}; +use std::thread; +use std::time::Duration; + +use et_core::proto::FlowControlMode; +use et_net::connection::{ConnError, Connection, WritePacketError}; +use prost::Message; + +use super::{ + session_flow::{FlowControl, FlowWriteResult}, + ActiveSession, SessionError, SessionWriteError, +}; + +const TEST_TIMEOUT: Duration = Duration::from_secs(3); + +fn connection_pair() -> (Connection, Connection) { + let listener = TcpListener::bind((Ipv4Addr::LOCALHOST, 0)).unwrap(); + let address = listener.local_addr().unwrap(); + let client = thread::spawn(move || TcpStream::connect(address).unwrap()); + let (server, _) = listener.accept().unwrap(); + let client = client.join().unwrap(); + let key = [7u8; 32]; + ( + Connection::new_server(server, &key), + Connection::new_client(client, &key), + ) +} + +#[test] +fn default_none_preserves_write_ownership_for_exact_once_retry() { + let (server, mut client) = connection_pair(); + let (terminal, _terminal_peer) = et_net::local::wake_pair().unwrap(); + let session = ActiveSession::new(server, &terminal, None).unwrap(); + let mut writes = 0; + + let error = session + .send_packet_owned_with(46, b"exactly-once", |_, _, _| { + writes += 1; + Err(WritePacketError::BeforeReplay(ConnError::Io( + std::io::Error::other("injected clone failure"), + ))) + }) + .unwrap_err(); + assert!(matches!(error, SessionWriteError::BeforeReplay(_))); + assert_eq!(writes, 1); + + session.send_packet_owned(46, b"exactly-once").unwrap(); + let packet = client.read_packet().unwrap(); + assert_eq!( + (packet.header(), packet.payload()), + (46, b"exactly-once".as_slice()) + ); + + let error = session + .send_packet_owned_with(47, b"replay-owned", |_, _, _| { + writes += 1; + Err(WritePacketError::ReplayOwned(ConnError::Io( + std::io::Error::other("injected post-admission failure"), + ))) + }) + .unwrap_err(); + assert!(matches!(error, SessionWriteError::ReplayOwned(_))); + assert_eq!(writes, 2); +} + +#[test] +fn graceful_hup_drains_and_joins_a_deliberately_blocked_writer() { + for mode in [FlowControlMode::Backpressure, FlowControlMode::Discard] { + // Given: a flow writer has removed the final packet from its queue but + // is deterministically blocked before taking the connection lock. + let (server, mut client) = connection_pair(); + let (terminal, _terminal_peer) = et_net::local::wake_pair().unwrap(); + let session = Arc::new(ActiveSession::new(server, &terminal, Some(mode as i32)).unwrap()); + session.start_flow_writer(); + let connection = session.connection.lock().unwrap(); + session.send_packet(47, b"final-packet").unwrap(); + let flow = session.flow_control.as_ref().unwrap(); + flow.wait_in_flight(); + + // When: terminal HUP requests graceful completion while the write is blocked. + let (finished_tx, finished_rx) = mpsc::sync_channel(0); + let finishing = Arc::clone(&session); + let worker = thread::spawn(move || finished_tx.send(finishing.finish_terminal()).unwrap()); + flow.wait_for_stop(true); + drop(connection); + + // Then: the retained packet arrives before the joined writer permits half-close. + assert!(finished_rx.recv_timeout(TEST_TIMEOUT).unwrap().is_ok()); + worker.join().unwrap(); + let packet = client.read_packet().unwrap(); + assert_eq!( + (packet.header(), packet.payload()), + (47, b"final-packet".as_slice()) + ); + } +} + +#[test] +fn terminal_finish_fails_and_joins_after_unrecoverable_before_replay() { + let (server, mut client) = connection_pair(); + client.set_io_timeout(Some(TEST_TIMEOUT)).unwrap(); + let (terminal, _terminal_peer) = et_net::local::wake_pair().unwrap(); + let session = Arc::new( + ActiveSession::new( + server, + &terminal, + Some(FlowControlMode::Backpressure as i32), + ) + .unwrap(), + ); + let flow = Arc::clone(session.flow_control.as_ref().unwrap()); + flow.enqueue(et_core::packet::Packet::new(55, b"retained".as_slice())) + .unwrap(); + let (attempt_tx, attempt_rx) = mpsc::sync_channel(0); + let (release_tx, release_rx) = mpsc::sync_channel(0); + let worker_session = Arc::clone(&session); + let worker_flow = Arc::clone(&flow); + let handle = thread::spawn(move || { + while let Some(packet) = worker_flow.next_packet() { + attempt_tx.send(()).unwrap(); + release_rx.recv().unwrap(); + let (result, connected) = super::session_flow::writer::write_packet_with( + &worker_session, + &worker_flow, + &packet, + |connection, packet| { + connection.prepare_write_packet_with(packet.header(), packet.payload(), |_| { + Err(std::io::Error::other("injected clone failure")) + }) + }, + ); + if !worker_flow.complete(packet, &result, connected) { + break; + } + } + }); + *session.flow_writer.lock().unwrap() = Some(handle); + attempt_rx.recv().unwrap(); + + let (finished_tx, finished_rx) = mpsc::sync_channel(0); + let finishing = Arc::clone(&session); + let finisher = thread::spawn(move || finished_tx.send(finishing.finish_terminal()).unwrap()); + flow.wait_for_stop(true); + release_tx.send(()).unwrap(); + + let error = finished_rx.recv_timeout(TEST_TIMEOUT).unwrap().unwrap_err(); + assert!(matches!(error, SessionError::Connection(ConnError::Io(_)))); + finisher.join().unwrap(); + assert!(attempt_rx.try_recv().is_err()); + assert!(session.flow_writer.lock().unwrap().is_none()); + assert!(client.read_packet().is_err()); +} + +#[test] +fn hard_shutdown_wakes_and_joins_a_deliberately_blocked_writer() { + // Given: a writer blocked after taking ownership of a queued packet. + let (server, _client) = connection_pair(); + let (terminal, _terminal_peer) = et_net::local::wake_pair().unwrap(); + let session = Arc::new( + ActiveSession::new(server, &terminal, Some(FlowControlMode::Discard as i32)).unwrap(), + ); + session.start_flow_writer(); + let connection = session.connection.lock().unwrap(); + session.send_packet(48, b"discardable").unwrap(); + let flow = session.flow_control.as_ref().unwrap(); + flow.wait_in_flight(); + + // When: hard shutdown is requested and wakes the worker. + let (finished_tx, finished_rx) = mpsc::sync_channel(0); + let stopping = Arc::clone(&session); + let worker = thread::spawn(move || finished_tx.send(stopping.shutdown()).unwrap()); + flow.wait_for_stop(false); + drop(connection); + + // Then: shutdown returns only after the writer has joined. + assert!(finished_rx.recv_timeout(TEST_TIMEOUT).unwrap().is_ok()); + worker.join().unwrap(); +} + +#[test] +fn before_replay_failure_waits_for_resume_then_delivers_once() { + // Given: the flow writer owns one packet and socket cloning fails before replay admission. + let (server, _client) = connection_pair(); + let (terminal, _terminal_peer) = et_net::local::wake_pair().unwrap(); + let session = ActiveSession::new( + server, + &terminal, + Some(FlowControlMode::Backpressure as i32), + ) + .unwrap(); + let flow = session.flow_control.as_ref().unwrap(); + flow.enqueue(et_core::packet::Packet::new(40, b"clone-retry".as_slice())) + .unwrap(); + let packet = flow.next_packet().unwrap(); + + // When: preparation injects a clone failure, then normal preparation retries it. + let (failed, connected) = super::session_flow::writer::write_packet_with( + &session, + flow, + &packet, + |connection, packet| { + connection.prepare_write_packet_with(packet.header(), packet.payload(), |_| { + Err(std::io::Error::other("injected clone failure")) + }) + }, + ); + assert!(matches!(failed, FlowWriteResult::BeforeReplay(_))); + assert!(!connected); + assert!(!session.connection.lock().unwrap().connected()); + assert!(flow.complete(packet, &failed, connected)); + let (restored_tx, restored_rx) = mpsc::sync_channel(0); + let flow_waiter = Arc::clone(session.flow_control.as_ref().unwrap()); + let waiting = thread::spawn(move || restored_tx.send(flow_waiter.next_packet()).unwrap()); + assert!(restored_rx.try_recv().is_err()); + let (recovered_server, mut recovered_client) = connection_pair(); + recovered_client.set_io_timeout(Some(TEST_TIMEOUT)).unwrap(); + *session.connection.lock().unwrap() = recovered_server; + flow.resume(true); + let restored = restored_rx.recv_timeout(TEST_TIMEOUT).unwrap().unwrap(); + waiting.join().unwrap(); + let (delivered, connected) = + super::session_flow::writer::write_packet(&session, flow, &restored); + assert!(flow.complete(restored, &delivered, connected)); + + // Then: the restored plaintext is encrypted under one sequence and delivered once. + let received = recovered_client.read_packet().unwrap(); + assert_eq!( + (received.header(), received.payload()), + (40, b"clone-retry".as_slice()) + ); + recovered_client + .set_io_timeout(Some(Duration::from_millis(50))) + .unwrap(); + assert!(recovered_client.read_packet().is_err()); +} + +#[test] +fn non_transport_before_replay_is_fatal_instead_of_pausing_as_disconnected() { + let (server, _client) = connection_pair(); + let (terminal, _terminal_peer) = et_net::local::wake_pair().unwrap(); + let session = ActiveSession::new( + server, + &terminal, + Some(FlowControlMode::Backpressure as i32), + ) + .unwrap(); + let flow = session.flow_control.as_ref().unwrap(); + flow.enqueue(et_core::packet::Packet::new( + 53, + b"semantic-failure".as_slice(), + )) + .unwrap(); + let packet = flow.next_packet().unwrap(); + let (result, connected) = + super::session_flow::writer::write_packet_with(&session, flow, &packet, |_, _| { + Err(WritePacketError::BeforeReplay(ConnError::Backpressure)) + }); + + assert!(matches!(result, FlowWriteResult::Fatal(_))); + assert!(connected); + assert!(!flow.complete(packet, &result, connected)); +} + +#[test] +fn before_replay_waits_for_explicit_resume_while_session_is_recoverable() { + let (server, _client) = connection_pair(); + let (terminal, _terminal_peer) = et_net::local::wake_pair().unwrap(); + let session = Arc::new( + ActiveSession::new( + server, + &terminal, + Some(FlowControlMode::Backpressure as i32), + ) + .unwrap(), + ); + let flow = Arc::clone(session.flow_control.as_ref().unwrap()); + flow.enqueue(et_core::packet::Packet::new( + 54, + b"graceful-retry".as_slice(), + )) + .unwrap(); + let (attempt_tx, attempt_rx) = mpsc::sync_channel(0); + let (result_tx, result_rx) = mpsc::sync_channel(0); + let (completed_tx, completed_rx) = mpsc::sync_channel(0); + let (done_tx, done_rx) = mpsc::sync_channel(0); + let worker_session = Arc::clone(&session); + let worker_flow = Arc::clone(&flow); + let worker = thread::spawn(move || { + let mut attempt = 0; + while let Some(packet) = worker_flow.next_packet() { + attempt += 1; + attempt_tx.send(attempt).unwrap(); + let fail = result_rx.recv().unwrap(); + let (result, connected) = match (attempt, fail) { + (1, true) => super::session_flow::writer::write_packet_with( + &worker_session, + &worker_flow, + &packet, + |connection, packet| { + connection.prepare_write_packet_with( + packet.header(), + packet.payload(), + |_| Err(std::io::Error::other("persistent clone failure")), + ) + }, + ), + (_, true) => ( + FlowWriteResult::BeforeReplay(SessionError::Connection(ConnError::Io( + std::io::Error::other("persistent clone failure"), + ))), + false, + ), + (_, false) => (FlowWriteResult::Delivered, true), + }; + if !worker_flow.complete(packet, &result, connected) { + break; + } + completed_tx.send(attempt).unwrap(); + } + done_tx.send(()).unwrap(); + }); + + assert_eq!(attempt_rx.recv().unwrap(), 1); + result_tx.send(true).unwrap(); + assert_eq!(completed_rx.recv().unwrap(), 1); + assert!(!session.connection.lock().unwrap().connected()); + assert!(attempt_rx.try_recv().is_err()); + assert!(done_rx.try_recv().is_err()); + + flow.resume(true); + assert_eq!(attempt_rx.recv().unwrap(), 2); + result_tx.send(true).unwrap(); + assert_eq!(completed_rx.recv().unwrap(), 2); + assert!(attempt_rx.try_recv().is_err()); + assert!(done_rx.try_recv().is_err()); + + // Installing/authorizing a usable transport and resuming permits exactly + // one final delivery, after which graceful join can complete. + flow.resume(true); + assert_eq!(attempt_rx.recv().unwrap(), 3); + result_tx.send(false).unwrap(); + assert_eq!(completed_rx.recv().unwrap(), 3); + flow.stop_gracefully(); + done_rx.recv().unwrap(); + worker.join().unwrap(); +} + +#[test] +fn disconnected_recovery_pause_transfers_graceful_output_to_candidate() { + // Given: the old transport is disconnected before recovery pauses the writer. + let flow = FlowControl::new(et_core::flow_control::FlowControlMode::Backpressure); + flow.disconnected(); + flow.pause().unwrap(); + flow.enqueue(et_core::packet::Packet::new(55, b"candidate".as_slice())) + .unwrap(); + + // When: terminal EOF requests graceful completion and recovery installs a candidate. + flow.stop_gracefully(); + flow.resume(true); + + // Then: the candidate owns and drains the final packet exactly once. + let packet = flow.next_packet().unwrap(); + assert_eq!( + (packet.header(), packet.payload()), + (55, b"candidate".as_slice()) + ); + assert!(flow.complete(packet, &FlowWriteResult::Delivered, true)); + assert!(flow.next_packet().is_none()); + assert!(!flow.unrecoverable()); +} + +#[test] +fn buffered_only_admission_cannot_complete_graceful_session() { + // Given: graceful completion owns one final packet. + let flow = FlowControl::new(et_core::flow_control::FlowControlMode::Backpressure); + flow.enqueue(et_core::packet::Packet::new( + 57, + b"buffered-only".as_slice(), + )) + .unwrap(); + flow.stop_gracefully(); + let packet = flow.next_packet().unwrap(); + + // When: replay admits it but no live transport receives it. + assert!(!flow.complete(packet, &FlowWriteResult::Delivered, false)); + + // Then: removing the session is a typed failure, not false delivery. + assert!(flow.unrecoverable()); +} + +#[test] +fn disconnected_recovery_pause_fails_gracefully_when_candidate_fails() { + // Given: graceful EOF is waiting on a candidate for a disconnected transport. + let flow = FlowControl::new(et_core::flow_control::FlowControlMode::Backpressure); + flow.disconnected(); + flow.pause().unwrap(); + flow.enqueue(et_core::packet::Packet::new(56, b"undelivered".as_slice())) + .unwrap(); + flow.stop_gracefully(); + + // When: candidate recovery fails and resumes without a live transport. + flow.resume(false); + + // Then: completion is bounded and reported as unrecoverable. + assert!(flow.next_packet().is_none()); + assert!(flow.unrecoverable()); +} + +#[test] +fn live_send_reset_disconnects_without_requeueing_replay_owned_packet() { + // Given: one in-flight packet followed by queued output. + let flow = FlowControl::new(et_core::flow_control::FlowControlMode::Backpressure); + flow.enqueue(et_core::packet::Packet::new(41, b"in-flight".as_slice())) + .unwrap(); + flow.enqueue(et_core::packet::Packet::new(42, b"queued".as_slice())) + .unwrap(); + let in_flight = flow.next_packet().unwrap(); + + // When: the live socket resets after replay accepted the packet. + let error = FlowWriteResult::ReplayOwned(SessionError::Connection( + et_net::connection::ConnError::Io(std::io::ErrorKind::ConnectionReset.into()), + )); + assert!(flow.complete(in_flight, &error, false)); + flow.resume(true); + + // Then: recovery sends only the still-queued packet; replay owns the first. + let queued = flow.next_packet().unwrap(); + assert_eq!( + (queued.header(), queued.payload()), + (42, b"queued".as_slice()) + ); +} + +#[test] +fn live_send_reset_recovers_replay_then_sends_queued_and_subsequent_once() { + // Given: replay has accepted one packet when its live socket resets, with + // another packet still reserved in the flow queue. + let (mut sender, mut receiver) = connection_pair(); + let flow = FlowControl::new(et_core::flow_control::FlowControlMode::Backpressure); + flow.enqueue(et_core::packet::Packet::new(51, b"replay".as_slice())) + .unwrap(); + flow.enqueue(et_core::packet::Packet::new(52, b"queued".as_slice())) + .unwrap(); + let replay = flow.next_packet().unwrap(); + let prepared = sender + .prepare_write_packet(replay.header(), replay.payload()) + .unwrap(); + drop(prepared); + sender.disconnect(); + let reset = FlowWriteResult::ReplayOwned(SessionError::Connection( + et_net::connection::ConnError::Io(std::io::ErrorKind::ConnectionReset.into()), + )); + assert!(flow.complete(replay, &reset, false)); + + // When: both peers recover on a replacement socket. + let listener = TcpListener::bind((Ipv4Addr::LOCALHOST, 0)).unwrap(); + let address = listener.local_addr().unwrap(); + let connector = thread::spawn(move || TcpStream::connect(address).unwrap()); + let (new_sender, _) = listener.accept().unwrap(); + let new_receiver = connector.join().unwrap(); + let (receiver_tx, receiver_rx) = mpsc::sync_channel(0); + let recovering_receiver = thread::spawn(move || { + receiver.recover(new_receiver).unwrap(); + receiver_tx.send(receiver).unwrap(); + }); + sender.recover(new_sender).unwrap(); + let mut receiver = receiver_rx.recv_timeout(TEST_TIMEOUT).unwrap(); + recovering_receiver.join().unwrap(); + flow.resume(true); + + // Then: replay delivers the failed live packet, while the still-queued + // packet and output produced after recovery each follow exactly once. + for (header, payload) in [(52, b"queued".as_slice()), (53, b"subsequent".as_slice())] { + if header == 53 { + flow.enqueue(et_core::packet::Packet::new(header, payload)) + .unwrap(); + } + let packet = flow.next_packet().unwrap(); + sender + .write_packet(packet.header(), packet.payload()) + .unwrap(); + assert!(flow.complete(packet, &FlowWriteResult::Delivered, sender.connected())); + } + for expected in [ + (51, b"replay".as_slice()), + (52, b"queued".as_slice()), + (53, b"subsequent".as_slice()), + ] { + let packet = receiver.read_packet().unwrap(); + assert_eq!((packet.header(), packet.payload()), expected); + } + receiver + .set_io_timeout(Some(Duration::from_millis(50))) + .unwrap(); + assert!( + receiver.read_packet().is_err(), + "recovered output was duplicated" + ); +} + +#[test] +fn recover_hold_post_admission_failure_replays_without_plaintext_duplicate() { + // Given: post-install held plaintext is accepted by replay, then live send fails. + let (server, mut client) = connection_pair(); + let (terminal, _terminal_peer) = et_net::local::wake_pair().unwrap(); + let session = ActiveSession::new(server, &terminal, None).unwrap(); + session + .recover_hold + .lock() + .unwrap() + .push((54, b"held-once".to_vec())); + let result = session.flush_recover_hold_with(|connection, header, payload| { + let prepared = connection.prepare_write_packet(header, payload)?; + drop(prepared); + connection.disconnect(); + Err(et_net::connection::WritePacketError::ReplayOwned( + et_net::connection::ConnError::Io(std::io::ErrorKind::ConnectionReset.into()), + )) + }); + assert!(result.is_err()); + assert!(session.recover_hold.lock().unwrap().is_empty()); + + // When: the installed connection recovers again. + let listener = TcpListener::bind((Ipv4Addr::LOCALHOST, 0)).unwrap(); + let address = listener.local_addr().unwrap(); + let connector = thread::spawn(move || TcpStream::connect(address).unwrap()); + let (new_server, _) = listener.accept().unwrap(); + let new_client = connector.join().unwrap(); + let (client_tx, client_rx) = mpsc::sync_channel(0); + let recovering_client = thread::spawn(move || { + client.recover(new_client).unwrap(); + client_tx.send(client).unwrap(); + }); + session + .connection + .lock() + .unwrap() + .recover(new_server) + .unwrap(); + let mut client = client_rx.recv_timeout(TEST_TIMEOUT).unwrap(); + recovering_client.join().unwrap(); + + // Then: replay supplies the packet exactly once; no plaintext duplicate was retained. + let packet = client.read_packet().unwrap(); + assert_eq!( + (packet.header(), packet.payload()), + (54, b"held-once".as_slice()) + ); + client + .set_io_timeout(Some(Duration::from_millis(50))) + .unwrap(); + assert!(client.read_packet().is_err()); +} + +#[test] +fn full_control_lane_rejects_nonblocking_while_discard_terminal_and_recovery_progress() { + // Given: discard mode's lossless control lane is full. + let flow = FlowControl::new(et_core::flow_control::FlowControlMode::Discard); + let control = et_core::packet::Packet::new(61, vec![1; 1024]); + loop { + match flow.enqueue(control.clone()) { + Ok(()) => {} + Err(SessionError::Connection(et_net::connection::ConnError::Backpressure)) => break, + Err(error) => panic!("unexpected control admission: {error}"), + } + } + + // When: terminal output arrives and a recovery pause/resume completes. + let terminal = et_core::packet::Packet::new( + et_core::proto::TerminalPacketType::TerminalBuffer as u8, + et_core::proto::TerminalBuffer { + buffer: Some(b"newest".to_vec()), + } + .encode_to_vec(), + ); + flow.enqueue(terminal).unwrap(); + flow.pause().unwrap(); + flow.resume(true); + + // Then: the bridge-facing full result was immediate and terminal remains serviceable. + let next = flow.next_packet().unwrap(); + assert_eq!( + next.header(), + et_core::proto::TerminalPacketType::TerminalBuffer as u8 + ); +} + +#[test] +fn oversized_control_fails_permanently_without_waiting() { + let flow = FlowControl::new(et_core::flow_control::FlowControlMode::Backpressure); + let oversized = et_core::packet::Packet::new(62, vec![0; 64 * 1024]); + + assert!(matches!( + flow.enqueue(oversized), + Err(SessionError::Connection( + et_net::connection::ConnError::PacketTooLarge + )) + )); +} + +#[test] +fn concurrent_hard_stop_cannot_be_downgraded_to_graceful_drain() { + // Given: a packet removed from the queue while transport preparation is blocked. + let (server, mut client) = connection_pair(); + client.set_io_timeout(Some(TEST_TIMEOUT)).unwrap(); + let (terminal, _terminal_peer) = et_net::local::wake_pair().unwrap(); + let session = Arc::new( + ActiveSession::new( + server, + &terminal, + Some(FlowControlMode::Backpressure as i32), + ) + .unwrap(), + ); + session.start_flow_writer(); + let connection = session.connection.lock().unwrap(); + session.send_packet(49, b"must-not-drain").unwrap(); + let flow = session.flow_control.as_ref().unwrap(); + flow.wait_in_flight(); + + // When: hard shutdown wins before a concurrent terminal HUP asks for graceful stop. + flow.stop_hard(); + flow.stop_gracefully(); + drop(connection); + session.join_flow_writer(true).unwrap(); + + // Then: the hard-stopped writer does not drain the packet after transport unblocks. + assert!(client.read_packet().is_err()); +} diff --git a/crates/et-server/src/session_flow_write.rs b/crates/et-server/src/session_flow_write.rs new file mode 100644 index 0000000..ad4d389 --- /dev/null +++ b/crates/et-server/src/session_flow_write.rs @@ -0,0 +1,99 @@ +use et_core::packet::Packet; +use et_net::connection::{ConnError, WritePacketError}; + +use super::FlowControl; +use crate::session::{ActiveSession, SessionError}; + +pub(crate) enum FlowWriteResult { + Delivered, + BeforeReplay(SessionError), + ReplayOwned(SessionError), + Fatal(SessionError), +} + +#[cfg(unix)] +pub(crate) fn write_packet( + session: &ActiveSession, + flow: &FlowControl, + packet: &Packet, +) -> (FlowWriteResult, bool) { + write_packet_with(session, flow, packet, |connection, packet| { + connection.prepare_write_packet(packet.header(), packet.payload()) + }) +} + +#[cfg(unix)] +pub(crate) fn write_packet_with( + session: &ActiveSession, + flow: &FlowControl, + packet: &Packet, + prepare: F, +) -> (FlowWriteResult, bool) +where + F: FnOnce( + &mut et_net::connection::Connection, + &Packet, + ) -> Result, +{ + let prepared = match session.connection.lock() { + Ok(_connection) if flow.is_hard_stopped() => { + return (FlowWriteResult::Fatal(SessionError::Unavailable), false); + } + Ok(mut connection) => prepare(&mut connection, packet), + Err(_) => return (FlowWriteResult::Fatal(SessionError::Unavailable), false), + }; + let result = match prepared.and_then(et_net::connection::PreparedWrite::send) { + Ok(()) => FlowWriteResult::Delivered, + Err(WritePacketError::BeforeReplay(ConnError::Io(error))) => { + FlowWriteResult::BeforeReplay(SessionError::Connection(ConnError::Io(error))) + } + Err(WritePacketError::BeforeReplay(error)) => { + FlowWriteResult::Fatal(SessionError::Connection(error)) + } + Err(WritePacketError::ReplayOwned(error)) => { + FlowWriteResult::ReplayOwned(SessionError::Connection(error)) + } + }; + match session.connection.lock() { + Ok(mut connection) => { + if matches!( + result, + FlowWriteResult::BeforeReplay(_) | FlowWriteResult::ReplayOwned(_) + ) { + connection.disconnect(); + } + (result, connection.connected()) + } + Err(_) => (FlowWriteResult::Fatal(SessionError::Unavailable), false), + } +} + +#[cfg(windows)] +pub(crate) fn write_packet( + session: &ActiveSession, + flow: &FlowControl, + packet: &Packet, +) -> (FlowWriteResult, bool) { + match session.connection.lock() { + Ok(_) if flow.is_hard_stopped() => { + (FlowWriteResult::Fatal(SessionError::Unavailable), false) + } + Ok(mut connection) => { + let result = match connection.write_packet_owned(packet.header(), packet.payload()) { + Ok(()) => FlowWriteResult::Delivered, + Err(WritePacketError::BeforeReplay(ConnError::Io(error))) => { + connection.disconnect(); + FlowWriteResult::BeforeReplay(SessionError::Connection(ConnError::Io(error))) + } + Err(WritePacketError::BeforeReplay(error)) => { + FlowWriteResult::Fatal(SessionError::Connection(error)) + } + Err(WritePacketError::ReplayOwned(error)) => { + FlowWriteResult::ReplayOwned(SessionError::Connection(error)) + } + }; + (result, connection.connected()) + } + Err(_) => (FlowWriteResult::Fatal(SessionError::Unavailable), false), + } +} diff --git a/crates/et-server/src/session_io.rs b/crates/et-server/src/session_io.rs new file mode 100644 index 0000000..211a0a0 --- /dev/null +++ b/crates/et-server/src/session_io.rs @@ -0,0 +1,182 @@ +use std::io::Write; +use std::net::TcpStream; +use std::sync::atomic::Ordering; +use std::time::{Duration, Instant}; + +use et_core::backed_writer::{ + MAX_BACKUP_PACKETS, MAX_DISCONNECT_PACKETS, MAX_RECOVERY_BACKUP_BYTES, +}; +use et_net::local::LocalStream; + +use super::{ActiveSession, SessionError}; + +impl ActiveSession { + pub(crate) fn take_wake_reader(&self) -> Result { + self.wake_reader + .lock() + .map_err(|_| SessionError::Unavailable)? + .take() + .ok_or(SessionError::Unavailable) + } + + #[cfg_attr(windows, allow(dead_code))] + pub(crate) fn try_clone_stream(&self) -> Result<(TcpStream, u64), SessionError> { + let connection = self + .connection + .lock() + .map_err(|_| SessionError::Unavailable)?; + let generation = self.connection_generation.load(Ordering::Acquire); + let stream = connection + .try_clone_stream() + .map_err(SessionError::Connection)?; + Ok((stream, generation)) + } + + pub(crate) fn try_read_packet(&self) -> Result, SessionError> { + if let Some(flow) = &self.flow_control { + flow.set_reader_waiting(true); + } + let result = (|| { + self.connection + .lock() + .map_err(|_| SessionError::Unavailable)? + .try_read_packet() + .map_err(SessionError::Connection) + })(); + if let Some(flow) = &self.flow_control { + flow.set_reader_waiting(false); + } + result + } + + pub(crate) fn note_bridge_generation(&self, generation: u64) -> Result<(), SessionError> { + let mut observed = self + .bridge_generation + .lock() + .map_err(|_| SessionError::Unavailable)?; + if generation > *observed { + *observed = generation; + self.bridge_changed.notify_all(); + } + Ok(()) + } + + pub(crate) fn wait_for_bridge_generation( + &self, + expected: u64, + timeout: Duration, + ) -> Result<(), SessionError> { + let deadline = Instant::now() + .checked_add(timeout) + .ok_or(SessionError::Unavailable)?; + let mut observed = self + .bridge_generation + .lock() + .map_err(|_| SessionError::Unavailable)?; + while *observed < expected { + let remaining = deadline + .checked_duration_since(Instant::now()) + .ok_or(SessionError::RecoverBusy)?; + let (next, result) = self + .bridge_changed + .wait_timeout(observed, remaining) + .map_err(|_| SessionError::Unavailable)?; + observed = next; + if result.timed_out() && *observed < expected { + return Err(SessionError::RecoverBusy); + } + } + Ok(()) + } + + pub(crate) fn connection_state(&self) -> Result<(bool, u64), SessionError> { + let connection = self + .connection + .lock() + .map_err(|_| SessionError::Unavailable)?; + Ok(( + connection.connected(), + self.connection_generation.load(Ordering::Acquire), + )) + } + + /// Soft-drop the encrypted client transport without killing the terminal. + /// + /// Used when the client TCP path dies (sleep, Wi-Fi, NAT) so terminal + /// output keeps buffering and a returning client can recover the same + /// session. Does not set the session shutdown flag or close the terminal. + pub(crate) fn mark_client_disconnected( + &self, + expected_generation: u64, + ) -> Result { + let mut connection = self + .connection + .lock() + .map_err(|_| SessionError::Unavailable)?; + if self.connection_generation.load(Ordering::Acquire) != expected_generation { + return Ok(false); + } + connection.disconnect(); + if let Some(state) = &self.flow_control { + state.disconnected(); + } + Ok(true) + } + + /// Apply a client delivery acknowledgement to the replay backup. + pub(crate) fn acknowledge_delivery(&self, sequence: i64) -> Result<(), SessionError> { + self.connection + .lock() + .map_err(|_| SessionError::Unavailable)? + .acknowledge_delivery(sequence); + Ok(()) + } + + /// Keep-alive payload acknowledging everything read from the client. + pub(crate) fn keepalive_ack( + &self, + ) -> Result<[u8; et_core::keepalive::ACK_PAYLOAD_LEN], SessionError> { + Ok(self + .connection + .lock() + .map_err(|_| SessionError::Unavailable)? + .keepalive_ack()) + } + + pub(crate) fn can_buffer_write(&self, bytes: i64) -> Result { + if let Some(state) = &self.flow_control { + let bytes = usize::try_from(bytes).map_err(|_| SessionError::Unavailable)?; + return state.can_accept_terminal(bytes); + } + let hold = self + .recover_hold + .lock() + .map_err(|_| SessionError::Unavailable)?; + if hold.len() >= MAX_BACKUP_PACKETS + MAX_DISCONNECT_PACKETS { + return Ok(false); + } + let held = + i64::try_from(self.recover_hold_bytes.load(Ordering::Acquire)).unwrap_or(i64::MAX); + let requested = held.checked_add(bytes).unwrap_or(i64::MAX); + if requested > MAX_RECOVERY_BACKUP_BYTES { + return Ok(false); + } + Ok(self + .connection + .lock() + .map_err(|_| SessionError::Unavailable)? + .can_buffer_write(requested)) + } + + pub(crate) fn is_shutting_down(&self) -> bool { + self.shutdown.load(Ordering::Acquire) + } + + pub(super) fn signal(&self) -> Result<(), SessionError> { + self.wake_writer + .lock() + .map_err(|_| SessionError::Unavailable)? + .write_all(&[1]) + .map_err(SessionError::Io) + } +} diff --git a/crates/et-server/src/session_recovery.rs b/crates/et-server/src/session_recovery.rs new file mode 100644 index 0000000..fba2fb7 --- /dev/null +++ b/crates/et-server/src/session_recovery.rs @@ -0,0 +1,205 @@ +use std::net::{Shutdown, TcpStream}; +use std::sync::atomic::Ordering; +use std::sync::{Mutex, MutexGuard, TryLockError}; +use std::time::{Duration, Instant}; + +use et_core::proto::TerminalPacketType; +use et_net::connection::{WritePacketError, DEFAULT_RECOVERY_TIMEOUT}; + +use super::{ActiveSession, SessionError, RECOVERY_LOCK_TIMEOUT}; + +impl ActiveSession { + /// Acquire the single-flight recover permit without speaking on the wire. + /// + /// Callers must send `ReturningClient` only after this succeeds, so a + /// concurrent recover does not commit the peer to sequence exchange and + /// then fail with `RecoverBusy`. The permit releases the flag on drop + /// (including panic unwind). + pub(crate) fn try_begin_recover(&self) -> Result, SessionError> { + if self + .recovering + .compare_exchange(false, true, Ordering::AcqRel, Ordering::Acquire) + .is_err() + { + return Err(SessionError::RecoverBusy); + } + if let Some(flow) = &self.flow_control { + if let Err(error) = flow.pause() { + self.recovering.store(false, Ordering::Release); + return Err(error); + } + } + Ok(RecoverPermit { session: self }) + } + + /// Prepare → network handshake off-lock → install → flush hold. + /// + /// The connection mutex is held only for soft-disconnect/snapshot and for + /// installing the new stream, not for sequence exchange or peer auth. + fn recover_body(&self, stream: TcpStream) -> Result<(), SessionError> { + // Phase 1: soft-disconnect and snapshot under a short lock. + let mut candidate = { + let connection = lock_timeout(&self.connection, RECOVERY_LOCK_TIMEOUT)?; + // Snapshot onto the new stream without closing or disconnecting + // the live victim socket (ET #784 / ANT-2026-VAMER5RC). Terminal + // output during the off-lock handshake is queued in + // `recover_hold` because `recovering` is set. A failed recover + // must leave the existing session intact. + connection.prepare_recovery_candidate(stream) + }; + + // Phase 2: recovery network I/O without the session connection lock. + candidate + .run_recovery_handshake(DEFAULT_RECOVERY_TIMEOUT) + .map_err(SessionError::Connection)?; + // Any packet that decrypts with the session key authenticates the + // returning client; it is requeued and handled by the session loop. + candidate + .authenticate_peer(DEFAULT_RECOVERY_TIMEOUT) + .map_err(SessionError::Connection)?; + if self.flow_control.is_some() { + candidate + .minimize_output_buffering() + .map_err(SessionError::Connection)?; + } + let ack = candidate.keepalive_ack(); + candidate + .write_packet_live(TerminalPacketType::KeepAlive as u8, &ack) + .map_err(SessionError::Connection)?; + let new_control = candidate + .try_clone_stream() + .map_err(SessionError::Connection)?; + + // Phase 3: install under a short lock. + { + let mut control = lock_timeout(&self.control, RECOVERY_LOCK_TIMEOUT)?; + let mut connection = lock_timeout(&self.connection, RECOVERY_LOCK_TIMEOUT)?; + let old_control = std::mem::replace(&mut *control, new_control); + let _ = old_control.shutdown(Shutdown::Both); + *connection = candidate; + self.connection_generation.fetch_add(1, Ordering::Release); + } + + // Phase 4: drain terminal output queued while the handshake ran. + // Still under `recovering` so concurrent send_packet keeps queuing + // until the permit drops; Drop flushes once more after clearing. + self.flush_recover_hold() + } + + fn flush_recover_hold(&self) -> Result<(), SessionError> { + self.flush_recover_hold_with(|connection, header, payload| { + connection.write_packet_owned(header, payload) + }) + } + + pub(super) fn flush_recover_hold_with(&self, mut write: F) -> Result<(), SessionError> + where + F: FnMut(&mut et_net::connection::Connection, u8, &[u8]) -> Result<(), WritePacketError>, + { + loop { + let batch = { + let mut hold = self + .recover_hold + .lock() + .map_err(|_| SessionError::Unavailable)?; + if hold.is_empty() { + return Ok(()); + } + std::mem::take(&mut *hold) + }; + let mut connection = lock_timeout(&self.connection, RECOVERY_LOCK_TIMEOUT)?; + let mut remaining = batch.into_iter(); + while let Some((header, payload)) = remaining.next() { + if let Err(error) = write(&mut connection, header, &payload) { + drop(connection); + let mut hold = self + .recover_hold + .lock() + .map_err(|_| SessionError::Unavailable)?; + let concurrent = std::mem::take(&mut *hold); + match error { + WritePacketError::BeforeReplay(error) => { + hold.push((header, payload)); + hold.extend(remaining); + hold.extend(concurrent); + return Err(SessionError::Connection(error)); + } + WritePacketError::ReplayOwned(error) => { + // Replay owns the failed packet; retain only the + // unwritten tail and concurrent plaintext. + self.recover_hold_bytes.fetch_sub( + u64::try_from(payload.len()) + .map_err(|_| SessionError::Unavailable)?, + Ordering::AcqRel, + ); + hold.extend(remaining); + hold.extend(concurrent); + return Err(SessionError::Connection(error)); + } + } + } + self.recover_hold_bytes.fetch_sub( + u64::try_from(payload.len()).map_err(|_| SessionError::Unavailable)?, + Ordering::AcqRel, + ); + } + } + } +} + +/// Single-flight recover permit. Dropping it (normally or on panic) always +/// clears [`ActiveSession::recovering`] and wakes the terminal bridge. +pub(crate) struct RecoverPermit<'a> { + session: &'a ActiveSession, +} + +impl RecoverPermit<'_> { + /// Run the recovery handshake and install the new stream. + pub(crate) fn complete(self, stream: TcpStream) -> Result<(), SessionError> { + // `self` drops after this returns (or panics), clearing `recovering` + // and flushing any straggler hold packets. + self.session.recover_body(stream) + } +} + +impl Drop for RecoverPermit<'_> { + fn drop(&mut self) { + // Flush while still marked recovering so send_packet keeps queuing + // rather than racing into a half-installed connection. + let _ = self.session.flush_recover_hold(); + self.session.recovering.store(false, Ordering::Release); + // Catch anything that observed `recovering` and queued after the first + // flush but before the flag cleared (re-check is under the hold lock). + let _ = self.session.flush_recover_hold(); + if let Some(state) = &self.session.flow_control { + let connected = self + .session + .connection + .lock() + .is_ok_and(|connection| connection.connected()); + state.resume(connected); + } + // Wake the bridge even on failure so it re-checks connection state. + let _ = self.session.signal(); + } +} + +/// Acquire a [`Mutex`] with a deadline so recover cannot park forever behind a +/// bridge thread blocked in a live write. +fn lock_timeout(mutex: &Mutex, timeout: Duration) -> Result, SessionError> { + let deadline = Instant::now() + .checked_add(timeout) + .ok_or(SessionError::RecoverBusy)?; + loop { + match mutex.try_lock() { + Ok(guard) => return Ok(guard), + Err(TryLockError::Poisoned(_)) => return Err(SessionError::Unavailable), + Err(TryLockError::WouldBlock) => { + if Instant::now() >= deadline { + return Err(SessionError::RecoverBusy); + } + std::thread::sleep(Duration::from_millis(5)); + } + } + } +} diff --git a/crates/et-server/src/terminal_bridge.rs b/crates/et-server/src/terminal_bridge.rs index ac11596..2fdf9a8 100644 --- a/crates/et-server/src/terminal_bridge.rs +++ b/crates/et-server/src/terminal_bridge.rs @@ -12,7 +12,7 @@ use rustix::event::{poll, PollFd, PollFlags}; #[cfg(unix)] use rustix::time::Timespec; -use crate::session::{ActiveSession, SessionError}; +use crate::session::{ActiveSession, SessionError, SessionWriteError}; const READ_BUFFER: usize = 16 * 1024; @@ -74,6 +74,11 @@ fn run_mode_poll( // An outbound forwarding packet that could not fit in the disconnected // replay buffer. Keep ownership until recovery restores write capacity. let mut pending_outbound: Option = None; + // A complete terminal packet already read from the local stream. Retain + // ownership across backpressure instead of dropping it or reading ahead. + let mut pending_terminal: Option = None; + let mut terminal_closing = false; + let mut terminal_eof = false; loop { let mut resume_outbound_drain = false; if session.is_shutting_down() { @@ -90,6 +95,10 @@ fn run_mode_poll( send_or_hold(&session, packet, &mut connected, &mut connection_generation)?; resume_outbound_drain = pending_outbound.is_none(); } + if let Some(packet) = pending_terminal.take() { + pending_terminal = + send_or_hold(&session, packet, &mut connected, &mut connection_generation)?; + } let (client, polled_generation) = if connected { match session.try_clone_stream() { Ok((stream, generation)) => (Some(stream), Some(generation)), @@ -108,7 +117,8 @@ fn run_mode_poll( } else { (None, None) }; - let accept_terminal = session.can_buffer_write((READ_BUFFER * 2) as i64)?; + let accept_terminal = + pending_terminal.is_none() && session.can_buffer_write((READ_BUFFER * 2) as i64)?; // When the client is down, poll with a short timeout so recovery wakes // (via the wake pipe) are still processed promptly and we re-check // `session.connected()` after recover installs a new stream. @@ -118,7 +128,10 @@ fn run_mode_poll( forwarder.wake().map_err(forward_error)?, client.as_ref(), accept_terminal, - pending_forward.is_some() || pending_outbound.is_some() || !connected, + pending_forward.is_some() + || pending_outbound.is_some() + || pending_terminal.is_some() + || !connected, )?; let client_events_are_stale = wake_events.intersects(PollFlags::IN | PollFlags::HUP); if client_events_are_stale { @@ -129,45 +142,34 @@ fn run_mode_poll( (connected, connection_generation) = session.connection_state()?; session.note_bridge_generation(connection_generation)?; } - let terminal_closed = terminal_events.intersects(PollFlags::HUP | PollFlags::ERR); - if terminal_closed || terminal_events.contains(PollFlags::IN) { - loop { - match read_terminal_packet(&mut terminal, &mut decoder) { - Ok(Some(packet)) => { - let packet = if mode == BridgeMode::Terminal { - validate_terminal_output(&packet)?; - packet - } else { - jumphost_terminal_packet(&session, packet)? - }; - decoder = LocalPacketDecoder::new(); - if let Err(error) = session.send_packet(packet.header(), packet.payload()) { - if !client_transport_error( - &error, - &session, - &mut connected, - &mut connection_generation, - )? { - return Err(error); - } - } - } - Ok(None) if terminal_closed => continue, - Ok(None) => break, - Err(SessionError::Io(error)) - if terminal_closed && error.kind() == io::ErrorKind::UnexpectedEof => - { - break; - } - Err(error) => return Err(error), + terminal_closing |= terminal_events.intersects(PollFlags::HUP | PollFlags::ERR); + if pending_terminal.is_none() + && !terminal_eof + && (terminal_closing || terminal_events.contains(PollFlags::IN)) + { + match read_terminal_packet(&mut terminal, &mut decoder) { + Ok(Some(packet)) => { + let packet = if mode == BridgeMode::Terminal { + validate_terminal_output(&packet)?; + packet + } else { + jumphost_terminal_packet(&session, packet)? + }; + decoder = LocalPacketDecoder::new(); + pending_terminal = + send_or_hold(&session, packet, &mut connected, &mut connection_generation)?; } - if !terminal_closed { - break; + Ok(None) => {} + Err(SessionError::Io(error)) + if terminal_closing && error.kind() == io::ErrorKind::UnexpectedEof => + { + terminal_eof = true; } + Err(error) => return Err(error), } - if terminal_closed { - return Ok(()); - } + } + if terminal_eof && pending_terminal.is_none() { + return Ok(()); } // Recovery authentication may read more than its proof packet into // BackedReader. Drain it after the wake even when the new socket no @@ -175,7 +177,7 @@ fn run_mode_poll( let client_data_ready = connected && (client_events_are_stale || client_events.contains(PollFlags::IN)); if client_data_ready { - while pending_forward.is_none() { + while pending_forward.is_none() && pending_outbound.is_none() { match session.try_read_packet() { // Jumphost relays every packet verbatim to the jump // terminal, which owns the destination connection. @@ -186,7 +188,18 @@ fn run_mode_poll( Ok(Some(packet)) if is_forward_packet(packet.header()) => { pending_forward = forwarder.try_receive(packet).map_err(forward_error)?; } - Ok(Some(packet)) => forward_client_packet(&session, &mut terminal, packet)?, + Ok(Some(packet)) => { + if let Some(control) = + forward_client_packet(&session, &mut terminal, packet)? + { + pending_outbound = send_or_hold( + &session, + control, + &mut connected, + &mut connection_generation, + )?; + } + } Ok(None) => break, Err(error) => { if client_transport_error( @@ -250,6 +263,7 @@ fn run_mode_windows( // further client packets are read so forwarding data stays ordered. let mut pending_forward: Option = None; let mut pending_outbound: Option = None; + let mut pending_terminal: Option = None; loop { if session.is_shutting_down() { return Ok(()); @@ -270,6 +284,11 @@ fn run_mode_windows( send_or_hold(&session, packet, &mut connected, &mut connection_generation)?; progress |= pending_outbound.is_none(); } + if let Some(packet) = pending_terminal.take() { + pending_terminal = + send_or_hold(&session, packet, &mut connected, &mut connection_generation)?; + progress |= pending_terminal.is_none(); + } // Connection state changes are announced through the wake channel. if drain_available(&mut wake)? { @@ -282,7 +301,7 @@ fn run_mode_windows( } // Terminal -> client, honouring the same write-buffer backpressure. - if session.can_buffer_write((READ_BUFFER * 2) as i64)? { + if pending_terminal.is_none() && session.can_buffer_write((READ_BUFFER * 2) as i64)? { match read_terminal_packet(&mut terminal, &mut decoder) { Ok(Some(packet)) => { progress = true; @@ -293,16 +312,8 @@ fn run_mode_windows( jumphost_terminal_packet(&session, packet)? }; decoder = LocalPacketDecoder::new(); - if let Err(error) = session.send_packet(packet.header(), packet.payload()) { - if !client_transport_error( - &error, - &session, - &mut connected, - &mut connection_generation, - )? { - return Err(error); - } - } + pending_terminal = + send_or_hold(&session, packet, &mut connected, &mut connection_generation)?; } Ok(None) => {} Err(error) => return Err(error), @@ -311,7 +322,7 @@ fn run_mode_windows( // Client -> terminal / forwarder. if connected { - while pending_forward.is_none() { + while pending_forward.is_none() && pending_outbound.is_none() { match session.try_read_packet() { Ok(Some(packet)) => { progress = true; @@ -321,8 +332,15 @@ fn run_mode_windows( } else if is_forward_packet(packet.header()) { pending_forward = forwarder.try_receive(packet).map_err(forward_error)?; - } else { - forward_client_packet(&session, &mut terminal, packet)?; + } else if let Some(control) = + forward_client_packet(&session, &mut terminal, packet)? + { + pending_outbound = send_or_hold( + &session, + control, + &mut connected, + &mut connection_generation, + )?; } } Ok(None) => break, @@ -382,10 +400,19 @@ fn send_or_hold( connected: &mut bool, connection_generation: &mut u64, ) -> Result, SessionError> { - match session.send_packet(packet.header(), packet.payload()) { + match session.send_packet_owned(packet.header(), packet.payload()) { Ok(()) => Ok(None), - Err(SessionError::Connection(ConnError::Backpressure)) => Ok(Some(packet)), - Err(error) => { + Err(SessionWriteError::BeforeReplay(SessionError::Connection(ConnError::Backpressure))) => { + Ok(Some(packet)) + } + Err(SessionWriteError::BeforeReplay(error)) => { + if client_transport_error(&error, session, connected, connection_generation)? { + Ok(Some(packet)) + } else { + Err(error) + } + } + Err(SessionWriteError::ReplayOwned(error)) => { if client_transport_error(&error, session, connected, connection_generation)? { Ok(None) } else { @@ -590,13 +617,15 @@ fn forward_client_packet( session: &ActiveSession, terminal: &mut LocalStream, packet: Packet, -) -> Result<(), SessionError> { +) -> Result, SessionError> { match packet.header() { value if value == TerminalPacketType::TerminalBuffer as u8 || value == TerminalPacketType::TerminalInfo as u8 => { - write_local_packet(terminal, &packet).map_err(SessionError::Io) + write_local_packet(terminal, &packet) + .map_err(SessionError::Io) + .map(|()| None) } header if header == TerminalPacketType::KeepAlive as u8 => { if let Some(ack) = et_core::keepalive::decode_ack(packet.payload()) { @@ -605,10 +634,10 @@ fn forward_client_packet( // The echo acknowledges everything read from the client, letting // an et.rs client trim its own replay backup. Legacy peers // (upstream C++, released et.rs) ignore the payload. - session.send_packet( + Ok(Some(Packet::new( TerminalPacketType::KeepAlive as u8, - &session.keepalive_ack()?, - ) + session.keepalive_ack()?.to_vec(), + ))) } _ => Err(SessionError::Io(io::Error::new( io::ErrorKind::InvalidData, diff --git a/crates/et-server/tests/flow_control_recovery.rs b/crates/et-server/tests/flow_control_recovery.rs new file mode 100644 index 0000000..3ad5755 --- /dev/null +++ b/crates/et-server/tests/flow_control_recovery.rs @@ -0,0 +1,249 @@ +#![forbid(unsafe_code)] + +mod runtime_support; +mod support; + +use std::io::{self, Read}; +use std::net::{Ipv4Addr, Shutdown, TcpListener, TcpStream}; +use std::sync::mpsc; +use std::thread; + +use et_core::keys::passkey_to_key; +use et_core::packet::Packet; +use et_core::proto::{ + ConnectStatus, FlowControlMode, SequenceHeader, TerminalBuffer, TerminalPacketType, +}; +use et_net::connection::Connection; +use et_net::framing_io::{read_proto_limited, write_proto}; +use et_net::local_packet::{read_local_packet, write_local_packet}; +use prost::Message; +use runtime_support::{default_payload, initialize, TestRuntime, ID_A, KEY_A, TIMEOUT}; + +struct RecoveryGate { + port: u16, + snapshot: mpsc::Receiver<()>, + release: mpsc::SyncSender<()>, + stop: mpsc::Receiver, + worker: Option>>, +} + +impl RecoveryGate { + fn start(mut server: TcpStream) -> Self { + let listener = TcpListener::bind((Ipv4Addr::LOCALHOST, 0)).unwrap(); + let port = listener.local_addr().unwrap().port(); + let (snapshot_tx, snapshot) = mpsc::sync_channel(0); + let (release, release_rx) = mpsc::sync_channel(0); + let (stop_tx, stop) = mpsc::sync_channel(1); + let worker = thread::spawn(move || { + let (mut client, _) = listener.accept()?; + stop_tx + .send(server.try_clone()?) + .map_err(|_| io::Error::other("proxy stop receiver closed"))?; + let mut client_read = client.try_clone()?; + let mut server_write = server.try_clone()?; + let upstream = thread::spawn(move || io::copy(&mut client_read, &mut server_write)); + + let sequence: SequenceHeader = read_proto_limited(&mut server, 4 * 1024)?; + snapshot_tx + .send(()) + .map_err(|_| io::Error::other("snapshot receiver closed"))?; + release_rx + .recv() + .map_err(|_| io::Error::other("recovery release sender closed"))?; + write_proto(&mut client, &sequence)?; + io::copy(&mut server, &mut client)?; + let _ = client.shutdown(Shutdown::Both); + upstream + .join() + .map_err(|_| io::Error::other("recovery upload worker panicked"))??; + Ok(()) + }); + Self { + port, + snapshot, + release, + stop, + worker: Some(worker), + } + } + + fn finish(&mut self) { + let stream = self.stop.recv_timeout(TIMEOUT).unwrap(); + if let Err(error) = stream.shutdown(Shutdown::Both) { + assert_eq!( + error.kind(), + io::ErrorKind::NotConnected, + "could not stop the recovery proxy" + ); + } + self.worker.take().unwrap().join().unwrap().unwrap(); + } +} + +impl Drop for RecoveryGate { + fn drop(&mut self) { + if let Ok(stream) = self.stop.try_recv() { + let _ = stream.shutdown(Shutdown::Both); + } + if let Some(worker) = self.worker.take() { + let _ = worker.join(); + } + } +} + +#[test] +fn recovery_snapshot_holds_flow_output_and_control_for_new_connection() { + let mut server = TestRuntime::start(); + let _terminal = server.register(ID_A, KEY_A); + let (stream, response) = server.handshake(ID_A); + assert_eq!(response.status, Some(ConnectStatus::NewClient as i32)); + let key = passkey_to_key(KEY_A).unwrap(); + let mut payload = default_payload(); + payload.flowcontrol = Some(FlowControlMode::Backpressure as i32); + let (mut client, initial) = initialize(stream, &key, payload); + assert_eq!(initial.error, None); + + let (returning, response) = server.handshake(ID_A); + assert_eq!(response.status, Some(ConnectStatus::ReturningClient as i32)); + let mut gate = RecoveryGate::start(returning); + let recovery_stream = TcpStream::connect((Ipv4Addr::LOCALHOST, gate.port)).unwrap(); + runtime_support::bound(&recovery_stream); + let (client_tx, client_rx) = mpsc::sync_channel::(0); + let recovery = thread::spawn(move || { + client.recover(recovery_stream).unwrap(); + client + .write_packet(TerminalPacketType::KeepAlive as u8, &[]) + .unwrap(); + client_tx.send(client).unwrap(); + }); + + gate.snapshot.recv_timeout(TIMEOUT).unwrap(); + server.handle.send_packet(ID_A, 41, b"output").unwrap(); + server + .handle + .send_packet(ID_A, TerminalPacketType::KeepAlive as u8, b"control") + .unwrap(); + gate.release.send(()).unwrap(); + + let mut client = client_rx.recv_timeout(TIMEOUT).unwrap(); + recovery.join().unwrap(); + let acknowledgement = client.read_packet().unwrap(); + assert_eq!( + acknowledgement.header(), + TerminalPacketType::KeepAlive as u8 + ); + let output = client.read_packet().unwrap(); + let control = client.read_packet().unwrap(); + assert_eq!( + (output.header(), output.payload()), + (41, b"output".as_slice()) + ); + assert_eq!( + (control.header(), control.payload()), + (TerminalPacketType::KeepAlive as u8, b"control".as_slice()) + ); + + client.shutdown().unwrap(); + gate.finish(); + server.runtime.shutdown().unwrap(); +} + +#[test] +fn terminal_hup_during_recovery_delivers_final_output_once_on_the_new_connection() { + let mut server = TestRuntime::start(); + let mut terminal = server.register(ID_A, KEY_A); + let (stream, response) = server.handshake(ID_A); + assert_eq!(response.status, Some(ConnectStatus::NewClient as i32)); + let key = passkey_to_key(KEY_A).unwrap(); + let mut payload = default_payload(); + payload.flowcontrol = Some(FlowControlMode::Backpressure as i32); + let (mut client, initial) = initialize(stream, &key, payload); + assert_eq!(initial.error, None); + let init = read_local_packet(&mut terminal).unwrap(); + assert_eq!(init.header(), TerminalPacketType::TerminalInit as u8); + + let mut old_stream = client.try_clone_stream().unwrap(); + old_stream.set_read_timeout(Some(TIMEOUT)).unwrap(); + let (returning, response) = server.handshake(ID_A); + assert_eq!(response.status, Some(ConnectStatus::ReturningClient as i32)); + let mut gate = RecoveryGate::start(returning); + let recovery_stream = TcpStream::connect((Ipv4Addr::LOCALHOST, gate.port)).unwrap(); + runtime_support::bound(&recovery_stream); + let (client_tx, client_rx) = mpsc::sync_channel::(0); + let recovery = thread::spawn(move || { + client.recover(recovery_stream).unwrap(); + client + .write_packet(TerminalPacketType::KeepAlive as u8, &[]) + .unwrap(); + client_tx.send(client).unwrap(); + }); + + // The server has taken its connection snapshot and the flow writer is + // paused. Queue the shell's final output while recovery remains gated. + gate.snapshot.recv_timeout(TIMEOUT).unwrap(); + let final_output = TerminalBuffer { + buffer: Some(b"final-before-hup".to_vec()), + }; + write_local_packet( + &mut terminal, + &Packet::new( + TerminalPacketType::TerminalBuffer as u8, + final_output.encode_to_vec(), + ), + ) + .unwrap(); + + // Release the exact sequence-header barrier and wait for the recovered + // keepalive. That event proves the candidate was installed and the old + // stream was retired; checking before it would race the recovery timeout. + gate.release.send(()).unwrap(); + let mut client = client_rx.recv_timeout(TIMEOUT).unwrap(); + recovery.join().unwrap(); + let acknowledgement = client.read_packet().unwrap(); + assert_eq!( + acknowledgement.header(), + TerminalPacketType::KeepAlive as u8 + ); + + // Graceful stop must preserve the recovery pause: retiring the old stream + // yields EOF/reset without even one byte of the final encrypted packet. + let mut byte = [0u8; 1]; + let old_read = old_stream.read(&mut byte); + assert!( + match &old_read { + Ok(0) => true, + Err(error) + if matches!( + error.kind(), + io::ErrorKind::ConnectionReset | io::ErrorKind::ConnectionAborted + ) => + { + true + } + Ok(_) | Err(_) => false, + }, + "final output reached the old connection during recovery: {old_read:?}" + ); + + let delivered = client.read_packet().unwrap(); + assert_eq!(delivered.header(), TerminalPacketType::TerminalBuffer as u8); + assert_eq!( + TerminalBuffer::decode(delivered.payload()).unwrap(), + final_output + ); + drop(terminal); + if let Ok(packet) = client.read_packet() { + assert_eq!( + packet.header(), + TerminalPacketType::KeepAlive as u8, + "final terminal output was duplicated before EOF" + ); + assert!( + client.read_packet().is_err(), + "terminal EOF must follow the recovery keepalive" + ); + } + + gate.finish(); + server.runtime.shutdown().unwrap(); +} diff --git a/crates/et-server/tests/runtime_adversarial.rs b/crates/et-server/tests/runtime_adversarial.rs index 1c1fe0a..c0a9931 100644 --- a/crates/et-server/tests/runtime_adversarial.rs +++ b/crates/et-server/tests/runtime_adversarial.rs @@ -14,8 +14,8 @@ use std::time::{Duration, Instant}; use et_core::keys::passkey_to_key; use et_core::proto::{ - ConnectRequest, ConnectResponse, ConnectStatus, InitialPayload, InitialResponse, - PortForwardSourceRequest, SocketEndpoint, + ConnectRequest, ConnectResponse, ConnectStatus, FlowControlMode, InitialPayload, + InitialResponse, PortForwardSourceRequest, SocketEndpoint, }; use et_net::connection::Connection; use et_net::framing_io::{read_proto_limited, write_proto}; @@ -114,6 +114,7 @@ fn delayed_valid_reverse_forward_times_out_rolls_back_and_resets_slot() { environmentvariable: None, }], environmentvariables: HashMap::new(), + flowcontrol: None, }; let (stream, response) = server.handshake(ID_A); assert_eq!(response.status, Some(ConnectStatus::NewClient as i32)); @@ -222,6 +223,7 @@ fn stalled_privileged_tcp_helper_honors_initialization_deadline_and_resets_slot( environmentvariable: None, }], environmentvariables: HashMap::new(), + flowcontrol: None, }; let (stream, response) = server.handshake(ID_A); assert_eq!(response.status, Some(ConnectStatus::NewClient as i32)); @@ -290,6 +292,7 @@ fn unbindable_reverse_tunnel_reports_an_error_and_resets_the_slot() { jumphost: Some(false), reversetunnels: vec![Default::default()], environmentvariables: HashMap::new(), + flowcontrol: None, }; let (stream, response) = server.handshake(ID_A); assert_eq!(response.status, Some(ConnectStatus::NewClient as i32)); @@ -325,6 +328,7 @@ fn occupied_reverse_row_is_fatal_and_rolls_back_sibling() { jumphost: Some(false), reversetunnels: vec![request(&occupied_path), request(&usable_path)], environmentvariables: HashMap::new(), + flowcontrol: None, }; let (stream, _) = server.handshake(ID_A); @@ -384,6 +388,7 @@ fn reverse_bind_failure_never_activates_the_session() { jumphost: Some(false), reversetunnels: vec![request(&occupied_path), request(&sibling_path)], environmentvariables: HashMap::new(), + flowcontrol: None, }; let (stream, _) = server.handshake(ID_A); @@ -423,6 +428,7 @@ fn reverse_failures_are_plain_fatal_errors() { jumphost: Some(false), reversetunnels: requests, environmentvariables: HashMap::new(), + flowcontrol: None, }; let (stream, _) = server.handshake(ID_A); @@ -459,6 +465,7 @@ fn reverse_listener_cap_is_prebind_transactional_on_server() { }) .collect(), environmentvariables: HashMap::new(), + flowcontrol: None, }; let (stream, _) = server.handshake(ID_A); @@ -497,6 +504,7 @@ fn obsolete_origin_marker_has_no_privileged_meaning() { environmentvariable: Some("ET_RS_SSH_CONFIG_REMOTE_FORWARD".to_owned()), }], environmentvariables: HashMap::new(), + flowcontrol: None, }; let (stream, _) = server.handshake(ID_A); @@ -516,6 +524,7 @@ fn jumphost_payload_is_relayed_to_the_registered_terminal() { jumphost: Some(true), reversetunnels: Vec::new(), environmentvariables: HashMap::new(), + flowcontrol: Some(FlowControlMode::Discard as i32), }; let (stream, response) = server.handshake(ID_A); assert_eq!(response.status, Some(ConnectStatus::NewClient as i32)); @@ -543,6 +552,7 @@ fn jumphost_payload_is_relayed_to_the_registered_terminal() { client_received_tx.send(()).unwrap(); let relayed = terminal_handshake.join().unwrap(); assert_eq!(relayed.jumphost, Some(true)); + assert_eq!(relayed.flowcontrol, Some(FlowControlMode::Discard as i32)); server.runtime.shutdown().unwrap(); } diff --git a/crates/et-server/tests/runtime_disconnect.rs b/crates/et-server/tests/runtime_disconnect.rs index b9c70a9..a25c145 100644 --- a/crates/et-server/tests/runtime_disconnect.rs +++ b/crates/et-server/tests/runtime_disconnect.rs @@ -8,8 +8,7 @@ use std::sync::mpsc; use std::time::Duration; use et_core::keys::passkey_to_key; -use et_core::proto::{ConnectStatus, SequenceHeader}; -use et_net::framing_io::read_proto_limited; +use et_core::proto::ConnectStatus; use et_server::SessionState; use runtime_support::{default_payload, initialize, TestRuntime, ID_A, KEY_A, TIMEOUT}; @@ -69,25 +68,24 @@ fn terminal_eof_removes_active_session_and_allows_fresh_registration() { } #[test] -fn terminal_eof_interrupts_blocked_returning_recovery() { +fn terminal_eof_preserves_returning_recovery_after_permit_acquisition() { let mut server = TestRuntime::start(); let terminal = server.register(ID_A, KEY_A); let (stream, response) = server.handshake(ID_A); assert_eq!(response.status, Some(ConnectStatus::NewClient as i32)); let key = passkey_to_key(KEY_A).unwrap(); - let (_active_client, initial) = initialize(stream, &key, default_payload()); + let (mut active_client, initial) = initialize(stream, &key, default_payload()); assert_eq!(initial.error, None); server .handle .wait_for_state(ID_A, SessionState::Active, TIMEOUT) .unwrap(); - let (mut returning, response) = server.handshake(ID_A); + let (returning, response) = server.handshake(ID_A); assert_eq!(response.status, Some(ConnectStatus::ReturningClient as i32)); - let _: SequenceHeader = read_proto_limited(&mut returning, 80 * 1024 * 1024).unwrap(); drop(terminal); server.handle.wait_disconnected(ID_A, TIMEOUT).unwrap(); - assert_prompt_eof_or_reset(&mut returning); + active_client.recover(returning).unwrap(); server.runtime.shutdown().unwrap(); } diff --git a/crates/et-server/tests/runtime_recovery.rs b/crates/et-server/tests/runtime_recovery.rs index deaad48..79671d4 100644 --- a/crates/et-server/tests/runtime_recovery.rs +++ b/crates/et-server/tests/runtime_recovery.rs @@ -10,12 +10,16 @@ use std::thread; use std::time::{Duration, Instant}; use et_core::keys::passkey_to_key; -use et_core::proto::{ConnectResponse, ConnectStatus, TerminalPacketType}; +use et_core::packet::Packet; +use et_core::proto::{ + ConnectResponse, ConnectStatus, FlowControlMode, TermInit, TerminalBuffer, TerminalPacketType, +}; use et_net::connection::DEFAULT_LIVE_WRITE_TIMEOUT; use et_net::framing_io::{read_proto_limited, write_proto}; use et_net::handshake::client_request; -use et_net::local_packet::read_local_packet; +use et_net::local_packet::{read_local_packet, write_local_packet}; use et_server::SessionState; +use prost::Message; use runtime_support::{default_payload, initialize, TestRuntime, ID_A, KEY_A, TIMEOUT}; // Deadlock watchdog only. Success is driven by exact bridge-generation and @@ -188,6 +192,62 @@ fn returning_client_receives_exact_buffered_server_catchup() { server.runtime.shutdown().unwrap(); } +#[test] +fn discard_flow_control_resumes_queued_terminal_output_after_recovery() { + let mut server = TestRuntime::start(); + let mut terminal = server.register(ID_A, KEY_A); + terminal.set_read_timeout(Some(TIMEOUT)).unwrap(); + let (stream, response) = server.handshake(ID_A); + assert_eq!(response.status, Some(ConnectStatus::NewClient as i32)); + let key = passkey_to_key(KEY_A).unwrap(); + let mut payload = default_payload(); + payload.flowcontrol = Some(FlowControlMode::Discard as i32); + let (mut client, initial) = initialize(stream, &key, payload); + assert_eq!(initial.error, None); + server + .handle + .wait_for_state(ID_A, SessionState::Active, TIMEOUT) + .unwrap(); + let init = read_local_packet(&mut terminal).unwrap(); + let term_init = TermInit::decode(init.payload()).unwrap(); + assert_eq!(term_init.flowcontrol, Some(FlowControlMode::Discard as i32)); + + client.shutdown().unwrap(); + let output = TerminalBuffer { + buffer: Some(b"newest-while-disconnected".to_vec()), + }; + write_local_packet( + &mut terminal, + &Packet::new( + TerminalPacketType::TerminalBuffer as u8, + output.encode_to_vec(), + ), + ) + .unwrap(); + + let (returning, response) = server.handshake(ID_A); + assert_eq!(response.status, Some(ConnectStatus::ReturningClient as i32)); + client.recover(returning).unwrap(); + client + .write_packet(TerminalPacketType::KeepAlive as u8, &[]) + .unwrap(); + let mut recovered_output = None; + for _ in 0..2 { + let recovered = client.read_packet().unwrap(); + match recovered.header() { + header if header == TerminalPacketType::KeepAlive as u8 => {} + header if header == TerminalPacketType::TerminalBuffer as u8 => { + recovered_output = Some(TerminalBuffer::decode(recovered.payload()).unwrap()); + break; + } + header => panic!("unexpected recovered packet type {header}"), + } + } + assert_eq!(recovered_output, Some(output)); + + server.runtime.shutdown().unwrap(); +} + #[test] fn hard_client_drop_keeps_session_for_returning_recover() { // Regression: laptop sleep / Wi-Fi drop aborts the TCP socket. The server @@ -335,19 +395,29 @@ fn recover_succeeds_while_old_peer_blackholes_writes() { let mut timed_out_live_write = false; for round in 0..1_024u32 { let send_started = Instant::now(); - server - .handle - .send_packet(ID_A, 40, &payload) - .unwrap_or_else(|error| panic!("flood round {round}: {error}")); + let result = server.handle.send_packet(ID_A, 40, &payload); let elapsed = send_started.elapsed(); assert!( elapsed < std::time::Duration::from_secs(8), "send_packet round {round} blocked for {:?} — live write timeout not applied", elapsed ); - if elapsed >= DEFAULT_LIVE_WRITE_TIMEOUT / 2 { - timed_out_live_write = true; - break; + match result { + Ok(()) if elapsed >= DEFAULT_LIVE_WRITE_TIMEOUT / 2 => { + timed_out_live_write = true; + break; + } + Ok(()) => {} + Err(error) => { + assert!( + error + .to_string() + .contains("io: live write deadline elapsed"), + "flood round {round} failed unexpectedly: {error}" + ); + timed_out_live_write = true; + break; + } } } assert!( diff --git a/crates/et-server/tests/runtime_support/mod.rs b/crates/et-server/tests/runtime_support/mod.rs index 014f5d1..485a3c4 100644 --- a/crates/et-server/tests/runtime_support/mod.rs +++ b/crates/et-server/tests/runtime_support/mod.rs @@ -117,6 +117,7 @@ pub fn default_payload() -> InitialPayload { jumphost: Some(false), reversetunnels: Vec::new(), environmentvariables: HashMap::new(), + flowcontrol: None, } } diff --git a/crates/et-server/tests/terminal_bridge.rs b/crates/et-server/tests/terminal_bridge.rs index a435c39..1b4d167 100644 --- a/crates/et-server/tests/terminal_bridge.rs +++ b/crates/et-server/tests/terminal_bridge.rs @@ -5,7 +5,9 @@ mod support; use et_core::keys::passkey_to_key; use et_core::packet::Packet; -use et_core::proto::{ConnectStatus, TermInit, TerminalBuffer, TerminalInfo, TerminalPacketType}; +use et_core::proto::{ + ConnectStatus, FlowControlMode, TermInit, TerminalBuffer, TerminalInfo, TerminalPacketType, +}; use et_net::local_packet::{read_local_packet, write_local_packet}; use prost::Message; use runtime_support::{default_payload, initialize, TestRuntime, ID_A, KEY_A, TIMEOUT}; @@ -157,6 +159,115 @@ fn terminal_hup_still_delivers_buffered_final_packet() { server.runtime.shutdown().unwrap(); } +#[test] +fn flow_control_mode_reaches_terminal_and_relays_output() { + let mut server = TestRuntime::start(); + let mut terminal = server.register(ID_A, KEY_A); + terminal.set_read_timeout(Some(TIMEOUT)).unwrap(); + let (stream, response) = server.handshake(ID_A); + assert_eq!(response.status, Some(ConnectStatus::NewClient as i32)); + let key = passkey_to_key(KEY_A).unwrap(); + let mut payload = default_payload(); + payload.flowcontrol = Some(FlowControlMode::Backpressure as i32); + let (mut client, initial) = initialize(stream, &key, payload); + assert_eq!(initial.error, None); + + let init = read_local_packet(&mut terminal).unwrap(); + let term_init = TermInit::decode(init.payload()).unwrap(); + assert_eq!( + term_init.flowcontrol, + Some(FlowControlMode::Backpressure as i32) + ); + + let output = TerminalBuffer { + buffer: Some(b"flow-controlled-output".to_vec()), + }; + write_local_packet( + &mut terminal, + &Packet::new( + TerminalPacketType::TerminalBuffer as u8, + output.encode_to_vec(), + ), + ) + .unwrap(); + + let received = client.read_packet().unwrap(); + assert_eq!(received.header(), TerminalPacketType::TerminalBuffer as u8); + assert_eq!(TerminalBuffer::decode(received.payload()).unwrap(), output); + + drop(terminal); + server.runtime.shutdown().unwrap(); +} + +#[test] +fn saturated_terminal_hup_retains_final_packets_by_mode() { + use std::sync::mpsc; + use std::thread; + + for mode in [FlowControlMode::Backpressure, FlowControlMode::Discard] { + // Given: terminal output fills the local/server path while the client is not reading. + let mut server = TestRuntime::start(); + let mut terminal = server.register(ID_A, KEY_A); + let (stream, response) = server.handshake(ID_A); + assert_eq!(response.status, Some(ConnectStatus::NewClient as i32)); + let key = passkey_to_key(KEY_A).unwrap(); + let mut payload = default_payload(); + payload.flowcontrol = Some(mode as i32); + let (mut client, initial) = initialize(stream, &key, payload); + assert_eq!(initial.error, None); + let _init = read_local_packet(&mut terminal).unwrap(); + let packets: Vec = (0u8..32) + .map(|value| TerminalBuffer { + buffer: Some(vec![value; 16 * 1024]), + }) + .collect(); + let sent_packets = packets.clone(); + let (written_tx, written_rx) = mpsc::sync_channel(32); + let producer = thread::spawn(move || { + for (index, packet) in sent_packets.iter().enumerate() { + write_local_packet( + &mut terminal, + &Packet::new( + TerminalPacketType::TerminalBuffer as u8, + packet.encode_to_vec(), + ), + ) + .unwrap(); + written_tx.send(index).unwrap(); + } + drop(terminal); + }); + for expected in 0..8 { + assert_eq!(written_rx.recv_timeout(TIMEOUT).unwrap(), expected); + } + + // When: the client starts draining after saturation and local HUP follows. + let reader = thread::spawn(move || { + let mut values = Vec::new(); + while let Ok(packet) = client.read_packet() { + if packet.header() == TerminalPacketType::TerminalBuffer as u8 { + let decoded = TerminalBuffer::decode(packet.payload()).unwrap(); + values.push(decoded.buffer.unwrap()[0]); + } + } + values + }); + producer.join().unwrap(); + let values = reader.join().unwrap(); + + // Then: backpressure is lossless; discard is ordered and retains newest output. + assert!( + values.windows(2).all(|pair| pair[0] < pair[1]), + "{values:?}" + ); + assert_eq!(values.last(), Some(&31)); + if mode == FlowControlMode::Backpressure { + assert_eq!(values, (0u8..32).collect::>()); + } + server.runtime.shutdown().unwrap(); + } +} + #[test] fn terminal_environment_is_forwarded_without_interpolation() { let mut server = TestRuntime::start();