Skip to main content

inferlab_runtime/
ssh.rs

1//! The one authority for the SSH client invocation shape, shared by every
2//! module that reaches a remote machine.
3
4use crate::shell::shell_quote;
5use std::process::{Command, Output};
6use thiserror::Error;
7
8const SSH_OPTIONS: &[&str] = &[
9    "-o",
10    "BatchMode=yes",
11    "-o",
12    "ConnectTimeout=10",
13    "-o",
14    "ServerAliveInterval=5",
15    "-o",
16    "ServerAliveCountMax=2",
17    "--",
18];
19
20/// The full SSH argv for bounded execution through an owning operation or
21/// cleanup deadline.
22pub fn ssh_argv(target: &str, script: &str) -> Vec<String> {
23    // Keep the remote account's interactive login initialization, which owns
24    // tool discovery and declared pass-through values, but replace that shell
25    // before the InferLab script runs. The non-login replacement preserves the
26    // initialized environment without running `.bash_logout`, whose exit status
27    // must not replace the remote operation's result.
28    let remote_command = format!("exec \"$BASH\" -c {}", shell_quote(script));
29    let mut argv: Vec<String> = ["ssh"]
30        .into_iter()
31        .map(str::to_owned)
32        .chain(SSH_OPTIONS.iter().map(|option| (*option).to_owned()))
33        .collect();
34    argv.extend([
35        target.to_owned(),
36        "bash".to_owned(),
37        "-lic".to_owned(),
38        shell_quote(&remote_command),
39    ]);
40    argv
41}
42
43#[derive(Debug, Error)]
44pub enum SshError {
45    #[error("failed to launch SSH for {target:?}: {source}")]
46    Launch {
47        target: String,
48        #[source]
49        source: std::io::Error,
50    },
51}
52
53pub fn ssh_output(target: &str, script: &str) -> Result<Output, SshError> {
54    let argv = ssh_argv(target, script);
55    Command::new(&argv[0])
56        .args(&argv[1..])
57        .output()
58        .map_err(|source| SshError::Launch {
59            target: target.to_owned(),
60            source,
61        })
62}
63
64#[cfg(test)]
65mod tests {
66    use super::ssh_argv;
67    use std::error::Error;
68    use std::fs;
69    use std::io;
70    use std::path::Path;
71    use std::process::{Command, Output};
72
73    fn run_remote_command(home: &Path, script: &str) -> Result<Output, Box<dyn Error>> {
74        let argv = ssh_argv("fixture-target", script);
75        let separator = argv
76            .iter()
77            .position(|argument| argument == "--")
78            .ok_or_else(|| io::Error::other("SSH argv has no option separator"))?;
79        let remote_command = argv
80            .get(separator + 2..)
81            .ok_or_else(|| io::Error::other("SSH argv has no remote command"))?
82            .join(" ");
83        Ok(Command::new("sh")
84            .args(["-c", &remote_command])
85            .env("HOME", home)
86            .output()?)
87    }
88
89    #[test]
90    fn login_shell_teardown_does_not_override_remote_script_result() -> Result<(), Box<dyn Error>> {
91        let home = tempfile::tempdir()?;
92        fs::write(
93            home.path().join(".bash_profile"),
94            "export INFERLAB_LOGIN_MARKER=loaded\n",
95        )?;
96        fs::write(
97            home.path().join(".bash_logout"),
98            "printf 'logout-ran\\n' >&2\nfalse\n",
99        )?;
100
101        let alive = run_remote_command(
102            home.path(),
103            "set -eu; printf '%s\\n' \"$INFERLAB_LOGIN_MARKER\"; printf 'probe-stderr\\n' >&2; exit 0",
104        )?;
105        assert!(alive.status.success());
106        assert_eq!(String::from_utf8(alive.stdout)?, "loaded\n");
107        let alive_stderr = String::from_utf8(alive.stderr)?;
108        assert!(alive_stderr.contains("probe-stderr"));
109        assert!(!alive_stderr.contains("logout-ran"));
110
111        let dead = run_remote_command(home.path(), "set -eu; printf 'dead\\n'; exit 3")?;
112        assert_eq!(dead.status.code(), Some(3));
113        assert_eq!(String::from_utf8(dead.stdout)?, "dead\n");
114        Ok(())
115    }
116}