use crate::{Result, QsshError};
use std::net::{IpAddr, SocketAddr};
use std::path::{Path, PathBuf};
pub struct ValidationLimits {
pub max_username_len: usize,
pub max_hostname_len: usize,
pub max_path_len: usize,
pub max_command_len: usize,
pub max_password_len: usize,
pub min_port: u16,
pub max_port: u16,
}
impl Default for ValidationLimits {
fn default() -> Self {
Self {
max_username_len: 256,
max_hostname_len: 253, max_path_len: 4096,
max_command_len: 32768,
max_password_len: 1024,
min_port: 1,
max_port: 65535,
}
}
}
pub struct InputValidator {
limits: ValidationLimits,
}
impl InputValidator {
pub fn new() -> Self {
Self {
limits: ValidationLimits::default(),
}
}
pub fn with_limits(limits: ValidationLimits) -> Self {
Self { limits }
}
pub fn validate_username(&self, username: &str) -> Result<String> {
if username.is_empty() {
return Err(QsshError::Config("Username cannot be empty".into()));
}
if username.len() > self.limits.max_username_len {
return Err(QsshError::Config(format!(
"Username too long (max {} characters)",
self.limits.max_username_len
)));
}
if !username.chars().all(|c| c.is_ascii_alphanumeric() || c == '_' || c == '-' || c == '.') {
return Err(QsshError::Config("Username contains invalid characters".into()));
}
if username == "root" && !cfg!(feature = "allow_root") {
log::warn!("Root login attempted - use 'allow_root' feature to enable");
}
Ok(username.to_string())
}
pub fn validate_hostname(&self, hostname: &str) -> Result<String> {
if hostname.is_empty() {
return Err(QsshError::Config("Hostname cannot be empty".into()));
}
if hostname.len() > self.limits.max_hostname_len {
return Err(QsshError::Config(format!(
"Hostname too long (max {} characters)",
self.limits.max_hostname_len
)));
}
if hostname.parse::<IpAddr>().is_ok() {
return Ok(hostname.to_string());
}
if !hostname.chars().all(|c| c.is_ascii_alphanumeric() || c == '.' || c == '-') {
return Err(QsshError::Config("Hostname contains invalid characters".into()));
}
if hostname.starts_with('.') || hostname.starts_with('-') ||
hostname.ends_with('.') || hostname.ends_with('-') {
return Err(QsshError::Config("Invalid hostname format".into()));
}
if hostname.contains("..") {
return Err(QsshError::Config("Hostname contains consecutive dots".into()));
}
Ok(hostname.to_string())
}
pub fn validate_port(&self, port: u16) -> Result<u16> {
if port < self.limits.min_port || port > self.limits.max_port {
return Err(QsshError::Config(format!(
"Port must be between {} and {}",
self.limits.min_port, self.limits.max_port
)));
}
if port < 1024 && !cfg!(feature = "allow_privileged_ports") {
log::warn!("Using privileged port {} - requires root/admin privileges", port);
}
Ok(port)
}
pub fn validate_socket_addr(&self, addr: &str) -> Result<SocketAddr> {
addr.parse::<SocketAddr>()
.map_err(|e| QsshError::Config(format!("Invalid socket address: {}", e)))
}
pub fn validate_path(&self, path: &str) -> Result<PathBuf> {
if path.is_empty() {
return Err(QsshError::Config("Path cannot be empty".into()));
}
if path.len() > self.limits.max_path_len {
return Err(QsshError::Config(format!(
"Path too long (max {} characters)",
self.limits.max_path_len
)));
}
if path.contains("../") || path.contains("..\\") {
return Err(QsshError::Config("Path contains directory traversal".into()));
}
if path.contains('\0') {
return Err(QsshError::Config("Path contains null byte".into()));
}
let path_buf = PathBuf::from(path);
let path_str = path_buf.to_string_lossy();
if path_str.contains("//") || path_str.contains("\\\\") {
log::warn!("Path contains double slashes: {}", path_str);
}
Ok(path_buf)
}
pub fn validate_command(&self, command: &str) -> Result<String> {
if command.len() > self.limits.max_command_len {
return Err(QsshError::Config(format!(
"Command too long (max {} characters)",
self.limits.max_command_len
)));
}
if command.contains('\0') {
return Err(QsshError::Config("Command contains null byte".into()));
}
let dangerous_patterns = [
"rm -rf",
"dd if=",
"mkfs",
"format",
"> /dev/",
":(){ :|:", ];
for pattern in &dangerous_patterns {
if command.contains(pattern) {
log::warn!("Potentially dangerous command pattern detected: {}", pattern);
}
}
Ok(command.to_string())
}
pub fn validate_password(&self, password: &str) -> Result<()> {
if password.len() > self.limits.max_password_len {
return Err(QsshError::Config(format!(
"Password too long (max {} characters)",
self.limits.max_password_len
)));
}
if password.contains('\0') {
return Err(QsshError::Config("Password contains null byte".into()));
}
Ok(())
}
pub fn validate_port_forward(&self, spec: &str) -> Result<(u16, String, u16)> {
let parts: Vec<&str> = spec.split(':').collect();
if parts.len() != 3 {
return Err(QsshError::Config(
"Port forward must be in format: local_port:remote_host:remote_port".into()
));
}
let local_port = parts[0].parse::<u16>()
.map_err(|_| QsshError::Config("Invalid local port number".into()))?;
self.validate_port(local_port)?;
let remote_host = self.validate_hostname(parts[1])?;
let remote_port = parts[2].parse::<u16>()
.map_err(|_| QsshError::Config("Invalid remote port number".into()))?;
self.validate_port(remote_port)?;
Ok((local_port, remote_host, remote_port))
}
pub fn validate_display_number(&self, display: u32) -> Result<u32> {
if display > 99 {
log::warn!("Unusual X11 display number: {}", display);
}
if display > 59535 { return Err(QsshError::Config("X11 display number too large".into()));
}
Ok(display)
}
pub fn validate_term_type(&self, term: &str) -> Result<String> {
let valid_terms = [
"xterm", "xterm-256color", "xterm-color",
"vt100", "vt220", "vt320",
"screen", "screen-256color",
"tmux", "tmux-256color",
"rxvt", "rxvt-unicode",
"linux", "dumb", "ansi",
];
let term_lower = term.to_lowercase();
let is_known = valid_terms.iter().any(|&t| term_lower.starts_with(t));
if !is_known {
log::warn!("Unknown terminal type: {}", term);
}
if term.contains('\0') || term.contains('\n') || term.contains('\r') {
return Err(QsshError::Config("Invalid characters in terminal type".into()));
}
Ok(term.to_string())
}
pub fn sanitize_env_var(&self, key: &str, value: &str) -> Result<(String, String)> {
if !key.chars().all(|c| c.is_ascii_alphanumeric() || c == '_') {
return Err(QsshError::Config("Invalid environment variable name".into()));
}
if value.contains('\0') {
return Err(QsshError::Config("Environment value contains null byte".into()));
}
let sensitive_vars = ["LD_PRELOAD", "LD_LIBRARY_PATH", "PATH", "PYTHONPATH"];
if sensitive_vars.contains(&key) {
log::warn!("Setting sensitive environment variable: {}", key);
}
Ok((key.to_string(), value.to_string()))
}
}
pub fn parse_user_host(destination: &str) -> Result<(Option<String>, String)> {
let validator = InputValidator::new();
if destination.contains('@') {
let parts: Vec<&str> = destination.split('@').collect();
if parts.len() != 2 {
return Err(QsshError::Config("Invalid user@host format".into()));
}
let username = validator.validate_username(parts[0])?;
let hostname = validator.validate_hostname(parts[1])?;
Ok((Some(username), hostname))
} else {
let hostname = validator.validate_hostname(destination)?;
Ok((None, hostname))
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_validate_username() {
let validator = InputValidator::new();
assert!(validator.validate_username("alice").is_ok());
assert!(validator.validate_username("user-123").is_ok());
assert!(validator.validate_username("test.user").is_ok());
assert!(validator.validate_username("").is_err());
assert!(validator.validate_username("user@domain").is_err());
assert!(validator.validate_username("user/path").is_err());
assert!(validator.validate_username("user\0null").is_err());
}
#[test]
fn test_validate_hostname() {
let validator = InputValidator::new();
assert!(validator.validate_hostname("example.com").is_ok());
assert!(validator.validate_hostname("192.168.1.1").is_ok());
assert!(validator.validate_hostname("::1").is_ok());
assert!(validator.validate_hostname("sub-domain.example.org").is_ok());
assert!(validator.validate_hostname("").is_err());
assert!(validator.validate_hostname(".example.com").is_err());
assert!(validator.validate_hostname("example..com").is_err());
assert!(validator.validate_hostname("example.com-").is_err());
}
#[test]
fn test_validate_port_forward() {
let validator = InputValidator::new();
assert!(validator.validate_port_forward("8080:localhost:80").is_ok());
assert!(validator.validate_port_forward("3000:192.168.1.1:3000").is_ok());
assert!(validator.validate_port_forward("8080:localhost").is_err());
assert!(validator.validate_port_forward("invalid:localhost:80").is_err());
assert!(validator.validate_port_forward("8080:invalid..host:80").is_err());
}
#[test]
fn test_path_traversal_prevention() {
let validator = InputValidator::new();
assert!(validator.validate_path("../etc/passwd").is_err());
assert!(validator.validate_path("../../secret").is_err());
assert!(validator.validate_path("path/../../../etc").is_err());
assert!(validator.validate_path("/home/user/file").is_ok());
assert!(validator.validate_path("relative/path/file").is_ok());
}
}