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 crate::{
6    container::{BoundedError, BoundedWait, CommandCleanupEvidence, run_with_bound},
7    operation_bound::OperationBound,
8};
9use std::process::Output;
10use thiserror::Error;
11
12const SSH_OPTIONS: &[&str] = &["-o", "BatchMode=yes", "--"];
13
14/// The full SSH argv for bounded execution through an owning operation or
15/// cleanup deadline.
16pub fn ssh_argv(target: &str, script: &str) -> Vec<String> {
17    // Keep the remote account's interactive login initialization, which owns
18    // tool discovery and declared pass-through values, but replace that shell
19    // before the InferLab script runs. The non-login replacement preserves the
20    // initialized environment without running `.bash_logout`, whose exit status
21    // must not replace the remote operation's result.
22    let remote_command = format!("exec \"$BASH\" -c {}", shell_quote(script));
23    let mut argv: Vec<String> = ["ssh"]
24        .into_iter()
25        .map(str::to_owned)
26        .chain(SSH_OPTIONS.iter().map(|option| (*option).to_owned()))
27        .collect();
28    argv.extend([
29        target.to_owned(),
30        "bash".to_owned(),
31        "-lic".to_owned(),
32        shell_quote(&remote_command),
33    ]);
34    argv
35}
36
37#[derive(Debug, Error)]
38pub enum SshError {
39    #[error("failed to launch SSH for {target:?}: {source}")]
40    Launch {
41        target: String,
42        #[source]
43        source: std::io::Error,
44    },
45    #[error("failed to {operation} for SSH target {target:?}: {source}")]
46    Io {
47        target: String,
48        operation: &'static str,
49        #[source]
50        source: std::io::Error,
51    },
52    #[error(
53        "failed while waiting for SSH target {target:?} after {operation_elapsed_ms} ms: {source}; child cleanup: {cleanup:?}"
54    )]
55    WaitCleanup {
56        target: String,
57        operation_elapsed_ms: u64,
58        #[source]
59        source: std::io::Error,
60        cleanup: Box<CommandCleanupEvidence>,
61    },
62    #[error(
63        "SSH for target {target:?} was interrupted after {operation_elapsed_ms} ms; child cleanup: {cleanup:?}"
64    )]
65    Interrupted {
66        target: String,
67        operation_elapsed_ms: u64,
68        cleanup: Box<CommandCleanupEvidence>,
69    },
70    #[error(
71        "failed to clean up interrupted SSH for target {target:?} after {operation_elapsed_ms} ms: {source}; child cleanup: {cleanup:?}"
72    )]
73    InterruptCleanup {
74        target: String,
75        operation_elapsed_ms: u64,
76        #[source]
77        source: std::io::Error,
78        cleanup: Box<CommandCleanupEvidence>,
79    },
80    #[error(
81        "SSH supervisor for target {target:?} unexpectedly exhausted an unbounded operation after {operation_elapsed_ms} ms; child cleanup: {cleanup:?}"
82    )]
83    UnexpectedDeadline {
84        target: String,
85        operation_elapsed_ms: u64,
86        cleanup: Option<Box<CommandCleanupEvidence>>,
87    },
88}
89
90pub fn ssh_output(target: &str, script: &str) -> Result<Output, SshError> {
91    run_ssh(target, script, None)
92}
93
94pub fn ssh_output_with_input(target: &str, script: &str, input: &[u8]) -> Result<Output, SshError> {
95    run_ssh(target, script, Some(input))
96}
97
98fn run_ssh(target: &str, script: &str, input: Option<&[u8]>) -> Result<Output, SshError> {
99    let argv = ssh_argv(target, script);
100    match run_with_bound(&argv, None, input, &OperationBound::unbounded(), None) {
101        Ok(BoundedWait::Exited {
102            status,
103            stdout,
104            stderr,
105        }) => Ok(Output {
106            status,
107            stdout,
108            stderr,
109        }),
110        Ok(BoundedWait::Expired {
111            kill,
112            operation_elapsed_ms,
113            cleanup,
114        }) => {
115            kill.map_err(|source| SshError::Io {
116                target: target.to_owned(),
117                operation: "clean up SSH after unexpected deadline",
118                source,
119            })?;
120            Err(SshError::UnexpectedDeadline {
121                target: target.to_owned(),
122                operation_elapsed_ms,
123                cleanup: cleanup.map(Box::new),
124            })
125        }
126        Ok(BoundedWait::Interrupted {
127            kill,
128            operation_elapsed_ms,
129            cleanup,
130        }) => match kill {
131            Ok(()) => Err(SshError::Interrupted {
132                target: target.to_owned(),
133                operation_elapsed_ms,
134                cleanup: Box::new(cleanup),
135            }),
136            Err(source) => Err(SshError::InterruptCleanup {
137                target: target.to_owned(),
138                operation_elapsed_ms,
139                source,
140                cleanup: Box::new(cleanup),
141            }),
142        },
143        Err(BoundedError::Launch(source)) => Err(SshError::Launch {
144            target: target.to_owned(),
145            source,
146        }),
147        Err(BoundedError::Stdin(source)) => Err(SshError::Io {
148            target: target.to_owned(),
149            operation: "write SSH stdin",
150            source,
151        }),
152        Err(BoundedError::Wait(source)) => Err(SshError::Io {
153            target: target.to_owned(),
154            operation: "wait for SSH",
155            source,
156        }),
157        Err(BoundedError::WaitCleanup {
158            source,
159            operation_elapsed_ms,
160            cleanup,
161        }) => Err(SshError::WaitCleanup {
162            target: target.to_owned(),
163            operation_elapsed_ms,
164            source,
165            cleanup: Box::new(cleanup),
166        }),
167    }
168}
169
170#[cfg(test)]
171mod tests {
172    use super::ssh_argv;
173    use std::error::Error;
174    use std::fs;
175    use std::io;
176    use std::path::Path;
177    use std::process::{Command, Output};
178
179    #[test]
180    fn ssh_argv_leaves_connection_policy_to_openssh_configuration() {
181        let argv = ssh_argv("fixture-target", "exit 0");
182
183        assert_eq!(
184            &argv[..5],
185            ["ssh", "-o", "BatchMode=yes", "--", "fixture-target"]
186        );
187        assert!(!argv.iter().any(|argument| {
188            argument.starts_with("ConnectTimeout=")
189                || argument.starts_with("ServerAliveInterval=")
190                || argument.starts_with("ServerAliveCountMax=")
191        }));
192    }
193
194    fn run_remote_command(home: &Path, script: &str) -> Result<Output, Box<dyn Error>> {
195        let argv = ssh_argv("fixture-target", script);
196        let separator = argv
197            .iter()
198            .position(|argument| argument == "--")
199            .ok_or_else(|| io::Error::other("SSH argv has no option separator"))?;
200        let remote_command = argv
201            .get(separator + 2..)
202            .ok_or_else(|| io::Error::other("SSH argv has no remote command"))?
203            .join(" ");
204        Ok(Command::new("sh")
205            .args(["-c", &remote_command])
206            .env("HOME", home)
207            .output()?)
208    }
209
210    #[test]
211    fn login_shell_teardown_does_not_override_remote_script_result() -> Result<(), Box<dyn Error>> {
212        let home = tempfile::tempdir()?;
213        fs::write(
214            home.path().join(".bash_profile"),
215            "export INFERLAB_LOGIN_MARKER=loaded\n",
216        )?;
217        fs::write(
218            home.path().join(".bash_logout"),
219            "printf 'logout-ran\\n' >&2\nfalse\n",
220        )?;
221
222        let alive = run_remote_command(
223            home.path(),
224            "set -eu; printf '%s\\n' \"$INFERLAB_LOGIN_MARKER\"; printf 'probe-stderr\\n' >&2; exit 0",
225        )?;
226        assert!(alive.status.success());
227        assert_eq!(String::from_utf8(alive.stdout)?, "loaded\n");
228        let alive_stderr = String::from_utf8(alive.stderr)?;
229        assert!(alive_stderr.contains("probe-stderr"));
230        assert!(!alive_stderr.contains("logout-ran"));
231
232        let dead = run_remote_command(home.path(), "set -eu; printf 'dead\\n'; exit 3")?;
233        assert_eq!(dead.status.code(), Some(3));
234        assert_eq!(String::from_utf8(dead.stdout)?, "dead\n");
235        Ok(())
236    }
237}