use ipnetwork::IpNetwork;
use rayon::prelude::*;
use std::net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr, TcpStream};
use std::sync::{Arc, Mutex};
use std::time::Duration;
use crate::utils::progress_bar;
pub const MAX_IPV6_HOSTS: u128 = 1 << 16;
pub fn scan_address(ip: IpAddr, timeout: Option<Duration>) -> Option<IpAddr> {
match TcpStream::connect_timeout(
&SocketAddr::new(ip, 80),
timeout.unwrap_or(crate::scanner::port::CONNECT_TIMEOUT),
) {
Ok(_) => Some(ip),
Err(_) => None,
}
}
fn scan_all<I>(addrs: I, total: u64, timeout: Option<Duration>, finish_msg: &str) -> Vec<IpAddr>
where
I: ParallelIterator<Item = IpAddr>,
{
let pb = progress_bar(total, "addresses scanned");
let available_ips = Arc::new(Mutex::new(Vec::new()));
addrs.for_each(|ip| {
if let Some(available_ip) = scan_address(ip, timeout)
&& let Ok(mut guard) = available_ips.lock()
{
guard.push(available_ip);
}
pb.inc(1);
});
pb.finish_with_message(finish_msg.to_string());
let mut result = available_ips.lock().unwrap();
result.sort();
result.clone()
}
pub fn scan_subnet(subnet: IpNetwork, timeout: Option<Duration>) -> Vec<IpAddr> {
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,
"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,
"Subnet scan completed",
)
}
None => Vec::new(),
},
_ => Vec::new(),
}
}
pub fn scan_ip_range(start: IpAddr, end: IpAddr, timeout: Option<Duration>) -> 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();
}
let total = u64::from(end - start) + 1;
scan_all(
(start..=end).into_par_iter().map(ipv4),
total,
timeout,
"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,
"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))
}
#[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();
assert!(scan_address(ip, TEST_TIMEOUT).is_some());
}
#[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);
assert!(results.contains(&"127.0.0.1".parse::<IpAddr>().unwrap()));
assert!(results.windows(2).all(|w| w[0] <= w[1]));
}
#[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);
assert!(results.contains(&"127.0.0.1".parse::<IpAddr>().unwrap()));
assert!(results.windows(2).all(|w| w[0] <= w[1]));
}
#[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");
}
}
}