use clap::Parser;
use std::ffi::OsStr;
use std::path::{Path, PathBuf};
use crate::error::{Result, SshMcpError};
use crate::ssh::HostKeyCheckMode;
pub const DEFAULT_TIMEOUT_MS: u64 = 300_000;
pub const DEFAULT_MAX_CHARS: Option<usize> = Some(64_000);
pub const CONNECTION_TIMEOUT_SECS: u64 = 30;
pub const DEFAULT_RECONNECT_RETRIES: u64 = 3;
pub const DEFAULT_RECONNECT_BACKOFF_MS: u64 = 250;
pub const DEFAULT_HEALTH_PROBE_TIMEOUT_MS: u64 = 1500;
pub const MAX_RECONNECT_RETRIES: u64 = 10;
pub const MIN_RECONNECT_BACKOFF_MS: u64 = 10;
pub const MAX_RECONNECT_BACKOFF_MS: u64 = 30_000;
pub const MIN_HEALTH_PROBE_TIMEOUT_MS: u64 = 100;
pub const MAX_HEALTH_PROBE_TIMEOUT_MS: u64 = 30_000;
#[derive(Parser, Debug, Clone)]
#[command(name = "ssh-mcp")]
#[command(author = "0FL01")]
#[command(version = env!("CARGO_PKG_VERSION"))]
#[command(about = env!("CARGO_PKG_DESCRIPTION"))]
pub struct Args {
#[arg(long, env = "SSH_MCP_HOST")]
pub host: String,
#[arg(long, default_value = "22", env = "SSH_MCP_PORT")]
pub port: u16,
#[arg(long, env = "SSH_MCP_USER")]
pub user: String,
#[arg(long, env = "SSH_MCP_PASSWORD")]
pub password: Option<String>,
#[arg(long, env = "SSH_MCP_KEY")]
pub key: Option<PathBuf>,
#[arg(long, env = "SSH_MCP_SPOOL_DIR")]
pub spool_dir: Option<PathBuf>,
#[arg(long, env = "SSH_MCP_SU_PASSWORD")]
pub su_password: Option<String>,
#[arg(long, env = "SSH_MCP_SUDO_PASSWORD")]
pub sudo_password: Option<String>,
#[arg(long, default_value = "300000", env = "SSH_MCP_TIMEOUT")]
pub timeout: u64,
#[arg(long = "maxChars", env = "SSH_MCP_MAX_CHARS")]
pub max_chars: Option<String>,
#[arg(long, default_value = "false", env = "SSH_MCP_DISABLE_SUDO")]
pub disable_sudo: bool,
#[arg(long = "max-output-tokens", env = "SSH_MCP_MAX_OUTPUT_TOKENS")]
pub max_output_tokens: Option<String>,
#[arg(long, default_value = "info", env = "SSH_MCP_LOG_LEVEL", value_parser = clap::builder::PossibleValuesParser::new(["trace", "debug", "info", "warn", "error"]))]
pub log_level: String,
#[arg(long, env = "SSH_MCP_LOG_FILE")]
pub log_file: Option<PathBuf>,
#[arg(long, default_value = "text", env = "SSH_MCP_LOG_FORMAT", value_parser = clap::builder::PossibleValuesParser::new(["text", "json"]))]
pub log_format: String,
#[arg(long, default_value = "daily", env = "SSH_MCP_LOG_ROTATION", value_parser = clap::builder::PossibleValuesParser::new(["daily", "hourly", "never"]))]
pub log_rotation: String,
#[arg(long, default_value = "30", env = "SSH_MCP_KEEPALIVE_INTERVAL")]
pub keepalive_interval: u64,
#[arg(long, default_value = "3", env = "SSH_MCP_KEEPALIVE_MAX")]
pub keepalive_max: u64,
#[arg(long, default_value = "3", env = "SSH_MCP_RECONNECT_RETRIES")]
pub reconnect_retries: u64,
#[arg(long, default_value = "250", env = "SSH_MCP_RECONNECT_BACKOFF_MS")]
pub reconnect_backoff_ms: u64,
#[arg(long, default_value = "1500", env = "SSH_MCP_HEALTH_PROBE_TIMEOUT_MS")]
pub health_probe_timeout_ms: u64,
#[arg(
long = "strict-host-key-checking",
env = "SSH_MCP_STRICT_HOST_KEY_CHECKING",
value_enum,
default_value_t = HostKeyCheckMode::AcceptNew
)]
pub strict_host_key_checking: HostKeyCheckMode,
#[arg(long = "known-hosts", env = "SSH_MCP_KNOWN_HOSTS")]
pub known_hosts: Option<PathBuf>,
}
#[derive(Debug, Clone)]
pub struct Config {
pub host: String,
pub port: u16,
pub user: String,
pub password: Option<String>,
pub key: Option<PathBuf>,
pub su_password: Option<String>,
pub sudo_password: Option<String>,
pub timeout_ms: u64,
pub max_chars: Option<usize>,
pub max_output_tokens: Option<usize>,
pub disable_sudo: bool,
pub keepalive_interval: u64,
pub keepalive_max: u64,
pub reconnect_retries: u64,
pub reconnect_backoff_ms: u64,
pub health_probe_timeout_ms: u64,
pub strict_host_key_checking: HostKeyCheckMode,
pub known_hosts: Option<PathBuf>,
}
impl Config {
pub fn from_args(args: Args) -> Result<Self> {
let home = std::env::var_os("HOME");
Self::from_args_with_home(args, home.as_deref())
}
fn from_args_with_home(mut args: Args, home: Option<&OsStr>) -> Result<Self> {
args.key = args
.key
.map(|path| expand_key_path(path, home))
.transpose()?;
validate_args(&args)?;
let max_chars = parse_max_chars(args.max_chars.as_deref());
let max_output_tokens = parse_max_output_tokens(args.max_output_tokens.as_deref());
Ok(Config {
host: args.host,
port: args.port,
user: args.user,
password: sanitize_password(args.password),
key: args.key,
su_password: sanitize_password(args.su_password),
sudo_password: sanitize_password(args.sudo_password),
timeout_ms: args.timeout,
max_chars,
max_output_tokens,
disable_sudo: args.disable_sudo,
keepalive_interval: args.keepalive_interval,
keepalive_max: args.keepalive_max,
reconnect_retries: args.reconnect_retries,
reconnect_backoff_ms: args.reconnect_backoff_ms,
health_probe_timeout_ms: args.health_probe_timeout_ms,
strict_host_key_checking: args.strict_host_key_checking,
known_hosts: args.known_hosts,
})
}
}
fn expand_key_path(path: PathBuf, home: Option<&OsStr>) -> Result<PathBuf> {
if !path.as_os_str().as_encoded_bytes().starts_with(b"~/") {
return Ok(path);
}
let home = home.filter(|value| !value.is_empty()).ok_or_else(|| {
SshMcpError::Config(format!(
"Cannot expand SSH key path {}: HOME is not set",
path.display()
))
})?;
let suffix = path
.strip_prefix("~")
.expect("leading ~/ path must have a tilde component");
Ok(Path::new(home).join(suffix))
}
fn validate_args(args: &Args) -> Result<()> {
let mut errors = Vec::new();
if args.host.is_empty() {
errors.push("Missing required --host".to_string());
}
if args.user.is_empty() {
errors.push("Missing required --user".to_string());
}
if args.password.is_none() && args.key.is_none() {
errors.push("Must provide either --password or --key".to_string());
}
if let Some(ref key_path) = args.key
&& !key_path.exists()
{
errors.push(format!("SSH key file not found: {}", key_path.display()));
}
if args.reconnect_retries > MAX_RECONNECT_RETRIES {
errors.push(format!(
"--reconnect-retries must be <= {MAX_RECONNECT_RETRIES}"
));
}
if !(MIN_RECONNECT_BACKOFF_MS..=MAX_RECONNECT_BACKOFF_MS).contains(&args.reconnect_backoff_ms) {
errors.push(format!(
"--reconnect-backoff-ms must be between {MIN_RECONNECT_BACKOFF_MS} and {MAX_RECONNECT_BACKOFF_MS}"
));
}
if !(MIN_HEALTH_PROBE_TIMEOUT_MS..=MAX_HEALTH_PROBE_TIMEOUT_MS)
.contains(&args.health_probe_timeout_ms)
{
errors.push(format!(
"--health-probe-timeout-ms must be between {MIN_HEALTH_PROBE_TIMEOUT_MS} and {MAX_HEALTH_PROBE_TIMEOUT_MS}"
));
}
if !errors.is_empty() {
return Err(SshMcpError::Config(format!(
"Configuration error:\n{}",
errors.join("\n")
)));
}
Ok(())
}
pub const DEFAULT_MAX_OUTPUT_TOKENS: Option<usize> = Some(16_000);
pub fn parse_max_chars(value: Option<&str>) -> Option<usize> {
match value {
None => DEFAULT_MAX_CHARS,
Some(s) => {
let lowered = s.to_lowercase();
if lowered == "none" {
return None;
}
match s.parse::<i64>() {
Ok(n) if n <= 0 => None,
Ok(n) => Some(n as usize),
Err(_) => DEFAULT_MAX_CHARS,
}
}
}
}
pub fn parse_max_output_tokens(value: Option<&str>) -> Option<usize> {
match value {
None => DEFAULT_MAX_OUTPUT_TOKENS,
Some(s) => {
let lowered = s.to_lowercase().replace(" ", "");
if lowered == "none" {
return None;
}
if lowered.ends_with('k') {
let num_part = &lowered[..lowered.len() - 1];
match num_part.parse::<i64>() {
Ok(n) if n <= 0 => None,
Ok(n) => Some((n as usize).saturating_mul(1_000)),
Err(_) => DEFAULT_MAX_OUTPUT_TOKENS,
}
} else {
match lowered.parse::<i64>() {
Ok(n) if n <= 0 => None,
Ok(n) => Some(n as usize),
Err(_) => DEFAULT_MAX_OUTPUT_TOKENS,
}
}
}
}
}
fn sanitize_password(password: Option<String>) -> Option<String> {
password.filter(|p| !p.is_empty())
}
#[cfg(test)]
mod tests {
use super::*;
fn base_args() -> Args {
Args {
host: "localhost".to_string(),
port: 22,
user: "test".to_string(),
password: Some("secret".to_string()),
key: None,
spool_dir: None,
su_password: None,
sudo_password: None,
timeout: DEFAULT_TIMEOUT_MS,
max_chars: None,
disable_sudo: false,
max_output_tokens: None,
log_level: "info".to_string(),
log_file: None,
log_format: "text".to_string(),
log_rotation: "daily".to_string(),
keepalive_interval: 30,
keepalive_max: 3,
reconnect_retries: DEFAULT_RECONNECT_RETRIES,
reconnect_backoff_ms: DEFAULT_RECONNECT_BACKOFF_MS,
health_probe_timeout_ms: DEFAULT_HEALTH_PROBE_TIMEOUT_MS,
strict_host_key_checking: HostKeyCheckMode::AcceptNew,
known_hosts: None,
}
}
#[test]
fn test_parse_max_chars_none_string() {
assert_eq!(parse_max_chars(Some("none")), None);
assert_eq!(parse_max_chars(Some("None")), None);
assert_eq!(parse_max_chars(Some("NONE")), None);
}
#[test]
fn test_parse_max_chars_zero_or_negative() {
assert_eq!(parse_max_chars(Some("0")), None);
assert_eq!(parse_max_chars(Some("-1")), None);
assert_eq!(parse_max_chars(Some("-100")), None);
}
#[test]
fn test_parse_max_chars_positive() {
assert_eq!(parse_max_chars(Some("500")), Some(500));
assert_eq!(parse_max_chars(Some("2000")), Some(2000));
}
#[test]
fn test_parse_max_chars_invalid() {
assert_eq!(parse_max_chars(Some("abc")), DEFAULT_MAX_CHARS);
assert_eq!(parse_max_chars(Some("")), DEFAULT_MAX_CHARS);
}
#[test]
fn test_parse_max_chars_not_provided() {
assert_eq!(parse_max_chars(None), DEFAULT_MAX_CHARS);
}
#[test]
fn test_config_from_args_uses_default_max_chars() {
let config = Config::from_args(base_args()).unwrap();
assert_eq!(config.max_chars, Some(64_000));
assert_eq!(config.strict_host_key_checking, HostKeyCheckMode::AcceptNew);
assert!(config.known_hosts.is_none());
}
#[test]
fn test_config_expands_tilde_key_before_validation() {
let home = tempfile::tempdir().unwrap();
let key_path = home.path().join(".ssh/id_ed25519");
std::fs::create_dir_all(key_path.parent().unwrap()).unwrap();
std::fs::write(&key_path, "test key").unwrap();
let mut args = base_args();
args.password = None;
args.key = Some(PathBuf::from("~/.ssh/id_ed25519"));
let config = Config::from_args_with_home(args, Some(home.path().as_os_str())).unwrap();
assert_eq!(config.key, Some(key_path));
}
#[test]
fn test_expand_key_path_only_expands_leading_home_prefix() {
let home = OsStr::new("/home/test");
let cases = [
("~", "~"),
("~user/key", "~user/key"),
("dir/~/key", "dir/~/key"),
("$HOME/key", "$HOME/key"),
(r"~\key", r"~\key"),
("/tmp/key", "/tmp/key"),
];
for (input, expected) in cases {
assert_eq!(
expand_key_path(PathBuf::from(input), Some(home)).unwrap(),
PathBuf::from(expected)
);
}
assert_eq!(
expand_key_path(PathBuf::from("~/.ssh/id_ed25519"), Some(home)).unwrap(),
PathBuf::from("/home/test/.ssh/id_ed25519")
);
}
#[test]
fn test_expand_key_path_requires_home() {
for home in [None, Some(OsStr::new(""))] {
let error = expand_key_path(PathBuf::from("~/.ssh/id_ed25519"), home).unwrap_err();
assert!(error.to_string().contains("HOME is not set"));
}
}
#[test]
fn test_args_parse_host_key_options() {
let args = Args::try_parse_from([
"ssh-mcp",
"--host",
"example.com",
"--user",
"alice",
"--password",
"secret",
"--strict-host-key-checking",
"yes",
"--known-hosts",
"/tmp/known_hosts",
])
.unwrap();
assert_eq!(args.strict_host_key_checking, HostKeyCheckMode::Yes);
assert_eq!(args.known_hosts, Some(PathBuf::from("/tmp/known_hosts")));
}
#[test]
fn test_args_parse_spool_dir() {
let args = Args::try_parse_from([
"ssh-mcp",
"--host",
"example.com",
"--user",
"alice",
"--password",
"secret",
"--spool-dir",
"/tmp/ssh-mcp-alice",
])
.unwrap();
assert_eq!(args.spool_dir, Some(PathBuf::from("/tmp/ssh-mcp-alice")));
}
#[test]
fn test_sanitize_password() {
assert_eq!(
sanitize_password(Some("secret".to_string())),
Some("secret".to_string())
);
assert_eq!(sanitize_password(Some(String::new())), None);
assert_eq!(sanitize_password(None), None);
}
#[test]
fn test_parse_max_output_tokens_none_string() {
assert_eq!(parse_max_output_tokens(Some("none")), None);
assert_eq!(parse_max_output_tokens(Some("None")), None);
assert_eq!(parse_max_output_tokens(Some("NONE")), None);
}
#[test]
fn test_parse_max_output_tokens_zero_or_negative() {
assert_eq!(parse_max_output_tokens(Some("0")), None);
assert_eq!(parse_max_output_tokens(Some("-1")), None);
assert_eq!(parse_max_output_tokens(Some("-100")), None);
}
#[test]
fn test_parse_max_output_tokens_positive() {
assert_eq!(parse_max_output_tokens(Some("500")), Some(500));
assert_eq!(parse_max_output_tokens(Some("12000")), Some(12_000));
}
#[test]
fn test_parse_max_output_tokens_with_k_suffix() {
assert_eq!(parse_max_output_tokens(Some("12k")), Some(12_000));
assert_eq!(parse_max_output_tokens(Some("5K")), Some(5_000));
assert_eq!(parse_max_output_tokens(Some("100k")), Some(100_000));
}
#[test]
fn test_parse_max_output_tokens_invalid() {
assert_eq!(
parse_max_output_tokens(Some("abc")),
DEFAULT_MAX_OUTPUT_TOKENS
);
assert_eq!(parse_max_output_tokens(Some("")), DEFAULT_MAX_OUTPUT_TOKENS);
}
#[test]
fn test_parse_max_output_tokens_not_provided() {
assert_eq!(parse_max_output_tokens(None), DEFAULT_MAX_OUTPUT_TOKENS);
}
#[test]
fn test_validate_args_rejects_reconnect_retries_out_of_range() {
let mut args = base_args();
args.reconnect_retries = MAX_RECONNECT_RETRIES.saturating_add(1);
let result = validate_args(&args);
assert!(result.is_err());
}
#[test]
fn test_validate_args_rejects_reconnect_backoff_out_of_range() {
let mut args = base_args();
args.reconnect_backoff_ms = MIN_RECONNECT_BACKOFF_MS.saturating_sub(1);
let result = validate_args(&args);
assert!(result.is_err());
}
#[test]
fn test_validate_args_rejects_health_probe_timeout_out_of_range() {
let mut args = base_args();
args.health_probe_timeout_ms = MAX_HEALTH_PROBE_TIMEOUT_MS.saturating_add(1);
let result = validate_args(&args);
assert!(result.is_err());
}
}