restore ssh_tunnel.rs file

This commit is contained in:
runstone 2026-06-12 08:53:54 +08:00
parent 91f1206783
commit ef1697d8a8
1 changed files with 75 additions and 111 deletions

View File

@ -5,10 +5,7 @@ use std::sync::Arc;
use base64::engine::general_purpose::STANDARD as BASE64_STANDARD;
use base64::Engine;
use russh::client::{self, Config, Handle};
use russh::keys::agent::{
client::{AgentClient, AgentStream},
AgentIdentity,
};
use russh::keys::agent::{client::AgentClient, AgentIdentity};
use russh::keys::{decode_secret_key, key::PrivateKeyWithHashAlg, PrivateKey};
use russh::ChannelMsg;
use tokio::io::{AsyncReadExt, AsyncWriteExt};
@ -46,12 +43,12 @@ impl client::Handler for SshClient {
}
async fn connect_and_authenticate(
ssh_host: String,
ssh_host: &str,
ssh_port: u16,
ssh_user: String,
ssh_password: String,
ssh_key_path: String,
ssh_key_passphrase: String,
ssh_user: &str,
ssh_password: &str,
ssh_key_path: &str,
ssh_key_passphrase: &str,
use_ssh_agent: bool,
connect_timeout_secs: u64,
) -> Result<Handle<SshClient>, String> {
@ -67,15 +64,15 @@ async fn connect_and_authenticate(
if !ssh_key_path.is_empty() {
// Validate SSH key file path
validate_file_path(&ssh_key_path, |_| false)?;
validate_file_path(ssh_key_path, |_| false)?;
let passphrase = if ssh_key_passphrase.is_empty() { None } else { Some(ssh_key_passphrase.as_str()) };
let passphrase = if ssh_key_passphrase.is_empty() { None } else { Some(ssh_key_passphrase) };
let key_pair =
load_ssh_private_key(&ssh_key_path, passphrase).map_err(|e| format!("Failed to load SSH key: {e}"))?;
load_ssh_private_key(ssh_key_path, passphrase).map_err(|e| format!("Failed to load SSH key: {e}"))?;
let auth_res = tokio::time::timeout(
connect_timeout,
session.authenticate_publickey(
&ssh_user,
ssh_user,
PrivateKeyWithHashAlg::new(
Arc::new(key_pair),
session.best_supported_rsa_hash().await.ok().flatten().flatten(),
@ -89,7 +86,7 @@ async fn connect_and_authenticate(
return Err("SSH public key authentication failed".to_string());
}
} else if !ssh_password.is_empty() {
let auth_res = tokio::time::timeout(connect_timeout, session.authenticate_password(&ssh_user, &ssh_password))
let auth_res = tokio::time::timeout(connect_timeout, session.authenticate_password(ssh_user, ssh_password))
.await
.map_err(|_| format!("SSH password auth timed out ({connect_timeout_secs}s)"))?
.map_err(|e| format!("SSH password auth failed: {e}"))?;
@ -97,7 +94,7 @@ async fn connect_and_authenticate(
return Err("SSH password authentication failed".to_string());
}
} else if use_ssh_agent {
match try_authenticate_with_agent(&mut session, &ssh_user, &connect_timeout).await {
match try_authenticate_with_agent(&mut session, ssh_user, &connect_timeout).await {
Ok(()) => {}
Err(agent_err) => return Err(agent_err),
}
@ -108,18 +105,6 @@ async fn connect_and_authenticate(
Ok(session)
}
/// Connect to the platform-appropriate SSH agent.
#[cfg(unix)]
async fn connect_ssh_agent() -> Result<AgentClient<Box<dyn AgentStream + Send + Unpin>>, String> {
AgentClient::connect_env().await.map(|c| c.dynamic()).map_err(|e| e.to_string())
}
/// Connect to the platform-appropriate SSH agent (Pageant on Windows).
#[cfg(windows)]
async fn connect_ssh_agent() -> Result<AgentClient<Box<dyn AgentStream + Send + Unpin>>, String> {
AgentClient::connect_pageant().await.map(|c| c.dynamic()).map_err(|e| e.to_string())
}
/// Try to authenticate using ssh-agent identities. Returns `Ok(())` on success,
/// or an error describing why agent auth failed (unavailable, no identities, all rejected).
async fn try_authenticate_with_agent(
@ -323,12 +308,7 @@ fn read_u32(bytes: &[u8], pos: &mut usize) -> Result<u32, String> {
/// Accept connections on the local listener and forward them through the SSH session.
/// Returns when the SSH session dies (listener error or session.is_closed()).
async fn forward_loop(
session: Arc<Handle<SshClient>>,
listener: Arc<TcpListener>,
remote_host: String,
remote_port: u16,
) {
async fn forward_loop(session: &Handle<SshClient>, listener: &TcpListener, remote_host: &str, remote_port: u16) {
let mut idle_check = tokio::time::interval(IDLE_SESSION_CHECK_INTERVAL);
idle_check.set_missed_tick_behavior(MissedTickBehavior::Delay);
@ -370,7 +350,7 @@ async fn forward_loop(
let mut channel = match session
.channel_open_direct_tcpip(
remote_host.clone(),
remote_host,
remote_port.into(),
peer_addr.ip().to_string(),
peer_addr.port().into(),
@ -427,8 +407,8 @@ async fn forward_loop(
/// Uses exponential backoff for reconnect attempts and gives up after
/// MAX_RECONNECT_ATTEMPTS to avoid log storms from permanent failures.
#[allow(clippy::too_many_arguments)]
fn tunnel_reconnect_loop(
session: Handle<SshClient>,
async fn tunnel_reconnect_loop(
mut session: Handle<SshClient>,
connect_host: String,
connect_port: u16,
ssh_user: String,
@ -437,60 +417,66 @@ fn tunnel_reconnect_loop(
ssh_key_passphrase: String,
use_ssh_agent: bool,
connect_timeout_secs: u64,
listener: Arc<TcpListener>,
listener: TcpListener,
remote_host: String,
remote_port: u16,
) -> std::pin::Pin<Box<dyn std::future::Future<Output = ()>>> {
let mut session = Arc::new(session);
let label = format!("{connect_host}:{connect_port} -> {remote_host}:{remote_port}");
Box::pin(async move {
) {
loop {
log::info!("SSH tunnel active: {}:{} -> {}:{}", connect_host, connect_port, remote_host, remote_port);
forward_loop(&session, &listener, &remote_host, remote_port).await;
log::warn!("SSH tunnel connection lost ({}:{}), reconnecting...", connect_host, connect_port);
// Reconnect with exponential backoff
let mut delay = INITIAL_RECONNECT_DELAY;
let mut attempts: u32 = 0;
loop {
log::info!("SSH tunnel active: {label}");
if attempts >= MAX_RECONNECT_ATTEMPTS {
log::error!(
"SSH tunnel ({connect_host}:{connect_port}): max reconnect attempts ({MAX_RECONNECT_ATTEMPTS}) exhausted, giving up"
);
return;
}
forward_loop(session.clone(), listener.clone(), remote_host.clone(), remote_port).await;
tokio::time::sleep(delay).await;
log::warn!("SSH tunnel connection lost ({label}), reconnecting...");
// Reconnect with exponential backoff
let mut delay = INITIAL_RECONNECT_DELAY;
let mut attempts: u32 = 0;
loop {
if attempts >= MAX_RECONNECT_ATTEMPTS {
log::error!(
"SSH tunnel ({label}): max reconnect attempts ({MAX_RECONNECT_ATTEMPTS}) exhausted, giving up"
match connect_and_authenticate(
&connect_host,
connect_port,
&ssh_user,
&ssh_password,
&ssh_key_path,
&ssh_key_passphrase,
use_ssh_agent,
connect_timeout_secs,
)
.await
{
Ok(new_session) => {
session = new_session;
log::info!(
"SSH tunnel reconnected to {}:{} (attempt {})",
connect_host,
connect_port,
attempts + 1
);
return;
break;
}
tokio::time::sleep(delay).await;
match connect_and_authenticate(
connect_host.clone(),
connect_port,
ssh_user.clone(),
ssh_password.clone(),
ssh_key_path.clone(),
ssh_key_passphrase.clone(),
use_ssh_agent,
connect_timeout_secs,
)
.await
{
Ok(new_session) => {
session = Arc::new(new_session);
log::info!("SSH tunnel reconnected to {label} (attempt {})", attempts + 1);
break;
}
Err(e) => {
attempts += 1;
log::error!("SSH reconnect failed ({label}, attempt {attempts}/{MAX_RECONNECT_ATTEMPTS}): {e}");
delay = std::cmp::min(delay * 2, MAX_RECONNECT_DELAY);
}
Err(e) => {
attempts += 1;
log::error!(
"SSH reconnect failed ({}:{}, attempt {attempts}/{MAX_RECONNECT_ATTEMPTS}): {e}",
connect_host,
connect_port,
);
// Exponential backoff: double the delay, cap at MAX_RECONNECT_DELAY
delay = std::cmp::min(delay * 2, MAX_RECONNECT_DELAY);
}
}
}
})
}
}
struct TunnelEntry {
@ -546,7 +532,7 @@ impl TunnelManager {
return Ok(port);
}
}
// Slow SSH connection 鈥?do this outside the lock.
// Slow SSH connection do this outside the lock.
let (handle, local_port) = spawn_tunnel(
ssh_host,
ssh_port,
@ -684,20 +670,18 @@ async fn spawn_tunnel(
// Initial connection: fail fast on bad credentials
let session = connect_and_authenticate(
connect_host.to_string(),
connect_host,
connect_port,
ssh_user.to_string(),
ssh_password.to_string(),
ssh_key_path.to_string(),
ssh_key_passphrase.to_string(),
ssh_user,
ssh_password,
ssh_key_path,
ssh_key_passphrase,
use_ssh_agent,
connect_timeout_secs,
)
.await?;
// Spawn the tunnel loop on a dedicated blocking thread to avoid
// a known Rust async Send inference limitation (rust-lang/rust#102211).
let tunnel_args = (
let handle = tokio::spawn(tunnel_reconnect_loop(
session,
connect_host.to_string(),
connect_port,
@ -707,30 +691,10 @@ async fn spawn_tunnel(
ssh_key_passphrase.to_string(),
use_ssh_agent,
connect_timeout_secs,
Arc::new(listener),
listener,
remote_host.to_string(),
remote_port,
);
let handle = tokio::task::spawn_blocking(move || {
let rt = tokio::runtime::Builder::new_current_thread()
.enable_all()
.build()
.expect("failed to create tunnel runtime");
rt.block_on(tunnel_reconnect_loop(
tunnel_args.0,
tunnel_args.1,
tunnel_args.2,
tunnel_args.3,
tunnel_args.4,
tunnel_args.5,
tunnel_args.6,
tunnel_args.7,
tunnel_args.8,
tunnel_args.9,
tunnel_args.10,
tunnel_args.11,
));
});
));
Ok((handle, local_port))
}