From 1e4bd36a459a5dde2ce4d22d5fec77f59ba04c43 Mon Sep 17 00:00:00 2001 From: Jeongkyu Shin Date: Fri, 28 Aug 2026 15:40:12 +0900 Subject: [PATCH 1/6] feat(ssh): enforce source binding policy --- Cargo.toml | 2 +- src/commands/interactive/connection.rs | 11 +- src/jump/chain.rs | 11 + src/ssh/client/command.rs | 9 + src/ssh/client/connection.rs | 10 + src/ssh/client/file_transfer.rs | 1 + src/ssh/mod.rs | 2 +- src/ssh/session_policy.rs | 17 + src/ssh/session_policy_tests.rs | 29 ++ src/ssh/ssh_config/ip_qos.rs | 166 ++++++++ src/ssh/ssh_config/mod.rs | 31 +- .../ssh_config/parser/options/connection.rs | 116 +----- src/ssh/ssh_config/parser/options/support.rs | 16 +- src/ssh/ssh_config/resolver.rs | 2 +- src/ssh/ssh_config/types.rs | 4 +- src/ssh/tokio_client/connection.rs | 385 +++++++++++++++++- src/ssh/tokio_client/connection_tests.rs | 148 ++++++- src/ssh/tokio_client/error.rs | 42 ++ 18 files changed, 862 insertions(+), 140 deletions(-) create mode 100644 src/ssh/ssh_config/ip_qos.rs diff --git a/Cargo.toml b/Cargo.toml index d90a2527..d8626808 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -71,7 +71,7 @@ smallvec = "1.15.1" lru = "0.18.0" uuid = { version = "1.23.1", features = ["v4"] } tokio-util = "0.7.18" -socket2 = "0.6" +socket2 = { version = "0.6", features = ["all"] } shell-words = "1.1.1" base64 = "0.22.1" hmac = "0.13.0" diff --git a/src/commands/interactive/connection.rs b/src/commands/interactive/connection.rs index be14602c..c90b7919 100644 --- a/src/commands/interactive/connection.rs +++ b/src/commands/interactive/connection.rs @@ -25,7 +25,7 @@ use zeroize::Zeroizing; use crate::jump::{JumpHostChain, parse_jump_hosts}; use crate::node::Node; use crate::ssh::{ - SessionRequest, + SessionPurpose, SessionRequest, known_hosts::get_check_method_for_target, tokio_client::{AuthMethod, Client, Error as SshError, ServerCheckMethod, SshConnectionConfig}, }; @@ -53,6 +53,9 @@ impl InteractiveCommand { ) -> Result { const SSH_CONNECT_TIMEOUT_SECS: u64 = 30; let connect_timeout = Duration::from_secs(SSH_CONNECT_TIMEOUT_SECS); + let ssh_config = ssh_config + .clone() + .with_session_purpose(SessionPurpose::Interactive); // SECURITY: Add a small delay before connection attempts to prevent rapid-fire attempts // This helps mitigate brute-force attacks and prevents triggering fail2ban too quickly @@ -72,7 +75,7 @@ impl InteractiveCommand { username, auth_method, check_method.clone(), - ssh_config, + &ssh_config, ), ) .await @@ -116,7 +119,7 @@ impl InteractiveCommand { username, password_auth, check_method, - ssh_config, + &ssh_config, ), ) .await @@ -316,6 +319,7 @@ impl InteractiveCommand { .with_connect_timeout(adjusted_timeout) .with_command_timeout(Duration::from_secs(300)) .with_ssh_connection_config(self.ssh_connection_config.clone()) + .with_session_purpose(SessionPurpose::Interactive) .with_ssh_password(self.ssh_password.clone()); // Connect through the chain @@ -476,6 +480,7 @@ impl InteractiveCommand { .with_connect_timeout(adjusted_timeout) .with_command_timeout(Duration::from_secs(300)) .with_ssh_connection_config(self.ssh_connection_config.clone()) + .with_session_purpose(SessionPurpose::Interactive) .with_ssh_password(self.ssh_password.clone()); // Connect through the chain diff --git a/src/jump/chain.rs b/src/jump/chain.rs index 2a8aeef4..3d0eab51 100644 --- a/src/jump/chain.rs +++ b/src/jump/chain.rs @@ -25,6 +25,7 @@ use super::connection::JumpHostConnection; use super::parser::{JumpHost, get_max_jump_hosts}; use super::rate_limiter::ConnectionRateLimiter; use crate::security::Password; +use crate::ssh::SessionPurpose; use crate::ssh::known_hosts::StrictHostKeyChecking; use crate::ssh::tokio_client::{ AuthMethod, Error as SshError, SshConnectionConfig, SshConnectionConfigResolver, @@ -70,6 +71,8 @@ pub struct JumpHostChain { ssh_connection_config: SshConnectionConfig, /// Per-host SSH connection configuration resolver. ssh_connection_config_resolver: Option, + /// Traffic profile propagated to every direct socket in the chain. + session_purpose: SessionPurpose, /// Pre-collected SSH password (from the dispatcher's single up-front prompt). /// When `use_password` is set on a per-call basis, this is consumed by every /// jump-host auth step instead of prompting per-call, which would otherwise @@ -111,6 +114,7 @@ impl JumpHostChain { max_connection_age: Duration::from_secs(1800), // 30 minutes ssh_connection_config: SshConnectionConfig::default(), ssh_connection_config_resolver: None, + session_purpose: SessionPurpose::Bulk, ssh_password: None, } } @@ -143,11 +147,18 @@ impl JumpHostChain { self } + #[must_use] + pub fn with_session_purpose(mut self, purpose: SessionPurpose) -> Self { + self.session_purpose = purpose; + self + } + fn connection_config_for_host(&self, host: &str) -> SshConnectionConfig { self.ssh_connection_config_resolver .as_ref() .map(|resolver| resolver.resolve_for_host(host)) .unwrap_or_else(|| self.ssh_connection_config.clone()) + .with_session_purpose(self.session_purpose) } /// Create a direct connection chain (no jump hosts) diff --git a/src/ssh/client/command.rs b/src/ssh/client/command.rs index e203d216..80c7e7ee 100644 --- a/src/ssh/client/command.rs +++ b/src/ssh/client/command.rs @@ -147,6 +147,9 @@ impl SshClient { config.ssh_connection_config, config.ssh_connection_config_resolver, config.ssh_password.clone(), + config + .session_policy + .map_or(crate::ssh::SessionPurpose::Bulk, |policy| policy.purpose()), ) .await?; @@ -284,6 +287,9 @@ impl SshClient { config.ssh_connection_config, config.ssh_connection_config_resolver, config.ssh_password.clone(), + config + .session_policy + .map_or(crate::ssh::SessionPurpose::Bulk, |policy| policy.purpose()), ) .await?; @@ -427,6 +433,9 @@ impl SshClient { config.ssh_connection_config, config.ssh_connection_config_resolver, config.ssh_password.clone(), + config + .session_policy + .map_or(crate::ssh::SessionPurpose::Bulk, |policy| policy.purpose()), ) .await?; diff --git a/src/ssh/client/connection.rs b/src/ssh/client/connection.rs index fafe8f6e..38ce8d48 100644 --- a/src/ssh/client/connection.rs +++ b/src/ssh/client/connection.rs @@ -15,6 +15,7 @@ use super::core::SshClient; use crate::jump::{JumpHostChain, parse_jump_hosts}; use crate::security::Password; +use crate::ssh::SessionPurpose; use crate::ssh::known_hosts::StrictHostKeyChecking; use crate::ssh::tokio_client::{ AuthMethod, Client, ProxyMode, SshConnectionConfig, SshConnectionConfigResolver, @@ -236,6 +237,7 @@ impl SshClient { ssh_connection_config: Option<&SshConnectionConfig>, ssh_connection_config_resolver: Option<&SshConnectionConfigResolver>, pre_collected_password: Option>, + session_purpose: SessionPurpose, ) -> Result { // Create jump host chain with user-specified or default connect timeout let connect_timeout = @@ -250,6 +252,7 @@ impl SshClient { if let Some(resolver) = ssh_connection_config_resolver { chain = chain.with_ssh_connection_config_resolver(resolver.clone()); } + chain = chain.with_session_purpose(session_purpose); // Connect through the chain let connection = chain @@ -293,7 +296,13 @@ impl SshClient { ssh_connection_config: Option<&SshConnectionConfig>, ssh_connection_config_resolver: Option<&SshConnectionConfigResolver>, pre_collected_password: Option>, + session_purpose: SessionPurpose, ) -> Result { + let selected_config = ssh_connection_config + .cloned() + .unwrap_or_default() + .with_session_purpose(session_purpose); + let ssh_connection_config = Some(&selected_config); let jump_hosts_spec = match ssh_connection_config.and_then(|config| config.proxy_mode.as_ref()) { Some(ProxyMode::Jump(jump)) => Some(jump.as_str()), @@ -340,6 +349,7 @@ impl SshClient { ssh_connection_config, ssh_connection_config_resolver, pre_collected_password, + session_purpose, ) .await } diff --git a/src/ssh/client/file_transfer.rs b/src/ssh/client/file_transfer.rs index 904cece5..c5c9837c 100644 --- a/src/ssh/client/file_transfer.rs +++ b/src/ssh/client/file_transfer.rs @@ -798,6 +798,7 @@ impl SshClient { Some(ssh_connection_config), ssh_connection_config_resolver, pre_collected_password, + crate::ssh::SessionPurpose::Bulk, ) .await } diff --git a/src/ssh/mod.rs b/src/ssh/mod.rs index d6d004bf..57844836 100644 --- a/src/ssh/mod.rs +++ b/src/ssh/mod.rs @@ -30,5 +30,5 @@ pub use client::SshClient; pub use config_cache::{CacheConfig, CacheStats, GLOBAL_CACHE, SshConfigCache}; pub use handler::BsshHandler; pub use pool::ConnectionPool; -pub use session_policy::{CliTtyMode, SessionPolicy, SessionRequest}; +pub use session_policy::{CliTtyMode, SessionPolicy, SessionPurpose, SessionRequest}; pub use ssh_config::{SshConfig, SshHostConfig}; diff --git a/src/ssh/session_policy.rs b/src/ssh/session_policy.rs index 4ddfaa37..5c3f9360 100644 --- a/src/ssh/session_policy.rs +++ b/src/ssh/session_policy.rs @@ -41,6 +41,14 @@ pub enum CliTtyMode { Disable, } +/// Traffic profile used by transport options such as `IPQoS`. +#[derive(Debug, Clone, Copy, Default, PartialEq, Eq)] +pub enum SessionPurpose { + Interactive, + #[default] + Bulk, +} + #[derive(Debug, Clone, PartialEq, Eq)] pub enum SessionRequest { Exec(String), @@ -74,6 +82,15 @@ struct TokenContext { } impl SessionPolicy { + #[must_use] + pub fn purpose(&self) -> SessionPurpose { + if self.request_pty { + SessionPurpose::Interactive + } else { + SessionPurpose::Bulk + } + } + pub fn resolve( config: &SshHostConfig, node: &Node, diff --git a/src/ssh/session_policy_tests.rs b/src/ssh/session_policy_tests.rs index 24f32728..8fbfa9c5 100644 --- a/src/ssh/session_policy_tests.rs +++ b/src/ssh/session_policy_tests.rs @@ -127,6 +127,35 @@ fn request_tty_obeys_cli_precedence_and_config_modes() { ); } +#[test] +fn session_purpose_tracks_the_resolved_pty_policy() { + let interactive = SessionPolicy::resolve( + &SshHostConfig { + request_tty: Some("force".into()), + ..Default::default() + }, + &node(), + Some("true"), + CliTtyMode::Default, + false, + ) + .unwrap(); + assert_eq!( + interactive.purpose(), + crate::ssh::SessionPurpose::Interactive + ); + + let bulk = SessionPolicy::resolve( + &SshHostConfig::default(), + &node(), + Some("true"), + CliTtyMode::Default, + false, + ) + .unwrap(); + assert_eq!(bulk.purpose(), crate::ssh::SessionPurpose::Bulk); +} + #[cfg(unix)] #[tokio::test] async fn local_command_runs_once_and_propagates_failure() { diff --git a/src/ssh/ssh_config/ip_qos.rs b/src/ssh/ssh_config/ip_qos.rs new file mode 100644 index 00000000..cc658cf5 --- /dev/null +++ b/src/ssh/ssh_config/ip_qos.rs @@ -0,0 +1,166 @@ +// Copyright 2025 Lablup Inc. and Jeongkyu Shin +// +// Licensed under the Apache License, Version 2.0 (the "License"); +// you may not use this file except in compliance with the License. +// You may obtain a copy of the License at +// +// http://www.apache.org/licenses/LICENSE-2.0 +// +// Unless required by applicable law or agreed to in writing, software +// distributed under the License is distributed on an "AS IS" BASIS, +// WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +// See the License for the specific language governing permissions and +// limitations under the License. + +//! Typed OpenSSH-compatible `IPQoS` policy. + +use thiserror::Error; + +/// A socket traffic-class value, or an explicit request to leave it unset. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub enum IpQosValue { + None, + Class(u8), +} + +/// Interactive and bulk traffic classes selected by `IPQoS`. +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +pub struct IpQosPolicy { + pub interactive: IpQosValue, + pub bulk: IpQosValue, +} + +impl Default for IpQosPolicy { + fn default() -> Self { + // OpenSSH 10.3 defaults to EF for interactive sessions and CS0 for + // bulk sessions (readconf.c fill_default_options()). + Self { + interactive: IpQosValue::Class(0xb8), + bulk: IpQosValue::Class(0x00), + } + } +} + +#[derive(Debug, Clone, PartialEq, Eq, Error)] +pub enum IpQosParseError { + #[error("at least one traffic class is required")] + MissingValue, + #[error("at most an interactive and a bulk traffic class are accepted")] + TooManyValues, + #[error("invalid traffic class '{0}'")] + InvalidValue(String), +} + +impl IpQosPolicy { + pub fn parse(values: &[String]) -> Result { + let Some(interactive) = values.first() else { + return Err(IpQosParseError::MissingValue); + }; + if values.len() > 2 { + return Err(IpQosParseError::TooManyValues); + } + let interactive = parse_value(interactive)?; + let bulk = values + .get(1) + .map_or(Ok(interactive), |value| parse_value(value))?; + Ok(Self { interactive, bulk }) + } +} + +fn parse_value(value: &str) -> Result { + let class = match value.to_ascii_lowercase().as_str() { + "none" => return Ok(IpQosValue::None), + "af11" => 0x28, + "af12" => 0x30, + "af13" => 0x38, + "af21" => 0x48, + "af22" => 0x50, + "af23" => 0x58, + "af31" => 0x68, + "af32" => 0x70, + "af33" => 0x78, + "af41" => 0x88, + "af42" => 0x90, + "af43" => 0x98, + "cs0" => 0x00, + "cs1" => 0x20, + "cs2" => 0x40, + "cs3" => 0x60, + "cs4" => 0x80, + "cs5" => 0xa0, + "cs6" => 0xc0, + "cs7" => 0xe0, + "ef" => 0xb8, + "le" => 0x04, + "va" => 0x2c, + // OpenSSH retains these names for compatibility but deliberately + // leaves the system traffic class unchanged. + "lowdelay" | "throughput" | "reliability" => return Ok(IpQosValue::None), + _ => value + .parse::() + .map_err(|_| IpQosParseError::InvalidValue(value.to_string()))?, + }; + Ok(IpQosValue::Class(class)) +} + +#[cfg(test)] +mod tests { + use super::*; + + fn values(parts: &[&str]) -> Vec { + parts.iter().map(|part| (*part).to_string()).collect() + } + + #[test] + fn parses_named_numeric_and_single_value_policies() { + assert_eq!( + IpQosPolicy::parse(&values(&["ef", "cs1"])), + Ok(IpQosPolicy { + interactive: IpQosValue::Class(0xb8), + bulk: IpQosValue::Class(0x20), + }) + ); + assert_eq!( + IpQosPolicy::parse(&values(&["af21"])), + Ok(IpQosPolicy { + interactive: IpQosValue::Class(0x48), + bulk: IpQosValue::Class(0x48), + }) + ); + assert_eq!( + IpQosPolicy::parse(&values(&["255", "0"])), + Ok(IpQosPolicy { + interactive: IpQosValue::Class(255), + bulk: IpQosValue::Class(0), + }) + ); + } + + #[test] + fn deprecated_and_none_values_leave_the_system_class_unchanged() { + for value in ["none", "lowdelay", "throughput", "reliability"] { + assert_eq!( + IpQosPolicy::parse(&values(&[value])), + Ok(IpQosPolicy { + interactive: IpQosValue::None, + bulk: IpQosValue::None, + }) + ); + } + } + + #[test] + fn rejects_missing_extra_hex_and_unknown_values() { + assert_eq!(IpQosPolicy::parse(&[]), Err(IpQosParseError::MissingValue)); + assert_eq!( + IpQosPolicy::parse(&values(&["ef", "cs0", "extra"])), + Err(IpQosParseError::TooManyValues) + ); + for value in ["0xff", "expedited", "256", "invalid"] { + assert!(matches!( + IpQosPolicy::parse(&values(&[value])), + Err(IpQosParseError::InvalidValue(_)) + )); + } + } +} diff --git a/src/ssh/ssh_config/mod.rs b/src/ssh/ssh_config/mod.rs index ea1bfb63..eb8ecece 100644 --- a/src/ssh/ssh_config/mod.rs +++ b/src/ssh/ssh_config/mod.rs @@ -25,6 +25,7 @@ mod env_cache; mod include; #[cfg(test)] mod integration_tests; +mod ip_qos; mod match_directive; mod parser; mod path; @@ -38,6 +39,7 @@ mod security_fix_tests; mod types; // Re-export public types +pub use ip_qos::{IpQosParseError, IpQosPolicy, IpQosValue}; pub use types::SshHostConfig; /// SSH configuration parser and resolver @@ -626,13 +628,25 @@ Host backup-server // Verify vpn-server config let host1 = &config.hosts[0]; assert_eq!(host1.bind_interface, Some("tun0".to_string())); - assert_eq!(host1.ipqos, Some("lowdelay throughput".to_string())); + assert_eq!( + host1.ipqos, + Some(IpQosPolicy { + interactive: IpQosValue::None, + bulk: IpQosValue::None, + }) + ); assert_eq!(host1.rekey_limit, Some("1G 1h".to_string())); // Verify backup-server config let host2 = &config.hosts[1]; assert_eq!(host2.bind_interface, Some("eth1".to_string())); - assert_eq!(host2.ipqos, Some("af21".to_string())); + assert_eq!( + host2.ipqos, + Some(IpQosPolicy { + interactive: IpQosValue::Class(0x48), + bulk: IpQosValue::Class(0x48), + }) + ); assert_eq!(host2.rekey_limit, Some("default none".to_string())); } @@ -717,7 +731,13 @@ Host web1.example.com assert_eq!(host_config.bind_interface, Some("eth0".to_string())); // IPQoS should be from *.example.com - assert_eq!(host_config.ipqos, Some("lowdelay".to_string())); + assert_eq!( + host_config.ipqos, + Some(IpQosPolicy { + interactive: IpQosValue::None, + bulk: IpQosValue::None, + }) + ); // RekeyLimit should be from web1.example.com (most specific) assert_eq!(host_config.rekey_limit, Some("1G 2h".to_string())); @@ -875,7 +895,10 @@ Host test let config = SshConfig::parse(config_content).unwrap(); assert_eq!( config.hosts[0].ipqos, - Some("lowdelay throughput".to_string()) + Some(IpQosPolicy { + interactive: IpQosValue::None, + bulk: IpQosValue::None, + }) ); // Test RekeyLimit - should reject invalid format diff --git a/src/ssh/ssh_config/parser/options/connection.rs b/src/ssh/ssh_config/parser/options/connection.rs index 080d9eb1..72c70572 100644 --- a/src/ssh/ssh_config/parser/options/connection.rs +++ b/src/ssh/ssh_config/parser/options/connection.rs @@ -17,6 +17,7 @@ //! Handles connection-related configuration options including keepalive //! settings, timeouts, compression, and network settings. +use crate::ssh::ssh_config::IpQosPolicy; use crate::ssh::ssh_config::parser::helpers::parse_yes_no; use crate::ssh::ssh_config::types::SshHostConfig; use anyhow::{Context, Result}; @@ -165,117 +166,12 @@ pub(super) fn parse_connection_option( if args.is_empty() { anyhow::bail!("IPQoS requires a value at line {line_number}"); } - // IPQoS can have one or two values (interactive and bulk) - // Valid values are: af11-af43, cs0-cs7, ef, lowdelay, throughput, reliability, or numeric (0-63 for DSCP, 0-255 for ToS) - if args.len() > 2 { - anyhow::bail!( - "IPQoS at line {} accepts at most 2 values (interactive and bulk), got {}", - line_number, - args.len() - ); + if args.iter().any(|value| value.len() > 20) { + anyhow::bail!("IPQoS value at line {line_number} is too long"); } - - // Validate each QoS value - let valid_qos_values = [ - "af11", - "af12", - "af13", - "af21", - "af22", - "af23", - "af31", - "af32", - "af33", - "af41", - "af42", - "af43", - "cs0", - "cs1", - "cs2", - "cs3", - "cs4", - "cs5", - "cs6", - "cs7", - "ef", - "lowdelay", - "throughput", - "reliability", - "none", - ]; - - // Additional mappings for common aliases - let qos_aliases = [ - ("expedited", "ef"), - ("assured", "af11"), - ("besteffort", "cs0"), - ("background", "cs1"), - ]; - - for value in args { - // Check if it's a known QoS value or alias - let lower_value = value.to_lowercase(); - let normalized = qos_aliases - .iter() - .find(|(alias, _)| *alias == lower_value.as_str()) - .map(|(_, canonical)| *canonical) - .unwrap_or(lower_value.as_str()); - - if !valid_qos_values.contains(&normalized) { - // Check if it's a numeric value - match value.parse::() { - Ok(num) => { - // DSCP values are 0-63 (6 bits) - // ToS values are 0-255 (8 bits) but only certain values are valid - if num > 63 { - // If it's a ToS value (0-255), check if it's a valid one - // Valid ToS values: 0x10 (lowdelay), 0x08 (throughput), 0x04 (reliability) - let valid_tos = [0x00, 0x04, 0x08, 0x10, 0xff]; - if !valid_tos.contains(&num) { - tracing::warn!( - "IPQoS value '{}' ({:#04x}) at line {} is not a standard DSCP (0-63) or ToS value", - value, - num, - line_number - ); - } - } - } - Err(_) => { - // Check for hex notation (0x prefix) - if value.starts_with("0x") || value.starts_with("0X") { - if let Ok(num) = u8::from_str_radix(&value[2..], 16) { - if num > 63 && ![0x00, 0x04, 0x08, 0x10, 0xff].contains(&num) { - tracing::warn!( - "IPQoS hex value '{}' at line {} is outside standard ranges", - value, - line_number - ); - } - } else { - anyhow::bail!( - "IPQoS value '{value}' at line {line_number} is not a valid hexadecimal number" - ); - } - } else { - anyhow::bail!( - "IPQoS value '{value}' at line {line_number} is not valid. \ - Valid values are: af11-af43, cs0-cs7, ef, lowdelay, throughput, \ - reliability, none, or numeric (0-63 for DSCP, specific ToS values)" - ); - } - } - } - } - } - - // Limit total length to prevent memory exhaustion - let combined = args.join(" "); - if combined.len() > 100 { - anyhow::bail!("IPQoS value at line {line_number} is too long (max 100 characters)"); - } - - host.ipqos = Some(combined); + host.ipqos = Some(IpQosPolicy::parse(args).map_err(|error| { + anyhow::anyhow!("Invalid IPQoS value at line {line_number}: {error}") + })?); } "rekeylimit" => { if args.is_empty() { diff --git a/src/ssh/ssh_config/parser/options/support.rs b/src/ssh/ssh_config/parser/options/support.rs index 8ce8da23..bf0075ad 100644 --- a/src/ssh/ssh_config/parser/options/support.rs +++ b/src/ssh/ssh_config/parser/options/support.rs @@ -192,9 +192,9 @@ pub(super) const ACCEPTED_KEYWORDS: &[(&str, &str, KeywordSupport)] = &[ ("compression", "compression", Runtime(Transport)), ("tcpkeepalive", "tcpkeepalive", Runtime(Transport)), ("addressfamily", "addressfamily", Runtime(Transport)), - ("bindaddress", "bindaddress", Delegated(300)), - ("bindinterface", "bindinterface", Delegated(300)), - ("ipqos", "ipqos", Delegated(300)), + ("bindaddress", "bindaddress", Runtime(Transport)), + ("bindinterface", "bindinterface", Runtime(Transport)), + ("ipqos", "ipqos", Runtime(Transport)), ("rekeylimit", "rekeylimit", Delegated(301)), ("proxyjump", "proxyjump", Runtime(Proxy)), ("proxycommand", "proxycommand", Runtime(Proxy)), @@ -320,6 +320,9 @@ mod tests { ("compression", Transport), ("tcpkeepalive", Transport), ("addressfamily", Transport), + ("bindaddress", Transport), + ("bindinterface", Transport), + ("ipqos", Transport), ("proxyjump", Proxy), ("proxycommand", Proxy), ("proxyusefdpass", Proxy), @@ -350,12 +353,7 @@ mod tests { #[test] fn first_wave_delegations_match_the_split_issue_dag() { - let expected = HashMap::from([ - ("ipqos", 300), - ("bindaddress", 300), - ("bindinterface", 300), - ("rekeylimit", 301), - ]); + let expected = HashMap::from([("rekeylimit", 301)]); let delegated = ACCEPTED_KEYWORDS .iter() .filter_map(|(keyword, canonical, support)| match support { diff --git a/src/ssh/ssh_config/resolver.rs b/src/ssh/ssh_config/resolver.rs index 66699f90..d20dc906 100644 --- a/src/ssh/ssh_config/resolver.rs +++ b/src/ssh/ssh_config/resolver.rs @@ -357,7 +357,7 @@ pub(super) fn merge_host_config(base: &mut SshHostConfig, overlay: &SshHostConfi base.bind_interface = overlay.bind_interface.clone(); } if base.ipqos.is_none() && overlay.ipqos.is_some() { - base.ipqos = overlay.ipqos.clone(); + base.ipqos = overlay.ipqos; } if base.rekey_limit.is_none() && overlay.rekey_limit.is_some() { base.rekey_limit = overlay.rekey_limit.clone(); diff --git a/src/ssh/ssh_config/types.rs b/src/ssh/ssh_config/types.rs index 4c071bd6..43a6c68d 100644 --- a/src/ssh/ssh_config/types.rs +++ b/src/ssh/ssh_config/types.rs @@ -20,6 +20,8 @@ use std::path::PathBuf; use crate::forwarding::ForwardingDirective; +use super::IpQosPolicy; + /// Configuration block type #[derive(Debug, Clone, PartialEq)] pub enum ConfigBlock { @@ -123,7 +125,7 @@ pub struct SshHostConfig { pub enable_ssh_keysign: Option, // Network & connection pub bind_interface: Option, - pub ipqos: Option, + pub ipqos: Option, pub rekey_limit: Option, // X11 forwarding pub forward_x11_timeout: Option, diff --git a/src/ssh/tokio_client/connection.rs b/src/ssh/tokio_client/connection.rs index 160477e9..9c6ce60e 100644 --- a/src/ssh/tokio_client/connection.rs +++ b/src/ssh/tokio_client/connection.rs @@ -19,7 +19,7 @@ use russh::client::{Config, DisconnectReason, Handle, Handler}; use std::borrow::Cow; -use std::net::SocketAddr; +use std::net::{IpAddr, SocketAddr}; use std::path::PathBuf; use std::sync::{ Arc, Mutex, @@ -36,7 +36,8 @@ use super::proxy_command::{ }; use crate::forwarding::remote::RemoteForwardRegistry; use crate::forwarding::{ForwardingDirective, ForwardingPlan, ForwardingRuntime}; -use crate::ssh::SshConfig; +use crate::ssh::ssh_config::{IpQosPolicy, IpQosValue}; +use crate::ssh::{SessionPurpose, SshConfig}; /// Default keepalive interval in seconds. /// @@ -112,6 +113,18 @@ pub struct SshConnectionConfig { /// This is independent from SSH protocol keepalives. pub tcp_keep_alive: bool, + /// Optional local address selected by ssh_config `BindAddress`. + pub bind_address: Option, + + /// Optional local interface selected by ssh_config `BindInterface`. + pub bind_interface: Option, + + /// Interactive and bulk socket traffic classes selected by `IPQoS`. + pub ip_qos: IpQosPolicy, + + /// Which side of the `IPQoS` policy applies to this connection. + pub session_purpose: SessionPurpose, + /// Raw `UserKnownHostsFile` values selected for this host. /// /// They stay unexpanded until the connection target supplies `%h`, `%p`, @@ -158,6 +171,10 @@ impl Default for SshConnectionConfig { address_family: AddressFamily::Any, connection_attempts: 1, tcp_keep_alive: true, + bind_address: None, + bind_interface: None, + ip_qos: IpQosPolicy::default(), + session_purpose: SessionPurpose::Bulk, user_known_hosts_files: None, global_known_hosts_files: None, host_key_alias: None, @@ -361,6 +378,16 @@ impl SshConnectionConfigResolver { .as_ref() .and_then(|config| config.tcp_keep_alive) .unwrap_or(true); + let bind_address = host_config + .as_ref() + .and_then(|config| config.bind_address.clone()); + let bind_interface = host_config + .as_ref() + .and_then(|config| config.bind_interface.clone()); + let ip_qos = host_config + .as_ref() + .and_then(|config| config.ipqos) + .unwrap_or_default(); let user_known_hosts_files = host_config .as_ref() .and_then(|config| config.user_known_hosts_file.clone()); @@ -537,6 +564,8 @@ impl SshConnectionConfigResolver { .with_address_family(address_family) .with_connection_attempts(connection_attempts) .with_tcp_keep_alive(tcp_keep_alive) + .with_source_binding(bind_address, bind_interface) + .with_ip_qos(ip_qos) .with_known_hosts_files(user_known_hosts_files, global_known_hosts_files) .with_host_key_alias(host_key_alias) .with_known_host_policy( @@ -712,6 +741,41 @@ impl SshConnectionConfig { self } + /// Set the local source address/interface policy for direct TCP carriers. + #[must_use] + pub fn with_source_binding( + mut self, + bind_address: Option, + bind_interface: Option, + ) -> Self { + self.bind_address = bind_address; + self.bind_interface = bind_interface; + self + } + + /// Set the interactive/bulk traffic-class policy. + #[must_use] + pub fn with_ip_qos(mut self, ip_qos: IpQosPolicy) -> Self { + self.ip_qos = ip_qos; + self + } + + /// Select which side of the traffic-class policy applies to this session. + #[must_use] + pub fn with_session_purpose(mut self, purpose: SessionPurpose) -> Self { + self.session_purpose = purpose; + self + } + + /// Select the traffic class that applies to this connection's session. + #[must_use] + pub fn selected_ip_qos(&self) -> IpQosValue { + match self.session_purpose { + SessionPurpose::Interactive => self.ip_qos.interactive, + SessionPurpose::Bulk => self.ip_qos.bulk, + } + } + #[must_use] pub fn with_proxy_mode(mut self, proxy_mode: Option) -> Self { self.proxy_mode = proxy_mode; @@ -875,6 +939,294 @@ pub(super) fn configure_tcp_keepalive( Ok(()) } +fn socket_family(address: SocketAddr) -> &'static str { + if address.is_ipv4() { "IPv4" } else { "IPv6" } +} + +pub(super) fn resolve_bind_address( + bind_address: &str, + destination: SocketAddr, +) -> Result { + let family = socket_family(destination); + let resolved = + std::net::ToSocketAddrs::to_socket_addrs(&(bind_address, 0)).map_err(|source| { + super::Error::BindAddressResolution { + bind_address: bind_address.to_string(), + family, + source, + } + })?; + resolved + .into_iter() + .find(|source| source.is_ipv4() == destination.is_ipv4()) + .ok_or_else(|| super::Error::BindAddressFamily { + bind_address: bind_address.to_string(), + family, + destination, + }) +} + +#[cfg(unix)] +pub(super) fn resolve_interface_address( + interface: &str, + destination: SocketAddr, +) -> Result { + struct InterfaceAddresses(*mut libc::ifaddrs); + impl Drop for InterfaceAddresses { + fn drop(&mut self) { + if !self.0.is_null() { + // SAFETY: the pointer is initialized only by a successful + // getifaddrs call and is freed exactly once by this guard. + unsafe { libc::freeifaddrs(self.0) }; + } + } + } + + let family = socket_family(destination); + let mut raw_addresses = std::ptr::null_mut(); + // SAFETY: getifaddrs initializes the out pointer on success. The guard + // owns the resulting list for the rest of this function. + if unsafe { libc::getifaddrs(&mut raw_addresses) } != 0 { + return Err(super::Error::BindInterface { + interface: interface.to_string(), + family, + source: io::Error::last_os_error(), + }); + } + let addresses = InterfaceAddresses(raw_addresses); + let mut fallback = None; + let mut current = addresses.0; + while !current.is_null() { + // SAFETY: current belongs to the live getifaddrs list and remains + // valid until the guard is dropped. + let entry = unsafe { &*current }; + current = entry.ifa_next; + if entry.ifa_addr.is_null() + || entry.ifa_name.is_null() + || entry.ifa_flags & libc::IFF_UP as libc::c_uint == 0 + { + continue; + } + // SAFETY: getifaddrs guarantees a NUL-terminated interface name. + if unsafe { std::ffi::CStr::from_ptr(entry.ifa_name) }.to_bytes() != interface.as_bytes() { + continue; + } + + let address = match (destination, unsafe { + (*entry.ifa_addr).sa_family as libc::c_int + }) { + (SocketAddr::V4(_), libc::AF_INET) => { + // SAFETY: sa_family identifies this address as sockaddr_in. + let address = unsafe { &*entry.ifa_addr.cast::() }; + let octets = address.sin_addr.s_addr.to_ne_bytes(); + SocketAddr::new(IpAddr::V4(std::net::Ipv4Addr::from(octets)), 0) + } + (SocketAddr::V6(_), libc::AF_INET6) => { + // SAFETY: sa_family identifies this address as sockaddr_in6. + let address = unsafe { &*entry.ifa_addr.cast::() }; + SocketAddr::V6(std::net::SocketAddrV6::new( + std::net::Ipv6Addr::from(address.sin6_addr.s6_addr), + 0, + address.sin6_flowinfo, + address.sin6_scope_id, + )) + } + _ => continue, + }; + let local = match address.ip() { + IpAddr::V4(address) => address.is_loopback() || address.is_link_local(), + IpAddr::V6(address) => address.is_loopback() || address.is_unicast_link_local(), + }; + if local { + fallback.get_or_insert(address); + } else { + return Ok(address); + } + } + fallback.ok_or_else(|| super::Error::BindInterface { + interface: interface.to_string(), + family, + source: io::Error::new( + io::ErrorKind::AddrNotAvailable, + "interface has no active address for the destination family", + ), + }) +} + +#[cfg(not(unix))] +pub(super) fn resolve_interface_address( + interface: &str, + destination: SocketAddr, +) -> Result { + Err(super::Error::BindInterface { + interface: interface.to_string(), + family: socket_family(destination), + source: io::Error::new( + io::ErrorKind::Unsupported, + "BindInterface is not supported on this platform", + ), + }) +} + +pub(super) fn apply_ip_qos( + socket: &socket2::Socket, + destination: SocketAddr, + value: IpQosValue, +) -> Result<(), super::Error> { + let IpQosValue::Class(value) = value else { + return Ok(()); + }; + let result = if destination.is_ipv4() { + socket.set_tos_v4(u32::from(value)) + } else { + #[cfg(any( + target_os = "android", + target_os = "dragonfly", + target_os = "freebsd", + target_os = "fuchsia", + target_os = "linux", + target_os = "macos", + target_os = "netbsd", + target_os = "openbsd", + target_os = "illumos", + ))] + { + socket.set_tclass_v6(u32::from(value)) + } + #[cfg(not(any( + target_os = "android", + target_os = "dragonfly", + target_os = "freebsd", + target_os = "fuchsia", + target_os = "linux", + target_os = "macos", + target_os = "netbsd", + target_os = "openbsd", + target_os = "illumos", + )))] + { + Err(io::Error::new( + io::ErrorKind::Unsupported, + "IPv6 traffic class is not supported on this platform", + )) + } + }; + result.map_err(|source| super::Error::IpQos { + value, + family: socket_family(destination), + destination, + source, + }) +} + +fn is_connect_in_progress(error: &io::Error) -> bool { + if error.kind() == io::ErrorKind::WouldBlock { + return true; + } + #[cfg(unix)] + return matches!( + error.raw_os_error(), + Some(code) if code == libc::EINPROGRESS || code == libc::EALREADY + ); + #[cfg(windows)] + return matches!( + error.raw_os_error(), + Some(code) if code == 10035 || code == 10036 || code == 10037 + ); + #[cfg(not(any(unix, windows)))] + false +} + +pub(super) async fn connect_direct_socket( + destination: SocketAddr, + target_host: &str, + target_port: u16, + bind_address: Option<&str>, + bind_interface: Option<&str>, + ip_qos: IpQosValue, +) -> Result { + let domain = if destination.is_ipv4() { + socket2::Domain::IPV4 + } else { + socket2::Domain::IPV6 + }; + let socket = socket2::Socket::new(domain, socket2::Type::STREAM, Some(socket2::Protocol::TCP)) + .map_err(|source| super::Error::TcpConnect { + host: target_host.to_string(), + port: target_port, + source, + })?; + socket + .set_nonblocking(true) + .map_err(|source| super::Error::TcpConnect { + host: target_host.to_string(), + port: target_port, + source, + })?; + apply_ip_qos(&socket, destination, ip_qos)?; + + let source_address = if let Some(bind_address) = bind_address { + Some(resolve_bind_address(bind_address, destination)?) + } else if let Some(bind_interface) = bind_interface { + Some(resolve_interface_address(bind_interface, destination)?) + } else { + None + }; + if let Some(source_address) = source_address { + socket + .bind(&source_address.into()) + .map_err(|source| super::Error::SourceBind { + source_address, + destination, + source, + })?; + } + + let pending = match socket.connect(&destination.into()) { + Ok(()) => false, + Err(error) if is_connect_in_progress(&error) => true, + Err(source) => { + return Err(super::Error::TcpConnect { + host: target_host.to_string(), + port: target_port, + source, + }); + } + }; + let stream = tokio::net::TcpStream::from_std(socket.into()).map_err(|source| { + super::Error::TcpConnect { + host: target_host.to_string(), + port: target_port, + source, + } + })?; + if pending { + stream + .writable() + .await + .map_err(|source| super::Error::TcpConnect { + host: target_host.to_string(), + port: target_port, + source, + })?; + if let Some(source) = stream + .take_error() + .map_err(|source| super::Error::TcpConnect { + host: target_host.to_string(), + port: target_port, + source, + })? + { + return Err(super::Error::TcpConnect { + host: target_host.to_string(), + port: target_port, + source, + }); + } + } + Ok(stream) +} + use super::ToSocketAddrsWithHostname; /// A ssh connection to a remote server. @@ -1033,6 +1385,9 @@ struct DirectCarrierOptions<'a> { tcp_keepalive: Option<&'a socket2::TcpKeepalive>, address_family: AddressFamily, connection_attempts: usize, + bind_address: Option<&'a str>, + bind_interface: Option<&'a str>, + ip_qos: IpQosValue, } impl Client { @@ -1184,6 +1539,9 @@ impl Client { tcp_keepalive: tcp_keepalive.as_ref(), address_family: ssh_config.address_family, connection_attempts: ssh_config.connection_attempts, + bind_address: ssh_config.bind_address.as_deref(), + bind_interface: ssh_config.bind_interface.as_deref(), + ip_qos: ssh_config.selected_ip_qos(), }, ) .await?; @@ -1294,6 +1652,9 @@ impl Client { tcp_keepalive: None, address_family: AddressFamily::Any, connection_attempts: 1, + bind_address: None, + bind_interface: None, + ip_qos: IpQosValue::None, }, ) .await @@ -1325,6 +1686,9 @@ impl Client { tcp_keepalive, address_family, connection_attempts, + bind_address, + bind_interface, + ip_qos, } = carrier; let connection_attempts = connection_attempts.max(1); let target_host = addr.hostname(); @@ -1368,17 +1732,22 @@ impl Client { }; } else { for socket_addr in socket_addrs { - match tokio::net::TcpStream::connect(socket_addr).await { + match connect_direct_socket( + socket_addr, + &target_host, + target_port, + bind_address, + bind_interface, + ip_qos, + ) + .await + { Ok(stream) => { carrier_res = Ok((socket_addr, stream)); break 'rounds; } Err(error) => { - carrier_res = Err(super::Error::TcpConnect { - host: target_host.clone(), - port: target_port, - source: error, - }); + carrier_res = Err(error); } } } diff --git a/src/ssh/tokio_client/connection_tests.rs b/src/ssh/tokio_client/connection_tests.rs index 6548fa9e..f0f81b40 100644 --- a/src/ssh/tokio_client/connection_tests.rs +++ b/src/ssh/tokio_client/connection_tests.rs @@ -28,10 +28,13 @@ use std::time::Duration; use super::address_family::AddressFamily; use super::authentication::{AuthMethod, ServerCheckMethod}; use super::connection::{ - Client, SshConnectionConfig, SshConnectionConfigResolver, configure_tcp_keepalive, + Client, SshConnectionConfig, SshConnectionConfigResolver, apply_ip_qos, + configure_tcp_keepalive, connect_direct_socket, resolve_bind_address, + resolve_interface_address, }; use super::proxy_command::{ProxyCommandConfig, ProxyMode}; -use crate::ssh::SshConfig; +use crate::ssh::ssh_config::{IpQosPolicy, IpQosValue}; +use crate::ssh::{SessionPurpose, SshConfig}; #[test] fn test_default_compression_advertises_none_only() { @@ -740,3 +743,144 @@ fn authentication_yes_no_defaults_and_overrides_are_preserved() { assert!(!policy.password_authentication); assert!(!policy.batch_mode); } + +#[test] +fn source_binding_and_ipqos_reach_the_connection_config() { + let ssh_config = SshConfig::parse( + "Host target\n BindAddress 127.0.0.1\n BindInterface lo\n IPQoS ef cs1\n", + ) + .unwrap(); + let config = SshConnectionConfigResolver::new() + .with_ssh_config(Some(ssh_config)) + .resolve_for_host("target"); + + assert_eq!(config.bind_address.as_deref(), Some("127.0.0.1")); + assert_eq!(config.bind_interface.as_deref(), Some("lo")); + assert_eq!( + config.ip_qos, + IpQosPolicy { + interactive: IpQosValue::Class(0xb8), + bulk: IpQosValue::Class(0x20), + } + ); + assert_eq!(config.session_purpose, SessionPurpose::Bulk); + assert_eq!(config.selected_ip_qos(), IpQosValue::Class(0x20)); + assert_eq!( + config + .clone() + .with_session_purpose(SessionPurpose::Interactive) + .selected_ip_qos(), + IpQosValue::Class(0xb8) + ); +} + +#[tokio::test] +async fn direct_socket_binds_the_requested_loopback_source() { + let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let destination = listener.local_addr().unwrap(); + let client = connect_direct_socket( + destination, + "localhost", + destination.port(), + Some("127.0.0.1"), + Some("definitely-not-an-interface"), + IpQosValue::Class(0x20), + ) + .await + .unwrap(); + let (_server, peer) = listener.accept().await.unwrap(); + + assert_eq!(peer.ip(), "127.0.0.1".parse::().unwrap()); + assert_eq!(client.local_addr().unwrap().ip(), peer.ip()); + assert_eq!(socket2::SockRef::from(&client).tos_v4().unwrap(), 0x20); +} + +#[tokio::test] +async fn nonlocal_bind_address_fails_without_unbound_fallback() { + let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap(); + let destination = listener.local_addr().unwrap(); + let error = connect_direct_socket( + destination, + "localhost", + destination.port(), + Some("192.0.2.123"), + None, + IpQosValue::None, + ) + .await + .unwrap_err(); + + assert!(matches!(error, super::Error::SourceBind { .. })); +} + +#[test] +fn incompatible_bind_address_family_is_explicit() { + let destination = "127.0.0.1:22".parse().unwrap(); + let error = resolve_bind_address("::1", destination).unwrap_err(); + assert!(matches!(error, super::Error::BindAddressFamily { .. })); + assert!(error.to_string().contains("no IPv4 address")); +} + +#[cfg(unix)] +#[test] +fn bind_interface_selects_a_same_family_loopback_address() { + #[cfg(any( + target_os = "macos", + target_os = "freebsd", + target_os = "netbsd", + target_os = "openbsd" + ))] + let interface = "lo0"; + #[cfg(not(any( + target_os = "macos", + target_os = "freebsd", + target_os = "netbsd", + target_os = "openbsd" + )))] + let interface = "lo"; + + let destination = "127.0.0.1:22".parse().unwrap(); + let source = resolve_interface_address(interface, destination).unwrap(); + assert!(source.is_ipv4()); + assert!(source.ip().is_loopback()); + + let error = resolve_interface_address("bssh-no-such-if", destination).unwrap_err(); + assert!(matches!(error, super::Error::BindInterface { .. })); +} + +#[test] +fn ipv4_qos_sets_the_exact_tos_byte() { + let socket = socket2::Socket::new( + socket2::Domain::IPV4, + socket2::Type::STREAM, + Some(socket2::Protocol::TCP), + ) + .unwrap(); + let destination = "127.0.0.1:22".parse().unwrap(); + apply_ip_qos(&socket, destination, IpQosValue::Class(0xb8)).unwrap(); + assert_eq!(socket.tos_v4().unwrap(), 0xb8); +} + +#[cfg(any( + target_os = "android", + target_os = "dragonfly", + target_os = "freebsd", + target_os = "fuchsia", + target_os = "linux", + target_os = "macos", + target_os = "netbsd", + target_os = "openbsd", + target_os = "illumos", +))] +#[test] +fn ipv6_qos_sets_the_exact_traffic_class_byte() { + let socket = socket2::Socket::new( + socket2::Domain::IPV6, + socket2::Type::STREAM, + Some(socket2::Protocol::TCP), + ) + .unwrap(); + let destination = "[::1]:22".parse().unwrap(); + apply_ip_qos(&socket, destination, IpQosValue::Class(0x48)).unwrap(); + assert_eq!(socket.tclass_v6().unwrap(), 0x48); +} diff --git a/src/ssh/tokio_client/error.rs b/src/ssh/tokio_client/error.rs index 1d66589e..b4b0c192 100644 --- a/src/ssh/tokio_client/error.rs +++ b/src/ssh/tokio_client/error.rs @@ -75,6 +75,43 @@ pub enum Error { #[source] source: io::Error, }, + #[error("failed to resolve BindAddress '{bind_address}' for {family}: {source}")] + BindAddressResolution { + bind_address: String, + family: &'static str, + #[source] + source: io::Error, + }, + #[error("BindAddress '{bind_address}' has no {family} address for destination {destination}")] + BindAddressFamily { + bind_address: String, + family: &'static str, + destination: std::net::SocketAddr, + }, + #[error("BindInterface '{interface}' cannot select a {family} source address: {source}")] + BindInterface { + interface: String, + family: &'static str, + #[source] + source: io::Error, + }, + #[error( + "failed to bind source address {source_address} for destination {destination}: {source}" + )] + SourceBind { + source_address: std::net::SocketAddr, + destination: std::net::SocketAddr, + #[source] + source: io::Error, + }, + #[error("failed to apply IPQoS {value:#04x} to {family} destination {destination}: {source}")] + IpQos { + value: u8, + family: &'static str, + destination: std::net::SocketAddr, + #[source] + source: io::Error, + }, #[error("connection failed after {attempts} carrier attempts: {source}")] ConnectionAttemptsExhausted { attempts: usize, @@ -242,6 +279,11 @@ impl Error { | Self::AddressInvalid(_) | Self::DnsResolution { .. } | Self::TcpConnect { .. } + | Self::BindAddressResolution { .. } + | Self::BindAddressFamily { .. } + | Self::BindInterface { .. } + | Self::SourceBind { .. } + | Self::IpQos { .. } | Self::ConnectionAttemptsExhausted { .. } | Self::ConnectionTimeout { .. } | Self::ProtocolNegotiation { .. } From a4e3cee7115a575a40c4a46aa6c3fd1db1e1df39 Mon Sep 17 00:00:00 2001 From: Jeongkyu Shin Date: Sun, 30 Aug 2026 14:49:16 +0900 Subject: [PATCH 2/6] fix(ssh): preserve per-host interactive transport policy Retain the SSH connection config resolver in interactive commands so ProxyJump bastions apply their own BindAddress, BindInterface, and IPQoS directives instead of inheriting the final target's socket policy. Derive direct and jump traffic purpose from SessionPolicy so no-PTY shells use bulk QoS, correct Voice-Admit encoding to 0xb0, and cover fixed-config compatibility plus IPv4 and IPv6 socket behavior. Validated with scoped format, check, clippy, per-hop policy, source-binding, and IPQoS tests. Refs #300 --- src/app/dispatcher.rs | 2 + src/commands/interactive/connection.rs | 155 ++++++++++++++++++++++--- src/commands/interactive/types.rs | 4 +- src/commands/interactive/utils.rs | 3 + src/jump/chain.rs | 2 +- src/ssh/ssh_config/ip_qos.rs | 13 ++- tests/interactive_integration_test.rs | 9 ++ tests/interactive_test.rs | 2 + tests/ssh_keepalive_test.rs | 3 + 9 files changed, 172 insertions(+), 21 deletions(-) diff --git a/src/app/dispatcher.rs b/src/app/dispatcher.rs index ee1f754d..6019d1d6 100644 --- a/src/app/dispatcher.rs +++ b/src/app/dispatcher.rs @@ -524,6 +524,7 @@ async fn handle_interactive_command( use_pty, session_policy: None, ssh_connection_config, + ssh_connection_config_resolver: Some(ssh_connection_config_resolver), }; let result = interactive_cmd.execute().await?; @@ -654,6 +655,7 @@ async fn handle_exec_command( use_pty, session_policy: Some(session_policy), ssh_connection_config, + ssh_connection_config_resolver: Some(ssh_connection_config_resolver), }; let result = interactive_cmd.execute().await?; diff --git a/src/commands/interactive/connection.rs b/src/commands/interactive/connection.rs index c90b7919..e9eaba3c 100644 --- a/src/commands/interactive/connection.rs +++ b/src/commands/interactive/connection.rs @@ -22,16 +22,41 @@ use std::io::{self, IsTerminal, Write}; use tokio::time::{Duration, timeout}; use zeroize::Zeroizing; -use crate::jump::{JumpHostChain, parse_jump_hosts}; +use crate::jump::{JumpHostChain, parse_jump_hosts, parser::JumpHost}; use crate::node::Node; use crate::ssh::{ - SessionPurpose, SessionRequest, + SessionPolicy, SessionPurpose, SessionRequest, known_hosts::get_check_method_for_target, - tokio_client::{AuthMethod, Client, Error as SshError, ServerCheckMethod, SshConnectionConfig}, + tokio_client::{ + AuthMethod, Client, Error as SshError, ServerCheckMethod, SshConnectionConfig, + SshConnectionConfigResolver, + }, }; use super::types::{InteractiveCommand, NodeSession}; +fn build_interactive_jump_chain( + jump_hosts: Vec, + adjusted_timeout: Duration, + ssh_connection_config: &SshConnectionConfig, + resolver: Option<&SshConnectionConfigResolver>, + session_purpose: SessionPurpose, +) -> JumpHostChain { + let mut chain = JumpHostChain::new(jump_hosts) + .with_connect_timeout(adjusted_timeout) + .with_command_timeout(Duration::from_secs(300)) + .with_ssh_connection_config(ssh_connection_config.clone()) + .with_session_purpose(session_purpose); + if let Some(resolver) = resolver { + chain = chain.with_ssh_connection_config_resolver(resolver.clone()); + } + chain +} + +fn interactive_session_purpose(session_policy: Option<&SessionPolicy>) -> SessionPurpose { + session_policy.map_or(SessionPurpose::Interactive, SessionPolicy::purpose) +} + impl InteractiveCommand { /// Helper function to establish SSH connection with proper error handling and rate limiting /// This eliminates code duplication across different connection paths and prevents brute-force attacks @@ -50,12 +75,11 @@ impl InteractiveCommand { port: u16, allow_password_fallback: bool, ssh_config: &SshConnectionConfig, + session_purpose: SessionPurpose, ) -> Result { const SSH_CONNECT_TIMEOUT_SECS: u64 = 30; let connect_timeout = Duration::from_secs(SSH_CONNECT_TIMEOUT_SECS); - let ssh_config = ssh_config - .clone() - .with_session_purpose(SessionPurpose::Interactive); + let ssh_config = ssh_config.clone().with_session_purpose(session_purpose); // SECURITY: Add a small delay before connection attempts to prevent rapid-fire attempts // This helps mitigate brute-force attacks and prevents triggering fail2ban too quickly @@ -147,6 +171,25 @@ impl InteractiveCommand { result } + fn session_purpose(&self) -> SessionPurpose { + interactive_session_purpose(self.session_policy.as_ref()) + } + + fn build_jump_chain( + &self, + jump_hosts: Vec, + adjusted_timeout: Duration, + ) -> JumpHostChain { + build_interactive_jump_chain( + jump_hosts, + adjusted_timeout, + &self.ssh_connection_config, + self.ssh_connection_config_resolver.as_ref(), + self.session_purpose(), + ) + .with_ssh_password(self.ssh_password.clone()) + } + /// Prompt for password with secure handling async fn prompt_password(username: &str, host: &str) -> Result> { let username = username.to_string(); @@ -288,6 +331,7 @@ impl InteractiveCommand { node.port, !self.use_password, // Allow fallback unless explicit password mode &self.ssh_connection_config, + self.session_purpose(), ) .await? } else { @@ -315,12 +359,7 @@ impl InteractiveCommand { // Pass SSH connection config to jump host chain for keepalive settings. // Also pass the dispatcher's pre-collected password so jump-host // authentication consumes it instead of re-prompting per call. See #200. - let chain = JumpHostChain::new(jump_hosts) - .with_connect_timeout(adjusted_timeout) - .with_command_timeout(Duration::from_secs(300)) - .with_ssh_connection_config(self.ssh_connection_config.clone()) - .with_session_purpose(SessionPurpose::Interactive) - .with_ssh_password(self.ssh_password.clone()); + let chain = self.build_jump_chain(jump_hosts, adjusted_timeout); // Connect through the chain let connection = timeout( @@ -372,6 +411,7 @@ impl InteractiveCommand { node.port, !self.use_password, // Allow fallback unless explicit password mode &self.ssh_connection_config, + self.session_purpose(), ) .await? }; @@ -449,6 +489,7 @@ impl InteractiveCommand { node.port, !self.use_password, // Allow fallback unless explicit password mode &self.ssh_connection_config, + self.session_purpose(), ) .await? } else { @@ -476,12 +517,7 @@ impl InteractiveCommand { // Pass SSH connection config to jump host chain for keepalive settings. // Also pass the dispatcher's pre-collected password so jump-host // authentication consumes it instead of re-prompting per call. See #200. - let chain = JumpHostChain::new(jump_hosts) - .with_connect_timeout(adjusted_timeout) - .with_command_timeout(Duration::from_secs(300)) - .with_ssh_connection_config(self.ssh_connection_config.clone()) - .with_session_purpose(SessionPurpose::Interactive) - .with_ssh_password(self.ssh_password.clone()); + let chain = self.build_jump_chain(jump_hosts, adjusted_timeout); // Connect through the chain let connection = timeout( @@ -533,6 +569,7 @@ impl InteractiveCommand { node.port, !self.use_password, // Allow fallback unless explicit password mode &self.ssh_connection_config, + self.session_purpose(), ) .await? }; @@ -614,6 +651,88 @@ pub fn is_auth_error_for_password_fallback(error: &SshError) -> bool { #[cfg(test)] mod tests { use super::*; + use crate::ssh::ssh_config::{IpQosPolicy, IpQosValue, SshConfig}; + + #[test] + fn no_pty_shell_uses_bulk_ipqos_for_direct_and_jump_connections() { + let policy = SessionPolicy { + environment: Vec::new(), + local_command: None, + request_pty: false, + request: SessionRequest::Shell, + }; + + assert_eq!( + interactive_session_purpose(Some(&policy)), + SessionPurpose::Bulk + ); + assert_eq!( + interactive_session_purpose(None), + SessionPurpose::Interactive + ); + let config = SshConnectionConfig::new() + .with_ip_qos(IpQosPolicy { + interactive: IpQosValue::Class(0xb8), + bulk: IpQosValue::Class(0x20), + }) + .with_session_purpose(interactive_session_purpose(Some(&policy))); + assert_eq!(config.selected_ip_qos(), IpQosValue::Class(0x20)); + } + + #[test] + fn interactive_jump_chain_keeps_distinct_bastion_and_target_socket_policies() { + let ssh_config = SshConfig::parse( + r#" +Host bastion + BindAddress 127.0.0.2 + BindInterface lo + IPQoS cs5 cs1 + +Host target + BindAddress 127.0.0.3 + BindInterface target0 + IPQoS ef cs2 +"#, + ) + .expect("valid ssh_config"); + let resolver = SshConnectionConfigResolver::new().with_ssh_config(Some(ssh_config)); + let chain = build_interactive_jump_chain( + vec![JumpHost::new("bastion".to_string(), None, None)], + Duration::from_secs(45), + &SshConnectionConfig::default(), + Some(&resolver), + SessionPurpose::Bulk, + ); + + let bastion = chain.connection_config_for_host("bastion"); + assert_eq!(bastion.bind_address.as_deref(), Some("127.0.0.2")); + assert_eq!(bastion.bind_interface.as_deref(), Some("lo")); + assert_eq!(bastion.session_purpose, SessionPurpose::Bulk); + assert_eq!(bastion.selected_ip_qos(), IpQosValue::Class(0x20)); + + let target = chain.connection_config_for_host("target"); + assert_eq!(target.bind_address.as_deref(), Some("127.0.0.3")); + assert_eq!(target.bind_interface.as_deref(), Some("target0")); + assert_eq!(target.session_purpose, SessionPurpose::Bulk); + assert_eq!(target.selected_ip_qos(), IpQosValue::Class(0x40)); + + let fixed_config = SshConnectionConfig::new() + .with_source_binding(Some("127.0.0.4".to_string()), Some("lo".to_string())) + .with_ip_qos(IpQosPolicy { + interactive: IpQosValue::Class(0xb8), + bulk: IpQosValue::Class(0x60), + }); + let fixed_chain = build_interactive_jump_chain( + vec![JumpHost::new("manual-bastion".to_string(), None, None)], + Duration::from_secs(45), + &fixed_config, + None, + SessionPurpose::Bulk, + ); + let manual_bastion = fixed_chain.connection_config_for_host("manual-bastion"); + assert_eq!(manual_bastion.bind_address.as_deref(), Some("127.0.0.4")); + assert_eq!(manual_bastion.selected_ip_qos(), IpQosValue::Class(0x60)); + } #[test] fn test_key_auth_failed_triggers_password_fallback() { diff --git a/src/commands/interactive/types.rs b/src/commands/interactive/types.rs index a3b7424e..53c98899 100644 --- a/src/commands/interactive/types.rs +++ b/src/commands/interactive/types.rs @@ -27,7 +27,7 @@ use crate::pty::PtyConfig; use crate::security::Password; use crate::ssh::SessionPolicy; use crate::ssh::known_hosts::StrictHostKeyChecking; -use crate::ssh::tokio_client::{Client, SshConnectionConfig}; +use crate::ssh::tokio_client::{Client, SshConnectionConfig, SshConnectionConfigResolver}; /// SSH output polling interval for responsive display /// - 10ms provides very responsive output display @@ -71,6 +71,8 @@ pub struct InteractiveCommand { pub session_policy: Option, // SSH connection configuration (keepalive settings) pub ssh_connection_config: SshConnectionConfig, + /// Per-host connection policy retained for ProxyJump bastions. + pub ssh_connection_config_resolver: Option, } /// Result of an interactive session diff --git a/src/commands/interactive/utils.rs b/src/commands/interactive/utils.rs index 550e01ee..ae8bc64f 100644 --- a/src/commands/interactive/utils.rs +++ b/src/commands/interactive/utils.rs @@ -110,6 +110,7 @@ mod tests { request: crate::ssh::SessionRequest::Shell, }), ssh_connection_config: SshConnectionConfig::default(), + ssh_connection_config_resolver: Default::default(), }; assert!(cmd.should_use_raw_session().unwrap()); @@ -141,6 +142,7 @@ mod tests { use_pty: None, session_policy: None, ssh_connection_config: SshConnectionConfig::default(), + ssh_connection_config_resolver: Default::default(), }; let path = PathBuf::from("~/test/file.txt"); @@ -177,6 +179,7 @@ mod tests { use_pty: None, session_policy: None, ssh_connection_config: SshConnectionConfig::default(), + ssh_connection_config_resolver: Default::default(), }; let node = Node::new(String::from("example.com"), 22, String::from("alice")); diff --git a/src/jump/chain.rs b/src/jump/chain.rs index 3d0eab51..c4ed6844 100644 --- a/src/jump/chain.rs +++ b/src/jump/chain.rs @@ -153,7 +153,7 @@ impl JumpHostChain { self } - fn connection_config_for_host(&self, host: &str) -> SshConnectionConfig { + pub(crate) fn connection_config_for_host(&self, host: &str) -> SshConnectionConfig { self.ssh_connection_config_resolver .as_ref() .map(|resolver| resolver.resolve_for_host(host)) diff --git a/src/ssh/ssh_config/ip_qos.rs b/src/ssh/ssh_config/ip_qos.rs index cc658cf5..5d68253d 100644 --- a/src/ssh/ssh_config/ip_qos.rs +++ b/src/ssh/ssh_config/ip_qos.rs @@ -92,7 +92,7 @@ fn parse_value(value: &str) -> Result { "cs7" => 0xe0, "ef" => 0xb8, "le" => 0x04, - "va" => 0x2c, + "va" => 0xb0, // OpenSSH retains these names for compatibility but deliberately // leaves the system traffic class unchanged. "lowdelay" | "throughput" | "reliability" => return Ok(IpQosValue::None), @@ -149,6 +149,17 @@ mod tests { } } + #[test] + fn voice_admit_uses_the_on_wire_dscp_encoding() { + assert_eq!( + IpQosPolicy::parse(&values(&["va"])), + Ok(IpQosPolicy { + interactive: IpQosValue::Class(0xb0), + bulk: IpQosValue::Class(0xb0), + }) + ); + } + #[test] fn rejects_missing_extra_hex_and_unknown_values() { assert_eq!(IpQosPolicy::parse(&[]), Err(IpQosParseError::MissingValue)); diff --git a/tests/interactive_integration_test.rs b/tests/interactive_integration_test.rs index e0f93341..a7f86f15 100644 --- a/tests/interactive_integration_test.rs +++ b/tests/interactive_integration_test.rs @@ -56,6 +56,7 @@ fn test_interactive_command_builder() { session_policy: None, jump_hosts: None, ssh_connection_config: SshConnectionConfig::default(), + ssh_connection_config_resolver: Default::default(), }; assert!(!cmd.single_node); @@ -93,6 +94,7 @@ fn test_history_file_handling() { session_policy: None, jump_hosts: None, ssh_connection_config: SshConnectionConfig::default(), + ssh_connection_config_resolver: Default::default(), }; assert_eq!(cmd.history_file, history_path); @@ -193,6 +195,7 @@ async fn test_interactive_with_unreachable_nodes() { session_policy: None, jump_hosts: None, ssh_connection_config: SshConnectionConfig::default(), + ssh_connection_config_resolver: Default::default(), }; // This should fail to connect @@ -230,6 +233,7 @@ async fn test_interactive_with_no_nodes() { session_policy: None, jump_hosts: None, ssh_connection_config: SshConnectionConfig::default(), + ssh_connection_config_resolver: Default::default(), }; let result = cmd.execute().await; @@ -277,6 +281,7 @@ fn test_mode_configuration() { session_policy: None, jump_hosts: None, ssh_connection_config: SshConnectionConfig::default(), + ssh_connection_config_resolver: Default::default(), }; assert!(single_cmd.single_node); @@ -305,6 +310,7 @@ fn test_mode_configuration() { session_policy: None, jump_hosts: None, ssh_connection_config: SshConnectionConfig::default(), + ssh_connection_config_resolver: Default::default(), }; assert!(!multi_cmd.single_node); @@ -336,6 +342,7 @@ fn test_working_directory_config() { session_policy: None, jump_hosts: None, ssh_connection_config: SshConnectionConfig::default(), + ssh_connection_config_resolver: Default::default(), }; assert_eq!(cmd_with_dir.work_dir, Some("/var/www".to_string())); @@ -362,6 +369,7 @@ fn test_working_directory_config() { session_policy: None, jump_hosts: None, ssh_connection_config: SshConnectionConfig::default(), + ssh_connection_config_resolver: Default::default(), }; assert_eq!(cmd_without_dir.work_dir, None); @@ -400,6 +408,7 @@ fn test_prompt_format() { session_policy: None, jump_hosts: None, ssh_connection_config: SshConnectionConfig::default(), + ssh_connection_config_resolver: Default::default(), }; assert_eq!(cmd.prompt_format, format); diff --git a/tests/interactive_test.rs b/tests/interactive_test.rs index 64c6958d..d0233d22 100644 --- a/tests/interactive_test.rs +++ b/tests/interactive_test.rs @@ -46,6 +46,7 @@ async fn test_interactive_command_creation() { session_policy: None, jump_hosts: None, ssh_connection_config: SshConnectionConfig::default(), + ssh_connection_config_resolver: Default::default(), }; assert!(!cmd.single_node); @@ -77,6 +78,7 @@ async fn test_interactive_with_no_nodes() { session_policy: None, jump_hosts: None, ssh_connection_config: SshConnectionConfig::default(), + ssh_connection_config_resolver: Default::default(), }; let result = cmd.execute().await; diff --git a/tests/ssh_keepalive_test.rs b/tests/ssh_keepalive_test.rs index b5b7cb95..a3bfc521 100644 --- a/tests/ssh_keepalive_test.rs +++ b/tests/ssh_keepalive_test.rs @@ -698,6 +698,7 @@ fn test_interactive_mode_ssh_connection_config_default() { use_pty: None, session_policy: None, ssh_connection_config: SshConnectionConfig::default(), + ssh_connection_config_resolver: Default::default(), }; // Verify default values are applied @@ -747,6 +748,7 @@ fn test_interactive_mode_ssh_connection_config_custom() { use_pty: None, session_policy: None, ssh_connection_config: custom_config, + ssh_connection_config_resolver: Default::default(), }; // Verify custom values are applied @@ -794,6 +796,7 @@ fn test_interactive_mode_ssh_connection_config_disabled_keepalive() { use_pty: None, session_policy: None, ssh_connection_config: disabled_config, + ssh_connection_config_resolver: Default::default(), }; // Verify keepalive is disabled From 780744f584ab05769bae1528d83c1a6f44bb543d Mon Sep 17 00:00:00 2001 From: Jeongkyu Shin Date: Sun, 30 Aug 2026 15:22:46 +0900 Subject: [PATCH 3/6] fix(ssh): preserve resolved destination policy Keep final-destination SSH policy separate from jump-host resolution so HostName expansion cannot replace source, QoS, host-key, or proxy settings selected by the original alias. Resolve every interactive node from node.config_host(), pass that target config through host verification, direct connections, and jump construction, and make command and transfer ProxyJump lookup use the original alias. Preserve explicit jump-host precedence and fixed-config callers while applying the per-host resolver only to bastion aliases. Refs #300 --- src/commands/interactive/connection.rs | 131 +++++++++++++++++++---- src/commands/interactive/types.rs | 4 +- src/executor/connection_manager.rs | 90 +++++++--------- src/jump/chain.rs | 43 ++++++-- src/ssh/client/connection.rs | 140 +++++++++++++++++++++---- 5 files changed, 306 insertions(+), 102 deletions(-) diff --git a/src/commands/interactive/connection.rs b/src/commands/interactive/connection.rs index e9eaba3c..cff491a3 100644 --- a/src/commands/interactive/connection.rs +++ b/src/commands/interactive/connection.rs @@ -28,7 +28,7 @@ use crate::ssh::{ SessionPolicy, SessionPurpose, SessionRequest, known_hosts::get_check_method_for_target, tokio_client::{ - AuthMethod, Client, Error as SshError, ServerCheckMethod, SshConnectionConfig, + AuthMethod, Client, Error as SshError, ProxyMode, ServerCheckMethod, SshConnectionConfig, SshConnectionConfigResolver, }, }; @@ -57,6 +57,27 @@ fn interactive_session_purpose(session_policy: Option<&SessionPolicy>) -> Sessio session_policy.map_or(SessionPurpose::Interactive, SessionPolicy::purpose) } +fn interactive_target_connection_config( + node: &Node, + fixed_config: &SshConnectionConfig, + resolver: Option<&SshConnectionConfigResolver>, +) -> SshConnectionConfig { + resolver.map_or_else( + || fixed_config.clone(), + |resolver| resolver.resolve_for_host(node.config_host()), + ) +} + +fn interactive_jump_spec<'a>( + target_config: &'a SshConnectionConfig, + fallback: Option<&'a str>, +) -> Option<&'a str> { + fallback.or(match target_config.proxy_mode.as_ref() { + Some(ProxyMode::Jump(jump)) => Some(jump.as_str()), + Some(ProxyMode::Command(_) | ProxyMode::Direct) | None => None, + }) +} + impl InteractiveCommand { /// Helper function to establish SSH connection with proper error handling and rate limiting /// This eliminates code duplication across different connection paths and prevents brute-force attacks @@ -179,11 +200,12 @@ impl InteractiveCommand { &self, jump_hosts: Vec, adjusted_timeout: Duration, + target_config: &SshConnectionConfig, ) -> JumpHostChain { build_interactive_jump_chain( jump_hosts, adjusted_timeout, - &self.ssh_connection_config, + target_config, self.ssh_connection_config_resolver.as_ref(), self.session_purpose(), ) @@ -297,11 +319,16 @@ impl InteractiveCommand { pub(super) async fn connect_to_node(&self, node: Node) -> Result { // Determine authentication method using the same logic as exec mode let auth_method = self.determine_auth_method(&node).await?; + let target_config = interactive_target_connection_config( + &node, + &self.ssh_connection_config, + self.ssh_connection_config_resolver.as_ref(), + ); // Set up host key checking using the configured strict mode let check_method = get_check_method_for_target( self.strict_mode, - &self.ssh_connection_config, + &target_config, &node.host, node.port, &node.username, @@ -311,7 +338,9 @@ impl InteractiveCommand { let addr = (node.host.as_str(), node.port); // Create client connection - either direct or through jump hosts - let client = if let Some(ref jump_spec) = self.jump_hosts { + let client = if let Some(jump_spec) = + interactive_jump_spec(&target_config, self.jump_hosts.as_deref()) + { // Parse jump hosts let jump_hosts = parse_jump_hosts(jump_spec).with_context(|| { format!("Failed to parse jump host specification: '{jump_spec}'") @@ -330,7 +359,7 @@ impl InteractiveCommand { &node.host, node.port, !self.use_password, // Allow fallback unless explicit password mode - &self.ssh_connection_config, + &target_config, self.session_purpose(), ) .await? @@ -359,7 +388,7 @@ impl InteractiveCommand { // Pass SSH connection config to jump host chain for keepalive settings. // Also pass the dispatcher's pre-collected password so jump-host // authentication consumes it instead of re-prompting per call. See #200. - let chain = self.build_jump_chain(jump_hosts, adjusted_timeout); + let chain = self.build_jump_chain(jump_hosts, adjusted_timeout, &target_config); // Connect through the chain let connection = timeout( @@ -410,7 +439,7 @@ impl InteractiveCommand { &node.host, node.port, !self.use_password, // Allow fallback unless explicit password mode - &self.ssh_connection_config, + &target_config, self.session_purpose(), ) .await? @@ -455,11 +484,16 @@ impl InteractiveCommand { pub(super) async fn connect_to_node_pty(&self, node: Node) -> Result<(Client, Channel)> { // Determine authentication method using the same logic as exec mode let auth_method = self.determine_auth_method(&node).await?; + let target_config = interactive_target_connection_config( + &node, + &self.ssh_connection_config, + self.ssh_connection_config_resolver.as_ref(), + ); // Set up host key checking using the configured strict mode let check_method = get_check_method_for_target( self.strict_mode, - &self.ssh_connection_config, + &target_config, &node.host, node.port, &node.username, @@ -469,7 +503,9 @@ impl InteractiveCommand { let addr = (node.host.as_str(), node.port); // Create client connection - either direct or through jump hosts - let client = if let Some(ref jump_spec) = self.jump_hosts { + let client = if let Some(jump_spec) = + interactive_jump_spec(&target_config, self.jump_hosts.as_deref()) + { // Parse jump hosts let jump_hosts = parse_jump_hosts(jump_spec).with_context(|| { format!("Failed to parse jump host specification: '{jump_spec}'") @@ -488,7 +524,7 @@ impl InteractiveCommand { &node.host, node.port, !self.use_password, // Allow fallback unless explicit password mode - &self.ssh_connection_config, + &target_config, self.session_purpose(), ) .await? @@ -517,7 +553,7 @@ impl InteractiveCommand { // Pass SSH connection config to jump host chain for keepalive settings. // Also pass the dispatcher's pre-collected password so jump-host // authentication consumes it instead of re-prompting per call. See #200. - let chain = self.build_jump_chain(jump_hosts, adjusted_timeout); + let chain = self.build_jump_chain(jump_hosts, adjusted_timeout, &target_config); // Connect through the chain let connection = timeout( @@ -568,7 +604,7 @@ impl InteractiveCommand { &node.host, node.port, !self.use_password, // Allow fallback unless explicit password mode - &self.ssh_connection_config, + &target_config, self.session_purpose(), ) .await? @@ -688,35 +724,85 @@ Host bastion BindInterface lo IPQoS cs5 cs1 -Host target +Host alpha + HostName effective-alpha + HostKeyAlias alpha-key BindAddress 127.0.0.3 - BindInterface target0 + BindInterface alpha0 IPQoS ef cs2 + ProxyJump bastion + +Host beta + HostName effective-beta + HostKeyAlias beta-key + BindAddress 127.0.0.4 + BindInterface beta0 + IPQoS cs6 cs3 + ProxyJump beta-bastion + +Host effective-alpha + HostKeyAlias wrong-key + BindAddress 127.0.0.9 + BindInterface wrong0 + IPQoS cs7 cs7 "#, ) .expect("valid ssh_config"); let resolver = SshConnectionConfigResolver::new().with_ssh_config(Some(ssh_config)); + let alpha = Node::new("effective-alpha".to_string(), 22, "user".to_string()) + .with_original_host("alpha".to_string()); + let beta = Node::new("effective-beta".to_string(), 22, "user".to_string()) + .with_original_host("beta".to_string()); + let fixed_config = SshConnectionConfig::default(); + let alpha_config = + interactive_target_connection_config(&alpha, &fixed_config, Some(&resolver)); + let beta_config = + interactive_target_connection_config(&beta, &fixed_config, Some(&resolver)); + + assert_eq!(alpha_config.bind_address.as_deref(), Some("127.0.0.3")); + assert_eq!(alpha_config.bind_interface.as_deref(), Some("alpha0")); + assert_eq!(alpha_config.host_key_alias.as_deref(), Some("alpha-key")); + assert_eq!(alpha_config.ip_qos.bulk, IpQosValue::Class(0x40)); + assert_eq!(beta_config.bind_address.as_deref(), Some("127.0.0.4")); + assert_eq!(beta_config.bind_interface.as_deref(), Some("beta0")); + assert_eq!(beta_config.host_key_alias.as_deref(), Some("beta-key")); + assert_eq!(beta_config.ip_qos.bulk, IpQosValue::Class(0x60)); + assert_eq!(interactive_jump_spec(&alpha_config, None), Some("bastion")); + assert_eq!( + interactive_jump_spec(&alpha_config, Some("manual-bastion")), + Some("manual-bastion") + ); + assert_eq!( + interactive_jump_spec(&beta_config, None), + Some("beta-bastion") + ); + let chain = build_interactive_jump_chain( vec![JumpHost::new("bastion".to_string(), None, None)], Duration::from_secs(45), - &SshConnectionConfig::default(), + &alpha_config, Some(&resolver), SessionPurpose::Bulk, ); - let bastion = chain.connection_config_for_host("bastion"); + let bastion = chain.connection_config_for_jump_host("bastion"); assert_eq!(bastion.bind_address.as_deref(), Some("127.0.0.2")); assert_eq!(bastion.bind_interface.as_deref(), Some("lo")); assert_eq!(bastion.session_purpose, SessionPurpose::Bulk); assert_eq!(bastion.selected_ip_qos(), IpQosValue::Class(0x20)); - let target = chain.connection_config_for_host("target"); + let target = chain.destination_connection_config(); assert_eq!(target.bind_address.as_deref(), Some("127.0.0.3")); - assert_eq!(target.bind_interface.as_deref(), Some("target0")); + assert_eq!(target.bind_interface.as_deref(), Some("alpha0")); + assert_eq!(target.host_key_alias.as_deref(), Some("alpha-key")); + assert!(matches!( + target.proxy_mode.as_ref(), + Some(ProxyMode::Jump(jump)) if jump == "bastion" + )); assert_eq!(target.session_purpose, SessionPurpose::Bulk); assert_eq!(target.selected_ip_qos(), IpQosValue::Class(0x40)); - let fixed_config = SshConnectionConfig::new() + let manual_config = SshConnectionConfig::new() .with_source_binding(Some("127.0.0.4".to_string()), Some("lo".to_string())) .with_ip_qos(IpQosPolicy { interactive: IpQosValue::Class(0xb8), @@ -725,13 +811,16 @@ Host target let fixed_chain = build_interactive_jump_chain( vec![JumpHost::new("manual-bastion".to_string(), None, None)], Duration::from_secs(45), - &fixed_config, + &manual_config, None, SessionPurpose::Bulk, ); - let manual_bastion = fixed_chain.connection_config_for_host("manual-bastion"); + let manual_bastion = fixed_chain.connection_config_for_jump_host("manual-bastion"); assert_eq!(manual_bastion.bind_address.as_deref(), Some("127.0.0.4")); assert_eq!(manual_bastion.selected_ip_qos(), IpQosValue::Class(0x60)); + let manual_target = fixed_chain.destination_connection_config(); + assert_eq!(manual_target.bind_address.as_deref(), Some("127.0.0.4")); + assert_eq!(manual_target.selected_ip_qos(), IpQosValue::Class(0x60)); } #[test] diff --git a/src/commands/interactive/types.rs b/src/commands/interactive/types.rs index 53c98899..84ceb58d 100644 --- a/src/commands/interactive/types.rs +++ b/src/commands/interactive/types.rs @@ -69,9 +69,9 @@ pub struct InteractiveCommand { pub use_pty: Option, // None = auto-detect, Some(true) = force, Some(false) = disable /// Resolved live ssh_config policy for the SSH-compatible interactive path. pub session_policy: Option, - // SSH connection configuration (keepalive settings) + /// Fixed connection-policy fallback for callers without a resolver. pub ssh_connection_config: SshConnectionConfig, - /// Per-host connection policy retained for ProxyJump bastions. + /// Per-host policy resolver for each interactive target and jump alias. pub ssh_connection_config_resolver: Option, } diff --git a/src/executor/connection_manager.rs b/src/executor/connection_manager.rs index 56b38081..26358b2b 100644 --- a/src/executor/connection_manager.rs +++ b/src/executor/connection_manager.rs @@ -66,19 +66,10 @@ pub(crate) async fn execute_on_node_with_jump_hosts( let key_path = config.key_path.map(Path::new); - // Determine effective jump hosts: CLI takes precedence, then SSH config - // Store the SSH config jump hosts String to extend its lifetime - let ssh_config_jump_hosts = config - .ssh_config - .and_then(|ssh_config| ssh_config.get_proxy_jump(&node.host)); - - let effective_jump_hosts = if config.jump_hosts.is_some() { - // CLI jump hosts specified - config.jump_hosts - } else { - // Fall back to SSH config ProxyJump for this specific host - ssh_config_jump_hosts.as_deref() - }; + // Resolve ProxyJump against the original ssh_config alias, before HostName + // expansion. CLI remains authoritative. + let effective_jump_hosts = + resolve_effective_jump_hosts(config.jump_hosts, config.ssh_config, node.config_host()); let session_policy = config .ssh_config @@ -90,7 +81,7 @@ pub(crate) async fn execute_on_node_with_jump_hosts( (!command.is_empty()).then_some(command), config.tty_mode, std::io::stdin().is_terminal(), - effective_jump_hosts, + effective_jump_hosts.as_deref(), ) }) .transpose()?; @@ -104,7 +95,7 @@ pub(crate) async fn execute_on_node_with_jump_hosts( use_keychain: config.use_keychain, timeout_seconds: config.timeout, connect_timeout_seconds: config.connect_timeout, - jump_hosts_spec: effective_jump_hosts, + jump_hosts_spec: effective_jump_hosts.as_deref(), ssh_connection_config: config.ssh_connection_config, ssh_connection_config_resolver: config.ssh_connection_config_resolver, session_policy: session_policy.as_ref(), @@ -169,15 +160,8 @@ pub(crate) async fn upload_to_node( let key_path = key_path.map(Path::new); - // Determine effective jump hosts: CLI takes precedence, then SSH config - let ssh_config_jump_hosts = - ssh_config.and_then(|ssh_config| ssh_config.get_proxy_jump(&node.host)); - - let effective_jump_hosts = if jump_hosts.is_some() { - jump_hosts - } else { - ssh_config_jump_hosts.as_deref() - }; + let effective_jump_hosts = + resolve_effective_jump_hosts(jump_hosts, ssh_config, node.config_host()); // Check if the local path is a directory if local_path.is_dir() { @@ -189,7 +173,7 @@ pub(crate) async fn upload_to_node( Some(strict_mode), use_agent, use_password, - effective_jump_hosts, + effective_jump_hosts.as_deref(), connect_timeout_seconds, pre_collected_password, ssh_connection_config, @@ -205,7 +189,7 @@ pub(crate) async fn upload_to_node( Some(strict_mode), use_agent, use_password, - effective_jump_hosts, + effective_jump_hosts.as_deref(), connect_timeout_seconds, pre_collected_password, ssh_connection_config, @@ -236,15 +220,8 @@ pub(crate) async fn download_from_node( let key_path = key_path.map(Path::new); - // Determine effective jump hosts: CLI takes precedence, then SSH config - let ssh_config_jump_hosts = - ssh_config.and_then(|ssh_config| ssh_config.get_proxy_jump(&node.host)); - - let effective_jump_hosts = if jump_hosts.is_some() { - jump_hosts - } else { - ssh_config_jump_hosts.as_deref() - }; + let effective_jump_hosts = + resolve_effective_jump_hosts(jump_hosts, ssh_config, node.config_host()); // This function handles both files and directories // The caller should check if it's a directory and use the appropriate method @@ -256,7 +233,7 @@ pub(crate) async fn download_from_node( Some(strict_mode), use_agent, use_password, - effective_jump_hosts, + effective_jump_hosts.as_deref(), connect_timeout_seconds, pre_collected_password, ssh_connection_config, @@ -288,15 +265,8 @@ pub async fn download_dir_from_node( let key_path = key_path.map(Path::new); - // Determine effective jump hosts: CLI takes precedence, then SSH config - let ssh_config_jump_hosts = - ssh_config.and_then(|ssh_config| ssh_config.get_proxy_jump(&node.host)); - - let effective_jump_hosts = if jump_hosts.is_some() { - jump_hosts - } else { - ssh_config_jump_hosts.as_deref() - }; + let effective_jump_hosts = + resolve_effective_jump_hosts(jump_hosts, ssh_config, node.config_host()); client .download_dir_with_jump_hosts( @@ -306,7 +276,7 @@ pub async fn download_dir_from_node( Some(strict_mode), use_agent, use_password, - effective_jump_hosts, + effective_jump_hosts.as_deref(), connect_timeout_seconds, pre_collected_password, ssh_connection_config, @@ -322,24 +292,44 @@ pub async fn download_dir_from_node( /// 2. SSH config ProxyJump for the specific host /// 3. None (direct connection) /// -/// This is extracted for testing purposes and used internally by all connection functions. -#[allow(dead_code)] // Used for testing #[inline] fn resolve_effective_jump_hosts( cli_jump_hosts: Option<&str>, ssh_config: Option<&SshConfig>, - hostname: &str, + config_host: &str, ) -> Option { if cli_jump_hosts.is_some() { return cli_jump_hosts.map(String::from); } - ssh_config.and_then(|config| config.get_proxy_jump(hostname)) + ssh_config.and_then(|config| config.get_proxy_jump(config_host)) } #[cfg(test)] mod tests { use super::*; + #[test] + fn command_and_transfer_proxy_jump_uses_original_host_alias() { + let ssh_config = SshConfig::parse( + r#" +Host target-alias + HostName effective-target + ProxyJump alias-bastion + +Host effective-target + ProxyJump wrong-bastion +"#, + ) + .expect("valid ssh_config"); + let node = Node::new("effective-target".to_string(), 22, "user".to_string()) + .with_original_host("target-alias".to_string()); + + assert_eq!( + resolve_effective_jump_hosts(None, Some(&ssh_config), node.config_host()), + Some("alias-bastion".to_string()) + ); + } + /// Test that CLI jump hosts take precedence over SSH config #[test] fn test_resolve_effective_jump_hosts_cli_precedence() { diff --git a/src/jump/chain.rs b/src/jump/chain.rs index c4ed6844..3c9eab52 100644 --- a/src/jump/chain.rs +++ b/src/jump/chain.rs @@ -67,9 +67,13 @@ pub struct JumpHostChain { max_idle_time: Duration, /// Maximum connection age before forced renewal (default: 30 minutes) max_connection_age: Duration, - /// SSH connection configuration (keepalive settings) + /// Pre-resolved final-destination configuration and fixed-config fallback. + /// + /// This is resolved from the destination's original ssh_config alias before + /// `HostName` expansion and must not be looked up again using the effective + /// network host. When no resolver is configured, jump hosts also use it. ssh_connection_config: SshConnectionConfig, - /// Per-host SSH connection configuration resolver. + /// Per-host SSH connection configuration resolver for jump hosts only. ssh_connection_config_resolver: Option, /// Traffic profile propagated to every direct socket in the chain. session_purpose: SessionPurpose, @@ -131,7 +135,9 @@ impl JumpHostChain { /// Set SSH connection configuration (keepalive settings) /// /// Configures keepalive interval and maximum attempts to prevent - /// idle connection timeouts during jump host operations. + /// idle connection timeouts during jump host operations. The same config + /// is retained as the pre-resolved destination policy, while a later + /// per-host resolver applies only to jump-host aliases. pub fn with_ssh_connection_config(mut self, config: SshConnectionConfig) -> Self { self.ssh_connection_config = config; self.ssh_connection_config_resolver = None; @@ -153,7 +159,7 @@ impl JumpHostChain { self } - pub(crate) fn connection_config_for_host(&self, host: &str) -> SshConnectionConfig { + pub(crate) fn connection_config_for_jump_host(&self, host: &str) -> SshConnectionConfig { self.ssh_connection_config_resolver .as_ref() .map(|resolver| resolver.resolve_for_host(host)) @@ -161,6 +167,12 @@ impl JumpHostChain { .with_session_purpose(self.session_purpose) } + pub(crate) fn destination_connection_config(&self) -> SshConnectionConfig { + self.ssh_connection_config + .clone() + .with_session_purpose(self.session_purpose) + } + /// Create a direct connection chain (no jump hosts) pub fn direct() -> Self { Self::new(Vec::new()) @@ -238,7 +250,7 @@ impl JumpHostChain { } if self.is_direct() { - let ssh_connection_config = self.connection_config_for_host(destination_host); + let ssh_connection_config = self.destination_connection_config(); chain_connection::connect_direct( destination_host, destination_port, @@ -310,7 +322,7 @@ impl JumpHostChain { // Step 2: Chain through intermediate jump hosts for (i, jump_host) in self.jump_hosts.iter().skip(1).enumerate() { let ssh_connection_config = self - .connection_config_for_host(&jump_host.host) + .connection_config_for_jump_host(&jump_host.host) .without_forwarding(); debug!( "Connecting to intermediate jump host {} of {}: {}", @@ -338,7 +350,7 @@ impl JumpHostChain { } // Step 3: Connect to final destination through the last jump host - let ssh_connection_config = self.connection_config_for_host(destination_host); + let ssh_connection_config = self.destination_connection_config(); let final_client = tunnel::connect_to_destination( ¤t_client, destination_host, @@ -388,7 +400,7 @@ impl JumpHostChain { ) -> Result { let jump_host = &self.jump_hosts[0]; let ssh_connection_config = self - .connection_config_for_host(&jump_host.host) + .connection_config_for_jump_host(&jump_host.host) .without_forwarding(); debug!( @@ -562,26 +574,35 @@ Host bastion ServerAliveInterval 11 ServerAliveCountMax 2 -Host target +Host target-alias + HostName effective-target AddressFamily inet6 Compression no ServerAliveInterval 22 ServerAliveCountMax 4 + +Host effective-target + AddressFamily inet + Compression yes + ServerAliveInterval 99 + ServerAliveCountMax 9 "#, ) .expect("valid ssh_config"); let resolver = SshConnectionConfigResolver::new().with_ssh_config(Some(ssh_config)); + let destination_config = resolver.resolve_for_host("target-alias"); let chain = JumpHostChain::new(vec![JumpHost::new("bastion".to_string(), None, None)]) + .with_ssh_connection_config(destination_config) .with_ssh_connection_config_resolver(resolver); - let bastion = chain.connection_config_for_host("bastion"); + let bastion = chain.connection_config_for_jump_host("bastion"); assert_eq!(bastion.address_family, AddressFamily::V4); assert!(bastion.compression); assert_eq!(bastion.keepalive_interval, Some(11)); assert_eq!(bastion.keepalive_max, 2); - let target = chain.connection_config_for_host("target"); + let target = chain.destination_connection_config(); assert_eq!(target.address_family, AddressFamily::V6); assert!(!target.compression); assert_eq!(target.keepalive_interval, Some(22)); diff --git a/src/ssh/client/connection.rs b/src/ssh/client/connection.rs index 38ce8d48..04917732 100644 --- a/src/ssh/client/connection.rs +++ b/src/ssh/client/connection.rs @@ -13,7 +13,7 @@ // limitations under the License. use super::core::SshClient; -use crate::jump::{JumpHostChain, parse_jump_hosts}; +use crate::jump::{JumpHostChain, parse_jump_hosts, parser::JumpHost}; use crate::security::Password; use crate::ssh::SessionPurpose; use crate::ssh::known_hosts::StrictHostKeyChecking; @@ -31,6 +31,37 @@ use std::time::Duration; // - Balances user patience with reliability on poor networks const SSH_CONNECT_TIMEOUT_SECS: u64 = 30; +fn build_client_jump_chain( + jump_hosts: &[JumpHost], + connect_timeout: Duration, + target_config: Option<&SshConnectionConfig>, + jump_host_resolver: Option<&SshConnectionConfigResolver>, + pre_collected_password: Option>, + session_purpose: SessionPurpose, +) -> JumpHostChain { + let mut chain = JumpHostChain::new(jump_hosts.to_vec()) + .with_connect_timeout(connect_timeout) + .with_command_timeout(Duration::from_secs(300)) + .with_ssh_password(pre_collected_password); + if let Some(config) = target_config { + chain = chain.with_ssh_connection_config(config.clone()); + } + if let Some(resolver) = jump_host_resolver { + chain = chain.with_ssh_connection_config_resolver(resolver.clone()); + } + chain.with_session_purpose(session_purpose) +} + +fn client_jump_spec<'a>( + target_config: &'a SshConnectionConfig, + requested_jump_hosts: Option<&'a str>, +) -> Option<&'a str> { + requested_jump_hosts.or(match target_config.proxy_mode.as_ref() { + Some(ProxyMode::Jump(jump)) => Some(jump.as_str()), + Some(ProxyMode::Command(_) | ProxyMode::Direct) | None => None, + }) +} + /// Build the friendly, outer-context message for a failed direct SSH /// connection attempt. /// @@ -242,17 +273,14 @@ impl SshClient { // Create jump host chain with user-specified or default connect timeout let connect_timeout = Duration::from_secs(connect_timeout_seconds.unwrap_or(SSH_CONNECT_TIMEOUT_SECS)); - let mut chain = JumpHostChain::new(jump_hosts.to_vec()) - .with_connect_timeout(connect_timeout) - .with_command_timeout(Duration::from_secs(300)) - .with_ssh_password(pre_collected_password); - if let Some(cfg) = ssh_connection_config { - chain = chain.with_ssh_connection_config(cfg.clone()); - } - if let Some(resolver) = ssh_connection_config_resolver { - chain = chain.with_ssh_connection_config_resolver(resolver.clone()); - } - chain = chain.with_session_purpose(session_purpose); + let chain = build_client_jump_chain( + jump_hosts, + connect_timeout, + ssh_connection_config, + ssh_connection_config_resolver, + pre_collected_password, + session_purpose, + ); // Connect through the chain let connection = chain @@ -303,12 +331,7 @@ impl SshClient { .unwrap_or_default() .with_session_purpose(session_purpose); let ssh_connection_config = Some(&selected_config); - let jump_hosts_spec = - match ssh_connection_config.and_then(|config| config.proxy_mode.as_ref()) { - Some(ProxyMode::Jump(jump)) => Some(jump.as_str()), - Some(ProxyMode::Command(_) | ProxyMode::Direct) => None, - None => jump_hosts_spec, - }; + let jump_hosts_spec = client_jump_spec(&selected_config, jump_hosts_spec); if let Some(jump_spec) = jump_hosts_spec { // Parse jump hosts @@ -370,6 +393,8 @@ impl SshClient { #[cfg(test)] mod tests { use super::*; + use crate::ssh::SshConfig; + use crate::ssh::ssh_config::{IpQosPolicy, IpQosValue}; use crate::test_helpers::EnvGuard; use serial_test::serial; use tempfile::TempDir; @@ -381,6 +406,85 @@ mod tests { std::fs::write(path, key.to_openssh(LineEnding::LF).unwrap().as_bytes()).unwrap(); } + #[test] + fn command_and_transfer_jump_chain_preserves_alias_target_policy() { + let ssh_config = SshConfig::parse( + r#" +Host bastion + BindAddress 127.0.0.2 + IPQoS cs5 cs1 + +Host target-alias + HostName effective-target + HostKeyAlias alias-key + BindAddress 127.0.0.3 + IPQoS ef cs2 + ProxyJump bastion + +Host effective-target + HostKeyAlias wrong-key + BindAddress 127.0.0.9 + IPQoS cs7 cs7 +"#, + ) + .expect("valid ssh_config"); + let resolver = SshConnectionConfigResolver::new().with_ssh_config(Some(ssh_config)); + let target_config = resolver.resolve_for_host("target-alias"); + let chain = build_client_jump_chain( + &[JumpHost::new("bastion".to_string(), None, None)], + Duration::from_secs(30), + Some(&target_config), + Some(&resolver), + None, + SessionPurpose::Bulk, + ); + + let bastion = chain.connection_config_for_jump_host("bastion"); + assert_eq!(bastion.bind_address.as_deref(), Some("127.0.0.2")); + assert_eq!(bastion.selected_ip_qos(), IpQosValue::Class(0x20)); + let destination = chain.destination_connection_config(); + assert_eq!(destination.bind_address.as_deref(), Some("127.0.0.3")); + assert_eq!(destination.host_key_alias.as_deref(), Some("alias-key")); + assert_eq!(destination.selected_ip_qos(), IpQosValue::Class(0x40)); + assert!(matches!( + destination.proxy_mode.as_ref(), + Some(ProxyMode::Jump(jump)) if jump == "bastion" + )); + assert_eq!(client_jump_spec(&destination, None), Some("bastion")); + assert_eq!( + client_jump_spec(&destination, Some("manual-bastion")), + Some("manual-bastion") + ); + + let fixed_config = SshConnectionConfig::new() + .with_source_binding(Some("127.0.0.4".to_string()), None) + .with_ip_qos(IpQosPolicy { + interactive: IpQosValue::Class(0xb8), + bulk: IpQosValue::Class(0x60), + }); + let fixed_chain = build_client_jump_chain( + &[JumpHost::new("fixed-bastion".to_string(), None, None)], + Duration::from_secs(30), + Some(&fixed_config), + None, + None, + SessionPurpose::Bulk, + ); + assert_eq!( + fixed_chain + .connection_config_for_jump_host("fixed-bastion") + .bind_address + .as_deref(), + Some("127.0.0.4") + ); + assert_eq!( + fixed_chain + .destination_connection_config() + .selected_ip_qos(), + IpQosValue::Class(0x60) + ); + } + #[tokio::test] async fn test_determine_auth_method_with_key() { let temp_dir = TempDir::new().unwrap(); From f209db9a69f48dde19d3fd600207aa42afe2403d Mon Sep 17 00:00:00 2001 From: Jeongkyu Shin Date: Sun, 30 Aug 2026 15:40:10 +0900 Subject: [PATCH 4/6] fix(ssh): normalize direct ProxyJump values ProxyJump none was passed through as a requested jump specification, so command and transfer connections parsed it as a bastion named `none` even though the resolved target policy correctly selected a direct transport. Normalize none, direct, and the internal empty sentinel before jump selection while preserving the explicit decision over lower-priority configuration. Real CLI jump hosts remain authoritative, and both typed and fixed-config callers defensively reject direct markers as hop names. Cover command, file and directory upload, file and directory download, interactive selection, alias policy preservation, and typed resolver fallback behavior with focused regressions. Refs #300 --- src/commands/interactive/connection.rs | 16 +- src/executor/connection_manager.rs | 178 +++++++++++++++++++++-- src/ssh/client/connection.rs | 22 ++- src/ssh/tokio_client/connection.rs | 4 +- src/ssh/tokio_client/connection_tests.rs | 7 +- src/ssh/tokio_client/mod.rs | 1 + src/ssh/tokio_client/proxy_command.rs | 9 ++ 7 files changed, 210 insertions(+), 27 deletions(-) diff --git a/src/commands/interactive/connection.rs b/src/commands/interactive/connection.rs index cff491a3..dfdbf909 100644 --- a/src/commands/interactive/connection.rs +++ b/src/commands/interactive/connection.rs @@ -29,7 +29,7 @@ use crate::ssh::{ known_hosts::get_check_method_for_target, tokio_client::{ AuthMethod, Client, Error as SshError, ProxyMode, ServerCheckMethod, SshConnectionConfig, - SshConnectionConfigResolver, + SshConnectionConfigResolver, is_direct_proxy_jump, }, }; @@ -72,10 +72,13 @@ fn interactive_jump_spec<'a>( target_config: &'a SshConnectionConfig, fallback: Option<&'a str>, ) -> Option<&'a str> { - fallback.or(match target_config.proxy_mode.as_ref() { - Some(ProxyMode::Jump(jump)) => Some(jump.as_str()), - Some(ProxyMode::Command(_) | ProxyMode::Direct) | None => None, - }) + if let Some(requested) = fallback { + return (!is_direct_proxy_jump(requested)).then_some(requested); + } + match target_config.proxy_mode.as_ref() { + Some(ProxyMode::Jump(jump)) if !is_direct_proxy_jump(jump) => Some(jump.as_str()), + Some(ProxyMode::Jump(_) | ProxyMode::Command(_) | ProxyMode::Direct) | None => None, + } } impl InteractiveCommand { @@ -772,6 +775,9 @@ Host effective-alpha interactive_jump_spec(&alpha_config, Some("manual-bastion")), Some("manual-bastion") ); + for direct in ["", "none", "direct", " NONE "] { + assert_eq!(interactive_jump_spec(&alpha_config, Some(direct)), None); + } assert_eq!( interactive_jump_spec(&beta_config, None), Some("beta-bastion") diff --git a/src/executor/connection_manager.rs b/src/executor/connection_manager.rs index 26358b2b..90d7870f 100644 --- a/src/executor/connection_manager.rs +++ b/src/executor/connection_manager.rs @@ -25,7 +25,7 @@ use crate::ssh::{ CliTtyMode, SessionPolicy, SshClient, SshConfig, client::{CommandResult, ConnectionConfig}, known_hosts::StrictHostKeyChecking, - tokio_client::{SshConnectionConfig, SshConnectionConfigResolver}, + tokio_client::{SshConnectionConfig, SshConnectionConfigResolver, is_direct_proxy_jump}, }; /// Configuration for node execution. @@ -298,10 +298,24 @@ fn resolve_effective_jump_hosts( ssh_config: Option<&SshConfig>, config_host: &str, ) -> Option { - if cli_jump_hosts.is_some() { - return cli_jump_hosts.map(String::from); + if let Some(jump_hosts) = cli_jump_hosts { + return Some(normalize_jump_hosts(jump_hosts)); + } + ssh_config + .and_then(|config| config.get_proxy_jump(config_host)) + .map(|jump_hosts| normalize_jump_hosts(&jump_hosts)) +} + +/// Preserve an explicit direct decision as an empty jump specification. +/// +/// `Some("")` remains authoritative over lower-priority ssh_config values, +/// while both the session-policy and SSH client layers interpret it as no hop. +fn normalize_jump_hosts(jump_hosts: &str) -> String { + if is_direct_proxy_jump(jump_hosts) { + String::new() + } else { + jump_hosts.to_string() } - ssh_config.and_then(|config| config.get_proxy_jump(config_host)) } #[cfg(test)] @@ -346,6 +360,15 @@ Host example.com "example.com", ); assert_eq!(result, Some("cli-bastion.example.com".to_string())); + + // CLI can explicitly disable a configured jump without allowing the + // ssh_config fallback to become active again. + for direct in ["none", "direct", " NONE "] { + assert_eq!( + resolve_effective_jump_hosts(Some(direct), Some(&ssh_config), "example.com"), + Some(String::new()) + ); + } } /// Test that SSH config ProxyJump is used when no CLI jump hosts @@ -456,23 +479,148 @@ Host *.internal.example.com assert_eq!(result, None); } - /// Test ProxyJump none value (disables jump) #[test] - fn test_resolve_effective_jump_hosts_none_value() { - let ssh_config_content = r#" + fn direct_proxy_jump_values_normalize_to_an_empty_chain() { + for directive in ["none", "direct", "NONE", "DIRECT"] { + let ssh_config_content = format!( + r#" Host direct.example.com - ProxyJump none + ProxyJump {directive} Host *.example.com ProxyJump gateway.example.com -"#; - let ssh_config = SshConfig::parse(ssh_config_content).unwrap(); +"# + ); + let ssh_config = SshConfig::parse(&ssh_config_content).unwrap(); + + let result = + resolve_effective_jump_hosts(None, Some(&ssh_config), "direct.example.com"); + assert_eq!(result.as_deref(), Some(""), "value={directive}"); + assert!( + crate::jump::parse_jump_hosts(result.as_deref().unwrap()) + .unwrap() + .is_empty(), + "value={directive}" + ); + } + } + + #[tokio::test] + async fn command_and_all_transfer_paths_never_dial_proxy_jump_none() { + let listener = std::net::TcpListener::bind(("127.0.0.1", 0)).unwrap(); + let port = listener.local_addr().unwrap().port(); + drop(listener); - // The explicit disable must be obtained before the wildcard fallback. - // Note: The actual handling of "none" as special value would be - // done by the connection layer, but the config should return it - let result = resolve_effective_jump_hosts(None, Some(&ssh_config), "direct.example.com"); - assert_eq!(result, Some("none".to_string())); + let ssh_config = SshConfig::parse( + r#" +Host direct-target + HostName 127.0.0.1 + ProxyJump none +"#, + ) + .unwrap(); + let resolver = SshConnectionConfigResolver::new().with_ssh_config(Some(ssh_config.clone())); + let target_config = resolver.resolve_for_host("direct-target"); + let node = Node::new("127.0.0.1".to_string(), port, "user".to_string()) + .with_original_host("direct-target".to_string()); + let password = Arc::new(Password::new("test-password".to_string()).unwrap()); + let temp_dir = tempfile::TempDir::new().unwrap(); + let upload_file = temp_dir.path().join("upload.txt"); + std::fs::write(&upload_file, b"test").unwrap(); + + let assert_direct_target = |path: &str, error: anyhow::Error| { + let message = format!("{error:#}"); + assert!(message.contains("127.0.0.1"), "path={path}: {message}"); + assert!(!message.contains("none:22"), "path={path}: {message}"); + assert!( + !message.contains("jump host none"), + "path={path}: {message}" + ); + }; + + let execution_config = ExecutionConfig { + key_path: None, + strict_mode: StrictHostKeyChecking::AcceptNew, + use_agent: false, + use_password: true, + #[cfg(target_os = "macos")] + use_keychain: false, + timeout: Some(1), + connect_timeout: Some(1), + jump_hosts: None, + sudo_password: None, + ssh_password: Some(password.clone()), + ssh_config: Some(&ssh_config), + tty_mode: CliTtyMode::Disable, + ssh_connection_config: Some(&target_config), + ssh_connection_config_resolver: Some(&resolver), + }; + let error = execute_on_node_with_jump_hosts(node.clone(), "true", &execution_config) + .await + .unwrap_err(); + assert_direct_target("command", error); + + for (path, local_path) in [ + ("upload-file", upload_file.as_path()), + ("upload-directory", temp_dir.path()), + ] { + let error = upload_to_node( + node.clone(), + local_path, + "/tmp/remote", + None, + StrictHostKeyChecking::AcceptNew, + false, + true, + None, + Some(1), + Some(&ssh_config), + Some(password.clone()), + &target_config, + &resolver, + ) + .await + .unwrap_err(); + assert_direct_target(path, error); + } + + let error = download_from_node( + node.clone(), + "/tmp/remote-file", + &temp_dir.path().join("download.txt"), + None, + StrictHostKeyChecking::AcceptNew, + false, + true, + None, + Some(1), + Some(&ssh_config), + Some(password.clone()), + &target_config, + &resolver, + ) + .await + .unwrap_err(); + assert_direct_target("download-file", error); + + let error = download_dir_from_node( + node, + "/tmp/remote-dir", + &temp_dir.path().join("download-dir"), + None, + StrictHostKeyChecking::AcceptNew, + false, + true, + None, + Some(1), + Some(&ssh_config), + Some(password), + &target_config, + &resolver, + ) + .await + .unwrap_err(); + assert_direct_target("download-directory", error); } /// Test complex multi-hop chain with user and ports diff --git a/src/ssh/client/connection.rs b/src/ssh/client/connection.rs index 04917732..7d684632 100644 --- a/src/ssh/client/connection.rs +++ b/src/ssh/client/connection.rs @@ -19,6 +19,7 @@ use crate::ssh::SessionPurpose; use crate::ssh::known_hosts::StrictHostKeyChecking; use crate::ssh::tokio_client::{ AuthMethod, Client, ProxyMode, SshConnectionConfig, SshConnectionConfigResolver, + is_direct_proxy_jump, }; use anyhow::{Context, Result}; use std::path::Path; @@ -56,10 +57,13 @@ fn client_jump_spec<'a>( target_config: &'a SshConnectionConfig, requested_jump_hosts: Option<&'a str>, ) -> Option<&'a str> { - requested_jump_hosts.or(match target_config.proxy_mode.as_ref() { - Some(ProxyMode::Jump(jump)) => Some(jump.as_str()), - Some(ProxyMode::Command(_) | ProxyMode::Direct) | None => None, - }) + if let Some(requested) = requested_jump_hosts { + return (!is_direct_proxy_jump(requested)).then_some(requested); + } + match target_config.proxy_mode.as_ref() { + Some(ProxyMode::Jump(jump)) if !is_direct_proxy_jump(jump) => Some(jump.as_str()), + Some(ProxyMode::Jump(_) | ProxyMode::Command(_) | ProxyMode::Direct) | None => None, + } } /// Build the friendly, outer-context message for a failed direct SSH @@ -455,6 +459,16 @@ Host effective-target client_jump_spec(&destination, Some("manual-bastion")), Some("manual-bastion") ); + for direct in ["", "none", "direct", " NONE "] { + assert_eq!(client_jump_spec(&destination, Some(direct)), None); + } + + let direct_destination = + SshConnectionConfig::new().with_proxy_mode(Some(ProxyMode::Direct)); + assert_eq!( + client_jump_spec(&direct_destination, Some("cli-bastion")), + Some("cli-bastion") + ); let fixed_config = SshConnectionConfig::new() .with_source_binding(Some("127.0.0.4".to_string()), None) diff --git a/src/ssh/tokio_client/connection.rs b/src/ssh/tokio_client/connection.rs index 9c6ce60e..a6b557a0 100644 --- a/src/ssh/tokio_client/connection.rs +++ b/src/ssh/tokio_client/connection.rs @@ -32,7 +32,7 @@ use super::address_family::AddressFamily; use super::auth_policy::SshAuthenticationPolicy; use super::authentication::{AuthMethod, ServerCheckMethod}; use super::proxy_command::{ - ProxyCommandConfig, ProxyCommandProcess, ProxyMode, spawn_proxy_command, + ProxyCommandConfig, ProxyCommandProcess, ProxyMode, is_direct_proxy_jump, spawn_proxy_command, }; use crate::forwarding::remote::RemoteForwardRegistry; use crate::forwarding::{ForwardingDirective, ForwardingPlan, ForwardingRuntime}; @@ -619,7 +619,7 @@ impl SshConnectionConfigResolver { } fn proxy_jump_mode(jump: &str) -> ProxyMode { - if jump.eq_ignore_ascii_case("none") || jump.is_empty() { + if is_direct_proxy_jump(jump) { ProxyMode::Direct } else { ProxyMode::Jump(jump.to_string()) diff --git a/src/ssh/tokio_client/connection_tests.rs b/src/ssh/tokio_client/connection_tests.rs index f0f81b40..67f49b89 100644 --- a/src/ssh/tokio_client/connection_tests.rs +++ b/src/ssh/tokio_client/connection_tests.rs @@ -347,7 +347,12 @@ Host original-host #[test] fn test_proxy_none_disables_yaml_jump_fallback() { - for directive in ["ProxyCommand none", "ProxyJump none"] { + for directive in [ + "ProxyCommand none", + "ProxyJump none", + "ProxyJump direct", + "ProxyJump DIRECT", + ] { let ssh_config = SshConfig::parse(&format!("Host target\n {directive}\n")).expect("valid ssh_config"); let config = SshConnectionConfigResolver::new() diff --git a/src/ssh/tokio_client/mod.rs b/src/ssh/tokio_client/mod.rs index 63e1da61..22747ae9 100644 --- a/src/ssh/tokio_client/mod.rs +++ b/src/ssh/tokio_client/mod.rs @@ -44,6 +44,7 @@ pub use connection::{ SshConnectionConfigResolver, }; pub use error::{Error, TransportIntegrityCause}; +pub(crate) use proxy_command::is_direct_proxy_jump; pub use proxy_command::{ProxyCommandConfig, ProxyMode}; pub use to_socket_addrs_with_hostname::ToSocketAddrsWithHostname; diff --git a/src/ssh/tokio_client/proxy_command.rs b/src/ssh/tokio_client/proxy_command.rs index 21ea1e42..78627aa2 100644 --- a/src/ssh/tokio_client/proxy_command.rs +++ b/src/ssh/tokio_client/proxy_command.rs @@ -29,6 +29,15 @@ pub enum ProxyMode { Jump(String), } +/// Return whether a ProxyJump value explicitly requests a direct connection. +/// +/// The empty value is the internal sentinel used to retain an explicit direct +/// decision while passing through APIs that represent jump hosts as a string. +pub(crate) fn is_direct_proxy_jump(value: &str) -> bool { + let value = value.trim(); + value.is_empty() || value.eq_ignore_ascii_case("none") || value.eq_ignore_ascii_case("direct") +} + /// The unexpanded `ProxyCommand` and the context needed by its percent tokens. #[derive(Debug, Clone, PartialEq, Eq)] pub struct ProxyCommandConfig { From a8064543e4e784d19abba3e50ca4944b75cf2f97 Mon Sep 17 00:00:00 2001 From: Jeongkyu Shin Date: Sun, 30 Aug 2026 16:00:26 +0900 Subject: [PATCH 5/6] fix(ssh): preserve jump source precedence Dispatcher paths pre-combined the YAML cluster jump with the explicit CLI field, causing downstream command, transfer, and interactive connections to mistake a fallback value for a CLI override and bypass per-host ssh_config ProxyJump decisions. Thread only explicit CLI jump input through command parameters and keep ssh_config plus YAML inside the per-node resolver, which remains the single owner of CLI > ssh_config > YAML precedence. Share final jump selection between client and interactive paths, and derive SSH-mode session tokens from the same resolved ProxyMode. Validate config direct and jump overrides, YAML fallback, CLI jump and direct overrides, original HostName aliases, and the command plus every file and directory transfer production path. Refs #300 --- src/app/dispatcher.rs | 85 ++++++----- src/commands/exec.rs | 2 + src/commands/interactive/connection.rs | 64 ++++++-- src/commands/interactive/types.rs | 3 +- src/commands/upload.rs | 3 +- src/executor/connection_manager.rs | 195 +++++++++++++++++++------ src/executor/parallel.rs | 3 +- src/ssh/client/connection.rs | 12 +- src/ssh/tokio_client/mod.rs | 2 +- src/ssh/tokio_client/proxy_command.rs | 15 ++ 10 files changed, 279 insertions(+), 105 deletions(-) diff --git a/src/app/dispatcher.rs b/src/app/dispatcher.rs index 6019d1d6..34e1b7e0 100644 --- a/src/app/dispatcher.rs +++ b/src/app/dispatcher.rs @@ -260,11 +260,9 @@ pub async fn dispatch_command(cli: &Cli, ctx: &AppContext) -> Result { let use_keychain = determine_use_keychain(&ctx.ssh_config, hostname_for_ssh_config.as_deref()); - // Resolve jump_hosts: CLI takes precedence, then config - let jump_hosts = cli.jump_hosts.clone().or_else(|| { - ctx.config - .get_cluster_jump_host(ctx.cluster_name.as_deref().or(cli.cluster.as_deref())) - }); + // Preserve provenance: the resolver owns ssh_config and YAML + // fallback, while this field carries only an explicit CLI choice. + let cli_jump_hosts = cli.jump_hosts.clone(); let ssh_connection_config_resolver = build_ssh_connection_config_resolver( cli, @@ -286,7 +284,7 @@ pub async fn dispatch_command(cli: &Cli, ctx: &AppContext) -> Result { use_keychain, cli.timeout, Some(cli.connect_timeout), - jump_hosts, + cli_jump_hosts, ssh_password.clone(), ssh_connection_config_resolver, ) @@ -307,11 +305,7 @@ pub async fn dispatch_command(cli: &Cli, ctx: &AppContext) -> Result { ctx.cluster_name.as_deref().or(cli.cluster.as_deref()), ); - // Resolve jump_hosts: CLI takes precedence, then config - let jump_hosts = cli.jump_hosts.clone().or_else(|| { - ctx.config - .get_cluster_jump_host(ctx.cluster_name.as_deref().or(cli.cluster.as_deref())) - }); + let cli_jump_hosts = cli.jump_hosts.clone(); let params = FileTransferParams { nodes: ctx.nodes.clone(), @@ -323,7 +317,7 @@ pub async fn dispatch_command(cli: &Cli, ctx: &AppContext) -> Result { ssh_password: ssh_password.clone(), recursive: *recursive, ssh_config: Some(&ctx.ssh_config), - jump_hosts, + jump_hosts: cli_jump_hosts, ssh_connection_config_resolver: build_ssh_connection_config_resolver( cli, ctx, @@ -346,11 +340,7 @@ pub async fn dispatch_command(cli: &Cli, ctx: &AppContext) -> Result { ctx.cluster_name.as_deref().or(cli.cluster.as_deref()), ); - // Resolve jump_hosts: CLI takes precedence, then config - let jump_hosts = cli.jump_hosts.clone().or_else(|| { - ctx.config - .get_cluster_jump_host(ctx.cluster_name.as_deref().or(cli.cluster.as_deref())) - }); + let cli_jump_hosts = cli.jump_hosts.clone(); let params = FileTransferParams { nodes: ctx.nodes.clone(), @@ -362,7 +352,7 @@ pub async fn dispatch_command(cli: &Cli, ctx: &AppContext) -> Result { ssh_password: ssh_password.clone(), recursive: *recursive, ssh_config: Some(&ctx.ssh_config), - jump_hosts, + jump_hosts: cli_jump_hosts, ssh_connection_config_resolver: build_ssh_connection_config_resolver( cli, ctx, @@ -486,11 +476,9 @@ async fn handle_interactive_command( #[cfg(target_os = "macos")] let use_keychain = determine_use_keychain(&ctx.ssh_config, hostname.as_deref()); - // Resolve jump_hosts: CLI takes precedence, then config - let jump_hosts = cli.jump_hosts.clone().or_else(|| { - ctx.config - .get_cluster_jump_host(ctx.cluster_name.as_deref().or(cli.cluster.as_deref())) - }); + // The resolver retains ssh_config and YAML provenance. Thread only the + // explicit CLI override through the interactive command. + let cli_jump_hosts = cli.jump_hosts.clone(); // Build SSH connection config with keepalive settings for interactive mode let effective_cluster_name = ctx.cluster_name.as_deref().or(cli.cluster.as_deref()); @@ -519,7 +507,7 @@ async fn handle_interactive_command( #[cfg(target_os = "macos")] use_keychain, strict_mode: ctx.strict_mode, - jump_hosts, + jump_hosts: cli_jump_hosts, pty_config, use_pty, session_policy: None, @@ -563,6 +551,16 @@ fn resolve_ssh_mode_interactive_policy( Ok(matches!(policy.request, SessionRequest::Shell).then_some(policy)) } +fn session_policy_jump_spec(proxy_mode: Option<&ProxyMode>) -> Option<&str> { + match proxy_mode { + Some(ProxyMode::Jump(jump)) => Some(jump.as_str()), + // An empty authoritative value suppresses SessionPolicy's raw + // ssh_config fallback for explicit direct and ProxyCommand modes. + Some(ProxyMode::Direct | ProxyMode::Command(_)) => Some(""), + None => None, + } +} + /// Handle exec command or SSH mode interactive session async fn handle_exec_command( cli: &Cli, @@ -580,8 +578,10 @@ async fn handle_exec_command( .context("SSH interactive mode requires a destination node")?; let effective = ctx.ssh_config.find_host_config(node.config_host()); let effective_cluster_name = ctx.cluster_name.as_deref().or(cli.cluster.as_deref()); - let yaml_jump = ctx.config.get_cluster_jump_host(effective_cluster_name); - let jump_spec = cli.jump_hosts.as_deref().or(yaml_jump.as_deref()); + let connection_config = + build_ssh_connection_config_resolver(cli, ctx, effective_cluster_name) + .resolve_for_host(node.config_host()); + let jump_spec = session_policy_jump_spec(connection_config.proxy_mode.as_ref()); resolve_ssh_mode_interactive_policy( &effective, node, @@ -616,11 +616,7 @@ async fn handle_exec_command( #[cfg(target_os = "macos")] let use_keychain = determine_use_keychain(&ctx.ssh_config, hostname.as_deref()); - // Resolve jump_hosts: CLI takes precedence, then config - let jump_hosts = cli.jump_hosts.clone().or_else(|| { - ctx.config - .get_cluster_jump_host(ctx.cluster_name.as_deref().or(cli.cluster.as_deref())) - }); + let cli_jump_hosts = cli.jump_hosts.clone(); // Build SSH connection config with keepalive settings for SSH mode interactive session let effective_cluster_name = ctx.cluster_name.as_deref().or(cli.cluster.as_deref()); @@ -650,7 +646,7 @@ async fn handle_exec_command( #[cfg(target_os = "macos")] use_keychain, strict_mode: ctx.strict_mode, - jump_hosts, + jump_hosts: cli_jump_hosts, pty_config, use_pty, session_policy: Some(session_policy), @@ -703,22 +699,22 @@ async fn handle_exec_command( None }; - // Resolve jump_hosts: CLI takes precedence, then config + // Preserve source precedence in the resolver. Downstream receives only + // the explicit CLI override and must not mistake YAML for CLI input. let effective_cluster_name = ctx.cluster_name.as_deref().or(cli.cluster.as_deref()); let config_jump_host = ctx.config.get_cluster_jump_host(effective_cluster_name); - let jump_hosts = cli.jump_hosts.clone().or(config_jump_host.clone()); + let cli_jump_hosts = cli.jump_hosts.clone(); // Debug logging for jump host resolution tracing::debug!( - "Jump host resolution: cli={:?}, config={:?}, effective={:?}, cluster={:?}", + "Jump host sources: cli={:?}, yaml={:?}, cluster={:?}", cli.jump_hosts, config_jump_host, - jump_hosts, effective_cluster_name ); - if let Some(ref jh) = jump_hosts { - tracing::info!("Using jump host: {}", jh); + if let Some(ref jump_hosts) = cli_jump_hosts { + tracing::info!("Using CLI jump host override: {jump_hosts}"); } // Build SSH connection config resolver for exec mode. Each executor @@ -744,7 +740,7 @@ async fn handle_exec_command( byte_transparent: cli.is_ssh_mode(), timeout, connect_timeout: Some(cli.connect_timeout), - jump_hosts: jump_hosts.as_deref(), + jump_hosts: cli_jump_hosts.as_deref(), require_all_success: cli.require_all_success, check_all_nodes: cli.check_all_nodes, sudo_password, @@ -772,6 +768,17 @@ mod tests { }) } + #[test] + fn ssh_mode_policy_uses_the_resolved_proxy_decision() { + let jump = ProxyMode::Jump("resolved-bastion".to_string()); + assert_eq!( + session_policy_jump_spec(Some(&jump)), + Some("resolved-bastion") + ); + assert_eq!(session_policy_jump_spec(Some(&ProxyMode::Direct)), Some("")); + assert_eq!(session_policy_jump_spec(None), None); + } + #[test] fn plain_ssh_shell_always_consumes_the_resolved_session_policy() { let node = bssh::node::Node::new("127.0.0.1".into(), 2222, "remote".into()) diff --git a/src/commands/exec.rs b/src/commands/exec.rs index a02dc3ad..6a71ba2d 100644 --- a/src/commands/exec.rs +++ b/src/commands/exec.rs @@ -48,6 +48,8 @@ pub struct ExecuteCommandParams<'a> { pub byte_transparent: bool, pub timeout: Option, pub connect_timeout: Option, + /// Explicit CLI jump-host override. ssh_config and YAML fallbacks remain + /// in `ssh_connection_config_resolver` so their precedence is preserved. pub jump_hosts: Option<&'a str>, pub require_all_success: bool, pub check_all_nodes: bool, diff --git a/src/commands/interactive/connection.rs b/src/commands/interactive/connection.rs index dfdbf909..e120081b 100644 --- a/src/commands/interactive/connection.rs +++ b/src/commands/interactive/connection.rs @@ -28,8 +28,8 @@ use crate::ssh::{ SessionPolicy, SessionPurpose, SessionRequest, known_hosts::get_check_method_for_target, tokio_client::{ - AuthMethod, Client, Error as SshError, ProxyMode, ServerCheckMethod, SshConnectionConfig, - SshConnectionConfigResolver, is_direct_proxy_jump, + AuthMethod, Client, Error as SshError, ServerCheckMethod, SshConnectionConfig, + SshConnectionConfigResolver, select_proxy_jump, }, }; @@ -72,13 +72,7 @@ fn interactive_jump_spec<'a>( target_config: &'a SshConnectionConfig, fallback: Option<&'a str>, ) -> Option<&'a str> { - if let Some(requested) = fallback { - return (!is_direct_proxy_jump(requested)).then_some(requested); - } - match target_config.proxy_mode.as_ref() { - Some(ProxyMode::Jump(jump)) if !is_direct_proxy_jump(jump) => Some(jump.as_str()), - Some(ProxyMode::Jump(_) | ProxyMode::Command(_) | ProxyMode::Direct) | None => None, - } + select_proxy_jump(fallback, target_config.proxy_mode.as_ref()) } impl InteractiveCommand { @@ -691,6 +685,7 @@ pub fn is_auth_error_for_password_fallback(error: &SshError) -> bool { mod tests { use super::*; use crate::ssh::ssh_config::{IpQosPolicy, IpQosValue, SshConfig}; + use crate::ssh::tokio_client::ProxyMode; #[test] fn no_pty_shell_uses_bulk_ipqos_for_direct_and_jump_connections() { @@ -718,6 +713,57 @@ mod tests { assert_eq!(config.selected_ip_qos(), IpQosValue::Class(0x20)); } + #[test] + fn interactive_jump_selection_is_cli_then_ssh_config_then_yaml() { + let node = Node::new("effective-target".to_string(), 22, "user".to_string()) + .with_original_host("target-alias".to_string()); + let resolve = |ssh_config: SshConfig, cli_jump: Option<&str>| { + SshConnectionConfigResolver::new() + .with_ssh_config(Some(ssh_config)) + .with_cli_proxy_jump(cli_jump.map(str::to_owned)) + .with_yaml_proxy_jump(Some("yaml-bastion".to_string())) + .resolve_for_host(node.config_host()) + }; + + let config_jump = SshConfig::parse( + r#" +Host target-alias + HostName effective-target + ProxyJump config-bastion +"#, + ) + .unwrap(); + let target = resolve(config_jump.clone(), None); + assert_eq!(interactive_jump_spec(&target, None), Some("config-bastion")); + assert_eq!( + interactive_jump_spec(&target, Some("cli-bastion")), + Some("cli-bastion") + ); + for cli_direct in ["none", "direct"] { + assert_eq!(interactive_jump_spec(&target, Some(cli_direct)), None); + } + + for config_direct in ["none", "direct"] { + let ssh_config = SshConfig::parse(&format!( + "Host target-alias\n HostName effective-target\n ProxyJump {config_direct}\n" + )) + .unwrap(); + let target = resolve(ssh_config, None); + assert_eq!(interactive_jump_spec(&target, None), None); + } + + let yaml_target = resolve(SshConfig::new(), None); + assert_eq!( + interactive_jump_spec(&yaml_target, None), + Some("yaml-bastion") + ); + let cli_target = resolve(config_jump, Some("cli-bastion")); + assert_eq!( + interactive_jump_spec(&cli_target, Some("cli-bastion")), + Some("cli-bastion") + ); + } + #[test] fn interactive_jump_chain_keeps_distinct_bastion_and_target_socket_policies() { let ssh_config = SshConfig::parse( diff --git a/src/commands/interactive/types.rs b/src/commands/interactive/types.rs index 84ceb58d..3b5e48b1 100644 --- a/src/commands/interactive/types.rs +++ b/src/commands/interactive/types.rs @@ -62,7 +62,8 @@ pub struct InteractiveCommand { #[cfg(target_os = "macos")] pub use_keychain: bool, pub strict_mode: StrictHostKeyChecking, - // Jump hosts + /// Explicit CLI jump-host override. Per-host ssh_config and YAML fallback + /// are selected by `ssh_connection_config_resolver`. pub jump_hosts: Option, // PTY configuration pub pty_config: PtyConfig, diff --git a/src/commands/upload.rs b/src/commands/upload.rs index fff1f2d3..d1998d3e 100644 --- a/src/commands/upload.rs +++ b/src/commands/upload.rs @@ -39,7 +39,8 @@ pub struct FileTransferParams<'a> { pub ssh_password: Option>, pub recursive: bool, pub ssh_config: Option<&'a SshConfig>, - /// Jump hosts specification for connections. + /// Explicit CLI jump-host override. ssh_config and YAML remain in the + /// resolver so a YAML fallback cannot override a per-host ProxyJump. pub jump_hosts: Option, /// Per-host SSH connection configuration resolver. pub ssh_connection_config_resolver: SshConnectionConfigResolver, diff --git a/src/executor/connection_manager.rs b/src/executor/connection_manager.rs index 90d7870f..b9d4a10e 100644 --- a/src/executor/connection_manager.rs +++ b/src/executor/connection_manager.rs @@ -39,6 +39,8 @@ pub(crate) struct ExecutionConfig<'a> { pub use_keychain: bool, pub timeout: Option, pub connect_timeout: Option, + /// Explicit CLI/caller override. ssh_config and YAML are resolved through + /// their own sources and must not be pre-combined into this value. pub jump_hosts: Option<&'a str>, pub sudo_password: Option>, /// Pre-collected SSH password (collected once by the dispatcher and shared @@ -321,6 +323,63 @@ fn normalize_jump_hosts(jump_hosts: &str) -> String { #[cfg(test)] mod tests { use super::*; + use crate::ssh::tokio_client::select_proxy_jump; + + fn selected_jump_with_yaml_fallback( + cli_jump: Option<&str>, + ssh_config: &SshConfig, + yaml_jump: &str, + ) -> Option { + let resolver = SshConnectionConfigResolver::new() + .with_ssh_config(Some(ssh_config.clone())) + .with_cli_proxy_jump(cli_jump.map(str::to_owned)) + .with_yaml_proxy_jump(Some(yaml_jump.to_string())); + let target_config = resolver.resolve_for_host("target-alias"); + let explicit = resolve_effective_jump_hosts(cli_jump, Some(ssh_config), "target-alias"); + select_proxy_jump(explicit.as_deref(), target_config.proxy_mode.as_ref()).map(str::to_owned) + } + + #[test] + fn command_and_transfer_selection_is_cli_then_ssh_config_then_yaml() { + let config_jump = SshConfig::parse( + r#" +Host target-alias + HostName effective-target + ProxyJump config-bastion +"#, + ) + .unwrap(); + assert_eq!( + selected_jump_with_yaml_fallback(None, &config_jump, "yaml-bastion"), + Some("config-bastion".to_string()) + ); + assert_eq!( + selected_jump_with_yaml_fallback(Some("cli-bastion"), &config_jump, "yaml-bastion"), + Some("cli-bastion".to_string()) + ); + for cli_direct in ["none", "direct"] { + assert_eq!( + selected_jump_with_yaml_fallback(Some(cli_direct), &config_jump, "yaml-bastion"), + None + ); + } + + for config_direct in ["none", "direct"] { + let ssh_config = SshConfig::parse(&format!( + "Host target-alias\n HostName effective-target\n ProxyJump {config_direct}\n" + )) + .unwrap(); + assert_eq!( + selected_jump_with_yaml_fallback(None, &ssh_config, "yaml-bastion"), + None + ); + } + + assert_eq!( + selected_jump_with_yaml_fallback(None, &SshConfig::new(), "yaml-bastion"), + Some("yaml-bastion".to_string()) + ); + } #[test] fn command_and_transfer_proxy_jump_uses_original_host_alias() { @@ -505,38 +564,18 @@ Host *.example.com } } - #[tokio::test] - async fn command_and_all_transfer_paths_never_dial_proxy_jump_none() { - let listener = std::net::TcpListener::bind(("127.0.0.1", 0)).unwrap(); - let port = listener.local_addr().unwrap().port(); - drop(listener); - - let ssh_config = SshConfig::parse( - r#" -Host direct-target - HostName 127.0.0.1 - ProxyJump none -"#, - ) - .unwrap(); - let resolver = SshConnectionConfigResolver::new().with_ssh_config(Some(ssh_config.clone())); - let target_config = resolver.resolve_for_host("direct-target"); - let node = Node::new("127.0.0.1".to_string(), port, "user".to_string()) - .with_original_host("direct-target".to_string()); + async fn all_connection_path_errors( + node: Node, + cli_jump: Option<&str>, + ssh_config: &SshConfig, + resolver: &SshConnectionConfigResolver, + ) -> Vec<(&'static str, anyhow::Error)> { + let target_config = resolver.resolve_for_host(node.config_host()); let password = Arc::new(Password::new("test-password".to_string()).unwrap()); let temp_dir = tempfile::TempDir::new().unwrap(); let upload_file = temp_dir.path().join("upload.txt"); std::fs::write(&upload_file, b"test").unwrap(); - - let assert_direct_target = |path: &str, error: anyhow::Error| { - let message = format!("{error:#}"); - assert!(message.contains("127.0.0.1"), "path={path}: {message}"); - assert!(!message.contains("none:22"), "path={path}: {message}"); - assert!( - !message.contains("jump host none"), - "path={path}: {message}" - ); - }; + let mut errors = Vec::with_capacity(5); let execution_config = ExecutionConfig { key_path: None, @@ -547,18 +586,18 @@ Host direct-target use_keychain: false, timeout: Some(1), connect_timeout: Some(1), - jump_hosts: None, + jump_hosts: cli_jump, sudo_password: None, ssh_password: Some(password.clone()), - ssh_config: Some(&ssh_config), + ssh_config: Some(ssh_config), tty_mode: CliTtyMode::Disable, ssh_connection_config: Some(&target_config), - ssh_connection_config_resolver: Some(&resolver), + ssh_connection_config_resolver: Some(resolver), }; let error = execute_on_node_with_jump_hosts(node.clone(), "true", &execution_config) .await .unwrap_err(); - assert_direct_target("command", error); + errors.push(("command", error)); for (path, local_path) in [ ("upload-file", upload_file.as_path()), @@ -572,16 +611,16 @@ Host direct-target StrictHostKeyChecking::AcceptNew, false, true, - None, + cli_jump, Some(1), - Some(&ssh_config), + Some(ssh_config), Some(password.clone()), &target_config, - &resolver, + resolver, ) .await .unwrap_err(); - assert_direct_target(path, error); + errors.push((path, error)); } let error = download_from_node( @@ -592,16 +631,16 @@ Host direct-target StrictHostKeyChecking::AcceptNew, false, true, - None, + cli_jump, Some(1), - Some(&ssh_config), + Some(ssh_config), Some(password.clone()), &target_config, - &resolver, + resolver, ) .await .unwrap_err(); - assert_direct_target("download-file", error); + errors.push(("download-file", error)); let error = download_dir_from_node( node, @@ -611,16 +650,84 @@ Host direct-target StrictHostKeyChecking::AcceptNew, false, true, - None, + cli_jump, Some(1), - Some(&ssh_config), + Some(ssh_config), Some(password), &target_config, - &resolver, + resolver, ) .await .unwrap_err(); - assert_direct_target("download-directory", error); + errors.push(("download-directory", error)); + errors + } + + #[tokio::test] + async fn command_and_all_transfer_paths_never_dial_proxy_jump_none() { + let listener = std::net::TcpListener::bind(("127.0.0.1", 0)).unwrap(); + let port = listener.local_addr().unwrap().port(); + drop(listener); + + let ssh_config = SshConfig::parse( + r#" +Host direct-target + HostName 127.0.0.1 + ProxyJump none +"#, + ) + .unwrap(); + let resolver = SshConnectionConfigResolver::new() + .with_ssh_config(Some(ssh_config.clone())) + .with_yaml_proxy_jump(Some("127.0.0.3:1".to_string())); + let node = Node::new("127.0.0.1".to_string(), port, "user".to_string()) + .with_original_host("direct-target".to_string()); + + for (path, error) in all_connection_path_errors(node, None, &ssh_config, &resolver).await { + let message = format!("{error:#}"); + assert!(message.contains("127.0.0.1"), "path={path}: {message}"); + assert!(!message.contains("none:22"), "path={path}: {message}"); + assert!(!message.contains("127.0.0.3"), "path={path}: {message}"); + assert!( + !message.contains("jump host none"), + "path={path}: {message}" + ); + } + } + + #[tokio::test] + async fn all_connection_paths_prefer_config_then_cli_over_yaml() { + let ssh_config = SshConfig::parse( + r#" +Host target-alias + HostName 127.0.0.1 + ProxyJump config-bastion:not-a-port +"#, + ) + .unwrap(); + let node = Node::new("127.0.0.1".to_string(), 1, "user".to_string()) + .with_original_host("target-alias".to_string()); + let yaml_jump = "yaml-bastion:not-a-port"; + let resolver = SshConnectionConfigResolver::new() + .with_ssh_config(Some(ssh_config.clone())) + .with_yaml_proxy_jump(Some(yaml_jump.to_string())); + for (path, error) in + all_connection_path_errors(node.clone(), None, &ssh_config, &resolver).await + { + let message = format!("{error:#}"); + assert!(message.contains("config-bastion"), "path={path}: {message}"); + assert!(!message.contains("yaml-bastion"), "path={path}: {message}"); + } + + let cli_jump = "cli-bastion:not-a-port"; + let resolver = resolver.with_cli_proxy_jump(Some(cli_jump.to_string())); + for (path, error) in + all_connection_path_errors(node, Some(cli_jump), &ssh_config, &resolver).await + { + let message = format!("{error:#}"); + assert!(message.contains("cli-bastion"), "path={path}: {message}"); + assert!(!message.contains("yaml-bastion"), "path={path}: {message}"); + } } /// Test complex multi-hop chain with user and ports diff --git a/src/executor/parallel.rs b/src/executor/parallel.rs index 92147d79..948a23d6 100644 --- a/src/executor/parallel.rs +++ b/src/executor/parallel.rs @@ -51,6 +51,7 @@ pub struct ParallelExecutor { pub(crate) use_keychain: bool, pub(crate) timeout: Option, pub(crate) connect_timeout: Option, + /// Explicit caller override; dispatcher callers pass only CLI provenance. pub(crate) jump_hosts: Option, pub(crate) sudo_password: Option>, /// SSH password collected once up-front by the dispatcher. @@ -212,7 +213,7 @@ impl ParallelExecutor { self } - /// Set jump hosts for connections. + /// Set an explicit jump-host override for connections. pub fn with_jump_hosts(mut self, jump_hosts: Option) -> Self { self.jump_hosts = jump_hosts; self diff --git a/src/ssh/client/connection.rs b/src/ssh/client/connection.rs index 7d684632..04f0f703 100644 --- a/src/ssh/client/connection.rs +++ b/src/ssh/client/connection.rs @@ -18,8 +18,7 @@ use crate::security::Password; use crate::ssh::SessionPurpose; use crate::ssh::known_hosts::StrictHostKeyChecking; use crate::ssh::tokio_client::{ - AuthMethod, Client, ProxyMode, SshConnectionConfig, SshConnectionConfigResolver, - is_direct_proxy_jump, + AuthMethod, Client, SshConnectionConfig, SshConnectionConfigResolver, select_proxy_jump, }; use anyhow::{Context, Result}; use std::path::Path; @@ -57,13 +56,7 @@ fn client_jump_spec<'a>( target_config: &'a SshConnectionConfig, requested_jump_hosts: Option<&'a str>, ) -> Option<&'a str> { - if let Some(requested) = requested_jump_hosts { - return (!is_direct_proxy_jump(requested)).then_some(requested); - } - match target_config.proxy_mode.as_ref() { - Some(ProxyMode::Jump(jump)) if !is_direct_proxy_jump(jump) => Some(jump.as_str()), - Some(ProxyMode::Jump(_) | ProxyMode::Command(_) | ProxyMode::Direct) | None => None, - } + select_proxy_jump(requested_jump_hosts, target_config.proxy_mode.as_ref()) } /// Build the friendly, outer-context message for a failed direct SSH @@ -399,6 +392,7 @@ mod tests { use super::*; use crate::ssh::SshConfig; use crate::ssh::ssh_config::{IpQosPolicy, IpQosValue}; + use crate::ssh::tokio_client::ProxyMode; use crate::test_helpers::EnvGuard; use serial_test::serial; use tempfile::TempDir; diff --git a/src/ssh/tokio_client/mod.rs b/src/ssh/tokio_client/mod.rs index 22747ae9..7d616274 100644 --- a/src/ssh/tokio_client/mod.rs +++ b/src/ssh/tokio_client/mod.rs @@ -44,8 +44,8 @@ pub use connection::{ SshConnectionConfigResolver, }; pub use error::{Error, TransportIntegrityCause}; -pub(crate) use proxy_command::is_direct_proxy_jump; pub use proxy_command::{ProxyCommandConfig, ProxyMode}; +pub(crate) use proxy_command::{is_direct_proxy_jump, select_proxy_jump}; pub use to_socket_addrs_with_hostname::ToSocketAddrsWithHostname; // Re-export russh types commonly used with this module diff --git a/src/ssh/tokio_client/proxy_command.rs b/src/ssh/tokio_client/proxy_command.rs index 78627aa2..cfeca02e 100644 --- a/src/ssh/tokio_client/proxy_command.rs +++ b/src/ssh/tokio_client/proxy_command.rs @@ -38,6 +38,21 @@ pub(crate) fn is_direct_proxy_jump(value: &str) -> bool { value.is_empty() || value.eq_ignore_ascii_case("none") || value.eq_ignore_ascii_case("direct") } +/// Select the actual jump chain from an explicit caller override and the +/// resolver's lower-priority transport decision. +pub(crate) fn select_proxy_jump<'a>( + explicit_jump: Option<&'a str>, + resolved_mode: Option<&'a ProxyMode>, +) -> Option<&'a str> { + if let Some(explicit) = explicit_jump { + return (!is_direct_proxy_jump(explicit)).then_some(explicit); + } + match resolved_mode { + Some(ProxyMode::Jump(jump)) if !is_direct_proxy_jump(jump) => Some(jump.as_str()), + Some(ProxyMode::Jump(_) | ProxyMode::Command(_) | ProxyMode::Direct) | None => None, + } +} + /// The unexpanded `ProxyCommand` and the context needed by its percent tokens. #[derive(Debug, Clone, PartialEq, Eq)] pub struct ProxyCommandConfig { From ec604e8ea2a6f4288b576ff9b92de20a02c9e30c Mon Sep 17 00:00:00 2001 From: Jeongkyu Shin Date: Sun, 30 Aug 2026 16:29:23 +0900 Subject: [PATCH 6/6] fix(ssh): resolve interactive authentication per host Interactive connections resolved transport and host-key policy per node but selected credentials from the dispatcher's shared fallback, allowing one host's ssh_config authentication settings to leak into other aliases. Resolve each node's target config before authentication in both multiplex and PTY paths, and feed that same config through authentication, host verification, and direct or jump transport. Add two-host alias regressions that preserve explicit CLI identity and agent settings while distinguishing per-host IdentitiesOnly, IdentityFile, PasswordAuthentication, and BatchMode. Validated with cargo fmt, full lib/bin/test cargo check, scoped Clippy with warnings denied, and focused #296/#300 authentication and transport tests. Refs #300 --- src/commands/interactive/connection.rs | 180 +++++++++++++++++++++++-- 1 file changed, 172 insertions(+), 8 deletions(-) diff --git a/src/commands/interactive/connection.rs b/src/commands/interactive/connection.rs index e120081b..ee6e6be2 100644 --- a/src/commands/interactive/connection.rs +++ b/src/commands/interactive/connection.rs @@ -225,8 +225,11 @@ impl InteractiveCommand { .with_context(|| "Password prompt task failed")? } - /// Determine authentication method based on node and config (same logic as exec mode) - pub(super) async fn determine_auth_method(&self, node: &Node) -> Result { + fn auth_context( + &self, + node: &Node, + target_config: &SshConnectionConfig, + ) -> Result { // Use centralized authentication logic from auth module let mut auth_ctx = crate::ssh::AuthContext::new(node.username.clone(), node.host.clone()) .with_context(|| { @@ -245,7 +248,7 @@ impl InteractiveCommand { .with_password(self.use_password) .with_password_fallback(!self.use_password) // Enable fallback only if not using explicit password .with_pre_collected_password(self.ssh_password.clone()); - auth_ctx = auth_ctx.with_policy(self.ssh_connection_config.auth_policy.clone()); + auth_ctx = auth_ctx.with_policy(target_config.auth_policy.clone()); // Set macOS Keychain integration if available #[cfg(target_os = "macos")] @@ -253,7 +256,18 @@ impl InteractiveCommand { auth_ctx = auth_ctx.with_keychain(self.use_keychain); } - auth_ctx.determine_method().await + Ok(auth_ctx) + } + + /// Determine authentication method based on node and config (same logic as exec mode) + pub(super) async fn determine_auth_method( + &self, + node: &Node, + target_config: &SshConnectionConfig, + ) -> Result { + self.auth_context(node, target_config)? + .determine_method() + .await } /// Select nodes to connect to based on configuration @@ -314,13 +328,15 @@ impl InteractiveCommand { /// Connect to a single node and establish an interactive shell pub(super) async fn connect_to_node(&self, node: Node) -> Result { - // Determine authentication method using the same logic as exec mode - let auth_method = self.determine_auth_method(&node).await?; let target_config = interactive_target_connection_config( &node, &self.ssh_connection_config, self.ssh_connection_config_resolver.as_ref(), ); + // Resolve from the node's original ssh_config alias before selecting + // credentials so authentication, host verification, and transport all + // consume the same destination policy. + let auth_method = self.determine_auth_method(&node, &target_config).await?; // Set up host key checking using the configured strict mode let check_method = get_check_method_for_target( @@ -479,13 +495,14 @@ impl InteractiveCommand { /// Connect to a single node and establish a PTY-enabled SSH channel pub(super) async fn connect_to_node_pty(&self, node: Node) -> Result<(Client, Channel)> { - // Determine authentication method using the same logic as exec mode - let auth_method = self.determine_auth_method(&node).await?; let target_config = interactive_target_connection_config( &node, &self.ssh_connection_config, self.ssh_connection_config_resolver.as_ref(), ); + // Keep PTY authentication on the same per-node alias policy used by + // host verification and the direct or jump transport. + let auth_method = self.determine_auth_method(&node, &target_config).await?; // Set up host key checking using the configured strict mode let check_method = get_check_method_for_target( @@ -684,8 +701,155 @@ pub fn is_auth_error_for_password_fallback(error: &SshError) -> bool { #[cfg(test)] mod tests { use super::*; + use crate::config::{Config, InteractiveConfig}; + use crate::pty::PtyConfig; + use crate::ssh::known_hosts::StrictHostKeyChecking; use crate::ssh::ssh_config::{IpQosPolicy, IpQosValue, SshConfig}; use crate::ssh::tokio_client::ProxyMode; + use std::path::PathBuf; + + fn alias_auth_command(fixed_alias: &str) -> (InteractiveCommand, Node, Node) { + let ssh_config = SshConfig::parse( + r#" +Host alpha + HostName effective-alpha + IdentityFile /alpha-identity + IdentitiesOnly yes + PreferredAuthentications password + PubkeyAuthentication no + PasswordAuthentication no + NumberOfPasswordPrompts 1 + BatchMode no + +Host beta + HostName effective-beta + IdentityFile /beta-identity + IdentitiesOnly no + PreferredAuthentications password + PubkeyAuthentication no + PasswordAuthentication yes + NumberOfPasswordPrompts 7 + BatchMode yes +"#, + ) + .expect("valid ssh_config"); + let resolver = SshConnectionConfigResolver::new() + .with_ssh_config(Some(ssh_config)) + .with_cli_identity_files(vec![PathBuf::from("/cli-identity")]); + let fixed_config = resolver.resolve_for_host(fixed_alias); + let alpha = Node::new("effective-alpha".to_string(), 22, "user".to_string()) + .with_original_host("alpha".to_string()); + let beta = Node::new("effective-beta".to_string(), 22, "user".to_string()) + .with_original_host("beta".to_string()); + let command = InteractiveCommand { + single_node: false, + multiplex: true, + prompt_format: String::new(), + history_file: PathBuf::new(), + work_dir: None, + nodes: vec![alpha.clone(), beta.clone()], + config: Config::default(), + interactive_config: InteractiveConfig::default(), + cluster_name: None, + key_path: Some(PathBuf::from("/explicit-identity")), + use_agent: true, + use_password: false, + ssh_password: None, + #[cfg(target_os = "macos")] + use_keychain: false, + strict_mode: StrictHostKeyChecking::No, + jump_hosts: None, + pty_config: PtyConfig::default(), + use_pty: None, + session_policy: None, + ssh_connection_config: fixed_config, + ssh_connection_config_resolver: Some(resolver), + }; + + (command, alpha, beta) + } + + fn assert_distinct_alias_auth_policies( + command: &InteractiveCommand, + alpha: &Node, + beta: &Node, + ) { + let resolver = command.ssh_connection_config_resolver.as_ref(); + let alpha_config = + interactive_target_connection_config(alpha, &command.ssh_connection_config, resolver); + let beta_config = + interactive_target_connection_config(beta, &command.ssh_connection_config, resolver); + let alpha_context = command.auth_context(alpha, &alpha_config).unwrap(); + let beta_context = command.auth_context(beta, &beta_config).unwrap(); + + for context in [&alpha_context, &beta_context] { + assert_eq!( + context.key_path.as_deref(), + Some(std::path::Path::new("/explicit-identity")) + ); + assert!(context.use_agent); + assert_eq!( + context.policy.cli_identity_files, + [PathBuf::from("/cli-identity")] + ); + } + assert_eq!( + alpha_context.policy.identity_files, + [PathBuf::from("/alpha-identity")] + ); + assert!(alpha_context.policy.identities_only); + assert!(!alpha_context.policy.password_authentication); + assert!(!alpha_context.policy.batch_mode); + assert_eq!(alpha_context.policy.number_of_password_prompts, 1); + + assert_eq!( + beta_context.policy.identity_files, + [PathBuf::from("/beta-identity")] + ); + assert!(!beta_context.policy.identities_only); + assert!(beta_context.policy.password_authentication); + assert!(beta_context.policy.batch_mode); + assert_eq!(beta_context.policy.number_of_password_prompts, 7); + } + + #[tokio::test] + async fn connect_to_node_resolves_each_alias_auth_policy_before_authentication() { + // The shared fallback intentionally carries alpha's policy. The beta + // connection must still stop with beta's BatchMode decision before any + // network connection is attempted. + let (command, alpha, beta) = alias_auth_command("alpha"); + assert_distinct_alias_auth_policies(&command, &alpha, &beta); + + let error = match command.connect_to_node(beta).await { + Ok(_) => panic!("beta authentication policy must reject all methods"), + Err(error) => error, + }; + let rendered = format!("{error:#}"); + assert!(rendered.contains("disabled by BatchMode"), "{rendered}"); + assert!( + !rendered.contains("disabled by PasswordAuthentication"), + "{rendered}" + ); + } + + #[tokio::test] + async fn connect_to_node_pty_resolves_each_alias_auth_policy_before_authentication() { + // Reverse the shared fallback. The alpha PTY path must use alpha's + // PasswordAuthentication policy instead of inheriting beta's BatchMode. + let (command, alpha, beta) = alias_auth_command("beta"); + assert_distinct_alias_auth_policies(&command, &alpha, &beta); + + let error = match command.connect_to_node_pty(alpha).await { + Ok(_) => panic!("alpha authentication policy must reject all methods"), + Err(error) => error, + }; + let rendered = format!("{error:#}"); + assert!( + rendered.contains("disabled by PasswordAuthentication"), + "{rendered}" + ); + assert!(!rendered.contains("disabled by BatchMode"), "{rendered}"); + } #[test] fn no_pty_shell_uses_bulk_ipqos_for_direct_and_jump_connections() {