use prt_core::core::ssh_config::{SshHost, SshHostSource};
use prt_core::core::ssh_tunnel::{ResolvedHost, SshTunnelSpec, TunnelKind};
use std::collections::VecDeque;
use std::process::{Child, Command, Stdio};
use std::thread;
use std::time::{Duration, Instant};
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum TunnelStatus {
#[default]
Starting,
Alive,
Failed,
}
const INITIAL_BACKOFF: Duration = Duration::from_secs(2);
const MAX_BACKOFF: Duration = Duration::from_secs(60);
const STABILITY_THRESHOLD: Duration = Duration::from_secs(30);
const MAX_RECONNECT_ATTEMPTS: u32 = 10;
const LISTENER_HISTORY_CAP: usize = 6;
const LISTENER_MIN_SAMPLES: usize = 4;
fn push_listener_sample(history: &mut VecDeque<bool>, present: bool) {
history.push_back(present);
if history.len() > LISTENER_HISTORY_CAP {
history.pop_front();
}
}
fn history_is_flapping(history: &VecDeque<bool>) -> bool {
history.len() >= LISTENER_MIN_SAMPLES && history.contains(&true) && history.contains(&false)
}
pub struct SshTunnel {
pub spec: SshTunnelSpec,
args: Vec<String>,
child: Child,
pub last_status: TunnelStatus,
started_at: Instant,
pub auto_reconnect: bool,
retry_backoff: Duration,
retry_count: u32,
next_retry_at: Option<Instant>,
listener_history: VecDeque<bool>,
}
impl SshTunnel {
fn from_child(spec: SshTunnelSpec, args: Vec<String>, child: Child) -> Self {
Self {
spec,
args,
child,
last_status: TunnelStatus::Starting,
started_at: Instant::now(),
auto_reconnect: true,
retry_backoff: INITIAL_BACKOFF,
retry_count: 0,
next_retry_at: None,
listener_history: VecDeque::with_capacity(LISTENER_HISTORY_CAP),
}
}
pub fn spawn(spec: SshTunnelSpec) -> Result<Self, String> {
spec.validate()?;
let args = spec.ssh_args();
let child = spawn_ssh_args(&args)?;
Ok(Self::from_child(spec, args, child))
}
pub fn spawn_with_host(spec: SshTunnelSpec, host: Option<&SshHost>) -> Result<Self, String> {
spec.validate()?;
let args = match host {
Some(h) if h.source == SshHostSource::PrtConfig => {
spec.ssh_args_with(&resolved_from(h))
}
_ => spec.ssh_args(),
};
let child = spawn_ssh_args(&args)?;
Ok(Self::from_child(spec, args, child))
}
pub fn new(local_port: u16, remote_host: &str, remote_port: u16) -> Result<Self, String> {
let spec = SshTunnelSpec {
name: None,
kind: TunnelKind::Local,
local_port,
remote_host: Some("localhost".into()),
remote_port: Some(remote_port),
host_alias: remote_host.to_string(),
};
Self::spawn(spec)
}
pub fn summary(&self) -> String {
self.spec.summary()
}
pub fn refresh_status(&mut self) -> TunnelStatus {
let new = match self.child.try_wait() {
Ok(None) => match self.last_status {
TunnelStatus::Starting => {
TunnelStatus::Alive
}
TunnelStatus::Alive => {
if self.started_at.elapsed() >= STABILITY_THRESHOLD {
self.retry_backoff = INITIAL_BACKOFF;
self.retry_count = 0;
self.next_retry_at = None;
self.auto_reconnect = true;
}
TunnelStatus::Alive
}
other => other,
},
Ok(Some(_)) => TunnelStatus::Failed,
Err(_) => TunnelStatus::Failed,
};
self.last_status = new;
new
}
pub fn kill(&mut self) {
let _ = self.child.kill();
let _ = self.child.wait();
}
fn respawn(&mut self, validate: bool) -> Result<(), String> {
self.kill();
self.child = if validate {
spawn_ssh_args(&self.args)?
} else {
spawn_ssh_args_nowait(&self.args)?
};
self.last_status = TunnelStatus::Starting;
self.started_at = Instant::now();
self.listener_history.clear();
Ok(())
}
pub fn restart(&mut self) -> Result<(), String> {
self.respawn(true)?;
self.retry_backoff = INITIAL_BACKOFF;
self.retry_count = 0;
self.next_retry_at = None;
self.auto_reconnect = true;
Ok(())
}
pub fn uptime(&self) -> Duration {
self.started_at.elapsed()
}
pub fn record_listener(&mut self, present: bool) {
push_listener_sample(&mut self.listener_history, present);
}
pub fn is_flapping(&self) -> bool {
history_is_flapping(&self.listener_history)
}
pub fn pid(&self) -> u32 {
self.child.id()
}
pub fn command_string(&self) -> String {
let mut cmd = String::from("ssh");
for arg in &self.args {
cmd.push(' ');
cmd.push_str(&shell_quote(arg));
}
cmd
}
}
fn shell_quote(arg: &str) -> String {
let safe = !arg.is_empty()
&& arg.bytes().all(|b| {
b.is_ascii_alphanumeric()
|| matches!(
b,
b'_' | b'-' | b'.' | b'/' | b':' | b'=' | b'@' | b',' | b'+'
)
});
if safe {
return arg.to_string();
}
let mut out = String::with_capacity(arg.len() + 2);
out.push('\'');
for ch in arg.chars() {
if ch == '\'' {
out.push_str("'\\''");
} else {
out.push(ch);
}
}
out.push('\'');
out
}
impl Drop for SshTunnel {
fn drop(&mut self) {
self.kill();
}
}
fn resolved_from(h: &SshHost) -> ResolvedHost<'_> {
ResolvedHost {
hostname: h.hostname.as_deref(),
user: h.user.as_deref(),
port: h.port,
identity_file: h.identity_file.as_deref().and_then(|p| p.to_str()),
}
}
fn spawn_ssh_args(args: &[String]) -> Result<Child, String> {
let mut child = Command::new("ssh")
.args(args)
.stdin(Stdio::null())
.stdout(Stdio::null())
.stderr(Stdio::piped())
.spawn()
.map_err(|e| format!("failed to start ssh: {e}"))?;
thread::sleep(Duration::from_millis(150));
if let Ok(Some(status)) = child.try_wait() {
use std::io::Read;
let mut stderr = String::new();
if let Some(mut err) = child.stderr.take() {
let _ = err.read_to_string(&mut stderr);
}
let stderr = stderr.trim();
let details = if stderr.is_empty() {
format!("ssh exited with status {status}")
} else {
stderr.to_string()
};
return Err(format!("failed to establish ssh tunnel: {details}"));
}
Ok(child)
}
fn spawn_ssh_args_nowait(args: &[String]) -> Result<Child, String> {
Command::new("ssh")
.args(args)
.stdin(Stdio::null())
.stdout(Stdio::null())
.stderr(Stdio::null())
.spawn()
.map_err(|e| format!("failed to start ssh: {e}"))
}
pub struct ForwardManager {
pub tunnels: Vec<SshTunnel>,
}
impl Default for ForwardManager {
fn default() -> Self {
Self::new()
}
}
impl ForwardManager {
pub fn new() -> Self {
Self {
tunnels: Vec::new(),
}
}
pub fn add(
&mut self,
local_port: u16,
remote_host: &str,
remote_port: u16,
) -> Result<usize, String> {
let tunnel = SshTunnel::new(local_port, remote_host, remote_port)?;
self.tunnels.push(tunnel);
Ok(self.tunnels.len() - 1)
}
pub fn add_spec_with_host(
&mut self,
spec: SshTunnelSpec,
host: Option<&SshHost>,
) -> Result<usize, String> {
let tunnel = SshTunnel::spawn_with_host(spec, host)?;
self.tunnels.push(tunnel);
Ok(self.tunnels.len() - 1)
}
pub fn cleanup(&mut self) {
for tunnel in &mut self.tunnels {
tunnel.refresh_status();
}
}
pub fn reconnect_failed(&mut self) -> usize {
let now = Instant::now();
let mut reconnected = 0;
for tunnel in &mut self.tunnels {
if tunnel.last_status != TunnelStatus::Failed || !tunnel.auto_reconnect {
continue;
}
match tunnel.next_retry_at {
None => tunnel.next_retry_at = Some(now + tunnel.retry_backoff),
Some(at) if at > now => {}
Some(_) => {
tunnel.retry_count += 1;
let outcome = tunnel.respawn(false);
tunnel.retry_backoff = (tunnel.retry_backoff * 2).min(MAX_BACKOFF);
match outcome {
Ok(()) => {
reconnected += 1;
tunnel.next_retry_at = None;
}
Err(_) => {
tunnel.next_retry_at = Some(now + tunnel.retry_backoff);
}
}
if tunnel.retry_count >= MAX_RECONNECT_ATTEMPTS {
tunnel.auto_reconnect = false;
}
}
}
}
reconnected
}
pub fn drop_failed(&mut self) {
self.tunnels
.retain(|t| t.last_status != TunnelStatus::Failed || t.auto_reconnect);
}
pub fn kill_at(&mut self, idx: usize) {
if idx < self.tunnels.len() {
self.tunnels[idx].kill();
self.tunnels.remove(idx);
}
}
pub fn replace_at(
&mut self,
idx: usize,
spec: SshTunnelSpec,
host: Option<&SshHost>,
) -> Result<(), String> {
if idx >= self.tunnels.len() {
return Err("no such tunnel".into());
}
let new_tunnel = SshTunnel::spawn_with_host(spec, host)?;
self.tunnels[idx] = new_tunnel;
Ok(())
}
pub fn restart_at(&mut self, idx: usize) -> Result<(), String> {
self.tunnels
.get_mut(idx)
.ok_or_else(|| "no such tunnel".to_string())?
.restart()
}
pub fn kill_all(&mut self) {
for tunnel in &mut self.tunnels {
tunnel.kill();
}
self.tunnels.clear();
}
pub fn count(&self) -> usize {
self.tunnels.len()
}
pub fn summaries(&self) -> Vec<String> {
self.tunnels.iter().map(|t| t.summary()).collect()
}
pub fn specs(&self) -> Vec<SshTunnelSpec> {
self.tunnels.iter().map(|t| t.spec.clone()).collect()
}
}
impl Drop for ForwardManager {
fn drop(&mut self) {
self.kill_all();
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn forward_manager_new_is_empty() {
let fm = ForwardManager::new();
assert_eq!(fm.count(), 0);
}
#[test]
fn forward_manager_default_is_empty() {
let fm = ForwardManager::default();
assert_eq!(fm.count(), 0);
}
#[test]
fn specs_snapshot_is_empty_when_no_tunnels() {
let fm = ForwardManager::new();
assert!(fm.specs().is_empty());
}
#[test]
fn shell_quote_leaves_safe_args_untouched() {
assert_eq!(shell_quote("-L"), "-L");
assert_eq!(shell_quote("8080:localhost:80"), "8080:localhost:80");
assert_eq!(
shell_quote("/home/user/.ssh/id_rsa"),
"/home/user/.ssh/id_rsa"
);
assert_eq!(shell_quote("user@host"), "user@host");
}
#[test]
fn shell_quote_wraps_paths_with_spaces() {
assert_eq!(
shell_quote("/Users/x/my keys/id_rsa"),
"'/Users/x/my keys/id_rsa'"
);
}
#[test]
fn shell_quote_escapes_embedded_single_quote() {
assert_eq!(shell_quote("a'b"), "'a'\\''b'");
}
#[test]
fn shell_quote_quotes_empty_arg() {
assert_eq!(shell_quote(""), "''");
}
fn history(samples: &[bool]) -> VecDeque<bool> {
let mut h = VecDeque::new();
for &s in samples {
push_listener_sample(&mut h, s);
}
h
}
#[test]
fn stable_present_is_not_flapping() {
assert!(!history_is_flapping(&history(&[
true, true, true, true, true, true
])));
}
#[test]
fn all_absent_is_not_flapping() {
assert!(!history_is_flapping(&history(&[
false, false, false, false, false, false
])));
}
#[test]
fn alternating_presence_is_flapping() {
assert!(history_is_flapping(&history(&[true, false, true, false])));
}
#[test]
fn insufficient_samples_is_not_flapping() {
assert!(!history_is_flapping(&history(&[true, false])));
}
#[test]
fn empty_history_is_not_flapping() {
assert!(!history_is_flapping(&VecDeque::new()));
}
#[test]
fn window_caps_and_evicts_old_samples() {
let mut h = history(&[false, false, false, false, false, false]);
assert!(!history_is_flapping(&h)); for _ in 0..LISTENER_HISTORY_CAP {
push_listener_sample(&mut h, true);
}
assert_eq!(h.len(), LISTENER_HISTORY_CAP);
assert!(!history_is_flapping(&h)); }
}