use std::io::{self, Read, Write};
use std::os::windows::io::{AsRawHandle, RawHandle};
use std::sync::mpsc as std_mpsc;
use std::sync::Arc;
use std::thread;
use rmux_ipc::BlockingLocalStream;
use rmux_proto::{encode_attach_message, AttachMessage, TerminalSize};
use tokio::sync::{mpsc, oneshot};
use windows_sys::Win32::Foundation::INVALID_HANDLE_VALUE;
use crate::ClientError;
#[path = "attach_windows/action.rs"]
mod action;
#[path = "attach_windows/console_coordination.rs"]
mod console_coordination;
#[path = "attach_windows/console_input_read.rs"]
mod console_input_read;
#[cfg(test)]
#[path = "attach_windows/console_mode_ownership_tests.rs"]
mod console_mode_ownership_tests;
#[path = "attach_windows/input.rs"]
mod input;
#[path = "attach_windows/metrics.rs"]
mod metrics;
#[path = "attach_windows/output.rs"]
mod output;
#[path = "attach_windows/output_worker.rs"]
mod output_worker;
#[path = "attach/screen.rs"]
mod screen;
#[path = "attach_windows/shell_command.rs"]
mod shell_command;
#[path = "attach_windows/stream.rs"]
mod stream;
#[path = "attach_windows/terminal.rs"]
mod terminal;
#[path = "attach/terminal_cleanup.rs"]
mod terminal_cleanup;
#[path = "attach_windows/vt_input_passthrough.rs"]
mod vt_input_passthrough;
#[path = "attach_windows/vt_mode_scanner.rs"]
mod vt_mode_scanner;
#[path = "attach_windows/windows_version.rs"]
mod windows_version;
use crate::attach_lock_state::AttachLockState;
use screen::AttachScreenTracker;
pub use terminal::{AttachError, RawTerminal, Result};
const READ_BUFFER_SIZE: usize = 8192;
const ATTACH_INPUT_QUEUE_CAPACITY: usize = 256;
pub fn attach_terminal(stream: BlockingLocalStream) -> std::result::Result<(), ClientError> {
attach_terminal_with_initial_bytes(stream, Vec::new())
}
pub fn attach_terminal_with_initial_bytes(
stream: BlockingLocalStream,
initial_bytes: Vec<u8>,
) -> std::result::Result<(), ClientError> {
attach_terminal_with_initial_bytes_and_windows_console_key(stream, initial_bytes, false)
}
pub fn attach_terminal_with_initial_bytes_and_windows_console_key(
stream: BlockingLocalStream,
initial_bytes: Vec<u8>,
windows_console_key_enabled: bool,
) -> std::result::Result<(), ClientError> {
let input = io::stdin();
let raw_terminal = RawTerminal::enter().map_err(ClientError::from)?;
let output = output::AttachStdout::for_managed_terminal(io::stdout());
attach_with_stdio_and_raw_terminal_with_output_fence(
stream,
initial_bytes,
raw_terminal,
input,
output,
windows_console_key_enabled,
output::AttachStdout::flush_output_fence,
)
}
pub fn attach_with_terminal<Terminal, Input, Output>(
stream: BlockingLocalStream,
_terminal: &Terminal,
input: Input,
output: Output,
) -> std::result::Result<(), ClientError>
where
Input: Read + AsRawHandle + Send + 'static,
Output: Write + Send + 'static,
{
attach_with_stdio(stream, Vec::new(), input, output, false)
}
fn attach_with_stdio<Input, Output>(
stream: BlockingLocalStream,
initial_bytes: Vec<u8>,
input: Input,
output: Output,
windows_console_key_enabled: bool,
) -> std::result::Result<(), ClientError>
where
Input: Read + AsRawHandle + Send + 'static,
Output: Write + Send + 'static,
{
let raw_terminal = RawTerminal::enter().map_err(ClientError::from)?;
attach_with_stdio_and_raw_terminal(
stream,
initial_bytes,
raw_terminal,
input,
output,
windows_console_key_enabled,
)
}
fn attach_with_stdio_and_raw_terminal<Input, Output>(
stream: BlockingLocalStream,
initial_bytes: Vec<u8>,
raw_terminal: RawTerminal,
input: Input,
output: Output,
windows_console_key_enabled: bool,
) -> std::result::Result<(), ClientError>
where
Input: Read + AsRawHandle + Send + 'static,
Output: Write + Send + 'static,
{
attach_with_stdio_and_raw_terminal_with_output_fence(
stream,
initial_bytes,
raw_terminal,
input,
output,
windows_console_key_enabled,
Write::flush,
)
}
fn attach_with_stdio_and_raw_terminal_with_output_fence<Input, Output, FenceFlush>(
stream: BlockingLocalStream,
initial_bytes: Vec<u8>,
raw_terminal: RawTerminal,
input: Input,
output: Output,
windows_console_key_enabled: bool,
output_fence_flush: FenceFlush,
) -> std::result::Result<(), ClientError>
where
Input: Read + AsRawHandle + Send + 'static,
Output: Write + Send + 'static,
FenceFlush: FnMut(&mut Output) -> io::Result<()> + Send + 'static,
{
let _ = raw_terminal.flush_pending_input();
let screen_tracker = AttachScreenTracker::default();
drive_attach_stream_with_terminal_state(
stream,
initial_bytes,
raw_terminal,
&screen_tracker,
input,
AttachOutputTarget {
output,
fence_flush: output_fence_flush,
},
windows_console_key_enabled,
)
}
struct AttachOutputTarget<Output, FenceFlush> {
output: Output,
fence_flush: FenceFlush,
}
fn drive_attach_stream_with_terminal_state<Input, Output, FenceFlush>(
mut stream: BlockingLocalStream,
initial_bytes: Vec<u8>,
raw_terminal: RawTerminal,
screen_tracker: &AttachScreenTracker,
input: Input,
output_target: AttachOutputTarget<Output, FenceFlush>,
windows_console_key_enabled: bool,
) -> std::result::Result<(), ClientError>
where
Input: Read + AsRawHandle + Send + 'static,
Output: Write + Send + 'static,
FenceFlush: FnMut(&mut Output) -> io::Result<()> + Send + 'static,
{
let AttachOutputTarget {
output,
fence_flush: output_fence_flush,
} = output_target;
let initial_size = terminal::current_size();
if let Some(size) = initial_size {
write_attach_message(&mut stream, AttachMessage::Resize(size))?;
}
let (resize_tx, resize_rx) = mpsc::unbounded_channel();
let _resize_watcher = terminal::ResizeWatcher::spawn(initial_size, resize_tx);
drive_attach_stream_inner(
stream,
initial_bytes,
screen_tracker.clone(),
input,
output,
AttachLoopInputs {
resize_rx,
actions: action::ManagedTerminalActions::new(raw_terminal),
windows_console_key_enabled,
error_cleanup: Some(terminal_cleanup::fallback_attach_stop_sequence(
&std::env::var("TERM").unwrap_or_default(),
)),
},
output_fence_flush,
)
}
pub fn drive_attach_stream<Input, Output>(
stream: BlockingLocalStream,
input: Input,
output: Output,
) -> std::result::Result<(), ClientError>
where
Input: Read + AsRawHandle + Send + 'static,
Output: Write + Send + 'static,
{
drive_attach_stream_inner(
stream,
Vec::new(),
AttachScreenTracker::default(),
input,
output,
AttachLoopInputs {
resize_rx: closed_resize_rx(),
actions: action::StreamOnlyActions,
windows_console_key_enabled: false,
error_cleanup: None,
},
Write::flush,
)
}
struct AttachLoopInputs<Actions> {
resize_rx: mpsc::UnboundedReceiver<TerminalSize>,
actions: Actions,
windows_console_key_enabled: bool,
error_cleanup: Option<Vec<u8>>,
}
fn drive_attach_stream_inner<Input, Output, Actions, FenceFlush>(
stream: BlockingLocalStream,
initial_bytes: Vec<u8>,
screen_tracker: AttachScreenTracker,
input: Input,
output: Output,
loop_inputs: AttachLoopInputs<Actions>,
output_fence_flush: FenceFlush,
) -> std::result::Result<(), ClientError>
where
Input: Read + AsRawHandle + Send + 'static,
Output: Write + Send + 'static,
Actions: action::AttachActionExecutor + Send + 'static,
FenceFlush: FnMut(&mut Output) -> io::Result<()> + Send + 'static,
{
let AttachLoopInputs {
resize_rx,
actions,
windows_console_key_enabled,
error_cleanup,
} = loop_inputs;
let input_join_policy = input_join_policy(input.as_raw_handle());
let (input_tx, input_rx) = mpsc::channel(ATTACH_INPUT_QUEUE_CAPACITY);
let lock_state = Arc::new(AttachLockState::default());
let input_lock_state = Arc::clone(&lock_state);
let (input_thread, input_completion_rx) = spawn_input_worker(input, input_tx, input_lock_state);
let (action_tx, action_rx) = std_mpsc::channel();
let (action_completion_tx, action_completion_rx) = mpsc::unbounded_channel();
let action_lock_state = Arc::clone(&lock_state);
let action_thread = thread::spawn(move || {
action_loop(actions, action_rx, action_completion_tx, action_lock_state)
});
let (pipe, runtime) = stream.into_async_parts();
let output_result = runtime.block_on(async {
let (reader, writer) = tokio::io::split(pipe);
stream::drive_async_attach_with_output_fence(
reader,
writer,
initial_bytes,
output,
screen_tracker,
stream::AttachAsyncChannels::new(
input_rx,
resize_rx,
action_tx,
action_completion_rx,
Arc::clone(&lock_state),
windows_console_key_enabled,
)
.with_input_completion(input_completion_rx)
.with_error_cleanup(error_cleanup),
output_fence_flush,
)
.await
});
lock_state.close();
let input_result = match input_join_policy {
InputJoinPolicy::JoinOnClose => {
join_attach_thread(input_thread)?;
Ok(())
}
InputJoinPolicy::DetachOnClose => Ok(()),
};
let action_result = join_attach_thread(action_thread)?;
output_result?;
action_result?;
input_result
}
fn spawn_input_worker<Input>(
input: Input,
input_tx: mpsc::Sender<input::AttachInput>,
lock_state: Arc<AttachLockState>,
) -> (
thread::JoinHandle<()>,
oneshot::Receiver<std::result::Result<(), ClientError>>,
)
where
Input: Read + AsRawHandle + Send + 'static,
{
let (completion_tx, completion_rx) = oneshot::channel();
let worker = thread::spawn(move || {
let result = input_loop(input, input_tx, lock_state);
let _ = completion_tx.send(result);
});
(worker, completion_rx)
}
fn action_loop<Actions>(
mut actions: Actions,
action_rx: std_mpsc::Receiver<action::AttachAction>,
completion_tx: mpsc::UnboundedSender<
std::result::Result<action::AttachActionOutcome, ClientError>,
>,
lock_state: Arc<AttachLockState>,
) -> std::result::Result<(), ClientError>
where
Actions: action::AttachActionExecutor,
{
while let Ok(request) = action_rx.recv() {
if request.requires_exclusive_input() {
lock_state.wait_until_input_idle();
}
let result = action::run_attach_action(&mut actions, request);
if completion_tx.send(result).is_err() {
return Ok(());
}
}
Ok(())
}
struct AttachInputReadLease<'a> {
state: &'a AttachLockState,
}
impl<'a> AttachInputReadLease<'a> {
fn acquire(state: &'a AttachLockState) -> Option<Self> {
state.begin_input_read().then_some(Self { state })
}
}
impl Drop for AttachInputReadLease<'_> {
fn drop(&mut self) {
self.state.finish_input_read();
}
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
enum PendingCtrlCForward {
None,
Sent,
InputClosed,
}
fn forward_pending_ctrl_c_event(
input_tx: &mpsc::Sender<input::AttachInput>,
lock_state: &AttachLockState,
) -> PendingCtrlCForward {
let Some(_input_read_lease) = AttachInputReadLease::acquire(lock_state) else {
return PendingCtrlCForward::None;
};
if !terminal::take_pending_ctrl_c_event() {
return PendingCtrlCForward::None;
}
if input_tx
.blocking_send(input::synthetic_ctrl_c_input())
.is_err()
{
return PendingCtrlCForward::InputClosed;
}
PendingCtrlCForward::Sent
}
fn input_loop<Input>(
mut input: Input,
input_tx: mpsc::Sender<input::AttachInput>,
lock_state: Arc<AttachLockState>,
) -> std::result::Result<(), ClientError>
where
Input: Read + AsRawHandle,
{
let mut read_buffer = [0_u8; READ_BUFFER_SIZE];
let input_handle = input.as_raw_handle();
if is_absent_input_handle(input_handle) {
lock_state.wait_until_closed();
return Ok(());
}
let mut console_input = input::ConsoleInputReader::from_handle(input_handle);
let mut exclusive_input_generation = lock_state.exclusive_input_generation();
loop {
if lock_state.is_closed() || input_tx.is_closed() {
return Ok(());
}
let current_generation = lock_state.exclusive_input_generation();
if current_generation != exclusive_input_generation {
if let Some(console_input) = console_input.as_mut() {
console_input.reset_after_exclusive_action();
}
exclusive_input_generation = current_generation;
}
if lock_state.is_locked() {
lock_state.wait_while_locked();
continue;
}
match forward_pending_ctrl_c_event(&input_tx, &lock_state) {
PendingCtrlCForward::None => {}
PendingCtrlCForward::Sent => continue,
PendingCtrlCForward::InputClosed => return Ok(()),
}
if !terminal::wait_for_key_input(input_handle, 50).map_err(ClientError::Io)? {
if lock_state.is_closed() || input_tx.is_closed() {
return Ok(());
}
match forward_pending_ctrl_c_event(&input_tx, &lock_state) {
PendingCtrlCForward::None | PendingCtrlCForward::Sent => {}
PendingCtrlCForward::InputClosed => return Ok(()),
}
continue;
}
if !lock_state.begin_input_read() {
continue;
}
let inputs = if let Some(console_input) = console_input.as_mut() {
match console_input.read_key_inputs() {
Ok(inputs) => Ok(inputs),
Err(error) if error.kind() == io::ErrorKind::Interrupted => Err(None),
Err(error) => Err(Some(ClientError::Io(error))),
}
} else {
let bytes_read = match input.read(&mut read_buffer) {
Ok(0) => 0,
Ok(bytes_read) => bytes_read,
Err(error) if error.kind() == io::ErrorKind::Interrupted => {
lock_state.finish_input_read();
continue;
}
Err(error) => {
lock_state.finish_input_read();
return Err(ClientError::Io(error));
}
};
if bytes_read == 0 {
lock_state.finish_input_read();
return Ok(());
}
Ok(vec![input::AttachInput::bytes(
read_buffer[..bytes_read].to_vec(),
)])
};
lock_state.finish_input_read();
let inputs = match inputs {
Ok(inputs) => inputs,
Err(None) => continue,
Err(Some(error)) => return Err(error),
};
if inputs.is_empty() || lock_state.is_locked() {
continue;
}
for input in inputs {
if input_tx.blocking_send(input).is_err() {
return Ok(());
}
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum InputJoinPolicy {
JoinOnClose,
DetachOnClose,
}
fn input_join_policy(handle: RawHandle) -> InputJoinPolicy {
if is_absent_input_handle(handle) {
InputJoinPolicy::JoinOnClose
} else {
InputJoinPolicy::DetachOnClose
}
}
fn is_absent_input_handle(handle: RawHandle) -> bool {
handle.is_null() || std::ptr::eq(handle, INVALID_HANDLE_VALUE as RawHandle)
}
fn write_attach_message(
stream: &mut BlockingLocalStream,
message: AttachMessage,
) -> std::result::Result<(), ClientError> {
let frame = encode_attach_message(&message).map_err(ClientError::from)?;
stream.write_all(&frame).map_err(ClientError::Io)
}
fn closed_resize_rx() -> mpsc::UnboundedReceiver<TerminalSize> {
let (resize_tx, resize_rx) = mpsc::unbounded_channel();
drop(resize_tx);
resize_rx
}
fn join_attach_thread<Output>(
thread: thread::JoinHandle<Output>,
) -> std::result::Result<Output, ClientError> {
thread
.join()
.map_err(|_| ClientError::Io(io::Error::other("attach thread panicked")))
}
#[cfg(test)]
mod tests {
use std::fs::File;
use std::io::{self, Read, Write};
use std::os::windows::io::{AsRawHandle, FromRawHandle, OwnedHandle, RawHandle};
use std::sync::Arc;
use tokio::sync::mpsc;
use windows_sys::Win32::Foundation::HANDLE;
use windows_sys::Win32::System::Pipes::CreatePipe;
use super::{
input_join_policy, input_loop, join_attach_thread, spawn_input_worker, AttachLockState,
ClientError, InputJoinPolicy,
};
struct InvalidHandleInput;
impl Read for InvalidHandleInput {
fn read(&mut self, _buffer: &mut [u8]) -> io::Result<usize> {
panic!("an invalid wait handle must fail before attempting a read")
}
}
impl AsRawHandle for InvalidHandleInput {
fn as_raw_handle(&self) -> RawHandle {
1_usize as RawHandle
}
}
#[test]
fn pipe_stdin_handles_are_detached_on_attach_close() {
let mut read: HANDLE = std::ptr::null_mut();
let mut write: HANDLE = std::ptr::null_mut();
let ok = unsafe {
CreatePipe(&mut read, &mut write, std::ptr::null_mut(), 0)
};
assert_ne!(ok, 0, "CreatePipe failed: {}", io::Error::last_os_error());
let read = unsafe {
OwnedHandle::from_raw_handle(read)
};
let _write = unsafe {
OwnedHandle::from_raw_handle(write)
};
assert_eq!(
input_join_policy(read.as_raw_handle()),
InputJoinPolicy::DetachOnClose
);
}
#[test]
fn console_stdin_handles_are_detached_on_attach_close() {
let console_like = 1_usize as std::os::windows::io::RawHandle;
assert_eq!(
input_join_policy(console_like),
InputJoinPolicy::DetachOnClose
);
}
#[test]
fn input_worker_publishes_wait_failure_before_exiting() {
let (input_tx, _input_rx) = mpsc::channel(1);
let lock_state = Arc::new(AttachLockState::default());
let (worker, completion_rx) = spawn_input_worker(InvalidHandleInput, input_tx, lock_state);
let completion = completion_rx
.blocking_recv()
.expect("input worker must publish its result");
join_attach_thread(worker).expect("input worker must not panic");
assert!(
matches!(completion, Err(ClientError::Io(_))),
"invalid wait handle must surface as an input I/O error: {completion:?}"
);
}
#[test]
fn pipe_stdin_input_loop_preserves_paste_bytes() -> Result<(), Box<dyn std::error::Error>> {
let mut read: HANDLE = std::ptr::null_mut();
let mut write: HANDLE = std::ptr::null_mut();
let ok = unsafe {
CreatePipe(&mut read, &mut write, std::ptr::null_mut(), 0)
};
assert_ne!(ok, 0, "CreatePipe failed: {}", io::Error::last_os_error());
let reader = unsafe {
File::from_raw_handle(read)
};
let mut writer = unsafe {
File::from_raw_handle(write)
};
let (input_tx, mut input_rx) = mpsc::channel(1);
let lock_state = Arc::new(AttachLockState::default());
let input_lock_state = Arc::clone(&lock_state);
let input_thread =
std::thread::spawn(move || input_loop(reader, input_tx, input_lock_state));
let payload = b"\x1b[200~ascii\r\n\x02\x1b[<64;2;2M\x1b[9;2u\x1b[200~\xe6\x9d\xb1\xe4\xba\xac cafe\xcc\x81\x1b[201~";
writer.write_all(payload)?;
writer.flush()?;
drop(writer);
let received = input_rx.blocking_recv().expect("input payload");
assert_eq!(received.payload(), payload.as_slice());
lock_state.close();
input_thread
.join()
.map_err(|_| "attach input thread panicked")??;
Ok(())
}
}