use std::fs;
use std::io::{Read, Write};
use std::os::unix::fs::PermissionsExt;
use std::os::unix::net::{UnixListener, UnixStream};
use std::path::{Path, PathBuf};
use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
use std::sync::{Arc, Mutex};
use std::thread;
use std::time::{Duration, Instant};
use crate::shared::{OwnedChannelStream, SharedClient};
use super::client::{ProbeOutcome, probe_master};
use super::codec::{Frame, PROTOCOL_VERSION};
use super::{read_frame, write_frame};
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
pub enum Persist {
No,
Yes,
Seconds(u64),
}
pub struct MasterConfig {
pub control_path: PathBuf,
pub persist: Persist,
}
struct MasterState {
live_clients: AtomicU64,
foreground_done: AtomicBool,
shutdown: AtomicBool,
last_idle: Mutex<Option<Instant>>,
}
fn bind_socket(path: &Path) -> Result<UnixListener, String> {
match probe_master(path) {
ProbeOutcome::Live => {
return Err(format!(
"ControlPath {} already has a live master; refusing to clobber it",
path.display()
));
}
ProbeOutcome::Stale => {
let _ = fs::remove_file(path);
}
ProbeOutcome::Absent => {}
}
if let Some(parent) = path.parent()
&& !parent.as_os_str().is_empty()
&& !parent.exists()
{
return Err(format!(
"ControlPath parent directory {} does not exist",
parent.display()
));
}
let listener = UnixListener::bind(path)
.map_err(|e| format!("bind control socket {}: {e}", path.display()))?;
if let Err(e) = fs::set_permissions(path, fs::Permissions::from_mode(0o600)) {
let _ = fs::remove_file(path);
return Err(format!("chmod 0600 {}: {e}", path.display()));
}
Ok(listener)
}
pub fn run_master<F>(cfg: MasterConfig, shared: SharedClient, foreground: F) -> Result<i32, String>
where
F: FnOnce(&SharedClient) -> i32 + Send + 'static,
{
let listener = bind_socket(&cfg.control_path)?;
let state = Arc::new(MasterState {
live_clients: AtomicU64::new(0),
foreground_done: AtomicBool::new(false),
shutdown: AtomicBool::new(false),
last_idle: Mutex::new(Some(Instant::now())),
});
let accept_shared = shared.clone();
let accept_state = state.clone();
let accept = thread::spawn(move || {
accept_loop(listener, accept_shared, accept_state);
});
let reaper_state = state.clone();
let persist = cfg.persist;
let reaper_path = cfg.control_path.clone();
let reaper = thread::spawn(move || {
reaper_loop(persist, reaper_state, &reaper_path);
});
let status = foreground(&shared);
state.foreground_done.store(true, Ordering::SeqCst);
match cfg.persist {
Persist::No => {
state.shutdown.store(true, Ordering::SeqCst);
wake_accept(&cfg.control_path);
let _ = fs::remove_file(&cfg.control_path);
let _ = accept.join();
let _ = reaper.join();
Ok(status)
}
Persist::Yes | Persist::Seconds(_) => {
drop(accept);
drop(reaper);
Ok(status)
}
}
}
pub fn run_master_daemon(cfg: MasterConfig, shared: SharedClient) -> Result<(), String> {
let listener = bind_socket(&cfg.control_path)?;
let state = Arc::new(MasterState {
live_clients: AtomicU64::new(0),
foreground_done: AtomicBool::new(true),
shutdown: AtomicBool::new(false),
last_idle: Mutex::new(Some(Instant::now())),
});
let accept_shared = shared.clone();
let accept_state = state.clone();
let accept = thread::spawn(move || {
accept_loop(listener, accept_shared, accept_state);
});
let reaper_state = state.clone();
let persist = cfg.persist;
let reaper_path = cfg.control_path.clone();
let reaper = thread::spawn(move || {
reaper_loop(persist, reaper_state, &reaper_path);
});
while !state.shutdown.load(Ordering::SeqCst) {
thread::sleep(Duration::from_millis(100));
}
wake_accept(&cfg.control_path);
let _ = fs::remove_file(&cfg.control_path);
let _ = accept.join();
let _ = reaper.join();
Ok(())
}
fn wake_accept(path: &Path) {
let _ = UnixStream::connect(path);
}
fn accept_loop(listener: UnixListener, shared: SharedClient, state: Arc<MasterState>) {
for conn in listener.incoming() {
if state.shutdown.load(Ordering::SeqCst) {
break;
}
let stream = match conn {
Ok(s) => s,
Err(_) => continue,
};
let s = shared.clone();
let st = state.clone();
thread::spawn(move || {
handle_client(stream, s, st);
});
}
}
fn reaper_loop(persist: Persist, state: Arc<MasterState>, path: &Path) {
loop {
if state.shutdown.load(Ordering::SeqCst) {
wake_accept(path);
let _ = fs::remove_file(path);
return;
}
thread::sleep(Duration::from_millis(200));
let live = state.live_clients.load(Ordering::SeqCst);
let fg_done = state.foreground_done.load(Ordering::SeqCst);
match persist {
Persist::Yes => {
}
Persist::No => {
}
Persist::Seconds(n) => {
if !fg_done {
continue;
}
if live > 0 {
*state.last_idle.lock().unwrap() = None;
continue;
}
let mut idle = state.last_idle.lock().unwrap();
let since = idle.get_or_insert_with(Instant::now);
if since.elapsed() >= Duration::from_secs(n) {
drop(idle);
state.shutdown.store(true, Ordering::SeqCst);
wake_accept(path);
let _ = fs::remove_file(path);
return;
}
}
}
}
}
fn handle_client(stream: UnixStream, shared: SharedClient, state: Arc<MasterState>) {
let mut ctrl = stream;
match read_frame(&mut ctrl) {
Ok(Some(Frame::Hello { version })) if version == PROTOCOL_VERSION => {}
_ => return, }
if write_frame(
&mut ctrl,
&Frame::Hello {
version: PROTOCOL_VERSION,
},
)
.is_err()
{
return;
}
let open = match read_frame(&mut ctrl) {
Ok(Some(f)) => f,
_ => return,
};
let (want_pty, term, cols, rows, env, command) = match open {
Frame::OpenSession {
want_pty,
term,
cols,
rows,
env,
command,
} => (want_pty, term, cols, rows, env, command),
Frame::ExitRequest => {
state.shutdown.store(true, Ordering::SeqCst);
return;
}
Frame::AliveCheck => {
let _ = write_frame(&mut ctrl, &Frame::AliveOk);
return;
}
Frame::OpenDirectTcpip {
dest_host,
dest_port,
orig_host,
orig_port,
} => {
handle_direct_tcpip(
ctrl, shared, state, &dest_host, dest_port, &orig_host, orig_port,
);
return;
}
_ => return,
};
if !env.is_empty() {
shared.with_session_env(env);
}
let stream_res = if let Some(cmd) = command.as_deref() {
shared.exec_stream(cmd)
} else if want_pty {
shared.shell_stream(&term, cols, rows, 0, 0, Vec::new())
} else {
shared.shell_stream_no_pty()
};
let chan = match stream_res {
Ok(c) => c,
Err(e) => {
let _ = write_frame(
&mut ctrl,
&Frame::StderrData(format!("mux: open session failed: {e}\n").into_bytes()),
);
let _ = write_frame(&mut ctrl, &Frame::ExitStatus { code: 255 });
return;
}
};
state.live_clients.fetch_add(1, Ordering::SeqCst);
*state.last_idle.lock().unwrap() = None;
splice(ctrl, chan, &shared);
state.live_clients.fetch_sub(1, Ordering::SeqCst);
if state.live_clients.load(Ordering::SeqCst) == 0 {
*state.last_idle.lock().unwrap() = Some(Instant::now());
}
}
#[allow(clippy::too_many_arguments)]
fn handle_direct_tcpip(
mut ctrl: UnixStream,
shared: SharedClient,
state: Arc<MasterState>,
dest_host: &str,
dest_port: u32,
orig_host: &str,
orig_port: u32,
) {
if dest_port > 0xFFFF || orig_port > 0xFFFF {
let _ = write_frame(
&mut ctrl,
&Frame::OpenFail {
reason: format!(
"direct-tcpip {dest_host}:{dest_port}: invalid port (out of range)"
),
},
);
return;
}
let chan =
match shared.open_direct_tcpip(dest_host, dest_port as u16, orig_host, orig_port as u16) {
Ok(c) => c,
Err(e) => {
let _ = write_frame(
&mut ctrl,
&Frame::OpenFail {
reason: format!("direct-tcpip {dest_host}:{dest_port}: {e}"),
},
);
return;
}
};
if write_frame(&mut ctrl, &Frame::OpenOk).is_err() {
return;
}
state.live_clients.fetch_add(1, Ordering::SeqCst);
*state.last_idle.lock().unwrap() = None;
splice(ctrl, chan, &shared);
state.live_clients.fetch_sub(1, Ordering::SeqCst);
if state.live_clients.load(Ordering::SeqCst) == 0 {
*state.last_idle.lock().unwrap() = Some(Instant::now());
}
}
fn splice(ctrl: UnixStream, chan: OwnedChannelStream, shared: &SharedClient) {
let channel_id = chan.channel_id();
let _ = shared.set_read_timeout(Some(Duration::from_millis(50)));
let ctrl_read = match ctrl.try_clone() {
Ok(c) => c,
Err(_) => return,
};
let ctrl_write = Arc::new(Mutex::new(ctrl));
let in_shared = shared.clone();
let mut ctrl_in = ctrl_read;
let t_in = thread::spawn(move || {
loop {
match read_frame(&mut ctrl_in) {
Ok(Some(Frame::StdinData(d))) => {
let mut off = 0;
while off < d.len() {
match in_shared.channel_send_data(channel_id, &d[off..]) {
Ok(0) | Err(_) => return,
Ok(n) => off += n,
}
}
}
Ok(Some(Frame::Eof)) => {
let _ = in_shared.channel_send_eof(channel_id);
}
Ok(Some(Frame::WindowChange { cols, rows })) => {
let _ = in_shared.send_window_change(channel_id, cols, rows, 0, 0);
}
Ok(Some(Frame::ExitRequest)) => return,
Ok(Some(_)) => { }
Ok(None) | Err(_) => return,
}
}
});
let err_shared = shared.clone();
let err_write = ctrl_write.clone();
let t_err = thread::spawn(move || {
let mut buf = [0u8; 32 * 1024];
loop {
match err_shared.channel_recv_stderr(channel_id, &mut buf) {
Ok(0) => break,
Ok(n) => {
let mut g = match err_write.lock() {
Ok(g) => g,
Err(_) => break,
};
if write_frame(&mut *g, &Frame::StderrData(buf[..n].to_vec())).is_err() {
break;
}
}
Err(_) => break,
}
}
});
let mut chan = chan;
let mut buf = [0u8; 32 * 1024];
loop {
match chan.read(&mut buf) {
Ok(0) => break,
Ok(n) => {
let mut g = match ctrl_write.lock() {
Ok(g) => g,
Err(_) => break,
};
if write_frame(&mut *g, &Frame::StdoutData(buf[..n].to_vec())).is_err() {
break;
}
}
Err(_) => break,
}
}
let _ = t_err.join();
if let Ok(mut g) = ctrl_write.lock() {
if let Some(sig) = chan.exit_signal() {
let _ = write_frame(&mut *g, &Frame::ExitSignal { name: sig });
}
let code = chan.exit_status().unwrap_or(0);
let _ = write_frame(&mut *g, &Frame::ExitStatus { code: code as u32 });
let _ = write_frame(&mut *g, &Frame::Eof);
let _ = g.flush();
let _ = g.shutdown(std::net::Shutdown::Both);
}
drop(t_in);
drop(chan);
}