use crate::backoff::Backoff;
use crate::config::Config;
use crate::monitor::Monitor;
use crate::pidfile::PidFile;
use crate::{log_debug, log_err, log_info};
use std::fmt;
use std::io;
use std::os::unix::process::ExitStatusExt;
use std::process::ExitStatus;
use std::time::Duration;
use tokio::process::{Child, Command};
use tokio::signal::unix::{Signal, SignalKind, signal};
use tokio::time::{Instant, MissedTickBehavior};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Verdict {
Restart,
ExitOk,
ExitErr,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Reason {
Signalled,
PrematureExit,
ConnectionLost,
CleanExit,
Failed,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Death {
Signal(i32),
Exit(i32),
}
impl Death {
pub fn from_status(s: ExitStatus) -> Self {
match s.signal() {
Some(sig) => Self::Signal(sig),
None => Self::Exit(s.code().unwrap_or(0)),
}
}
}
pub fn classify(
death: Death,
start_count: u64,
uptime: Duration,
gate: Duration,
) -> (Verdict, Reason) {
let code = match death {
Death::Signal(_) => return (Verdict::Restart, Reason::Signalled),
Death::Exit(c) => c,
};
if start_count == 1 && !gate.is_zero() && uptime <= gate {
return (Verdict::ExitErr, Reason::PrematureExit);
}
match code {
255 => (Verdict::Restart, Reason::ConnectionLost),
0 => (Verdict::ExitOk, Reason::CleanExit),
1 | 2 if start_count > 1 || gate.is_zero() => (Verdict::Restart, Reason::ConnectionLost),
_ => (Verdict::ExitErr, Reason::Failed),
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum Sig {
Term,
Int,
Quit,
Hup,
Usr1,
Usr2,
}
impl fmt::Display for Sig {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
let s = match self {
Self::Term => "SIGTERM",
Self::Int => "SIGINT",
Self::Quit => "SIGQUIT",
Self::Hup => "SIGHUP",
Self::Usr1 => "SIGUSR1",
Self::Usr2 => "SIGUSR2",
};
f.write_str(s)
}
}
pub struct Signals {
term: Signal,
int: Signal,
quit: Signal,
hup: Signal,
usr1: Signal,
usr2: Signal,
}
impl Signals {
pub fn new() -> io::Result<Self> {
Ok(Self {
term: signal(SignalKind::terminate())?,
int: signal(SignalKind::interrupt())?,
quit: signal(SignalKind::quit())?,
hup: signal(SignalKind::hangup())?,
usr1: signal(SignalKind::user_defined1())?,
usr2: signal(SignalKind::user_defined2())?,
})
}
pub async fn next(&mut self) -> Sig {
tokio::select! {
_ = self.term.recv() => Sig::Term,
_ = self.int.recv() => Sig::Int,
_ = self.quit.recv() => Sig::Quit,
_ = self.hup.recv() => Sig::Hup,
_ = self.usr1.recv() => Sig::Usr1,
_ = self.usr2.recv() => Sig::Usr2,
}
}
}
pub async fn run(cfg: &Config, pid_file: Option<&PidFile>) -> Verdict {
let mut sigs = match Signals::new() {
Ok(s) => s,
Err(e) => {
log_err!("cannot install signal handlers: {e}");
return Verdict::ExitErr;
}
};
let monitor = match Monitor::bind(cfg).await {
Ok(m) => m,
Err(e) => {
log_err!("cannot open monitor socket: {e}");
return Verdict::ExitErr;
}
};
let ctx = Ctx {
cfg,
monitor: &monitor,
pid_file,
deadline: cfg.max_lifetime.map(|d| Instant::now() + d),
};
let deadline = ctx.deadline;
let mut backoff = Backoff::default();
let mut start_count: u64 = 0;
let mut last_start: Option<Instant> = None;
loop {
if cfg.max_start >= 0 && start_count >= cfg.max_start as u64 {
log_info!("max start count reached; exiting");
return Verdict::ExitOk;
}
if deadline.is_some_and(|d| Instant::now() >= d) {
log_info!("exceeded maximum time to live, shutting down");
return Verdict::ExitOk;
}
let uptime = last_start.map_or(Duration::MAX, |t| t.elapsed());
let delay = backoff.next_delay(uptime, cfg.poll);
log_debug!("checking for grace period, tries = {}", backoff.tries());
if !delay.is_zero() {
log_debug!("sleeping for grace time {} secs", delay.as_secs());
tokio::select! {
_ = tokio::time::sleep(delay) => {}
sig = sigs.next() => {
if let Some(v) = exiting(sig) {
log_info!("received signal to exit ({sig})");
return v;
}
log_debug!("{sig} during backoff; retrying now");
}
}
}
start_count += 1;
if cfg.max_start < 0 {
log_info!("starting ssh (count {start_count})");
} else {
log_info!("starting ssh (count {start_count} of {})", cfg.max_start);
}
let mut child = match spawn(cfg, &monitor) {
Ok(c) => c,
Err(e) => {
log_err!("{}: {e}", cfg.ssh_path.display());
return Verdict::ExitErr;
}
};
let started = Instant::now();
last_start = Some(started);
log_info!("ssh child pid is {}", child.id().unwrap_or(0));
let verdict = watch(&mut child, &ctx, start_count, started, &mut sigs).await;
if verdict != Verdict::Restart {
return verdict;
}
}
}
struct Ctx<'a> {
cfg: &'a Config,
monitor: &'a Monitor,
pid_file: Option<&'a PidFile>,
deadline: Option<Instant>,
}
enum Event {
Exited(io::Result<ExitStatus>),
Tick,
Deadline,
Signal(Sig),
}
async fn watch(
child: &mut Child,
ctx: &Ctx<'_>,
start_count: u64,
started: Instant,
sigs: &mut Signals,
) -> Verdict {
let cfg = ctx.cfg;
let mut tick = tokio::time::interval_at(Instant::now() + cfg.first_poll, cfg.poll);
tick.set_missed_tick_behavior(MissedTickBehavior::Delay);
loop {
let event = tokio::select! {
status = child.wait() => Event::Exited(status),
_ = tick.tick() => Event::Tick,
_ = at(ctx.deadline) => Event::Deadline,
sig = sigs.next() => Event::Signal(sig),
};
match event {
Event::Exited(Err(e)) => {
log_err!("waiting on ssh: {e}");
return Verdict::ExitErr;
}
Event::Exited(Ok(status)) => {
let death = Death::from_status(status);
let (verdict, reason) =
classify(death, start_count, started.elapsed(), cfg.gate_time);
report(death, reason);
return verdict;
}
Event::Tick => {
log_debug!("check on child {}", child.id().unwrap_or(0));
if ctx.monitor.enabled() && !ctx.monitor.probe(cfg).await {
log_info!("port down, restarting ssh");
kill(child, cfg).await;
return Verdict::Restart;
}
if cfg.touch_pid_file
&& let Some(p) = ctx.pid_file
&& let Err(e) = p.touch()
{
log_err!("could not touch pid file: {e}");
}
}
Event::Deadline => {
log_info!("exceeded maximum time to live, shutting down");
kill(child, cfg).await;
return Verdict::ExitOk;
}
Event::Signal(sig) => match sig {
Sig::Term | Sig::Int | Sig::Quit => {
log_info!("received signal to exit ({sig})");
kill(child, cfg).await;
return Verdict::ExitErr;
}
Sig::Usr1 => {
log_info!("signalled to kill and restart ssh");
kill(child, cfg).await;
return Verdict::Restart;
}
Sig::Hup | Sig::Usr2 => log_debug!("woken by {sig}"),
},
}
}
}
fn spawn(cfg: &Config, monitor: &Monitor) -> io::Result<Child> {
let argv = cfg.ssh_argv(monitor.next_forwards());
Command::new(&cfg.ssh_path)
.args(&argv)
.kill_on_drop(false)
.spawn()
}
async fn kill(child: &mut Child, cfg: &Config) {
let Some(pid) = child.id() else {
return; };
log_debug!("sending SIGTERM to {pid}");
unsafe { libc::kill(pid as libc::pid_t, libc::SIGTERM) };
match tokio::time::timeout(cfg.kill_timeout, child.wait()).await {
Ok(Ok(_)) => {}
Ok(Err(e)) => log_err!("waitpid() not successful: {e}"),
Err(_) => {
log_err!(
"ssh {pid} ignored SIGTERM after {}s; sending SIGKILL",
cfg.kill_timeout.as_secs()
);
let _ = child.start_kill();
if let Err(e) = child.wait().await {
log_err!("waitpid() not successful: {e}");
}
}
}
}
async fn at(deadline: Option<Instant>) {
match deadline {
Some(d) => tokio::time::sleep_until(d).await,
None => std::future::pending().await,
}
}
fn exiting(sig: Sig) -> Option<Verdict> {
match sig {
Sig::Term | Sig::Int | Sig::Quit => Some(Verdict::ExitErr),
Sig::Hup | Sig::Usr1 | Sig::Usr2 => None,
}
}
fn report(death: Death, reason: Reason) {
match (death, reason) {
(Death::Signal(s), _) => log_info!("ssh exited on signal {s}, restarting ssh"),
(Death::Exit(c), Reason::PrematureExit) => {
log_err!("ssh exited prematurely with status {c}; rash exiting")
}
(Death::Exit(c), Reason::ConnectionLost) => {
log_info!("ssh exited with error status {c}; restarting ssh")
}
(Death::Exit(c), Reason::CleanExit | Reason::Failed) => {
log_info!("ssh exited with status {c}; rash exiting")
}
(Death::Exit(c), Reason::Signalled) => log_info!("ssh exited with status {c}"),
}
}