From 9e60973d1aa43d358414456a4fc8756576c74961 Mon Sep 17 00:00:00 2001 From: Jeongkyu Shin Date: Tue, 1 Sep 2026 02:52:41 +0900 Subject: [PATCH] feat(ssh): implement OpenSSH compatibility flags --- src/app/dispatcher.rs | 64 +++++- src/app/query.rs | 12 +- src/cli/bssh.rs | 90 ++++++++- src/cli/mod.rs | 2 +- src/cli/pdsh.rs | 3 + src/cli/ssh_args.rs | 237 ++++++++++++++++++++--- src/commands/interactive/connection.rs | 1 + src/commands/interactive/utils.rs | 1 + src/main.rs | 16 +- src/ssh/client/command.rs | 52 ++++- src/ssh/session_policy.rs | 5 +- src/ssh/session_policy_tests.rs | 16 ++ src/ssh/tokio_client/algorithms.rs | 95 ++++++--- src/ssh/tokio_client/channel_manager.rs | 83 +++++++- src/ssh/tokio_client/connection.rs | 14 +- src/ssh/tokio_client/connection_tests.rs | 25 +++ src/ssh/tokio_client/mod.rs | 1 + tests/forwarding_live_test.rs | 199 +++++++++++++++++++ tests/openssh-regress/run.py | 4 +- tests/openssh-regress/selection.tsv | 2 +- tests/openssh-regress/test_run.py | 2 +- tests/ssh_compat_output_test.rs | 43 ++++ tests/ssh_config_dump_test.rs | 4 +- 23 files changed, 891 insertions(+), 80 deletions(-) diff --git a/src/app/dispatcher.rs b/src/app/dispatcher.rs index 34e1b7e0..1c81785b 100644 --- a/src/app/dispatcher.rs +++ b/src/app/dispatcher.rs @@ -30,7 +30,8 @@ use bssh::{ pty::PtyConfig, security::{Password, get_password, get_sudo_password}, ssh::{ - CliTtyMode, SessionPolicy, SessionRequest, + CliTtyMode, SessionPolicy, SessionRequest, SshClient, + client::ConnectionConfig, tokio_client::{AddressFamily, ProxyMode, SshConnectionConfigResolver}, }, }; @@ -75,6 +76,7 @@ fn build_ssh_connection_config_resolver( cli.remote_forwards.clone(), cli.dynamic_forwards.clone(), ) + .with_stdio_forward(cli.stdio_forward.is_some()) } /// Decide whether `-S` (sudo-password) is meaningful for the given dispatch path. @@ -548,7 +550,7 @@ fn resolve_ssh_mode_interactive_policy( stdin_is_terminal, jump_spec, )?; - Ok(matches!(policy.request, SessionRequest::Shell).then_some(policy)) + Ok((matches!(policy.request, SessionRequest::Shell) && !policy.stdin_null).then_some(policy)) } fn session_policy_jump_spec(proxy_mode: Option<&ProxyMode>) -> Option<&str> { @@ -568,6 +570,48 @@ async fn handle_exec_command( command: &str, ssh_password: Option>, ) -> Result<()> { + if let Some(target) = &cli.stdio_forward { + anyhow::ensure!( + cli.is_ssh_mode() && ctx.nodes.len() == 1, + "-W requires exactly one SSH destination" + ); + let node = ctx + .nodes + .first() + .context("-W requires an SSH destination node")?; + let effective_cluster_name = ctx.cluster_name.as_deref().or(cli.cluster.as_deref()); + let resolver = build_ssh_connection_config_resolver(cli, ctx, effective_cluster_name); + let resolved = resolver.resolve_for_host(node.config_host()); + let key_path = determine_ssh_key_path( + cli, + &ctx.config, + &ctx.ssh_config, + Some(node.config_host()), + effective_cluster_name, + ); + #[cfg(target_os = "macos")] + let use_keychain = determine_use_keychain(&ctx.ssh_config, Some(node.config_host())); + let config = ConnectionConfig { + key_path: key_path.as_deref(), + strict_mode: Some(ctx.strict_mode), + use_agent: cli.use_agent, + use_password: cli.password, + #[cfg(target_os = "macos")] + use_keychain, + timeout_seconds: None, + connect_timeout_seconds: Some(cli.connect_timeout), + jump_hosts_spec: cli.jump_hosts.as_deref(), + ssh_connection_config: Some(&resolved), + ssh_connection_config_resolver: Some(&resolver), + session_policy: None, + ssh_password, + }; + let mut client = SshClient::new(node.host.clone(), node.port, node.username.clone()); + return client + .connect_and_forward_stdio((target.host.clone(), target.port), &config) + .await; + } + // Resolve policy even for a plain ssh-compatible shell. Remote/subsystem/ // none requests stay on the command executor; shell requests retain the // existing interactive stdin, PTY resize, and byte-stream implementation. @@ -821,6 +865,22 @@ mod tests { ); assert_eq!(resolved.local_command.as_deref(), Some("true")); + configured.remote_command = None; + configured.stdin_null = Some(true); + assert!( + resolve_ssh_mode_interactive_policy( + &configured, + &node, + CliTtyMode::Default, + true, + None, + ) + .unwrap() + .is_none(), + "StdinNull shells must use the EOF-capable raw executor" + ); + configured.stdin_null = None; + configured.remote_command = Some("true".into()); assert!( resolve_ssh_mode_interactive_policy( diff --git a/src/app/query.rs b/src/app/query.rs index e677d857..9a8e85c0 100644 --- a/src/app/query.rs +++ b/src/app/query.rs @@ -36,16 +36,20 @@ pub fn is_supported_query(query: &str) -> bool { pub fn handle_query(query: &str) { match query { "cipher" => { - println!("aes128-ctr\naes192-ctr\naes256-ctr"); - println!("aes128-gcm@openssh.com\naes256-gcm@openssh.com"); - println!("chacha20-poly1305@openssh.com"); + println!( + "{}", + bssh::ssh::tokio_client::supported_cipher_names().join("\n") + ); } "cipher-auth" => { println!("aes128-gcm@openssh.com\naes256-gcm@openssh.com"); println!("chacha20-poly1305@openssh.com"); } "mac" => { - println!("hmac-sha2-256\nhmac-sha2-512\nhmac-sha1"); + println!( + "{}", + bssh::ssh::tokio_client::supported_mac_names().join("\n") + ); } "kex" => { println!("curve25519-sha256\ncurve25519-sha256@libssh.org"); diff --git a/src/cli/bssh.rs b/src/cli/bssh.rs index f3fbd272..1bdbc8c5 100644 --- a/src/cli/bssh.rs +++ b/src/cli/bssh.rs @@ -17,6 +17,8 @@ use anyhow::{Context, Result}; use clap::{Parser, Subcommand}; use std::path::PathBuf; +use super::ssh_args::StdioForwardTarget; + #[derive(Parser, Debug)] #[command( name = "bssh", @@ -271,6 +273,8 @@ pub struct Cli { short = 'c', long = "cipher", value_name = "cipher_spec", + allow_hyphen_values = true, + overrides_with = "cipher", help = "Select SSH transport ciphers (OpenSSH-compatible -c)" )] pub cipher: Option, @@ -279,10 +283,35 @@ pub struct Cli { short = 'm', long = "macs", value_name = "mac_spec", + allow_hyphen_values = true, + overrides_with = "macs", help = "Select SSH MAC algorithms (OpenSSH-compatible -m)" )] pub macs: Option, + #[arg( + short = 's', + long = "subsystem", + help = "Invoke the remote command as an SSH subsystem" + )] + pub subsystem: bool, + + #[arg( + short = 'n', + long = "stdin-null", + help = "Redirect stdin from /dev/null" + )] + pub stdin_null: bool, + + #[arg( + short = 'W', + long = "stdio-forward", + value_name = "host:port", + overrides_with = "stdio_forward", + help = "Forward standard input and output to host:port over SSH" + )] + pub stdio_forward: Option, + #[arg( short = 'F', long = "ssh-config", @@ -646,7 +675,9 @@ impl Cli { let mut options = Vec::with_capacity( self.ssh_options.len() + usize::from(self.cipher.is_some()) - + usize::from(self.macs.is_some()), + + usize::from(self.macs.is_some()) + + usize::from(self.subsystem) + + usize::from(self.stdin_null), ); if let Some(cipher) = &self.cipher { options.push(format!("Ciphers={cipher}")); @@ -654,6 +685,12 @@ impl Cli { if let Some(macs) = &self.macs { options.push(format!("MACs={macs}")); } + if self.subsystem { + options.push("SessionType=subsystem".to_string()); + } + if self.stdin_null { + options.push("StdinNull=yes".to_string()); + } options.extend(self.ssh_options.iter().cloned()); options } @@ -908,6 +945,57 @@ mod tests { ] ); } + + #[test] + fn compatibility_flags_parse_repetition_modifiers_and_session_overrides() { + let cli = Cli::try_parse_from([ + "bssh", + "-c", + "aes128-ctr", + "-c", + "-aes128-cbc", + "-m", + "hmac-sha1", + "-m", + "+hmac-sha2-256", + "-sn", + "-W[::1]:443", + "target", + "sftp", + ]) + .unwrap(); + + assert_eq!(cli.cipher.as_deref(), Some("-aes128-cbc")); + assert_eq!(cli.macs.as_deref(), Some("+hmac-sha2-256")); + assert!(cli.subsystem && cli.stdin_null); + assert_eq!(cli.stdio_forward.as_ref().unwrap().host, "::1"); + assert_eq!( + cli.ssh_config_overrides(), + [ + "Ciphers=-aes128-cbc", + "MACs=+hmac-sha2-256", + "SessionType=subsystem", + "StdinNull=yes", + ] + ); + } + + #[test] + fn second_option_pass_preserves_hyphen_leading_command_after_double_dash() { + let argv = ["bssh", "host", "-s", "--", "-literal-command"] + .map(str::to_string) + .to_vec(); + let first = Cli::try_parse_from(&argv).unwrap(); + let normalized = crate::cli::normalize_ssh_option_pass( + &argv, + first.destination.as_deref().unwrap(), + first.command_args.len(), + ); + let parsed = Cli::try_parse_from(normalized).unwrap(); + assert!(parsed.subsystem); + assert_eq!(parsed.destination.as_deref(), Some("host")); + assert_eq!(parsed.command_args, ["-literal-command"]); + } #[test] fn openssh_cipher_flag_rejects_unsupported_and_empty_policies() { for value in ["not-a-supported-cipher", ""] { diff --git a/src/cli/mod.rs b/src/cli/mod.rs index 0445f69c..bbc6db10 100644 --- a/src/cli/mod.rs +++ b/src/cli/mod.rs @@ -42,7 +42,7 @@ mod mode_detection_tests; // Re-export main CLI types from bssh module pub use bssh::{Cli, Commands}; -pub use ssh_args::SshDumpInvocation; +pub use ssh_args::{SshDumpInvocation, StdioForwardTarget, normalize_ssh_option_pass}; // Re-export pdsh compatibility utilities pub use pdsh::{ diff --git a/src/cli/pdsh.rs b/src/cli/pdsh.rs index cbb73020..d39423bd 100644 --- a/src/cli/pdsh.rs +++ b/src/cli/pdsh.rs @@ -319,6 +319,9 @@ impl PdshCli { ssh_options: Vec::new(), cipher: None, macs: None, + subsystem: false, + stdin_null: false, + stdio_forward: None, ssh_config: None, quiet: false, force_tty: false, diff --git a/src/cli/ssh_args.rs b/src/cli/ssh_args.rs index fe0565b4..1341ee7a 100644 --- a/src/cli/ssh_args.rs +++ b/src/cli/ssh_args.rs @@ -8,10 +8,122 @@ //! Order-preserving extraction of ssh_config command-line options. +use std::net::Ipv6Addr; use std::path::PathBuf; +use std::str::FromStr; use anyhow::{Context, Result}; +/// Destination of an OpenSSH-compatible `-W host:port` stdio forwarding +/// request. +#[derive(Debug, Clone, PartialEq, Eq)] +pub struct StdioForwardTarget { + pub host: String, + pub port: u16, +} + +impl FromStr for StdioForwardTarget { + type Err = String; + + fn from_str(value: &str) -> std::result::Result { + let (raw_host, raw_port) = value + .rsplit_once(':') + .ok_or_else(|| format!("Invalid -W target '{value}'; expected host:port"))?; + let host = if let Some(bracketed) = raw_host.strip_prefix('[') { + let host = bracketed + .strip_suffix(']') + .ok_or_else(|| format!("Invalid -W target '{value}'; expected [IPv6]:port"))?; + host.parse::() + .map_err(|_| format!("Invalid -W target '{value}'; expected [IPv6]:port"))?; + host.to_string() + } else if raw_host.is_empty() || raw_host.contains([':', '[', ']']) { + return Err(format!("Invalid -W target '{value}'; expected host:port")); + } else { + raw_host.to_string() + }; + let port = raw_port + .parse::() + .ok() + .filter(|port| *port > 0) + .or_else(|| service_port(raw_port)) + .ok_or_else(|| format!("Invalid -W target '{value}'; expected host:port"))?; + Ok(Self { host, port }) + } +} + +/// Move the OpenSSH second option pass in front of the destination so clap can +/// parse it. `trailing_var_arg` intentionally captures everything after the +/// destination; OpenSSH accepts another option group there until the first +/// remote-command argument. +pub fn normalize_ssh_option_pass( + args: &[String], + destination: &str, + trailing_count: usize, +) -> Vec { + if trailing_count == 0 || args.len() < trailing_count + 2 { + return args.to_vec(); + } + let search_end = args.len() - trailing_count; + let Some(destination_index) = args[..search_end] + .iter() + .rposition(|argument| argument == destination) + else { + return args.to_vec(); + }; + let trailing = &args[destination_index + 1..]; + let mut consumed = 0usize; + while consumed < trailing.len() { + let argument = &trailing[consumed]; + if argument == "--" { + consumed += 1; + break; + } + let Some(width) = scoped_second_pass_width(argument, trailing.get(consumed + 1)) else { + break; + }; + consumed += width; + } + if consumed == 0 { + return args.to_vec(); + } + + let mut normalized = Vec::with_capacity(args.len()); + normalized.extend_from_slice(&args[..destination_index]); + normalized.extend_from_slice(&trailing[..consumed]); + normalized.push(args[destination_index].clone()); + normalized.extend_from_slice(&trailing[consumed..]); + normalized +} + +fn scoped_second_pass_width(argument: &str, next: Option<&String>) -> Option { + if let Some(long) = argument.strip_prefix("--") { + let (name, attached) = long + .split_once('=') + .map_or((long, false), |(name, _)| (name, true)); + return match name { + "subsystem" | "stdin-null" => Some(1), + "cipher" | "macs" | "stdio-forward" | "ssh-config" | "option" => { + Some(usize::from(!attached && next.is_some()) + 1) + } + _ => None, + }; + } + let shorts = argument + .strip_prefix('-') + .filter(|value| !value.is_empty())?; + for (position, short) in shorts.char_indices() { + match short { + 's' | 'n' => {} + 'c' | 'm' | 'W' | 'F' | 'o' => { + let attached = position + short.len_utf8() < shorts.len(); + return Some(usize::from(!attached && next.is_some()) + 1); + } + _ => return None, + } + } + Some(1) +} + /// Inputs needed by `ssh -G`, in the order OpenSSH obtains them. #[derive(Debug, Clone, PartialEq, Eq)] pub struct SshDumpInvocation { @@ -359,7 +471,9 @@ fn apply_value( "remote-forward" => config_option("RemoteForward", value)?, "dynamic-forward" => config_option("DynamicForward", value)?, "stdio-forward" => { - validate_stdio_forward(value)?; + value + .parse::() + .map_err(anyhow::Error::msg)?; return Ok(()); } "diagnostic-file" => { @@ -407,55 +521,38 @@ fn set_priority( } } -fn validate_stdio_forward(value: &str) -> Result<()> { - let (host, port) = value.rsplit_once(':').context("-W requires host:port")?; - let valid_host = !host.is_empty() - && if host.starts_with('[') { - host.ends_with(']') && host.len() > 2 - } else { - !host.contains(':') && !host.ends_with(']') - }; - let valid_port = match port.parse::() { - Ok(port) => port > 0, - Err(_) => service_exists(port), - }; - if !valid_host || !valid_port { - anyhow::bail!("Invalid -W target '{value}'; expected host:port"); - } - Ok(()) -} - #[cfg(unix)] -fn service_exists(name: &str) -> bool { +fn service_port(name: &str) -> Option { if name.is_empty() || !name .chars() .all(|ch| ch.is_ascii_alphanumeric() || matches!(ch, '-' | '_')) { - return false; + return None; } std::fs::read_to_string("/etc/services") .ok() - .is_some_and(|services| { - services.lines().any(|line| { + .and_then(|services| { + services.lines().find_map(|line| { let fields = line .split('#') .next() .unwrap_or_default() .split_whitespace() .collect::>(); - fields.get(1).is_some_and(|port| port.ends_with("/tcp")) - && fields - .iter() - .enumerate() - .any(|(index, field)| index != 1 && *field == name) + let port = fields.get(1)?.strip_suffix("/tcp")?.parse().ok()?; + fields + .iter() + .enumerate() + .any(|(index, field)| index != 1 && *field == name) + .then_some(port) }) }) } #[cfg(not(unix))] -fn service_exists(_name: &str) -> bool { - false +fn service_port(_name: &str) -> Option { + None } fn scan_for_dump_flag(args: &[String]) -> bool { @@ -638,7 +735,7 @@ fn literal_user(value: &str) -> Result { #[cfg(test)] mod tests { - use super::SshDumpInvocation; + use super::{SshDumpInvocation, StdioForwardTarget, normalize_ssh_option_pass}; fn args(values: &[&str]) -> Vec { values.iter().map(|value| (*value).to_string()).collect() @@ -662,6 +759,84 @@ mod tests { ); } + #[test] + fn normalizes_scoped_second_option_pass_for_clap() { + let argv = args(&[ + "bssh", + "host", + "-sn", + "-c", + "-aes128-cbc", + "-W[::1]:443", + "remote-command", + ]); + assert_eq!( + normalize_ssh_option_pass(&argv, "host", 5), + args(&[ + "bssh", + "-sn", + "-c", + "-aes128-cbc", + "-W[::1]:443", + "host", + "remote-command", + ]) + ); + + let terminated = args(&["bssh", "host", "-s", "--", "-literal-command"]); + assert_eq!( + normalize_ssh_option_pass(&terminated, "host", 3), + args(&["bssh", "-s", "--", "host", "-literal-command"]) + ); + + let with_generic_option = args(&[ + "bssh", + "host", + "-oCiphers=aes128-ctr", + "-c", + "aes256-ctr", + "command", + ]); + assert_eq!( + normalize_ssh_option_pass(&with_generic_option, "host", 4), + args(&[ + "bssh", + "-oCiphers=aes128-ctr", + "-c", + "aes256-ctr", + "host", + "command", + ]) + ); + + let missing_value = args(&["bssh", "host", "-c"]); + assert_eq!( + normalize_ssh_option_pass(&missing_value, "host", 1), + args(&["bssh", "-c", "host"]) + ); + } + + #[test] + fn parses_numeric_service_and_bracketed_ipv6_stdio_targets() { + assert_eq!( + "example.com:443".parse::().unwrap(), + StdioForwardTarget { + host: "example.com".into(), + port: 443, + } + ); + assert_eq!( + "[2001:db8::1]:22".parse::().unwrap(), + StdioForwardTarget { + host: "2001:db8::1".into(), + port: 22, + } + ); + for invalid in ["host", ":22", "host:0", "2001:db8::1:22", "[bad]:22"] { + assert!(invalid.parse::().is_err(), "{invalid}"); + } + } + #[test] fn destination_values_are_last_and_ipv6_is_unwrapped() { let argv = args(&["bssh", "-Gp2200", "user@[::1]:2300"]); diff --git a/src/commands/interactive/connection.rs b/src/commands/interactive/connection.rs index ee6e6be2..9e89737f 100644 --- a/src/commands/interactive/connection.rs +++ b/src/commands/interactive/connection.rs @@ -857,6 +857,7 @@ Host beta environment: Vec::new(), local_command: None, request_pty: false, + stdin_null: false, request: SessionRequest::Shell, }; diff --git a/src/commands/interactive/utils.rs b/src/commands/interactive/utils.rs index ae8bc64f..d9fbff11 100644 --- a/src/commands/interactive/utils.rs +++ b/src/commands/interactive/utils.rs @@ -107,6 +107,7 @@ mod tests { environment: vec![("POLICY".into(), "value".into())], local_command: None, request_pty: false, + stdin_null: false, request: crate::ssh::SessionRequest::Shell, }), ssh_connection_config: SshConnectionConfig::default(), diff --git a/src/main.rs b/src/main.rs index 0705765f..07cbd920 100644 --- a/src/main.rs +++ b/src/main.rs @@ -308,7 +308,19 @@ async fn run_bssh_mode(args: &[String]) -> Result<()> { std::process::exit(0); } - let mut cli = Cli::parse(); + let mut cli = Cli::parse_from(args); + let effective_args = if cli.is_ssh_mode() { + bssh::cli::normalize_ssh_option_pass( + args, + cli.destination.as_deref().unwrap_or_default(), + cli.command_args.len(), + ) + } else { + args.to_vec() + }; + if effective_args != args { + cli = Cli::parse_from(&effective_args); + } bssh::ui::configure_color(cli.color); if cli.version { @@ -349,7 +361,7 @@ async fn run_bssh_mode(args: &[String]) -> Result<()> { // Initialize the application and load all configurations. A failure here is // a pre-connection failure, which `ping` reports as 255. - let init_result = initialize_app(&mut cli, args).await; + let init_result = initialize_app(&mut cli, &effective_args).await; let ctx = match init_result { Ok(ctx) => ctx, Err(e) => return Err(map_hard_failure(&cli.command, cli.is_ssh_mode(), e)), diff --git a/src/ssh/client/command.rs b/src/ssh/client/command.rs index ea715362..e7c3c105 100644 --- a/src/ssh/client/command.rs +++ b/src/ssh/client/command.rs @@ -56,6 +56,56 @@ impl SshClient { } } + /// Connect and expose a remote `direct-tcpip` channel on local stdio. + /// + /// This is intentionally separate from command/session execution: `-W` + /// creates no session channel, PTY, RemoteCommand, or command timeout. + pub async fn connect_and_forward_stdio( + &mut self, + target: (String, u16), + config: &ConnectionConfig<'_>, + ) -> Result<()> { + let auth_method = self + .determine_auth_method( + config.key_path, + config.use_agent, + config.use_password, + #[cfg(target_os = "macos")] + config.use_keychain, + config.ssh_password.clone(), + config.ssh_connection_config, + ) + .await?; + let strict_mode = config + .strict_mode + .unwrap_or(StrictHostKeyChecking::AcceptNew); + let client = self + .establish_connection( + &auth_method, + strict_mode, + config.jump_hosts_spec, + config.key_path, + config.use_agent, + config.use_password, + config.connect_timeout_seconds, + config.ssh_connection_config, + config.ssh_connection_config_resolver, + config.ssh_password.clone(), + crate::ssh::SessionPurpose::Bulk, + ) + .await?; + let address_family = config + .ssh_connection_config + .map_or(crate::ssh::tokio_client::AddressFamily::Any, |value| { + value.address_family + }); + let operation = client + .forward_stdio(target, address_family) + .await + .context("Failed to forward standard I/O over SSH"); + self.finish_with_disconnect(&client, operation).await + } + /// Execute a command on the remote host with basic configuration pub async fn connect_and_execute( &mut self, @@ -347,7 +397,7 @@ impl SshClient { forward_stdin: bool, ) -> Result { match session_policy { - Some(policy) if forward_stdin => { + Some(policy) if forward_stdin && !policy.stdin_null => { client .execute_session_streaming_with_stdin(policy, output_sender) .await diff --git a/src/ssh/session_policy.rs b/src/ssh/session_policy.rs index 5c3f9360..9e509ccf 100644 --- a/src/ssh/session_policy.rs +++ b/src/ssh/session_policy.rs @@ -62,6 +62,7 @@ pub struct SessionPolicy { pub environment: Vec<(String, String)>, pub local_command: Option, pub request_pty: bool, + pub stdin_null: bool, pub request: SessionRequest, } @@ -227,10 +228,11 @@ impl SessionPolicy { value => anyhow::bail!("Unsupported SessionType value: {value}"), }; let interactive = matches!(request, SessionRequest::Shell); + let stdin_null = config.stdin_null.unwrap_or(false); let request_pty = resolve_request_pty( cli_tty, config.request_tty.as_deref(), - stdin_is_terminal, + stdin_is_terminal && !stdin_null, interactive, )?; @@ -238,6 +240,7 @@ impl SessionPolicy { environment, local_command, request_pty, + stdin_null, request, }) } diff --git a/src/ssh/session_policy_tests.rs b/src/ssh/session_policy_tests.rs index 8fbfa9c5..5b850640 100644 --- a/src/ssh/session_policy_tests.rs +++ b/src/ssh/session_policy_tests.rs @@ -127,6 +127,22 @@ fn request_tty_obeys_cli_precedence_and_config_modes() { ); } +#[test] +fn stdin_null_disables_automatic_pty_but_not_forced_pty() { + let config = SshHostConfig { + stdin_null: Some(true), + request_tty: Some("yes".into()), + ..Default::default() + }; + let policy = SessionPolicy::resolve(&config, &node(), None, CliTtyMode::Default, true).unwrap(); + assert!(policy.stdin_null); + assert!(!policy.request_pty); + + let forced = SessionPolicy::resolve(&config, &node(), None, CliTtyMode::Force, true).unwrap(); + assert!(forced.stdin_null); + assert!(forced.request_pty); +} + #[test] fn session_purpose_tracks_the_resolved_pty_policy() { let interactive = SessionPolicy::resolve( diff --git a/src/ssh/tokio_client/algorithms.rs b/src/ssh/tokio_client/algorithms.rs index 4e24dcd0..0a3d2018 100644 --- a/src/ssh/tokio_client/algorithms.rs +++ b/src/ssh/tokio_client/algorithms.rs @@ -60,26 +60,6 @@ fn expand_supported( Ok(expanded) } -fn remove_matching( - defaults: &[T], - patterns: &[&str], - name: impl Fn(&T) -> &str, -) -> Result, String> { - let patterns = patterns - .iter() - .map(|pattern| Pattern::new(pattern).map_err(|error| error.to_string())) - .collect::, _>>()?; - Ok(defaults - .iter() - .filter(|algorithm| { - !patterns - .iter() - .any(|pattern| pattern.matches(name(algorithm))) - }) - .cloned() - .collect()) -} - fn combine(defaults: &[T], configured: Vec, mode: ListMode) -> Vec { match mode { ListMode::Replace => configured, @@ -123,8 +103,19 @@ fn resolve_policy( name: impl Fn(&T) -> &str + Copy, ) -> Result, String> { let (mode, names) = list_mode(values); + if names.is_empty() { + return Err(format!( + "empty {kind} policy; supported values: {}", + supported.iter().map(name).collect::>().join(",") + )); + } let resolved = if mode == ListMode::Remove { - remove_matching(defaults, &names, name)? + let removed = expand_supported(kind, &names, supported, name)?; + defaults + .iter() + .filter(|algorithm| !removed.contains(algorithm)) + .cloned() + .collect() } else { let configured = expand_supported(kind, &names, supported, name)?; combine(defaults, configured, mode) @@ -135,11 +126,40 @@ fn resolve_policy( Ok(resolved) } -pub(crate) fn resolve_ciphers(values: &[String]) -> Result, String> { - let supported = cipher::ALL_CIPHERS +fn selectable_ciphers() -> Vec { + cipher::ALL_CIPHERS .iter() .map(|value| **value) - .collect::>(); + .filter(|value| !matches!(value.as_ref(), "clear" | "none")) + .collect() +} + +fn selectable_macs() -> Vec { + mac::ALL_MAC_ALGORITHMS + .iter() + .map(|value| **value) + .filter(|value| value.as_ref() != "none") + .collect() +} + +#[must_use] +pub fn supported_cipher_names() -> Vec { + selectable_ciphers() + .iter() + .map(|value| value.as_ref().to_string()) + .collect() +} + +#[must_use] +pub fn supported_mac_names() -> Vec { + selectable_macs() + .iter() + .map(|value| value.as_ref().to_string()) + .collect() +} + +pub(crate) fn resolve_ciphers(values: &[String]) -> Result, String> { + let supported = selectable_ciphers(); resolve_policy( "cipher", values, @@ -150,10 +170,7 @@ pub(crate) fn resolve_ciphers(values: &[String]) -> Result, St } pub(crate) fn resolve_macs(values: &[String]) -> Result, String> { - let supported = mac::ALL_MAC_ALGORITHMS - .iter() - .map(|value| **value) - .collect::>(); + let supported = selectable_macs(); resolve_policy( "MAC", values, @@ -273,6 +290,28 @@ mod tests { assert!(error.contains("selects no supported algorithms")); } + #[test] + fn removal_and_internal_algorithms_fail_closed() { + for policy in ["-not-real", "+", "-", "^", "clear", "none"] { + let error = resolve_ciphers(&[policy.to_string()]).unwrap_err(); + assert!(error.contains("supported values"), "{policy}: {error}"); + } + assert!(resolve_macs(&["none".to_string()]).is_err()); + assert!( + !supported_cipher_names() + .iter() + .any(|name| name == "clear" || name == "none") + ); + assert!(!supported_mac_names().iter().any(|name| name == "none")); + + // A supported non-default removal is valid and simply leaves the + // default preference unchanged. + assert_eq!( + resolve_ciphers(&["-aes128-cbc".to_string()]).unwrap(), + russh::Preferred::DEFAULT.cipher.as_ref() + ); + } + #[test] fn configured_kex_preserves_protocol_extension_markers() { let resolved = resolve_kex(&[String::from("curve25519-sha256")]).unwrap(); diff --git a/src/ssh/tokio_client/channel_manager.rs b/src/ssh/tokio_client/channel_manager.rs index ae5236c8..1ce2da66 100644 --- a/src/ssh/tokio_client/channel_manager.rs +++ b/src/ssh/tokio_client/channel_manager.rs @@ -25,7 +25,7 @@ use russh::Channel; use russh::client::Msg; use std::io; use std::net::SocketAddr; -use tokio::io::{AsyncRead, AsyncReadExt}; +use tokio::io::{AsyncRead, AsyncReadExt, AsyncWrite, AsyncWriteExt}; use tokio::sync::mpsc::{Receiver, Sender, channel}; use tokio::task::JoinHandle; @@ -315,6 +315,87 @@ impl Client { Err(connect_err) } + /// Forward process stdin/stdout over a `direct-tcpip` channel. + /// + /// Local EOF half-closes only the sending side. The remote side is still + /// drained to stdout until the SSH channel closes, which is required for + /// protocols that send their final response after consuming request EOF. + pub async fn forward_stdio( + &self, + target: (String, u16), + address_family: AddressFamily, + ) -> Result<(), super::Error> { + self.forward_stdio_with_io( + target, + address_family, + tokio::io::stdin(), + tokio::io::stdout(), + ) + .await + } + + /// Forward arbitrary asynchronous input/output over a `direct-tcpip` + /// channel. The two directions are pumped independently so SSH flow + /// control in one direction cannot block progress in the other. + pub async fn forward_stdio_with_io( + &self, + target: (String, u16), + address_family: AddressFamily, + mut input: R, + mut output: W, + ) -> Result<(), super::Error> + where + R: AsyncRead + Unpin, + W: AsyncWrite + Unpin, + { + let channel = self + .open_direct_tcpip_channel_with_family(target, None, address_family) + .await?; + let (mut channel_read, channel_write) = channel.split(); + + let upload = async { + let mut writer = channel_write.make_writer(); + tokio::io::copy(&mut input, &mut writer) + .await + .map_err(super::Error::IoError)?; + writer.flush().await.map_err(super::Error::IoError)?; + drop(writer); + channel_write.eof().await?; + Ok::<(), super::Error>(()) + }; + let download = async { + while let Some(message) = channel_read.wait().await { + match message { + russh::ChannelMsg::Data { data } => { + output + .write_all(&data) + .await + .map_err(super::Error::IoError)?; + output.flush().await.map_err(super::Error::IoError)?; + } + // Remote EOF only half-closes remote-to-local traffic. + // Keep pumping stdin until local EOF or a full Close. + russh::ChannelMsg::Eof => {} + russh::ChannelMsg::Close => break, + // direct-tcpip is a single byte stream. Extended data is + // not part of that transport and must never contaminate + // stdout. + _ => {} + } + } + output.flush().await.map_err(super::Error::IoError) + }; + tokio::pin!(upload, download); + tokio::select! { + biased; + result = &mut download => result, + result = &mut upload => { + result?; + download.await + } + } + } + /// Execute a remote command via the ssh connection with streaming output. /// /// This method sends command output in real-time to the provided sender channel. diff --git a/src/ssh/tokio_client/connection.rs b/src/ssh/tokio_client/connection.rs index ba8f44eb..d05e6f9e 100644 --- a/src/ssh/tokio_client/connection.rs +++ b/src/ssh/tokio_client/connection.rs @@ -223,6 +223,7 @@ pub struct SshConnectionConfigResolver { cli_remote_forwards: Vec, cli_dynamic_forwards: Vec, cli_forwarding_order: Vec, + stdio_forward: bool, } impl SshConnectionConfigResolver { @@ -328,6 +329,15 @@ impl SshConnectionConfigResolver { self } + /// Apply OpenSSH's `-W` defaults after all explicit ssh_config values + /// have been resolved. Explicit `ClearAllForwardings=no` and + /// `ExitOnForwardFailure=no` therefore keep their documented precedence. + #[must_use] + pub fn with_stdio_forward(mut self, enabled: bool) -> Self { + self.stdio_forward = enabled; + self + } + pub fn resolve_for_host(&self, hostname: &str) -> SshConnectionConfig { if let Some(config) = &self.fixed_config { return config.clone(); @@ -549,11 +559,11 @@ impl SshConnectionConfigResolver { clear_all: host_config .as_ref() .and_then(|config| config.clear_all_forwardings) - .unwrap_or(false), + .unwrap_or(self.stdio_forward), exit_on_failure: host_config .as_ref() .and_then(|config| config.exit_on_forward_failure) - .unwrap_or(false), + .unwrap_or(self.stdio_forward), address_family, }; let rekey_limit = host_config diff --git a/src/ssh/tokio_client/connection_tests.rs b/src/ssh/tokio_client/connection_tests.rs index 536baef7..470a20dc 100644 --- a/src/ssh/tokio_client/connection_tests.rs +++ b/src/ssh/tokio_client/connection_tests.rs @@ -204,6 +204,31 @@ Host target assert_eq!(config.keepalive_max, 9); } +#[test] +fn stdio_forward_defaults_clear_forwards_but_preserves_explicit_no() { + let implicit = + SshConfig::parse("Host target\n LocalForward 127.0.0.1:2200 example.com:22\n").unwrap(); + let implicit = SshConnectionConfigResolver::new() + .with_ssh_config(Some(implicit)) + .with_stdio_forward(true) + .resolve_for_host("target"); + assert!(implicit.forwarding_plan.clear_all); + assert!(implicit.forwarding_plan.exit_on_failure); + assert!(implicit.forwarding_plan.parse().unwrap().is_empty()); + + let explicit = SshConfig::parse( + "Host target\n ClearAllForwardings no\n ExitOnForwardFailure no\n LocalForward 127.0.0.1:2200 example.com:22\n", + ) + .unwrap(); + let explicit = SshConnectionConfigResolver::new() + .with_ssh_config(Some(explicit)) + .with_stdio_forward(true) + .resolve_for_host("target"); + assert!(!explicit.forwarding_plan.clear_all); + assert!(!explicit.forwarding_plan.exit_on_failure); + assert_eq!(explicit.forwarding_plan.parse().unwrap().len(), 1); +} + /// The empty-after-filter path must fail with a specific error naming the host /// and the requested family, not the generic "could not resolve to any /// addresses". Passing a `SocketAddr` makes the candidate list exactly one diff --git a/src/ssh/tokio_client/mod.rs b/src/ssh/tokio_client/mod.rs index 7d616274..33d7c19d 100644 --- a/src/ssh/tokio_client/mod.rs +++ b/src/ssh/tokio_client/mod.rs @@ -36,6 +36,7 @@ mod to_socket_addrs_with_hostname; // Re-export public API types for backward compatibility pub use address_family::AddressFamily; +pub use algorithms::{supported_cipher_names, supported_mac_names}; pub use auth_policy::SshAuthenticationPolicy; pub use authentication::{AuthKeyboardInteractive, AuthMethod, ServerCheckMethod}; pub use channel_manager::{CommandExecutedResult, CommandOutput}; diff --git a/tests/forwarding_live_test.rs b/tests/forwarding_live_test.rs index cf3dcccb..80149ce8 100644 --- a/tests/forwarding_live_test.rs +++ b/tests/forwarding_live_test.rs @@ -1,5 +1,6 @@ use std::collections::HashMap; use std::net::{Ipv4Addr, SocketAddr}; +use std::process::Stdio; use std::sync::atomic::{AtomicBool, AtomicU16, AtomicUsize, Ordering}; use std::sync::{Arc, Mutex}; use std::time::Duration; @@ -19,6 +20,7 @@ use russh::server::{self, Msg, Server, Session}; use russh::{Channel, ChannelId, ChannelOpenFailure}; use tokio::io::{AsyncReadExt, AsyncWriteExt}; use tokio::net::{TcpListener, TcpStream}; +use tokio::process::Command; use tokio::sync::Mutex as AsyncMutex; use tokio::task::JoinHandle; use tokio::time::{sleep, timeout}; @@ -140,6 +142,7 @@ impl server::Handler for ForwardingServer { reply: server::ChannelOpenHandle, _session: &mut Session, ) -> Result<(), Self::Error> { + self.state.record("direct-tcpip"); let Ok(tcp) = TcpStream::connect((host_to_connect, port_to_connect as u16)).await else { reply.reject(ChannelOpenFailure::ConnectFailed).await; return Ok(()); @@ -239,6 +242,202 @@ impl server::Handler for ForwardingServer { } } +#[tokio::test] +async fn stdio_forward_is_byte_transparent_and_preserves_half_close() { + let ssh = TestSshServer::start().await; + let echo = EchoServer::start().await; + let config = SshConnectionConfig::new().with_address_family(AddressFamily::V4); + let client = Client::connect_with_ssh_config( + ssh.address, + "test", + AuthMethod::with_password("test"), + ServerCheckMethod::NoCheck, + &config, + ) + .await + .expect("connect stdio forwarding client"); + + let payload = b"stdio\0forward\nbytes".to_vec(); + let (mut input_writer, input_reader) = tokio::io::duplex(64); + let (output_writer, mut output_reader) = tokio::io::duplex(64); + let write_payload = payload.clone(); + let writer = async move { + input_writer.write_all(&write_payload).await.unwrap(); + input_writer.shutdown().await.unwrap(); + }; + let forward = client.forward_stdio_with_io( + (echo.address.ip().to_string(), echo.address.port()), + AddressFamily::V4, + input_reader, + output_writer, + ); + let reader = async { + let mut output = Vec::new(); + output_reader.read_to_end(&mut output).await.unwrap(); + output + }; + let (_, forwarded, output) = timeout(TEST_TIMEOUT, async { + tokio::join!(writer, forward, reader) + }) + .await + .expect("stdio forward timed out"); + forwarded.expect("stdio forward failed"); + assert_eq!(output, payload); + assert_eq!(ssh.state.authentications.load(Ordering::SeqCst), 1); + assert!( + ssh.state + .events() + .iter() + .any(|event| event == "direct-tcpip") + ); + assert!( + !ssh.state + .events() + .iter() + .any(|event| matches!(event.as_str(), "session" | "pty" | "exec")) + ); + + client.disconnect().await.expect("disconnect stdio client"); + echo.shutdown().await; + ssh.shutdown().await; +} + +#[tokio::test] +async fn stdio_forward_keeps_uploading_after_remote_half_close() { + let listener = TcpListener::bind((Ipv4Addr::LOCALHOST, 0)) + .await + .expect("bind half-close target"); + let target = listener.local_addr().expect("half-close target address"); + let (remote_eof_tx, remote_eof_rx) = tokio::sync::oneshot::channel(); + let (received_tx, received_rx) = tokio::sync::oneshot::channel(); + let target_task = tokio::spawn(async move { + let (mut stream, _) = listener.accept().await.expect("accept forwarded stream"); + stream.write_all(b"remote-eof").await.unwrap(); + stream.shutdown().await.unwrap(); + remote_eof_tx.send(()).unwrap(); + let mut received = Vec::new(); + stream.read_to_end(&mut received).await.unwrap(); + received_tx.send(received).unwrap(); + }); + + let ssh = TestSshServer::start().await; + let client = Client::connect_with_ssh_config( + ssh.address, + "test", + AuthMethod::with_password("test"), + ServerCheckMethod::NoCheck, + &SshConnectionConfig::new().with_address_family(AddressFamily::V4), + ) + .await + .expect("connect half-close forwarding client"); + let payload = b"upload-after-remote-eof".to_vec(); + let (mut input_writer, input_reader) = tokio::io::duplex(64); + let (output_writer, mut output_reader) = tokio::io::duplex(64); + let forward = client.forward_stdio_with_io( + (target.ip().to_string(), target.port()), + AddressFamily::V4, + input_reader, + output_writer, + ); + let drive = async { + remote_eof_rx.await.expect("target sent remote EOF"); + input_writer.write_all(&payload).await.unwrap(); + input_writer.shutdown().await.unwrap(); + let mut output = Vec::new(); + output_reader.read_to_end(&mut output).await.unwrap(); + let received = received_rx.await.expect("target received upload"); + (output, received) + }; + let (forwarded, (output, received)) = + timeout(TEST_TIMEOUT, async { tokio::join!(forward, drive) }) + .await + .expect("remote half-close forwarding timed out"); + forwarded.expect("remote half-close forwarding failed"); + assert_eq!(output, b"remote-eof"); + assert_eq!(received, payload); + + target_task.await.expect("half-close target task"); + client.disconnect().await.expect("disconnect stdio client"); + ssh.shutdown().await; +} + +#[tokio::test] +async fn bssh_stdio_forward_is_a_working_proxy_command_transport() { + let bastion = TestSshServer::start().await; + let target = TestSshServer::start().await; + let binary = env!("CARGO_BIN_EXE_bssh"); + let mut config_file = tempfile::NamedTempFile::new().expect("temporary ssh config"); + let config = format!( + "Host target\n\ + HostName 127.0.0.1\n\ + Port {}\n\ + User test\n\ + StrictHostKeyChecking no\n\ + UserKnownHostsFile /dev/null\n\ + ProxyCommand {binary} --password -q -p {} -oStrictHostKeyChecking=no -oUserKnownHostsFile=/dev/null -W %h:%p test@127.0.0.1\n", + target.address.port(), + bastion.address.port(), + ); + std::io::Write::write_all(config_file.as_file_mut(), config.as_bytes()) + .expect("write ssh config"); + + let output = timeout( + TEST_TIMEOUT, + Command::new(binary) + .args([ + "--password", + "-q", + "-F", + config_file.path().to_str().unwrap(), + "target", + "true", + ]) + .env("BSSH_PASSWORD", "test") + .kill_on_drop(true) + .stdin(Stdio::null()) + .stdout(Stdio::piped()) + .stderr(Stdio::piped()) + .output(), + ) + .await + .unwrap_or_else(|_| { + panic!( + "nested ProxyCommand timed out; bastion={:?}, target={:?}", + bastion.state.events(), + target.state.events() + ) + }) + .expect("run bssh through ProxyCommand"); + assert!( + output.status.success(), + "status={:?}, stdout={}, stderr={}", + output.status.code(), + String::from_utf8_lossy(&output.stdout), + String::from_utf8_lossy(&output.stderr), + ); + assert!(output.stdout.is_empty(), "ProxyCommand polluted stdout"); + assert!( + bastion + .state + .events() + .iter() + .any(|event| event == "direct-tcpip") + ); + assert!( + !bastion + .state + .events() + .iter() + .any(|event| matches!(event.as_str(), "session" | "pty" | "exec")) + ); + let target_events = target.state.events(); + assert!(target_events.iter().any(|event| event == "session")); + assert!(target_events.iter().any(|event| event == "exec")); + + target.shutdown().await; + bastion.shutdown().await; +} + struct TestSshServer { address: SocketAddr, state: Arc, diff --git a/tests/openssh-regress/run.py b/tests/openssh-regress/run.py index 98f84b2d..7d759774 100644 --- a/tests/openssh-regress/run.py +++ b/tests/openssh-regress/run.py @@ -102,9 +102,9 @@ def read_selection(path: Path) -> list[Selection]: if len(names) != len(set(names)): raise ValueError("selection.tsv contains duplicate test names") candidate_count = sum(row.disposition in {"run", "skip"} for row in rows) - if candidate_count != 89: + if candidate_count != 90: raise ValueError( - f"selection.tsv must contain 89 candidate tests, found {candidate_count}" + f"selection.tsv must contain 90 candidate tests, found {candidate_count}" ) return rows diff --git a/tests/openssh-regress/selection.tsv b/tests/openssh-regress/selection.tsv index d8b6e8b4..0a6c710c 100644 --- a/tests/openssh-regress/selection.tsv +++ b/tests/openssh-regress/selection.tsv @@ -64,7 +64,7 @@ krl skip Permanent candidate skip: KRL generation is outside the compatibility t limit-keytype run localcommand run login-timeout run -match-subsystem exclude sshd-side subsystem matching test; outside the client candidate set. +match-subsystem run Exercises the client -s subsystem request and preserves remote exit status. multiplex run multipubkey run password run diff --git a/tests/openssh-regress/test_run.py b/tests/openssh-regress/test_run.py index 267d2b0f..274db07d 100755 --- a/tests/openssh-regress/test_run.py +++ b/tests/openssh-regress/test_run.py @@ -24,7 +24,7 @@ class ManifestTests(unittest.TestCase): def test_committed_manifest_is_valid(self) -> None: selection = openssh_regress.read_selection(openssh_regress.DEFAULT_SELECTION) - self.assertEqual(sum(row.disposition == "run" for row in selection), 78) + self.assertEqual(sum(row.disposition == "run" for row in selection), 79) self.assertEqual(sum(row.disposition == "skip" for row in selection), 11) self.assertTrue(all(row.reason for row in selection if row.disposition != "run")) self.assertEqual( diff --git a/tests/ssh_compat_output_test.rs b/tests/ssh_compat_output_test.rs index 07cbbab1..66eeb761 100644 --- a/tests/ssh_compat_output_test.rs +++ b/tests/ssh_compat_output_test.rs @@ -41,6 +41,49 @@ fn version_matches_openssh_stream_contract() { ); } +#[test] +fn algorithm_queries_match_the_selectable_transport_surface() { + let ciphers = bssh().args(["-Q", "cipher"]).output().unwrap(); + let macs = bssh().args(["-Q", "mac"]).output().unwrap(); + assert!(ciphers.status.success() && macs.status.success()); + let ciphers = String::from_utf8(ciphers.stdout).unwrap(); + let macs = String::from_utf8(macs.stdout).unwrap(); + assert!(ciphers.lines().any(|value| value == "aes128-cbc")); + assert!( + macs.lines() + .any(|value| value == "hmac-sha2-256-etm@openssh.com") + ); + assert!( + !ciphers + .lines() + .any(|value| matches!(value, "clear" | "none")) + ); + assert!(!macs.lines().any(|value| value == "none")); +} + +#[test] +fn unsupported_algorithm_policies_fail_before_connecting_with_supported_values() { + for (flag, policy, kind) in [ + ("-c", "-definitely-not-a-cipher", "cipher"), + ("-m", "-definitely-not-a-mac", "mac"), + ("-c", "+", "cipher"), + ("-m", "^", "mac"), + ] { + let output = bssh() + .args([flag, policy, "unresolvable.invalid", "true"]) + .output() + .unwrap(); + let stderr = String::from_utf8_lossy(&output.stderr).to_ascii_lowercase(); + assert_eq!(output.status.code(), Some(1), "{flag} {policy}: {stderr}"); + assert!(output.stdout.is_empty()); + assert!(stderr.contains(kind), "{flag} {policy}: {stderr}"); + assert!( + stderr.contains("supported values"), + "{flag} {policy}: {stderr}" + ); + } +} + #[test] fn explicit_and_environment_color_controls_apply_to_stdout() { let never = bssh() diff --git a/tests/ssh_config_dump_test.rs b/tests/ssh_config_dump_test.rs index 194393c1..3eaa144e 100644 --- a/tests/ssh_config_dump_test.rs +++ b/tests/ssh_config_dump_test.rs @@ -332,9 +332,9 @@ fn raw_dump_dispatch_treats_bssh_subcommand_names_as_destinations() { } #[test] -fn normal_mode_stdio_forward_fails_closed_before_connecting() { +fn normal_mode_stdio_forward_connection_failure_is_ssh_level_and_stdout_clean() { let output = run(&["-W", "localhost:22", "host"]); - assert!(!output.status.success()); + assert_eq!(output.status.code(), Some(255)); assert!(output.stdout.is_empty()); }