use std::io::Read;
use std::path::{Path, PathBuf};
use std::time::Duration;
use ssh2_config::{ParseRule, SshConfig};
use super::{check_host, check_user, file_name, SshTarget, TargetSource};
use crate::Error;
const MAX_CONFIG_BYTES: u64 = 1024 * 1024;
const MAX_INCLUDE_DEPTH: usize = 16;
const MAX_CONFIG_FILES: usize = 64;
pub(crate) struct Resolved {
pub(crate) host: String,
pub(crate) port: u16,
pub(crate) user: String,
pub(crate) identity_files: Vec<PathBuf>,
pub(crate) connect_timeout: Option<Duration>,
}
pub(crate) fn home() -> Option<PathBuf> {
std::env::home_dir().filter(|p| !p.as_os_str().is_empty())
}
pub(crate) fn read_limited(path: &Path, limit: u64, what: &str) -> Result<String, Error> {
let file = std::fs::File::open(path).map_err(|e| Error::InvalidRequest(format!("{what} `{}` cannot be read: {e}", file_name(path))))?;
let mut text = String::new();
file.take(limit.saturating_add(1))
.read_to_string(&mut text)
.map_err(|e| Error::InvalidRequest(format!("{what} `{}` cannot be read: {e}", file_name(path))))?;
if u64::try_from(text.len()).unwrap_or(u64::MAX) > limit {
return Err(Error::InvalidRequest(format!("{what} `{}` is larger than {limit} bytes", file_name(path))));
}
Ok(text)
}
#[derive(Default)]
struct Budget {
bytes: u64,
files: usize,
}
fn expand(path: &Path, depth: usize, budget: &mut Budget, out: &mut String) -> Result<(), Error> {
if depth > MAX_INCLUDE_DEPTH {
return Err(Error::InvalidRequest(format!("ssh_config `Include` nested deeper than {MAX_INCLUDE_DEPTH} (an include loop?)")));
}
budget.files = budget.files.saturating_add(1);
if budget.files > MAX_CONFIG_FILES {
return Err(Error::InvalidRequest(format!("the ssh_config includes more than {MAX_CONFIG_FILES} files")));
}
let remaining = MAX_CONFIG_BYTES.saturating_sub(budget.bytes);
let text = read_limited(path, remaining, "the ssh_config file").map_err(|e| match e {
Error::InvalidRequest(why) if why.contains("larger than") => {
Error::InvalidRequest(format!("the ssh_config files together are larger than {MAX_CONFIG_BYTES} bytes"))
}
other => other,
})?;
budget.bytes = budget.bytes.saturating_add(u64::try_from(text.len()).unwrap_or(u64::MAX));
for line in text.lines() {
match include_args(line) {
None => {
out.push_str(line);
out.push('\n');
}
Some(patterns) => {
for pattern in patterns {
for included in include_paths(&pattern)? {
expand(&included, depth.saturating_add(1), budget, out)?;
}
}
}
}
}
Ok(())
}
fn include_args(line: &str) -> Option<Vec<String>> {
let mut stripped = String::new();
let mut in_quotes = false;
for c in line.trim().chars() {
match c {
'"' => {
in_quotes = !in_quotes;
stripped.push(c);
}
'#' if !in_quotes => break,
_ => stripped.push(c),
}
}
let trimmed = stripped.trim();
let split = trimmed.find(|c: char| c == '=' || c.is_whitespace()).unwrap_or(trimmed.len());
let (field, rest) = trimmed.split_at(split);
if field.to_lowercase() != "include" {
return None;
}
let rest = rest.trim().trim_start_matches('=').trim();
let mut args = Vec::new();
let mut current = String::new();
let mut quoted = false;
for c in rest.chars() {
match c {
'"' => quoted = !quoted,
c if c.is_whitespace() && !quoted => {
if !current.is_empty() {
args.push(std::mem::take(&mut current));
}
}
c => current.push(c),
}
}
if !current.is_empty() {
args.push(current);
}
Some(args)
}
fn include_paths(pattern: &str) -> Result<Vec<PathBuf>, Error> {
let expanded = if let Some(rest) = pattern.strip_prefix("~/") {
home().ok_or_else(|| Error::InvalidRequest("no home directory for an ssh_config `Include`".into()))?.join(rest)
} else if Path::new(pattern).is_absolute() {
PathBuf::from(pattern)
} else {
home().ok_or_else(|| Error::InvalidRequest("no home directory for an ssh_config `Include`".into()))?.join(".ssh").join(pattern)
};
let text = expanded.to_string_lossy().into_owned();
let paths = glob::glob(&text).map_err(|e| Error::InvalidRequest(format!("a bad ssh_config `Include` pattern ({e})")))?;
let mut found: Vec<PathBuf> = paths.filter_map(Result::ok).filter(|p| p.is_file()).collect();
found.sort();
if found.len() > MAX_CONFIG_FILES {
return Err(Error::InvalidRequest(format!("an ssh_config `Include` matches more than {MAX_CONFIG_FILES} files")));
}
Ok(found)
}
pub(crate) fn resolve(target: &SshTarget) -> Result<Resolved, Error> {
let resolved = match &target.source {
TargetSource::Direct => Resolved {
host: target.host.clone(),
port: target.port.unwrap_or(22),
user: target.user.clone().unwrap_or_default(),
identity_files: Vec::new(),
connect_timeout: None,
},
TargetSource::Config(path) => {
let path = match path {
Some(path) => path.clone(),
None => home()
.map(|home| home.join(".ssh").join("config"))
.ok_or_else(|| Error::InvalidRequest("no home directory to find ~/.ssh/config in".into()))?,
};
let mut budget = Budget::default();
let mut text = String::new();
expand(&path, 0, &mut budget, &mut text)?;
let config = SshConfig::default()
.parse(&mut text.as_bytes(), ParseRule::ALLOW_UNKNOWN_FIELDS | ParseRule::ALLOW_UNSUPPORTED_FIELDS)
.map_err(|e| Error::InvalidRequest(format!("the ssh_config file `{}` could not be parsed: {e}", file_name(&path))))?;
let params = config.query(&target.host);
let user = target.user.clone().or(params.user).ok_or_else(|| {
Error::InvalidRequest(format!("no user for `{}`: set `User` in the ssh_config file or call SshTarget::with_user", target.host))
})?;
Resolved {
host: params.host_name.unwrap_or_else(|| target.host.clone()),
port: target.port.or(params.port).unwrap_or(22),
user,
identity_files: params.identity_file.unwrap_or_default(),
connect_timeout: params.connect_timeout,
}
}
};
check_host(&resolved.host)?;
check_user(&resolved.user)?;
if resolved.port == 0 {
return Err(Error::InvalidRequest("the SSH port must not be 0".into()));
}
Ok(resolved)
}
#[cfg(test)]
mod tests {
use std::path::PathBuf;
use std::time::Duration;
use super::resolve;
use crate::ssh::SshTarget;
fn config_file(name: &str, text: &str) -> PathBuf {
let dir = PathBuf::from(env!("CARGO_MANIFEST_DIR")).join("..").join("..").join("target").join("tmp").join("client-ssh-config-tests");
std::fs::create_dir_all(&dir).unwrap_or_else(|e| panic!("{e}"));
let path = dir.join(name);
std::fs::write(&path, text).unwrap_or_else(|e| panic!("{e}"));
path
}
#[test]
fn an_alias_takes_host_port_user_and_identity_files_from_the_file() {
let path = config_file(
"config-a",
"Host build\n HostName build.example.com\n Port 2222\n User deploy\n IdentityFile /keys/id_test\n ConnectTimeout 5\n\nHost *\n User nobody\n",
);
let resolved = resolve(&SshTarget::from_ssh_config_file(&path, "build")).unwrap_or_else(|e| panic!("{e}"));
assert_eq!((resolved.host.as_str(), resolved.port, resolved.user.as_str()), ("build.example.com", 2222, "deploy"));
assert_eq!(resolved.identity_files, vec![PathBuf::from("/keys/id_test")]);
assert_eq!(resolved.connect_timeout, Some(Duration::from_secs(5)));
let resolved = resolve(&SshTarget::from_ssh_config_file(&path, "build").with_port(22).with_user("admin")).unwrap_or_else(|e| panic!("{e}"));
assert_eq!((resolved.port, resolved.user.as_str()), (22, "admin"));
let resolved = resolve(&SshTarget::from_ssh_config_file(&path, "other.example.com")).unwrap_or_else(|e| panic!("{e}"));
assert_eq!((resolved.host.as_str(), resolved.user.as_str()), ("other.example.com", "nobody"));
}
#[test]
fn includes_are_expanded_with_limits_and_loops_are_errors() {
let dir = PathBuf::from(env!("CARGO_MANIFEST_DIR")).join("..").join("..").join("target").join("tmp").join("client-ssh-config-tests");
let fwd = |p: &PathBuf| p.to_string_lossy().replace('\\', "/");
let inc = config_file("inc-hosts.conf", "Host build\n HostName included.example.com\n User deploy\n");
let main = config_file("config-inc", &format!("INCLUDE=\"{}\" # comment\nHost *\n Port 2200\n", fwd(&dir.join("inc-host*.conf"))));
let resolved = resolve(&SshTarget::from_ssh_config_file(&main, "build")).unwrap_or_else(|e| panic!("{e}"));
assert_eq!((resolved.host.as_str(), resolved.user.as_str(), resolved.port), ("included.example.com", "deploy", 2200));
let _ = inc;
let looped = dir.join("config-loop");
std::fs::write(&looped, format!("Include {}\nHost x\n User u\n", fwd(&looped))).unwrap_or_else(|e| panic!("{e}"));
let error = resolve(&SshTarget::from_ssh_config_file(&looped, "x")).err().map(|e| e.to_string()).unwrap_or_default();
assert!(error.contains("deeper than 16"), "{error}");
let (a, b) = (dir.join("config-a-b"), dir.join("config-b-a"));
std::fs::write(&a, format!("include {}\n", fwd(&b))).unwrap_or_else(|e| panic!("{e}"));
std::fs::write(&b, format!(" Include\t{}\n", fwd(&a))).unwrap_or_else(|e| panic!("{e}"));
assert!(resolve(&SshTarget::from_ssh_config_file(&a, "x")).is_err());
let big = config_file("config-big-part", &"# filler line\n".repeat(50_000));
let top = config_file("config-big", &format!("Include {}\nInclude {}\nHost x\n User u\n", fwd(&big), fwd(&big)));
let error = resolve(&SshTarget::from_ssh_config_file(&top, "x")).err().map(|e| e.to_string()).unwrap_or_default();
assert!(error.contains("together are larger"), "{error}");
let none = config_file("config-none", &format!("Include {}\nHost x\n User u\n", fwd(&dir.join("no-such-*.conf"))));
assert!(resolve(&SshTarget::from_ssh_config_file(&none, "x")).is_ok());
}
#[test]
fn bad_config_files_are_errors_not_panics() {
for (name, text) in
[("config-b", "Host\n"), ("config-c", "Port notanumber\n"), ("config-d", "\u{0}\u{1}\u{2}"), ("config-e", "Host x\n HostName bad host\n")]
{
let path = config_file(name, text);
let _ = resolve(&SshTarget::from_ssh_config_file(&path, "x"));
}
let missing = PathBuf::from(env!("CARGO_MANIFEST_DIR")).join("..").join("..").join("target").join("tmp").join("no-such-ssh-config");
assert!(resolve(&SshTarget::from_ssh_config_file(missing, "x")).is_err());
}
}