diff --git a/Cargo.lock b/Cargo.lock index 63ec200..bdfcd9d 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -401,7 +401,7 @@ checksum = "2b15c43186be67a4fd63bee50d0303afffcef381492ebe2c5d87f324e1b8815c" [[package]] name = "rexec" -version = "1.5.0" +version = "1.5.1" dependencies = [ "brace-expand", "clap", diff --git a/Cargo.toml b/Cargo.toml index 4014bc1..50c9544 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -1,6 +1,6 @@ [package] name = "rexec" -version = "1.5.1" +version = "1.5.2" readme = "https://github.com/house-of-vanity/rexec#readme" edition = "2021" description = "Parallel SSH executor" diff --git a/src/main.rs b/src/main.rs index f90680d..69289c3 100644 --- a/src/main.rs +++ b/src/main.rs @@ -2,8 +2,9 @@ extern crate log; use std::fs::read_to_string; use std::hash::Hash; -use std::io::{BufRead, BufReader}; +use std::io::{self, BufRead, BufReader}; use std::net::IpAddr; +use std::path::{Path, PathBuf}; use std::process::{self, Command, Stdio}; use std::sync::{Arc, Mutex}; use std::thread; @@ -157,28 +158,80 @@ fn shorten_hostname(hostname: &str, common_suffix: &Option) -> String { } } -/// Read and parse the SSH known_hosts file to extract server names -/// -/// # Returns -/// * `Vec` - List of hosts found in the known_hosts file -fn read_known_hosts() -> Vec { +/// Resolve the default SSH known_hosts path for the current user. +fn default_known_hosts_path() -> PathBuf { + let home = std::env::var_os("HOME") + .map(PathBuf::from) + .filter(|path| !path.as_os_str().is_empty()) + .unwrap_or_else(|| PathBuf::from(format!("/home/{}", whoami::username()))); + + home.join(".ssh").join("known_hosts") +} + +fn normalize_known_host_entry(entry: &str) -> Option { + if entry.is_empty() || entry.starts_with('|') { + return None; + } + + let hostname = entry + .strip_prefix('[') + .and_then(|value| value.split_once("]:").map(|(host, _)| host)) + .unwrap_or(entry); + + if hostname.is_empty() { + None + } else { + Some(hostname.to_string()) + } +} + +fn parse_known_hosts(content: &str) -> Vec { let mut result: Vec = Vec::new(); - // Read known_hosts file from the user's home directory - for line in read_to_string(format!("/home/{}/.ssh/known_hosts", whoami::username())) - .unwrap() - .lines() - { - let line = line.split(" ").collect::>(); - let hostname = line[0]; - result.push(Host { - name: hostname.to_string(), - ip: None, - }) + for line in content.lines() { + let line = line.trim(); + + if line.is_empty() || line.starts_with('#') { + continue; + } + + let mut fields = line.split_whitespace(); + let first = match fields.next() { + Some(field) => field, + None => continue, + }; + + let hostnames = if first.starts_with('@') { + match fields.next() { + Some(field) => field, + None => continue, + } + } else { + first + }; + + for hostname in hostnames.split(',').filter_map(normalize_known_host_entry) { + result.push(Host { + name: hostname, + ip: None, + }); + } } + result } +/// Read and parse the SSH known_hosts file to extract server names +/// +/// # Arguments +/// * `path` - Path to the known_hosts file +/// +/// # Returns +/// * `io::Result>` - List of hosts found in the known_hosts file +fn read_known_hosts(path: &Path) -> io::Result> { + read_to_string(path).map(|content| parse_known_hosts(&content)) +} + /// Expand a numeric range in the format [start:end] to a list of strings /// /// # Arguments @@ -441,8 +494,19 @@ fn main() { // Build the list of target hosts based on user selection method let hosts = if args.known_hosts { // Use regex pattern matching against known_hosts file - info!("Using ~/.ssh/known_hosts to build server list."); - let known_hosts = read_known_hosts(); + let known_hosts_path = default_known_hosts_path(); + info!("Using {} to build server list.", known_hosts_path.display()); + let known_hosts = match read_known_hosts(&known_hosts_path) { + Ok(hosts) => hosts, + Err(e) => { + error!( + "Failed to read known_hosts file {}: {}", + known_hosts_path.display(), + e + ); + process::exit(1); + } + }; let mut all_hosts = Vec::new(); for expression in args.expression.iter() { let re = match Regex::new(expression) { @@ -624,3 +688,41 @@ fn main() { processed = end; } } + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn parse_known_hosts_extracts_usable_hosts() { + let content = r#" +# comment +admin.example.com ssh-ed25519 AAAA +web.example.com,192.0.2.10 ecdsa-sha2-nistp256 AAAA +@cert-authority *.example.com ssh-rsa AAAA +|1|salt|hash ssh-ed25519 AAAA +[custom.example.com]:2222 ssh-rsa AAAA +"#; + + let hosts: Vec<_> = parse_known_hosts(content) + .iter() + .map(|host| host.name.as_str()) + .collect(); + + assert_eq!( + hosts, + vec![ + "admin.example.com", + "web.example.com", + "192.0.2.10", + "*.example.com", + "custom.example.com" + ] + ); + } + + #[test] + fn normalize_known_host_entry_skips_hashed_entries() { + assert_eq!(normalize_known_host_entry("|1|salt|hash"), None); + } +}