use ipnetwork::IpNetwork;
use rayon::prelude::*;
use std::io::ErrorKind;
use std::net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr};
use std::time::{Duration, Instant};
use crate::iface;
use crate::utils::progress_bar;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct HostHit {
pub ip: IpAddr,
pub latency: Duration,
}
pub const PROBE_PORTS: [u16; 4] = [80, 443, 22, 3389];
pub const MAX_IPV6_HOSTS: u128 = 1 << 16;
pub fn scan_address(ip: IpAddr, timeout: Option<Duration>) -> Option<HostHit> {
scan_address_with_retries(ip, timeout, 0)
}
pub fn scan_address_with_retries(
ip: IpAddr,
timeout: Option<Duration>,
retries: u32,
) -> Option<HostHit> {
let timeout = timeout.unwrap_or(crate::scanner::port::CONNECT_TIMEOUT);
for _ in 0..=retries {
if let Some(hit) = probe_address(ip, timeout) {
return Some(hit);
}
}
None
}
fn probe_address(ip: IpAddr, timeout: Duration) -> Option<HostHit> {
PROBE_PORTS
.iter()
.find_map(|&port| probe_port(ip, port, timeout))
}
fn probe_port(ip: IpAddr, port: u16, timeout: Duration) -> Option<HostHit> {
crate::rate::gate();
let start = Instant::now();
match iface::tcp_connect_timeout(SocketAddr::new(ip, port), timeout) {
Ok(_) => Some(HostHit {
ip,
latency: start.elapsed(),
}),
Err(e)
if matches!(
e.kind(),
ErrorKind::ConnectionRefused | ErrorKind::ConnectionReset
) =>
{
Some(HostHit {
ip,
latency: start.elapsed(),
})
}
Err(_) => None,
}
}
pub fn scan_hosts(addrs: Vec<IpAddr>, timeout: Option<Duration>, retries: u32) -> Vec<HostHit> {
let total = addrs.len() as u64;
scan_all(
addrs.into_par_iter(),
total,
timeout,
retries,
"Scan completed",
)
}
fn scan_all<I>(
addrs: I,
total: u64,
timeout: Option<Duration>,
retries: u32,
finish_msg: &str,
) -> Vec<HostHit>
where
I: ParallelIterator<Item = IpAddr>,
{
let pb = progress_bar(total, "addresses scanned");
let mut result: Vec<HostHit> = addrs
.filter_map(|ip| {
let available = scan_address_with_retries(ip, timeout, retries);
pb.inc(1);
available
})
.collect();
pb.finish_with_message(finish_msg.to_string());
result.sort_by_key(|h| h.ip);
result
}
pub fn scan_subnet(subnet: IpNetwork, timeout: Option<Duration>) -> Vec<HostHit> {
scan_subnet_with_retries(subnet, timeout, 0)
}
pub fn scan_subnet_with_retries(
subnet: IpNetwork,
timeout: Option<Duration>,
retries: u32,
) -> Vec<HostHit> {
match (subnet.network(), subnet.broadcast()) {
(IpAddr::V4(network), IpAddr::V4(broadcast)) => {
let start = u32::from(network);
let end = u32::from(broadcast);
let total = u64::from(end - start) + 1;
scan_all(
(start..=end).into_par_iter().map(ipv4),
total,
timeout,
retries,
"Subnet scan completed",
)
}
(IpAddr::V6(network), IpAddr::V6(broadcast)) => match ipv6_hosts(network, broadcast) {
Some(hosts) => {
let total = hosts.len() as u64;
scan_all(
hosts.into_par_iter(),
total,
timeout,
retries,
"Subnet scan completed",
)
}
None => Vec::new(),
},
_ => Vec::new(),
}
}
pub fn scan_ip_range(start: IpAddr, end: IpAddr, timeout: Option<Duration>) -> Vec<HostHit> {
scan_ip_range_with_retries(start, end, timeout, 0)
}
pub fn scan_ip_range_with_retries(
start: IpAddr,
end: IpAddr,
timeout: Option<Duration>,
retries: u32,
) -> Vec<HostHit> {
match (start, end) {
(IpAddr::V4(start), IpAddr::V4(end)) => {
let start = u32::from(start);
let end = u32::from(end);
if start > end {
return Vec::new();
}
let total = u64::from(end - start) + 1;
scan_all(
(start..=end).into_par_iter().map(ipv4),
total,
timeout,
retries,
"Range scan completed",
)
}
(IpAddr::V6(start), IpAddr::V6(end)) => {
if start > end {
return Vec::new();
}
match ipv6_hosts(start, end) {
Some(hosts) => {
let total = hosts.len() as u64;
scan_all(
hosts.into_par_iter(),
total,
timeout,
retries,
"Range scan completed",
)
}
None => Vec::new(),
}
}
_ => {
eprintln!("Range start and end must be the same IP family");
Vec::new()
}
}
}
fn ipv6_hosts(start: Ipv6Addr, end: Ipv6Addr) -> Option<Vec<IpAddr>> {
let start = u128::from(start);
let end = u128::from(end);
let count = end - start + 1;
if count > MAX_IPV6_HOSTS {
eprintln!(
"Refusing to scan {} IPv6 addresses (limit is {}); narrow the range or prefix",
count, MAX_IPV6_HOSTS
);
return None;
}
Some(
(start..=end)
.map(|n| IpAddr::V6(Ipv6Addr::from(n)))
.collect(),
)
}
fn ipv4(n: u32) -> IpAddr {
IpAddr::V4(Ipv4Addr::from(n))
}
pub fn subnet_addresses(subnet: IpNetwork) -> Vec<IpAddr> {
match (subnet.network(), subnet.broadcast()) {
(IpAddr::V4(network), IpAddr::V4(broadcast)) => {
let start = u32::from(network);
let end = u32::from(broadcast);
(start..=end).map(ipv4).collect()
}
(IpAddr::V6(network), IpAddr::V6(broadcast)) => {
ipv6_hosts(network, broadcast).unwrap_or_default()
}
_ => Vec::new(),
}
}
pub fn range_addresses(start: IpAddr, end: IpAddr) -> Vec<IpAddr> {
match (start, end) {
(IpAddr::V4(start), IpAddr::V4(end)) => {
let start = u32::from(start);
let end = u32::from(end);
if start > end {
return Vec::new();
}
(start..=end).map(ipv4).collect()
}
(IpAddr::V6(start), IpAddr::V6(end)) => {
if start > end {
return Vec::new();
}
ipv6_hosts(start, end).unwrap_or_default()
}
_ => Vec::new(),
}
}
#[cfg(test)]
mod tests {
use super::*;
use std::time::Duration;
const TEST_TIMEOUT: Option<Duration> = Some(Duration::from_millis(100));
fn is_localhost_available() -> bool {
scan_address("127.0.0.1".parse().unwrap(), TEST_TIMEOUT).is_some()
}
#[test]
fn test_scan_address_localhost() {
if !is_localhost_available() {
println!("Skipping test_scan_address_localhost: localhost is not available");
return;
}
let ip: IpAddr = "127.0.0.1".parse().unwrap();
let hit = scan_address(ip, TEST_TIMEOUT).expect("localhost should be up");
assert_eq!(hit.ip, ip);
}
#[test]
fn test_scan_address_unavailable() {
let ip: IpAddr = "192.168.255.255".parse().unwrap();
assert!(scan_address(ip, TEST_TIMEOUT).is_none());
}
#[test]
fn test_scan_subnet() {
if !is_localhost_available() {
println!("Skipping test_scan_subnet: localhost is not available");
return;
}
let subnet = "127.0.0.0/30".parse::<IpNetwork>().unwrap();
let results = scan_subnet(subnet, TEST_TIMEOUT);
let localhost: IpAddr = "127.0.0.1".parse().unwrap();
assert!(results.iter().any(|h| h.ip == localhost));
assert!(results.windows(2).all(|w| w[0].ip <= w[1].ip));
}
#[test]
fn test_scan_ip_range() {
if !is_localhost_available() {
println!("Skipping test_scan_ip_range: localhost is not available");
return;
}
let start: IpAddr = "127.0.0.1".parse().unwrap();
let end: IpAddr = "127.0.0.3".parse().unwrap();
let results = scan_ip_range(start, end, TEST_TIMEOUT);
let localhost: IpAddr = "127.0.0.1".parse().unwrap();
assert!(results.iter().any(|h| h.ip == localhost));
assert!(results.windows(2).all(|w| w[0].ip <= w[1].ip));
}
#[test]
fn test_scan_empty_range() {
let start: IpAddr = "127.0.0.10".parse().unwrap();
let end: IpAddr = "127.0.0.1".parse().unwrap();
let results = scan_ip_range(start, end, TEST_TIMEOUT);
assert!(results.is_empty());
}
#[test]
fn test_scan_ip_range_family_mismatch() {
let start: IpAddr = "127.0.0.1".parse().unwrap();
let end: IpAddr = "::1".parse().unwrap();
assert!(scan_ip_range(start, end, TEST_TIMEOUT).is_empty());
}
#[test]
fn test_ipv6_hosts_within_limit() {
let start: Ipv6Addr = "2001:db8::".parse().unwrap();
let end: Ipv6Addr = "2001:db8::3".parse().unwrap();
let hosts = ipv6_hosts(start, end).expect("small span should be enumerated");
assert_eq!(hosts.len(), 4);
}
#[test]
fn test_ipv6_hosts_over_limit() {
let net = "2001:db8::/64".parse::<IpNetwork>().unwrap();
if let (IpAddr::V6(network), IpAddr::V6(broadcast)) = (net.network(), net.broadcast()) {
assert!(ipv6_hosts(network, broadcast).is_none());
} else {
panic!("expected IPv6 network");
}
}
#[test]
fn probe_reports_up_at_the_first_responding_port() {
use std::net::TcpListener;
let listener = TcpListener::bind("127.0.0.1:0").expect("bind loopback");
let port = listener.local_addr().unwrap().port();
let ip: IpAddr = "127.0.0.1".parse().unwrap();
let hit = probe_port(ip, port, Duration::from_millis(200));
assert!(hit.is_some(), "an open port should mark the host up");
assert_eq!(hit.unwrap().ip, ip);
}
#[test]
fn probe_ports_spans_more_than_just_port_80() {
assert!(PROBE_PORTS.contains(&443));
assert!(PROBE_PORTS.len() > 1);
}
#[test]
fn subnet_addresses_enumerates_every_host_without_probing() {
let subnet = "127.0.0.0/30".parse::<IpNetwork>().unwrap();
let addrs = subnet_addresses(subnet);
assert_eq!(addrs.len(), 4);
assert!(addrs.contains(&"127.0.0.1".parse().unwrap()));
}
#[test]
fn range_addresses_enumerates_inclusive_and_rejects_backwards() {
let start: IpAddr = "10.0.0.1".parse().unwrap();
let end: IpAddr = "10.0.0.5".parse().unwrap();
assert_eq!(range_addresses(start, end).len(), 5);
assert!(range_addresses(end, start).is_empty());
}
}