use std::collections::BTreeSet;
use super::constants::DEFAULT_PORTS;
use crate::error::{Error, Result};
pub const MAX_PORTS_PER_SCAN: usize = 4096;
pub fn validate_ports(ports: &[u16]) -> Result<()> {
if ports.contains(&0) {
return Err(Error::invalid_input(
"port 0 is invalid (not a connectable TCP port)",
));
}
if ports.len() > MAX_PORTS_PER_SCAN {
return Err(Error::invalid_input(format!(
"too many ports requested ({} > {})",
ports.len(),
MAX_PORTS_PER_SCAN
)));
}
Ok(())
}
pub fn parse_ports(input: Option<&str>) -> Option<Vec<u16>> {
input.map(|s| {
s.split(',')
.filter_map(|part| parse_port_token(part).ok())
.flatten()
.collect()
})
}
fn parse_port_token(part: &str) -> Result<Vec<u16>> {
let part = part.trim();
if part.is_empty() {
return Ok(Vec::new());
}
if let Some((start, end)) = part.split_once('-') {
let start_trim = start.trim();
let end_trim = end.trim();
let s: u16 = start_trim
.parse()
.map_err(|_| Error::invalid_input(format!("invalid port in range: '{start_trim}'")))?;
let e: u16 = end_trim
.parse()
.map_err(|_| Error::invalid_input(format!("invalid port in range: '{end_trim}'")))?;
if s > e {
return Err(Error::invalid_input(format!(
"inverted port range: {s}-{e}"
)));
}
Ok((s..=e).collect())
} else {
let port: u16 = part
.parse()
.map_err(|_| Error::invalid_input(format!("invalid port: '{part}'")))?;
Ok(vec![port])
}
}
pub fn parse_ports_checked(input: Option<&str>) -> Result<Option<Vec<u16>>> {
let Some(raw) = input else {
return Ok(None);
};
let raw = raw.trim();
if raw.is_empty() {
return Ok(None);
}
let mut ports: BTreeSet<u16> = BTreeSet::new();
for part in raw.split(',') {
let expanded = parse_port_token(part).map_err(|e| {
let detail = match e {
Error::InvalidInput(detail) => detail,
other => other.to_string(),
};
Error::invalid_input(format!("invalid port list '{raw}': {detail}"))
})?;
ports.extend(expanded);
if ports.len() > MAX_PORTS_PER_SCAN {
let so_far: Vec<u16> = ports.iter().copied().collect();
validate_ports(&so_far)?;
}
}
if ports.is_empty() {
return Err(Error::invalid_input(format!("invalid port list: {raw}")));
}
let ports: Vec<u16> = ports.into_iter().collect();
validate_ports(&ports)?;
Ok(Some(ports))
}
pub fn default_ports() -> Vec<u16> {
DEFAULT_PORTS.to_vec()
}
#[cfg(test)]
mod tests;