mod known_hosts;
mod session;
#[cfg(feature = "sftp")]
mod sftp;
mod ssh_config;
use std::fmt;
use std::path::{Path, PathBuf};
use std::sync::Arc;
use std::time::Duration;
use net_backend_protocol::Secret;
pub use session::{SshChunk, SshRun, SshSession};
#[cfg(feature = "sftp")]
#[cfg_attr(docsrs, doc(cfg(feature = "sftp")))]
pub use sftp::{SftpEntry, SftpEntryKind};
use crate::{Error, MAX_TIMEOUT};
pub const DEFAULT_SSH_CONNECT_TIMEOUT: Duration = Duration::from_secs(15);
pub const DEFAULT_SSH_COMMAND_TIMEOUT: Duration = Duration::from_secs(60);
pub const DEFAULT_SSH_MAX_OUTPUT_BYTES: u64 = 8 * 1024 * 1024;
pub const DEFAULT_SFTP_TIMEOUT: Duration = Duration::from_secs(300);
pub const DEFAULT_SFTP_MAX_BYTES: u64 = 256 * 1024 * 1024;
pub const MAX_SSH_COMMAND_BYTES: usize = 64 * 1024;
#[derive(Clone)]
pub struct SshAuth(pub(crate) AuthKind);
#[derive(Clone)]
pub(crate) enum AuthKind {
KeyFile { path: PathBuf, passphrase: Option<Secret> },
Agent,
Password(Secret),
KeyboardInteractive(Arc<dyn SshPromptResponder>),
}
impl SshAuth {
pub fn key_file(path: impl Into<PathBuf>) -> Self {
Self(AuthKind::KeyFile { path: path.into(), passphrase: None })
}
pub fn key_file_with_passphrase(path: impl Into<PathBuf>, passphrase: impl Into<Secret>) -> Self {
Self(AuthKind::KeyFile { path: path.into(), passphrase: Some(passphrase.into()) })
}
pub fn agent() -> Self {
Self(AuthKind::Agent)
}
pub fn password(password: impl Into<Secret>) -> Self {
Self(AuthKind::Password(password.into()))
}
pub fn keyboard_interactive(responder: impl SshPromptResponder) -> Self {
Self(AuthKind::KeyboardInteractive(Arc::new(responder)))
}
}
impl fmt::Debug for SshAuth {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
match &self.0 {
AuthKind::KeyFile { path, passphrase } => {
f.debug_struct("KeyFile").field("file", &file_name(path)).field("passphrase", &passphrase.as_ref().map(|_| "<redacted>")).finish()
}
AuthKind::Agent => f.write_str("Agent"),
AuthKind::Password(_) => f.write_str("Password(<redacted>)"),
AuthKind::KeyboardInteractive(_) => f.write_str("KeyboardInteractive"),
}
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
#[non_exhaustive]
pub struct SshPrompt {
pub text: String,
pub echo: bool,
}
impl SshPrompt {
pub fn new(text: impl Into<String>, echo: bool) -> Self {
Self { text: text.into(), echo }
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
#[non_exhaustive]
pub struct SshPromptRequest {
pub name: String,
pub instructions: String,
pub prompts: Vec<SshPrompt>,
}
impl SshPromptRequest {
pub fn new(prompts: Vec<SshPrompt>) -> Self {
Self { name: String::new(), instructions: String::new(), prompts }
}
}
pub trait SshPromptResponder: Send + Sync + 'static {
fn respond(&self, request: &SshPromptRequest) -> Option<Vec<Secret>>;
}
#[derive(Clone, Default)]
pub struct SshPromptAnswers {
answers: Vec<(String, Secret)>,
}
impl fmt::Debug for SshPromptAnswers {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("SshPromptAnswers").field("words", &self.answers.iter().map(|(w, _)| w.as_str()).collect::<Vec<_>>()).finish()
}
}
impl SshPromptAnswers {
pub fn new() -> Self {
Self::default()
}
pub fn answer_containing(mut self, word: impl Into<String>, answer: impl Into<Secret>) -> Self {
self.answers.push((word.into().to_lowercase(), answer.into()));
self
}
}
impl SshPromptResponder for SshPromptAnswers {
fn respond(&self, request: &SshPromptRequest) -> Option<Vec<Secret>> {
request
.prompts
.iter()
.map(|prompt| {
let text = prompt.text.to_lowercase();
self.answers.iter().find(|(word, _)| text.contains(word.as_str())).map(|(_, answer)| answer.clone())
})
.collect()
}
}
pub(crate) fn file_name(path: &Path) -> String {
path.file_name().map_or_else(|| "<no file name>".to_string(), |n| n.to_string_lossy().into_owned())
}
#[derive(Clone)]
pub(crate) enum TargetSource {
Direct,
Config(Option<PathBuf>),
}
#[derive(Clone)]
pub struct SshTarget {
pub(crate) host: String,
pub(crate) source: TargetSource,
pub(crate) port: Option<u16>,
pub(crate) user: Option<String>,
pub(crate) auth: Vec<SshAuth>,
pub(crate) known_hosts: Vec<PathBuf>,
pub(crate) pinned: Vec<String>,
pub(crate) connect_timeout: Duration,
pub(crate) keepalive_interval: Duration,
pub(crate) keepalive_max: u32,
pub(crate) command_timeout: Duration,
pub(crate) max_output_bytes: u64,
pub(crate) sftp_timeout: Duration,
pub(crate) max_transfer_bytes: u64,
pub(crate) max_channels: usize,
pub(crate) allow_terrapin_vulnerable: bool,
pub(crate) allow_in_release: bool,
}
impl fmt::Debug for SshTarget {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("SshTarget")
.field("host", &self.host)
.field("from_ssh_config", &matches!(self.source, TargetSource::Config(_)))
.field("port", &self.port)
.field("user", &self.user)
.field("auth", &self.auth)
.field("known_hosts_files", &self.known_hosts.len())
.field("pinned_fingerprints", &self.pinned.len())
.field("connect_timeout", &self.connect_timeout)
.field("keepalive", &(self.keepalive_interval, self.keepalive_max))
.field("command_timeout", &self.command_timeout)
.field("max_output_bytes", &self.max_output_bytes)
.field("sftp_timeout", &self.sftp_timeout)
.field("max_transfer_bytes", &self.max_transfer_bytes)
.field("max_channels", &self.max_channels)
.field("allow_terrapin_vulnerable", &self.allow_terrapin_vulnerable)
.field("allow_in_release", &self.allow_in_release)
.finish()
}
}
impl SshTarget {
pub fn new(host: impl Into<String>, user: impl Into<String>) -> Self {
Self::blank(host.into(), TargetSource::Direct, Some(user.into()))
}
pub fn from_ssh_config(alias: impl Into<String>) -> Self {
Self::blank(alias.into(), TargetSource::Config(None), None)
}
pub fn from_ssh_config_file(path: impl Into<PathBuf>, alias: impl Into<String>) -> Self {
Self::blank(alias.into(), TargetSource::Config(Some(path.into())), None)
}
fn blank(host: String, source: TargetSource, user: Option<String>) -> Self {
Self {
host,
source,
port: None,
user,
auth: Vec::new(),
known_hosts: Vec::new(),
pinned: Vec::new(),
connect_timeout: DEFAULT_SSH_CONNECT_TIMEOUT,
keepalive_interval: Duration::from_secs(15),
keepalive_max: 3,
command_timeout: DEFAULT_SSH_COMMAND_TIMEOUT,
max_output_bytes: DEFAULT_SSH_MAX_OUTPUT_BYTES,
sftp_timeout: DEFAULT_SFTP_TIMEOUT,
max_transfer_bytes: DEFAULT_SFTP_MAX_BYTES,
max_channels: 8,
allow_terrapin_vulnerable: false,
allow_in_release: false,
}
}
pub fn with_port(mut self, port: u16) -> Self {
self.port = Some(port);
self
}
pub fn with_user(mut self, user: impl Into<String>) -> Self {
self.user = Some(user.into());
self
}
pub fn with_auth(mut self, auth: SshAuth) -> Self {
self.auth.push(auth);
self
}
pub fn with_known_hosts_file(mut self, path: impl Into<PathBuf>) -> Self {
self.known_hosts.push(path.into());
self
}
pub fn trust_host_key_fingerprint(mut self, fingerprint: impl Into<String>) -> Self {
self.pinned.push(fingerprint.into().trim().to_string());
self
}
pub fn with_connect_timeout(mut self, timeout: Duration) -> Self {
self.connect_timeout = timeout.clamp(Duration::from_secs(1), MAX_TIMEOUT);
self
}
pub fn with_keepalive(mut self, interval: Duration, max_missed: u32) -> Self {
self.keepalive_interval = interval.clamp(Duration::from_secs(1), MAX_TIMEOUT);
self.keepalive_max = max_missed.clamp(1, 100);
self
}
pub fn with_command_timeout(mut self, timeout: Duration) -> Self {
self.command_timeout = timeout.clamp(Duration::from_millis(1), MAX_TIMEOUT);
self
}
pub fn with_max_output_bytes(mut self, bytes: u64) -> Self {
self.max_output_bytes = bytes.max(1024);
self
}
pub fn with_sftp_timeout(mut self, timeout: Duration) -> Self {
self.sftp_timeout = timeout.clamp(Duration::from_millis(1), MAX_TIMEOUT);
self
}
pub fn with_max_transfer_bytes(mut self, bytes: u64) -> Self {
self.max_transfer_bytes = bytes.max(1024);
self
}
pub fn with_max_channels(mut self, channels: usize) -> Self {
self.max_channels = channels.clamp(1, 64);
self
}
pub fn allow_terrapin_vulnerable(mut self, allow: bool) -> Self {
self.allow_terrapin_vulnerable = allow;
self
}
pub fn allow_in_release(mut self, allow: bool) -> Self {
self.allow_in_release = allow;
self
}
pub fn host(&self) -> &str {
&self.host
}
pub fn port(&self) -> Option<u16> {
self.port
}
pub fn user(&self) -> Option<&str> {
self.user.as_deref()
}
pub fn is_allowed(&self) -> bool {
allowed(cfg!(debug_assertions), self.allow_in_release)
}
pub fn validate(&self) -> Result<(), Error> {
check_host(&self.host)?;
if let Some(user) = &self.user {
check_user(user)?;
}
if self.port == Some(0) {
return Err(Error::invalid("the SSH port must not be 0"));
}
for pin in &self.pinned {
check_fingerprint(pin)?;
}
match self.source {
TargetSource::Direct if self.user.is_none() => Err(Error::invalid("no SSH user given")),
TargetSource::Direct if self.auth.is_empty() => Err(Error::invalid("no SSH authentication method given (SshAuth::key_file or SshAuth::agent)")),
_ => Ok(()),
}
}
}
pub(crate) fn allowed(debug_build: bool, allow_in_release: bool) -> bool {
debug_build || allow_in_release
}
pub(crate) fn check_host(host: &str) -> Result<(), Error> {
if host.is_empty() || host.len() > 255 || host.starts_with('-') || host.chars().any(|c| c.is_whitespace() || c.is_control()) {
return Err(Error::invalid("the SSH host is not a valid host name or address"));
}
Ok(())
}
pub(crate) fn check_user(user: &str) -> Result<(), Error> {
if user.is_empty() || user.len() > 255 || user.chars().any(|c| c.is_whitespace() || c.is_control()) {
return Err(Error::invalid("the SSH user name is empty or has whitespace / control characters"));
}
Ok(())
}
pub(crate) fn check_fingerprint(pin: &str) -> Result<(), Error> {
let ok = pin.strip_prefix("SHA256:").is_some_and(|b64| b64.len() == 43 && b64.bytes().all(|b| b.is_ascii_alphanumeric() || b == b'+' || b == b'/'));
if ok {
Ok(())
} else {
Err(Error::invalid("a pinned host key fingerprint must look like `SHA256:` + 43 base64 characters (as `ssh-keygen -lf` prints it)"))
}
}
#[derive(Clone)]
pub struct SshCommand {
pub(crate) command: String,
pub(crate) timeout: Option<Duration>,
pub(crate) max_output_bytes: Option<u64>,
pub(crate) stdin: Option<Vec<u8>>,
}
impl fmt::Debug for SshCommand {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("SshCommand")
.field("command_bytes", &self.command.len())
.field("timeout", &self.timeout)
.field("max_output_bytes", &self.max_output_bytes)
.field("stdin_bytes", &self.stdin.as_ref().map(Vec::len))
.finish()
}
}
impl SshCommand {
pub fn new(command: impl Into<String>) -> Self {
Self { command: command.into(), timeout: None, max_output_bytes: None, stdin: None }
}
pub fn with_timeout(mut self, timeout: Duration) -> Self {
self.timeout = Some(timeout.clamp(Duration::from_millis(1), MAX_TIMEOUT));
self
}
pub fn with_max_output_bytes(mut self, bytes: u64) -> Self {
self.max_output_bytes = Some(bytes.max(1024));
self
}
pub fn with_stdin(mut self, stdin: impl Into<Vec<u8>>) -> Self {
self.stdin = Some(stdin.into());
self
}
pub fn command(&self) -> &str {
&self.command
}
pub(crate) fn validate(&self) -> Result<(), Error> {
if self.command.trim().is_empty() {
return Err(Error::invalid("the SSH command is empty"));
}
if self.command.contains('\0') {
return Err(Error::invalid("the SSH command contains a NUL byte"));
}
if self.command.len() > MAX_SSH_COMMAND_BYTES {
return Err(Error::RequestTooLarge { limit: MAX_SSH_COMMAND_BYTES as u64, size: self.command.len() as u64 });
}
Ok(())
}
}
impl From<&str> for SshCommand {
fn from(command: &str) -> Self {
Self::new(command)
}
}
impl From<String> for SshCommand {
fn from(command: String) -> Self {
Self::new(command)
}
}
#[derive(Clone, Debug, PartialEq, Eq)]
#[non_exhaustive]
pub struct SshExit {
pub status: Option<u32>,
pub signal: Option<String>,
pub stdout_bytes: u64,
pub stderr_bytes: u64,
}
impl SshExit {
pub fn success(&self) -> bool {
self.status == Some(0)
}
}
#[derive(Clone, Copy, Debug, PartialEq, Eq, Hash)]
#[non_exhaustive]
pub enum SshStream {
Stdout,
Stderr,
}
#[derive(Clone, PartialEq, Eq)]
#[non_exhaustive]
pub struct SshOutput {
pub exit: SshExit,
pub stdout: Vec<u8>,
pub stderr: Vec<u8>,
}
impl fmt::Debug for SshOutput {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("SshOutput").field("exit", &self.exit).field("stdout_bytes", &self.stdout.len()).field("stderr_bytes", &self.stderr.len()).finish()
}
}
impl SshOutput {
pub fn stdout_text(&self) -> String {
String::from_utf8_lossy(&self.stdout).into_owned()
}
pub fn stderr_text(&self) -> String {
String::from_utf8_lossy(&self.stderr).into_owned()
}
pub fn success(&self) -> bool {
self.exit.success()
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn prompt_answers_match_words_and_give_up_on_unknown_prompts() {
let answers = SshPromptAnswers::new().answer_containing("Password", "fake-pw-1").answer_containing("code", "123456");
let request = SshPromptRequest::new(vec![SshPrompt::new("Password: ", false), SshPrompt::new("Verification code: ", true)]);
let got: Option<Vec<String>> = answers.respond(&request).map(|a| a.iter().map(|s| s.expose().to_string()).collect());
assert_eq!(got, Some(vec!["fake-pw-1".to_string(), "123456".to_string()]));
assert!(answers.respond(&SshPromptRequest::new(vec![SshPrompt::new("Favourite colour?", true)])).is_none());
let debug = format!("{answers:?} {:?} {:?}", SshAuth::password("fake-pw-1"), SshAuth::keyboard_interactive(answers.clone()));
assert!(!debug.contains("fake-pw-1") && !debug.contains("123456"), "{debug}");
}
#[test]
fn the_release_guard() {
assert!(allowed(true, false), "debug builds always");
assert!(!allowed(false, false), "release builds refuse by default");
assert!(allowed(false, true), "release builds with allow_in_release");
assert_eq!(SshTarget::new("h", "u").is_allowed(), cfg!(debug_assertions));
assert!(SshTarget::new("h", "u").allow_in_release(true).is_allowed());
}
#[test]
fn targets_are_validated() {
let ok = SshTarget::new("host.example.com", "deploy").with_auth(SshAuth::agent());
assert!(ok.validate().is_ok());
assert!(ok.clone().with_port(0).validate().is_err());
assert!(SshTarget::new("host.example.com", "deploy").validate().is_err(), "no auth");
assert!(SshTarget::new("", "deploy").with_auth(SshAuth::agent()).validate().is_err());
assert!(SshTarget::new("-oProxyCommand=x", "deploy").with_auth(SshAuth::agent()).validate().is_err());
assert!(SshTarget::new("host", "de ploy").with_auth(SshAuth::agent()).validate().is_err());
let pin = format!("SHA256:{}", "A".repeat(43));
assert!(ok.clone().trust_host_key_fingerprint(pin).validate().is_ok());
assert!(ok.clone().trust_host_key_fingerprint("SHA256:short").validate().is_err());
assert!(SshTarget::from_ssh_config("build").validate().is_ok());
}
#[test]
fn debug_output_never_shows_secrets_or_full_paths() {
let target = SshTarget::new("host", "deploy")
.with_auth(SshAuth::key_file_with_passphrase("/home/someone/.ssh/id_ed25519", "fake-pass-1234"))
.with_known_hosts_file("/home/someone/.ssh/known_hosts");
let debug = format!("{target:?} {:?}", SshCommand::new("echo fake-secret-arg"));
assert!(!debug.contains("fake-pass") && !debug.contains("someone") && !debug.contains("fake-secret"), "{debug}");
assert!(debug.contains("id_ed25519"));
}
#[test]
fn commands_are_checked() {
assert!(SshCommand::new(" ").validate().is_err());
assert!(SshCommand::new("a\0b").validate().is_err());
assert!(matches!(SshCommand::new("x".repeat(MAX_SSH_COMMAND_BYTES + 1)).validate(), Err(Error::RequestTooLarge { .. })));
let command = SshCommand::new("x").with_timeout(Duration::MAX).with_max_output_bytes(1);
assert_eq!((command.timeout, command.max_output_bytes), (Some(MAX_TIMEOUT), Some(1024)));
}
}