use prt_core::core::ssh_config::{SshHost, SshHostSource};
use prt_core::core::ssh_tunnel::{ResolvedHost, SshTunnelSpec, TunnelKind};
use std::process::{Child, Command, Stdio};
use std::thread;
use std::time::Duration;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
pub enum TunnelStatus {
#[default]
Starting,
Alive,
Failed,
}
pub struct SshTunnel {
pub spec: SshTunnelSpec,
args: Vec<String>,
child: Child,
pub last_status: TunnelStatus,
}
impl SshTunnel {
pub fn spawn(spec: SshTunnelSpec) -> Result<Self, String> {
spec.validate()?;
let args = spec.ssh_args();
let child = spawn_ssh_args(&args)?;
Ok(Self {
spec,
args,
child,
last_status: TunnelStatus::Starting,
})
}
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 {
spec,
args,
child,
last_status: TunnelStatus::Starting,
})
}
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
}
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();
}
pub fn restart(&mut self) -> Result<(), String> {
self.kill();
self.child = spawn_ssh_args(&self.args)?;
self.last_status = TunnelStatus::Starting;
Ok(())
}
}
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)
}
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 drop_failed(&mut self) {
self.tunnels
.retain(|t| t.last_status != TunnelStatus::Failed);
}
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());
}
}