From 55a678707014e8c6cada0dceebab93f0e9f7e2bd Mon Sep 17 00:00:00 2001 From: Jonathan Liebig Date: Thu, 17 Sep 2026 02:44:55 +0200 Subject: [PATCH] fix(windows): recover VT input without changing readers refs #4251 --- src/client/input.rs | 2 +- src/client/input/windows_vti.rs | 11 +- src/client/mod.rs | 28 +++-- src/client/terminal_setup.rs | 190 +++++++++++++++++++++++++++----- 4 files changed, 185 insertions(+), 46 deletions(-) diff --git a/src/client/input.rs b/src/client/input.rs index e77d51fa..848d0156 100644 --- a/src/client/input.rs +++ b/src/client/input.rs @@ -330,7 +330,7 @@ fn windows_stdin_reader_loop( windows_crossterm_reader_loop(event_tx, should_quit); } else { match windows_vti::console_input_handle() { - Ok(handle) if crate::platform::windows_virtual_terminal_input_active() => { + Ok(handle) => { windows_vti::raw_console_reader_loop(handle, event_tx, should_quit); } _ => windows_crossterm_reader_loop(event_tx, should_quit), diff --git a/src/client/input/windows_vti.rs b/src/client/input/windows_vti.rs index e9cfa4e6..48fe8220 100644 --- a/src/client/input/windows_vti.rs +++ b/src/client/input/windows_vti.rs @@ -86,14 +86,17 @@ fn push_platform_input_events( #[cfg(windows)] pub(super) fn console_input_handle() -> std::io::Result { use windows_sys::Win32::Foundation::{HANDLE, INVALID_HANDLE_VALUE}; - use windows_sys::Win32::System::Console::{GetStdHandle, STD_INPUT_HANDLE}; + use windows_sys::Win32::System::Console::{GetConsoleMode, GetStdHandle, STD_INPUT_HANDLE}; let handle: HANDLE = unsafe { GetStdHandle(STD_INPUT_HANDLE) }; if handle.is_null() || handle == INVALID_HANDLE_VALUE { - Err(std::io::Error::last_os_error()) - } else { - Ok(handle) + return Err(std::io::Error::last_os_error()); } + let mut mode = 0; + if unsafe { GetConsoleMode(handle, &mut mode) } == 0 { + return Err(std::io::Error::last_os_error()); + } + Ok(handle) } #[cfg(windows)] diff --git a/src/client/mod.rs b/src/client/mod.rs index cb465096..1cef87a7 100644 --- a/src/client/mod.rs +++ b/src/client/mod.rs @@ -71,7 +71,7 @@ use terminal_geometry::{ use terminal_geometry::{reported_cell_size_from_events, store_reported_cell_size}; use terminal_setup::{ effective_mouse_capture, effective_sgr_pixel_mouse, set_mouse_capture, - setup_direct_attach_terminal, setup_terminal, should_draw_host_cursor, + setup_direct_attach_terminal, setup_terminal, should_draw_host_cursor, TerminalGuard, }; fn refresh_host_mouse_capture(enabled: bool, sgr_pixels: bool) { @@ -80,9 +80,7 @@ fn refresh_host_mouse_capture(enabled: bool, sgr_pixels: bool) { } } #[cfg(windows)] -use terminal_setup::{ - enable_windows_virtual_terminal_input, is_ssh_session, windows_vti_input_backend_enabled, -}; +use terminal_setup::{is_ssh_session, windows_vti_input_backend_enabled}; #[cfg(test)] use terminal_setup::{ should_enable_host_color_scheme_reports, windows_virtual_terminal_input_mode, @@ -327,6 +325,7 @@ fn run_client_with_mode( should_quit, loop_config, attach_escape, + &terminal_guard, ) .await }); @@ -377,6 +376,7 @@ async fn run_client_loop( should_quit: Arc, config: ClientLoopConfig, attach_escape: Option, + _terminal_guard: &TerminalGuard, ) -> Result<(), ClientError> { #[cfg(windows)] let _ = config.mouse_scroll_lines; @@ -1827,17 +1827,21 @@ async fn run_client_loop( ); let mouse_mode_changed = enabled != state.mouse_capture_active || next_sgr_pixels != host_sgr_pixels_active.load(Ordering::Acquire); + #[cfg(windows)] + if enabled && windows_vti_input_backend_enabled() && is_ssh_session() { + _terminal_guard + .recover_windows_virtual_terminal_input() + .map_err(ClientError::ConnectionFailed)?; + } if mouse_mode_changed { - #[cfg(windows)] - if enabled && windows_vti_input_backend_enabled() && is_ssh_session() { - let _ = enable_windows_virtual_terminal_input(); - } set_mouse_capture(enabled, next_sgr_pixels) .map_err(ClientError::ConnectionFailed)?; - #[cfg(windows)] - if enabled && windows_vti_input_backend_enabled() && !is_ssh_session() { - let _ = enable_windows_virtual_terminal_input(); - } + } + #[cfg(windows)] + if enabled && windows_vti_input_backend_enabled() && !is_ssh_session() { + _terminal_guard + .recover_windows_virtual_terminal_input() + .map_err(ClientError::ConnectionFailed)?; } state.mouse_capture_active = enabled; host_mouse_capture_active.store(enabled, Ordering::Release); diff --git a/src/client/terminal_setup.rs b/src/client/terminal_setup.rs index 5ea7001f..1e072a4f 100644 --- a/src/client/terminal_setup.rs +++ b/src/client/terminal_setup.rs @@ -3,6 +3,8 @@ use std::io::{self, Write as _}; use std::sync::atomic::{AtomicBool, Ordering}; use std::sync::Arc; +#[cfg(windows)] +use std::sync::{Mutex, MutexGuard}; use crossterm::event::{ DisableBracketedPaste, DisableFocusChange, DisableMouseCapture, EnableBracketedPaste, @@ -48,25 +50,23 @@ pub(super) fn setup_terminal_with_capabilities( restore_claimed: Arc::new(AtomicBool::new(false)), restored: false, #[cfg(windows)] - restore_windows_input_mode: None, + restore_windows_input_mode: Arc::new(WindowsInputModeRestore::default()), }; crate::terminal_modes::clear_host_mouse_reporting(&mut io::stdout())?; let host_color_scheme_reports = should_enable_host_color_scheme_reports(enable_client_protocols); - #[cfg(windows)] let windows_ssh_session = is_ssh_session(); #[cfg(windows)] let mut windows_virtual_terminal_input = if windows_vti_input_backend_enabled() && windows_ssh_session { - enable_windows_virtual_terminal_input() + enable_windows_virtual_terminal_input( + &terminal_guard.restore_claimed, + &terminal_guard.restore_windows_input_mode, + ) } else { WindowsVirtualTerminalInputSetup::default() }; - #[cfg(windows)] - { - terminal_guard.restore_windows_input_mode = windows_virtual_terminal_input.restore_mode; - } if enable_client_protocols { set_mouse_capture(mouse_capture, false)?; @@ -87,8 +87,10 @@ pub(super) fn setup_terminal_with_capabilities( #[cfg(windows)] if enable_client_protocols && windows_vti_input_backend_enabled() && !windows_ssh_session { - windows_virtual_terminal_input = enable_windows_virtual_terminal_input(); - terminal_guard.restore_windows_input_mode = windows_virtual_terminal_input.restore_mode; + windows_virtual_terminal_input = enable_windows_virtual_terminal_input( + &terminal_guard.restore_claimed, + &terminal_guard.restore_windows_input_mode, + ); } #[cfg(windows)] @@ -126,7 +128,7 @@ pub(super) struct TerminalGuard { restore_claimed: Arc, restored: bool, #[cfg(windows)] - restore_windows_input_mode: Option, + restore_windows_input_mode: Arc, } pub(super) fn write_host_color_scheme_report_mode( @@ -171,10 +173,59 @@ pub(super) fn should_draw_host_cursor(mode: crate::config::HostCursorModeConfig) pub(super) struct WindowsVirtualTerminalInputSetup { active: bool, restore_mode: Option, + warning: Option<&'static str>, } #[cfg(windows)] -pub(super) fn enable_windows_virtual_terminal_input() -> WindowsVirtualTerminalInputSetup { +#[derive(Default)] +struct WindowsInputModeRestore { + mode: Mutex>, +} + +#[cfg(windows)] +impl WindowsInputModeRestore { + fn lock(&self) -> MutexGuard<'_, Option> { + match self.mode.lock() { + Ok(mode) => mode, + Err(poisoned) => poisoned.into_inner(), + } + } + + fn activate( + &self, + restore_claimed: &AtomicBool, + operation: impl FnOnce() -> WindowsVirtualTerminalInputSetup, + ) -> WindowsVirtualTerminalInputSetup { + let mut restore_mode = self.lock(); + if restore_claimed.load(Ordering::Acquire) { + return WindowsVirtualTerminalInputSetup::default(); + } + let setup = operation(); + if restore_mode.is_none() { + *restore_mode = setup.restore_mode; + } + setup + } + + fn take(&self) -> Option { + self.lock().take() + } +} + +#[cfg(windows)] +fn enable_windows_virtual_terminal_input( + restore_claimed: &AtomicBool, + restore_mode: &WindowsInputModeRestore, +) -> WindowsVirtualTerminalInputSetup { + let setup = restore_mode.activate(restore_claimed, enable_windows_virtual_terminal_input_inner); + if let Some(warning) = setup.warning { + tracing::warn!("{warning}"); + } + setup +} + +#[cfg(windows)] +fn enable_windows_virtual_terminal_input_inner() -> WindowsVirtualTerminalInputSetup { use windows_sys::Win32::Foundation::{HANDLE, INVALID_HANDLE_VALUE}; use windows_sys::Win32::System::Console::{ GetConsoleMode, GetStdHandle, SetConsoleMode, ENABLE_VIRTUAL_TERMINAL_INPUT, @@ -183,44 +234,65 @@ pub(super) fn enable_windows_virtual_terminal_input() -> WindowsVirtualTerminalI let handle: HANDLE = unsafe { GetStdHandle(STD_INPUT_HANDLE) }; if handle.is_null() || handle == INVALID_HANDLE_VALUE { - tracing::warn!("failed to get Windows console input handle for VT input"); - return WindowsVirtualTerminalInputSetup::default(); + return WindowsVirtualTerminalInputSetup { + warning: Some("failed to get Windows console input handle for VT input"), + ..WindowsVirtualTerminalInputSetup::default() + }; } let mut mode = 0; if unsafe { GetConsoleMode(handle, &mut mode) } == 0 { - tracing::warn!("failed to read Windows console input mode for VT input"); - return WindowsVirtualTerminalInputSetup::default(); + return WindowsVirtualTerminalInputSetup { + warning: Some("failed to read Windows console input mode for VT input"), + ..WindowsVirtualTerminalInputSetup::default() + }; } let desired = windows_virtual_terminal_input_mode(mode); if desired == mode { return WindowsVirtualTerminalInputSetup { active: true, - restore_mode: None, + ..WindowsVirtualTerminalInputSetup::default() }; } if unsafe { SetConsoleMode(handle, desired) } == 0 { - tracing::warn!("failed to enable Windows virtual terminal input"); - return WindowsVirtualTerminalInputSetup::default(); + return WindowsVirtualTerminalInputSetup { + warning: Some("failed to enable Windows virtual terminal input"), + ..WindowsVirtualTerminalInputSetup::default() + }; } let mut applied = 0; if unsafe { GetConsoleMode(handle, &mut applied) } == 0 { - tracing::warn!("failed to verify Windows virtual terminal input mode"); - let _ = unsafe { SetConsoleMode(handle, mode) }; - return WindowsVirtualTerminalInputSetup::default(); + let rollback_failed = unsafe { SetConsoleMode(handle, mode) } == 0; + return WindowsVirtualTerminalInputSetup { + restore_mode: rollback_failed.then_some(mode), + warning: Some(if rollback_failed { + "failed to verify or restore Windows virtual terminal input mode" + } else { + "failed to verify Windows virtual terminal input mode" + }), + ..WindowsVirtualTerminalInputSetup::default() + }; } if applied & ENABLE_VIRTUAL_TERMINAL_INPUT == 0 { - tracing::warn!("Windows virtual terminal input bit did not stick"); - let _ = unsafe { SetConsoleMode(handle, mode) }; - return WindowsVirtualTerminalInputSetup::default(); + let rollback_failed = unsafe { SetConsoleMode(handle, mode) } == 0; + return WindowsVirtualTerminalInputSetup { + restore_mode: rollback_failed.then_some(mode), + warning: Some(if rollback_failed { + "Windows virtual terminal input bit did not stick and the prior mode could not be restored" + } else { + "Windows virtual terminal input bit did not stick" + }), + ..WindowsVirtualTerminalInputSetup::default() + }; } WindowsVirtualTerminalInputSetup { active: true, restore_mode: Some(mode), + warning: None, } } @@ -340,11 +412,13 @@ fn restore_terminal_state_once( reset_keyboard_enhancements: bool, reset_modify_other_keys: bool, reset_host_color_scheme_reports: bool, - #[cfg(windows)] restore_windows_input_mode: Option, + #[cfg(windows)] restore_windows_input_mode: &WindowsInputModeRestore, ) -> io::Result<()> { if restore_claimed.swap(true, Ordering::AcqRel) { return Ok(()); } + #[cfg(windows)] + let restore_windows_input_mode = restore_windows_input_mode.take(); restore_terminal_state( reset_keyboard_enhancements, reset_modify_other_keys, @@ -443,6 +517,19 @@ fn disable_windows_win32_input_mode(writer: &mut impl std::io::Write) -> io::Res } impl TerminalGuard { + #[cfg(windows)] + pub(super) fn recover_windows_virtual_terminal_input(&self) -> io::Result<()> { + let active = enable_windows_virtual_terminal_input( + &self.restore_claimed, + &self.restore_windows_input_mode, + ) + .active; + if active && windows_win32_input_mode_enabled() { + enable_windows_win32_input_mode(&mut io::stdout())?; + } + Ok(()) + } + /// Captures the restoration state for use by the process panic hook. pub(super) fn panic_restore(&self) -> impl Fn() + Send + Sync + 'static { let restore_claimed = self.restore_claimed.clone(); @@ -450,7 +537,7 @@ impl TerminalGuard { let reset_modify_other_keys = self.reset_modify_other_keys; let reset_host_color_scheme_reports = self.reset_host_color_scheme_reports; #[cfg(windows)] - let restore_windows_input_mode = self.restore_windows_input_mode; + let restore_windows_input_mode = self.restore_windows_input_mode.clone(); move || { let _ = restore_terminal_state_once( &restore_claimed, @@ -458,7 +545,7 @@ impl TerminalGuard { reset_modify_other_keys, reset_host_color_scheme_reports, #[cfg(windows)] - restore_windows_input_mode, + &restore_windows_input_mode, ); } } @@ -471,7 +558,7 @@ impl TerminalGuard { self.reset_modify_other_keys, self.reset_host_color_scheme_reports, #[cfg(windows)] - self.restore_windows_input_mode, + &self.restore_windows_input_mode, ) } } @@ -485,7 +572,7 @@ impl Drop for TerminalGuard { self.reset_modify_other_keys, self.reset_host_color_scheme_reports, #[cfg(windows)] - self.restore_windows_input_mode, + &self.restore_windows_input_mode, ); } } @@ -495,6 +582,51 @@ impl Drop for TerminalGuard { mod tests { use super::*; + #[cfg(windows)] + #[test] + fn windows_input_mode_restore_tracks_recovery_after_cleanup_callback_creation() { + let restore_claimed = Arc::new(AtomicBool::new(false)); + let restore_mode = Arc::new(WindowsInputModeRestore::default()); + let callback_claimed = restore_claimed.clone(); + let callback_mode = restore_mode.clone(); + let cleanup = move || { + if callback_claimed.swap(true, Ordering::AcqRel) { + None + } else { + callback_mode.take() + } + }; + + let failed = + restore_mode.activate(&restore_claimed, WindowsVirtualTerminalInputSetup::default); + assert!(!failed.active); + let already_active = + restore_mode.activate(&restore_claimed, || WindowsVirtualTerminalInputSetup { + active: true, + ..WindowsVirtualTerminalInputSetup::default() + }); + assert!(already_active.active); + restore_mode.activate(&restore_claimed, || WindowsVirtualTerminalInputSetup { + active: true, + restore_mode: Some(152), + warning: None, + }); + restore_mode.activate(&restore_claimed, || WindowsVirtualTerminalInputSetup { + active: true, + restore_mode: Some(999), + warning: None, + }); + + assert_eq!(cleanup(), Some(152)); + assert_eq!(cleanup(), None); + let called = AtomicBool::new(false); + restore_mode.activate(&restore_claimed, || { + called.store(true, Ordering::Release); + WindowsVirtualTerminalInputSetup::default() + }); + assert!(!called.load(Ordering::Acquire)); + } + #[derive(Clone, Default)] struct SharedOutput(std::rc::Rc>>);