1use super::host::HostSpec;
4use super::process::{CommandOutput, ProcessRunner, StdioMode};
5use crate::error::{HortoError, Result};
6use std::path::{Path, PathBuf};
7
8#[derive(Debug, Clone, Default)]
10pub struct SshEnv {
11 pub force_askpass: bool,
13}
14
15impl SshEnv {
16 #[must_use]
18 pub fn as_pairs(&self) -> Vec<(String, String)> {
19 let mut out = Vec::new();
20 if self.force_askpass {
21 out.push(("SSH_ASKPASS_REQUIRE".into(), "force".into()));
22 if let Ok(ask) = super::askpass::resolve_askpass() {
23 out.push(("SSH_ASKPASS".into(), ask.display().to_string()));
24 out.push((
25 "DISPLAY".into(),
26 std::env::var("DISPLAY").unwrap_or_else(|_| ":0".into()),
27 ));
28 }
29 }
30 out
31 }
32
33 fn as_refs(pairs: &[(String, String)]) -> Vec<(&str, &str)> {
34 pairs
35 .iter()
36 .map(|(k, v)| (k.as_str(), v.as_str()))
37 .collect()
38 }
39}
40
41#[derive(Debug, Clone)]
43pub struct SshSession {
44 pub host: HostSpec,
46 pub env: SshEnv,
48 pub config_file: Option<PathBuf>,
50}
51
52fn require_ok(program: &str, out: &CommandOutput) -> Result<()> {
53 if out.success() {
54 return Ok(());
55 }
56 let detail = if out.stderr.trim().is_empty() {
57 out.stdout.trim().to_owned()
58 } else {
59 out.stderr.trim().to_owned()
60 };
61 Err(HortoError::command(
62 program,
63 format!("exit {}: {detail}", out.status),
64 ))
65}
66
67#[must_use]
69pub fn remote_install_key_banner(host: &str, pub_path: &Path) -> String {
70 format!(
71 "[horto remote] PC → box '{host}': install pubkey {} onto box (ssh-copy-id; may ask password)",
72 pub_path.display()
73 )
74}
75
76fn default_identity_pubkey() -> Result<PathBuf> {
78 let home = std::env::var_os("HOME")
79 .map(PathBuf::from)
80 .ok_or_else(|| HortoError::msg("HOME unset; cannot find default SSH identity"))?;
81 let ssh_dir = home.join(".ssh");
82 for name in ["id_ed25519.pub", "id_ecdsa.pub", "id_rsa.pub"] {
83 let path = ssh_dir.join(name);
84 if path.is_file() {
85 return Ok(path);
86 }
87 }
88 Err(HortoError::msg(format!(
89 "no default SSH pubkey in {} (tried id_ed25519.pub, id_ecdsa.pub, id_rsa.pub)",
90 ssh_dir.display()
91 )))
92}
93
94fn identity_private_key(pub_path: &Path) -> Option<PathBuf> {
96 let stem = pub_path.file_stem()?;
97 let priv_path = pub_path.parent()?.join(stem);
98 priv_path.is_file().then_some(priv_path)
99}
100
101impl SshSession {
102 fn with_config_prefix(&self, rest: &[&str]) -> Vec<String> {
103 let mut owned = Vec::new();
104 if let Some(cfg) = &self.config_file {
105 owned.push("-F".into());
106 owned.push(cfg.display().to_string());
107 }
108 for a in rest {
109 owned.push((*a).to_owned());
110 }
111 owned
112 }
113
114 fn pubkey_already_authorized(
116 &self,
117 runner: &dyn ProcessRunner,
118 pub_path: &Path,
119 ) -> Result<bool> {
120 let Some(priv_path) = identity_private_key(pub_path) else {
121 return Ok(false);
122 };
123 let priv_s = priv_path.display().to_string();
124 let pairs = self.env.as_pairs();
125 let env = SshEnv::as_refs(&pairs);
126 let owned = self.with_config_prefix(&[
129 "-o",
130 "BatchMode=yes",
131 "-o",
132 "IdentitiesOnly=yes",
133 "-o",
134 "StrictHostKeyChecking=accept-new",
135 "-o",
136 "ConnectTimeout=10",
137 "-i",
138 &priv_s,
139 &self.host.raw,
140 "true",
141 ]);
142 let refs: Vec<&str> = owned.iter().map(String::as_str).collect();
143 let out = runner.run("ssh", &refs, &env, StdioMode::Capture)?;
144 Ok(out.success())
145 }
146
147 pub fn exec(
156 &self,
157 runner: &dyn ProcessRunner,
158 remote_cmd: &str,
159 stdio: StdioMode,
160 ) -> Result<CommandOutput> {
161 let pairs = self.env.as_pairs();
162 let env = SshEnv::as_refs(&pairs);
163 let owned = if stdio == StdioMode::Inherit {
166 self.with_config_prefix(&[
167 "-tt",
168 "-o",
169 "BatchMode=no",
170 "-o",
171 "StrictHostKeyChecking=accept-new",
172 &self.host.raw,
173 remote_cmd,
174 ])
175 } else {
176 self.with_config_prefix(&[
177 "-o",
178 "BatchMode=no",
179 "-o",
180 "StrictHostKeyChecking=accept-new",
181 &self.host.raw,
182 remote_cmd,
183 ])
184 };
185 let refs: Vec<&str> = owned.iter().map(String::as_str).collect();
186 let out = runner.run("ssh", &refs, &env, stdio)?;
187 require_ok("ssh", &out)?;
188 Ok(out)
189 }
190
191 pub fn exec_stdin(
197 &self,
198 runner: &dyn ProcessRunner,
199 remote_cmd: &str,
200 stdin: &[u8],
201 ) -> Result<CommandOutput> {
202 self.exec_stdin_inner(runner, remote_cmd, stdin, false)
203 }
204
205 pub fn exec_stdin_reboot(
211 &self,
212 runner: &dyn ProcessRunner,
213 remote_cmd: &str,
214 stdin: &[u8],
215 ) -> Result<CommandOutput> {
216 self.exec_stdin_inner(runner, remote_cmd, stdin, true)
217 }
218
219 fn exec_stdin_inner(
220 &self,
221 runner: &dyn ProcessRunner,
222 remote_cmd: &str,
223 stdin: &[u8],
224 reboot_timeouts: bool,
225 ) -> Result<CommandOutput> {
226 let pairs = self.env.as_pairs();
227 let env = SshEnv::as_refs(&pairs);
228 let owned = if reboot_timeouts {
229 self.with_config_prefix(&[
230 "-o",
231 "BatchMode=no",
232 "-o",
233 "StrictHostKeyChecking=accept-new",
234 "-o",
235 "ConnectTimeout=15",
236 "-o",
237 "ServerAliveInterval=2",
238 "-o",
239 "ServerAliveCountMax=2",
240 &self.host.raw,
241 remote_cmd,
242 ])
243 } else {
244 self.with_config_prefix(&[
245 "-o",
246 "BatchMode=no",
247 "-o",
248 "StrictHostKeyChecking=accept-new",
249 &self.host.raw,
250 remote_cmd,
251 ])
252 };
253 let refs: Vec<&str> = owned.iter().map(String::as_str).collect();
254 let out = runner.run_with_stdin("ssh", &refs, &env, stdin)?;
255 require_ok("ssh", &out)?;
256 Ok(out)
257 }
258
259 pub fn scp_to(
265 &self,
266 runner: &dyn ProcessRunner,
267 local: &Path,
268 remote_path: &str,
269 ) -> Result<()> {
270 let local_s = local
271 .to_str()
272 .ok_or_else(|| HortoError::msg("non-utf8 local path for scp"))?;
273 let dest = format!("{}:{remote_path}", self.host.raw);
274 let pairs = self.env.as_pairs();
275 let env = SshEnv::as_refs(&pairs);
276 let owned = self.with_config_prefix(&[
277 "-q",
278 "-o",
279 "StrictHostKeyChecking=accept-new",
280 local_s,
281 &dest,
282 ]);
283 let refs: Vec<&str> = owned.iter().map(String::as_str).collect();
284 let out = runner.run("scp", &refs, &env, StdioMode::Capture)?;
286 require_ok("scp", &out)
287 }
288
289 pub fn install_ssh_key(&self, runner: &dyn ProcessRunner) -> Result<()> {
299 let pub_path = default_identity_pubkey()?;
300 if self.pubkey_already_authorized(runner, &pub_path)? {
301 tracing::info!(
302 "[horto remote] pubkey {} already authorized on {}; skip ssh-copy-id",
303 pub_path.display(),
304 self.host.raw
305 );
306 return Ok(());
307 }
308 let pub_s = pub_path.display().to_string();
309 tracing::info!("{}", remote_install_key_banner(&self.host.raw, &pub_path));
310 let pairs = self.env.as_pairs();
311 let env = SshEnv::as_refs(&pairs);
312 let owned = self.with_config_prefix(&[
313 "-i",
314 &pub_s,
315 "-o",
316 "StrictHostKeyChecking=accept-new",
317 &self.host.raw,
318 ]);
319 let refs: Vec<&str> = owned.iter().map(String::as_str).collect();
320 let out = runner.run("ssh-copy-id", &refs, &env, StdioMode::Capture)?;
321 require_ok("ssh-copy-id", &out)
322 }
323
324 pub fn remote_has_rsync(&self, runner: &dyn ProcessRunner) -> bool {
326 self.exec(
327 runner,
328 "command -v rsync >/dev/null 2>&1",
329 StdioMode::Capture,
330 )
331 .is_ok()
332 }
333}
334
335#[cfg(test)]
336pub mod tests {
337 use super::*;
338 use crate::remote::host::parse_host_spec;
339 use crate::remote::process::ScriptedRunner;
340
341 #[test]
342 fn exec_builds_ssh_args() {
343 let runner = ScriptedRunner::default();
344 runner.push("ssh", ScriptedRunner::ok("x86_64\n"));
345 let session = SshSession {
346 host: parse_host_spec("box").unwrap(),
347 env: SshEnv::default(),
348 config_file: None,
349 };
350 let out = session
351 .exec(&runner, "uname -m", StdioMode::Capture)
352 .unwrap();
353 assert_eq!(out.stdout.trim(), "x86_64");
354 let calls = runner.calls.lock().unwrap();
355 assert_eq!(calls[0].0, "ssh");
356 assert!(calls[0].1.iter().any(|a| a == "box"));
357 assert!(calls[0].1.iter().any(|a| a == "uname -m"));
358 assert!(!calls[0].1.iter().any(|a| a == "-tt"));
359 drop(calls);
360 }
361
362 #[test]
363 fn remote_install_key_banner_includes_host_and_path() {
364 let msg = remote_install_key_banner("horto", Path::new("/home/u/.ssh/id_ed25519.pub"));
365 assert!(msg.contains("[horto remote]"));
366 assert!(msg.contains("horto"));
367 assert!(msg.contains("id_ed25519.pub"));
368 assert!(msg.contains("ssh-copy-id"));
369 }
370
371 #[test]
372 fn exec_inherit_forces_remote_tty() {
373 let runner = ScriptedRunner::default();
374 runner.push("ssh", ScriptedRunner::ok(""));
375 let session = SshSession {
376 host: parse_host_spec("box").unwrap(),
377 env: SshEnv::default(),
378 config_file: None,
379 };
380 session
381 .exec(&runner, "sudo true", StdioMode::Inherit)
382 .unwrap();
383 let args = &runner.calls.lock().unwrap()[0].1;
384 assert_eq!(args[0], "-tt");
385 assert!(args.iter().any(|a| a == "sudo true"));
386 }
387
388 use std::sync::MutexGuard;
389
390 fn home_lock() -> MutexGuard<'static, ()> {
392 crate::remote::ENV_LOCK
393 .lock()
394 .unwrap_or_else(std::sync::PoisonError::into_inner)
395 }
396
397 fn with_temp_home<R>(setup: impl FnOnce(&Path), f: impl FnOnce() -> R) -> R {
399 let _guard = home_lock();
400
401 let dir = tempfile::TempDir::new().unwrap();
402 setup(dir.path());
403
404 let prev_home = std::env::var_os("HOME");
405 std::env::set_var("HOME", dir.path());
406 let out = f();
407 match prev_home {
408 Some(h) => std::env::set_var("HOME", h),
409 None => std::env::remove_var("HOME"),
410 }
411 out
412 }
413
414 pub fn with_fake_default_pubkey<R>(f: impl FnOnce(PathBuf) -> R) -> R {
416 let _guard = home_lock();
417 let dir = tempfile::TempDir::new().unwrap();
418 let ssh = dir.path().join(".ssh");
419 std::fs::create_dir_all(&ssh).unwrap();
420 let pub_path = ssh.join("id_ed25519.pub");
421 std::fs::write(&pub_path, "ssh-ed25519 AAAATEST test@ci\n").unwrap();
422 std::fs::write(ssh.join("id_ed25519"), b"PRIVATE\n").unwrap();
423 let prev_home = std::env::var_os("HOME");
424 std::env::set_var("HOME", dir.path());
425 let out = f(pub_path);
426 match prev_home {
427 Some(h) => std::env::set_var("HOME", h),
428 None => std::env::remove_var("HOME"),
429 }
430 out
431 }
432
433 #[test]
434 fn default_identity_errors_when_home_unset() {
435 let _guard = home_lock();
436 let prev = std::env::var_os("HOME");
437 std::env::remove_var("HOME");
438 let err = default_identity_pubkey().unwrap_err();
439 if let Some(h) = prev {
440 std::env::set_var("HOME", h);
441 }
442 assert!(err.to_string().contains("HOME unset"));
443 }
444
445 #[test]
446 fn default_identity_errors_when_no_pubkey() {
447 with_temp_home(
448 |home| {
449 std::fs::create_dir_all(home.join(".ssh")).unwrap();
450 },
451 || {
452 let err = default_identity_pubkey().unwrap_err();
453 assert!(err.to_string().contains("no default SSH pubkey"));
454 },
455 );
456 }
457
458 #[test]
459 fn default_identity_falls_back_to_ecdsa_then_rsa() {
460 with_temp_home(
461 |home| {
462 let ssh = home.join(".ssh");
463 std::fs::create_dir_all(&ssh).unwrap();
464 std::fs::write(ssh.join("id_ecdsa.pub"), "ecdsa-sha2-nistp256 AAAA ecdsa\n")
465 .unwrap();
466 },
467 || {
468 let path = default_identity_pubkey().unwrap();
469 assert!(path.ends_with("id_ecdsa.pub"));
470 },
471 );
472 with_temp_home(
473 |home| {
474 let ssh = home.join(".ssh");
475 std::fs::create_dir_all(&ssh).unwrap();
476 std::fs::write(ssh.join("id_rsa.pub"), "ssh-rsa AAAA rsa\n").unwrap();
477 },
478 || {
479 let path = default_identity_pubkey().unwrap();
480 assert!(path.ends_with("id_rsa.pub"));
481 },
482 );
483 }
484
485 #[test]
486 fn install_key_passes_i_pubkey() {
487 with_fake_default_pubkey(|pub_path| {
488 let runner = ScriptedRunner::default();
489 runner.push("ssh", ScriptedRunner::fail(255, "Permission denied"));
491 runner.push("ssh-copy-id", ScriptedRunner::ok(""));
492 let session = SshSession {
493 host: parse_host_spec("box").unwrap(),
494 env: SshEnv::default(),
495 config_file: None,
496 };
497 session.install_ssh_key(&runner).unwrap();
498 let calls = runner.calls.lock().unwrap();
499 assert_eq!(calls[0].0, "ssh");
500 assert!(calls[0].1.iter().any(|a| a == "BatchMode=yes"));
501 assert_eq!(calls[1].0, "ssh-copy-id");
502 let args = &calls[1].1;
503 let i = args.iter().position(|a| a == "-i").expect("-i missing");
504 assert_eq!(Path::new(&args[i + 1]), pub_path.as_path());
505 assert!(args.iter().any(|a| a == "box"));
506 drop(calls);
507 });
508 }
509
510 #[test]
511 fn install_key_skips_when_already_authorized() {
512 with_fake_default_pubkey(|_| {
513 let runner = ScriptedRunner::default();
514 runner.push("ssh", ScriptedRunner::ok(""));
515 let session = SshSession {
516 host: parse_host_spec("box").unwrap(),
517 env: SshEnv::default(),
518 config_file: None,
519 };
520 session.install_ssh_key(&runner).unwrap();
521 let programs: Vec<_> = runner
522 .calls
523 .lock()
524 .unwrap()
525 .iter()
526 .map(|(p, _, _, _)| p.clone())
527 .collect();
528 assert_eq!(programs, vec!["ssh"]);
529 assert!(!programs.iter().any(|p| p == "ssh-copy-id"));
530 });
531 }
532
533 #[test]
534 fn install_key_without_private_key_still_runs_copy_id() {
535 with_temp_home(
536 |home| {
537 let ssh = home.join(".ssh");
538 std::fs::create_dir_all(&ssh).unwrap();
539 std::fs::write(ssh.join("id_ed25519.pub"), "ssh-ed25519 AAAATEST test@ci\n")
541 .unwrap();
542 },
543 || {
544 let runner = ScriptedRunner::default();
545 runner.push("ssh-copy-id", ScriptedRunner::ok(""));
546 let session = SshSession {
547 host: parse_host_spec("box").unwrap(),
548 env: SshEnv::default(),
549 config_file: None,
550 };
551 session.install_ssh_key(&runner).unwrap();
552 let calls = runner.calls.lock().unwrap();
553 assert_eq!(calls.len(), 1);
554 assert_eq!(calls[0].0, "ssh-copy-id");
555 drop(calls);
556 },
557 );
558 }
559
560 #[test]
561 fn install_key_errors_when_no_default_pubkey() {
562 with_temp_home(
563 |home| {
564 std::fs::create_dir_all(home.join(".ssh")).unwrap();
565 },
566 || {
567 let runner = ScriptedRunner::default();
568 let session = SshSession {
569 host: parse_host_spec("box").unwrap(),
570 env: SshEnv::default(),
571 config_file: None,
572 };
573 let err = session.install_ssh_key(&runner).unwrap_err();
574 assert!(err.to_string().contains("no default SSH pubkey"));
575 },
576 );
577 }
578
579 #[test]
580 fn scp_failure_surfaces() {
581 let runner = ScriptedRunner::default();
582 runner.push("scp", ScriptedRunner::fail(1, "Permission denied"));
583 let session = SshSession {
584 host: parse_host_spec("box").unwrap(),
585 env: SshEnv::default(),
586 config_file: None,
587 };
588 let err = session
589 .scp_to(&runner, Path::new("/tmp/x"), "/tmp/x")
590 .unwrap_err();
591 assert!(err.to_string().contains("scp"));
592 }
593
594 #[test]
595 fn config_file_adds_f_flag() {
596 let runner = ScriptedRunner::default();
597 runner.push("ssh", ScriptedRunner::ok("ok\n"));
598 let session = SshSession {
599 host: parse_host_spec("box").unwrap(),
600 env: SshEnv::default(),
601 config_file: Some(PathBuf::from("/tmp/ssh_config")),
602 };
603 session.exec(&runner, "true", StdioMode::Capture).unwrap();
604 let args = &runner.calls.lock().unwrap()[0].1;
605 assert_eq!(args[0], "-F");
606 assert_eq!(args[1], "/tmp/ssh_config");
607 }
608
609 #[test]
610 fn force_askpass_sets_env() {
611 std::env::set_var("SSH_ASKPASS", "/usr/bin/ssh-askpass");
612 let env = SshEnv {
613 force_askpass: true,
614 };
615 let pairs = env.as_pairs();
616 assert!(pairs.iter().any(|(k, _)| k == "SSH_ASKPASS_REQUIRE"));
617 assert!(pairs.iter().any(|(k, _)| k == "SSH_ASKPASS"));
618 std::env::remove_var("SSH_ASKPASS");
619 }
620
621 #[test]
622 fn require_ok_prefers_stdout_when_stderr_blank() {
623 let runner = ScriptedRunner::default();
624 runner.push(
625 "ssh",
626 CommandOutput {
627 status: 1,
628 stdout: "denied-stdout".into(),
629 stderr: String::new(),
630 },
631 );
632 let session = SshSession {
633 host: parse_host_spec("box").unwrap(),
634 env: SshEnv::default(),
635 config_file: None,
636 };
637 let err = session
638 .exec(&runner, "false", StdioMode::Capture)
639 .unwrap_err();
640 assert!(err.to_string().contains("denied-stdout"));
641 }
642
643 #[test]
644 fn remote_has_rsync_true_and_false() {
645 let runner = ScriptedRunner::default();
646 runner.push("ssh", ScriptedRunner::ok(""));
647 let session = SshSession {
648 host: parse_host_spec("box").unwrap(),
649 env: SshEnv::default(),
650 config_file: None,
651 };
652 assert!(session.remote_has_rsync(&runner));
653 runner.push("ssh", ScriptedRunner::fail(1, "no"));
654 assert!(!session.remote_has_rsync(&runner));
655 }
656
657 #[test]
658 fn scp_ok_path() {
659 let runner = ScriptedRunner::default();
660 runner.push("scp", ScriptedRunner::ok(""));
661 let session = SshSession {
662 host: parse_host_spec("box").unwrap(),
663 env: SshEnv::default(),
664 config_file: None,
665 };
666 session
667 .scp_to(&runner, Path::new("/tmp/x"), "/tmp/x")
668 .unwrap();
669 let call = &runner.calls.lock().unwrap()[0];
670 assert_eq!(call.0, "scp");
671 assert_eq!(call.3, StdioMode::Capture);
672 assert!(call.1.iter().any(|a| a == "-q"));
673 }
674
675 #[test]
676 fn exec_stdin_and_reboot_keepalive_flags() {
677 let runner = ScriptedRunner::default();
678 runner.push("ssh", ScriptedRunner::ok(""));
679 runner.push("ssh", ScriptedRunner::ok(""));
680 let session = SshSession {
681 host: parse_host_spec("box").unwrap(),
682 env: SshEnv::default(),
683 config_file: None,
684 };
685 session
686 .exec_stdin(&runner, "cat >/dev/null", b"pw\n")
687 .unwrap();
688 session
689 .exec_stdin_reboot(&runner, "sudo -S reboot", b"pw\n")
690 .unwrap();
691 let calls = runner.calls.lock().unwrap();
692 assert!(!calls[0].1.iter().any(|a| a.contains("ServerAliveInterval")));
693 assert!(calls[1].1.iter().any(|a| a == "ServerAliveInterval=2"));
694 drop(calls);
695 }
696}