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/app/dispatcher.rs b/src/app/dispatcher.rs index ee1f754d..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,11 +507,12 @@ 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, ssh_connection_config, + ssh_connection_config_resolver: Some(ssh_connection_config_resolver), }; let result = interactive_cmd.execute().await?; @@ -562,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, @@ -579,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, @@ -615,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()); @@ -649,11 +646,12 @@ 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), ssh_connection_config, + ssh_connection_config_resolver: Some(ssh_connection_config_resolver), }; let result = interactive_cmd.execute().await?; @@ -701,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 @@ -742,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, @@ -770,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 be14602c..ee6e6be2 100644 --- a/src/commands/interactive/connection.rs +++ b/src/commands/interactive/connection.rs @@ -22,16 +22,59 @@ 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::{ - 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, select_proxy_jump, + }, }; 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) +} + +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> { + select_proxy_jump(fallback, target_config.proxy_mode.as_ref()) +} + 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,9 +93,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(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 @@ -72,7 +117,7 @@ impl InteractiveCommand { username, auth_method, check_method.clone(), - ssh_config, + &ssh_config, ), ) .await @@ -116,7 +161,7 @@ impl InteractiveCommand { username, password_auth, check_method, - ssh_config, + &ssh_config, ), ) .await @@ -144,6 +189,26 @@ 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, + target_config: &SshConnectionConfig, + ) -> JumpHostChain { + build_interactive_jump_chain( + jump_hosts, + adjusted_timeout, + target_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(); @@ -160,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(|| { @@ -180,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")] @@ -188,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 @@ -249,13 +328,20 @@ 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( self.strict_mode, - &self.ssh_connection_config, + &target_config, &node.host, node.port, &node.username, @@ -265,7 +351,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}'") @@ -284,7 +372,8 @@ 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? } else { @@ -312,11 +401,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_ssh_password(self.ssh_password.clone()); + let chain = self.build_jump_chain(jump_hosts, adjusted_timeout, &target_config); // Connect through the chain let connection = timeout( @@ -367,7 +452,8 @@ 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? }; @@ -409,13 +495,19 @@ 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( self.strict_mode, - &self.ssh_connection_config, + &target_config, &node.host, node.port, &node.username, @@ -425,7 +517,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}'") @@ -444,7 +538,8 @@ 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? } else { @@ -472,11 +567,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_ssh_password(self.ssh_password.clone()); + let chain = self.build_jump_chain(jump_hosts, adjusted_timeout, &target_config); // Connect through the chain let connection = timeout( @@ -527,7 +618,8 @@ 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? }; @@ -609,6 +701,343 @@ 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() { + 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_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( + r#" +Host bastion + BindAddress 127.0.0.2 + BindInterface lo + IPQoS cs5 cs1 + +Host alpha + HostName effective-alpha + HostKeyAlias alpha-key + BindAddress 127.0.0.3 + 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") + ); + 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") + ); + + let chain = build_interactive_jump_chain( + vec![JumpHost::new("bastion".to_string(), None, None)], + Duration::from_secs(45), + &alpha_config, + Some(&resolver), + 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.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.destination_connection_config(); + assert_eq!(target.bind_address.as_deref(), Some("127.0.0.3")); + 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 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), + bulk: IpQosValue::Class(0x60), + }); + let fixed_chain = build_interactive_jump_chain( + vec![JumpHost::new("manual-bastion".to_string(), None, None)], + Duration::from_secs(45), + &manual_config, + None, + SessionPurpose::Bulk, + ); + 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] fn test_key_auth_failed_triggers_password_fallback() { diff --git a/src/commands/interactive/types.rs b/src/commands/interactive/types.rs index a3b7424e..3b5e48b1 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 @@ -62,15 +62,18 @@ 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, 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 policy resolver for each interactive target and jump alias. + 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/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 56b38081..b9d4a10e 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. @@ -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 @@ -66,19 +68,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 +83,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 +97,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 +162,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 +175,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 +191,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 +222,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 +235,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 +267,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 +278,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,23 +294,114 @@ 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); + 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(hostname)) } #[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() { + 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] @@ -356,6 +419,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 @@ -466,23 +538,196 @@ 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}" + ); + } + } + + 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 mut errors = Vec::with_capacity(5); + + 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: cli_jump, + 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(); + errors.push(("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, + cli_jump, + Some(1), + Some(ssh_config), + Some(password.clone()), + &target_config, + resolver, + ) + .await + .unwrap_err(); + errors.push((path, error)); + } - // 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 error = download_from_node( + node.clone(), + "/tmp/remote-file", + &temp_dir.path().join("download.txt"), + None, + StrictHostKeyChecking::AcceptNew, + false, + true, + cli_jump, + Some(1), + Some(ssh_config), + Some(password.clone()), + &target_config, + resolver, + ) + .await + .unwrap_err(); + errors.push(("download-file", error)); + + let error = download_dir_from_node( + node, + "/tmp/remote-dir", + &temp_dir.path().join("download-dir"), + None, + StrictHostKeyChecking::AcceptNew, + false, + true, + cli_jump, + Some(1), + Some(ssh_config), + Some(password), + &target_config, + resolver, + ) + .await + .unwrap_err(); + 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/jump/chain.rs b/src/jump/chain.rs index 2a8aeef4..3c9eab52 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, @@ -66,10 +67,16 @@ 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, /// 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 +118,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, } } @@ -127,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; @@ -143,11 +153,24 @@ impl JumpHostChain { self } - fn connection_config_for_host(&self, host: &str) -> SshConnectionConfig { + #[must_use] + pub fn with_session_purpose(mut self, purpose: SessionPurpose) -> Self { + self.session_purpose = purpose; + self + } + + 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)) .unwrap_or_else(|| self.ssh_connection_config.clone()) + .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) @@ -227,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, @@ -299,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 {}: {}", @@ -327,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, @@ -377,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!( @@ -551,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/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..04f0f703 100644 --- a/src/ssh/client/connection.rs +++ b/src/ssh/client/connection.rs @@ -13,11 +13,12 @@ // 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; use crate::ssh::tokio_client::{ - AuthMethod, Client, ProxyMode, SshConnectionConfig, SshConnectionConfigResolver, + AuthMethod, Client, SshConnectionConfig, SshConnectionConfigResolver, select_proxy_jump, }; use anyhow::{Context, Result}; use std::path::Path; @@ -30,6 +31,34 @@ 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> { + select_proxy_jump(requested_jump_hosts, target_config.proxy_mode.as_ref()) +} + /// Build the friendly, outer-context message for a failed direct SSH /// connection attempt. /// @@ -236,20 +265,19 @@ 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 = 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()); - } + 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 @@ -293,13 +321,14 @@ impl SshClient { ssh_connection_config: Option<&SshConnectionConfig>, ssh_connection_config_resolver: Option<&SshConnectionConfigResolver>, pre_collected_password: Option>, + session_purpose: SessionPurpose, ) -> Result { - 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 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 = client_jump_spec(&selected_config, jump_hosts_spec); if let Some(jump_spec) = jump_hosts_spec { // Parse jump hosts @@ -340,6 +369,7 @@ impl SshClient { ssh_connection_config, ssh_connection_config_resolver, pre_collected_password, + session_purpose, ) .await } @@ -360,6 +390,9 @@ impl SshClient { #[cfg(test)] 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; @@ -371,6 +404,95 @@ 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") + ); + 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) + .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(); 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..5d68253d --- /dev/null +++ b/src/ssh/ssh_config/ip_qos.rs @@ -0,0 +1,177 @@ +// 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" => 0xb0, + // 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 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)); + 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..a6b557a0 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, @@ -32,11 +32,12 @@ 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}; -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( @@ -590,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()) @@ -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..67f49b89 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() { @@ -344,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() @@ -740,3 +748,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 { .. } diff --git a/src/ssh/tokio_client/mod.rs b/src/ssh/tokio_client/mod.rs index 63e1da61..7d616274 100644 --- a/src/ssh/tokio_client/mod.rs +++ b/src/ssh/tokio_client/mod.rs @@ -45,6 +45,7 @@ pub use connection::{ }; pub use error::{Error, TransportIntegrityCause}; 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 21ea1e42..cfeca02e 100644 --- a/src/ssh/tokio_client/proxy_command.rs +++ b/src/ssh/tokio_client/proxy_command.rs @@ -29,6 +29,30 @@ 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") +} + +/// 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 { 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