use std::collections::{HashMap, VecDeque};
use std::io::ErrorKind;
use std::net::Ipv4Addr;
use std::sync::OnceLock;
use std::sync::atomic::{AtomicUsize, Ordering};
use std::time::Duration;
use tokio::net::TcpStream;
use tokio::sync::Semaphore;
use tokio::task::JoinSet;
use tokio::time::{Instant, timeout};
const LIVENESS: &[u16] = &[80, 443, 22, 445, 62078, 7000, 8008, 9100];
const SERVICES: &[u16] = &[
21, 25, 53, 110, 111, 135, 139, 143, 993, 995, 1433, 1521, 1883, 3306, 3389, 5001, 5060, 5432,
5672, 6379, 8000, 8001, 8006, 8080, 8081, 8096, 8123, 8443, 8888, 9090, 9091, 9443, 27017,
32400,
];
pub const AWAKE_WAIT: Duration = Duration::from_millis(250);
pub async fn scan(
targets: &[Ipv4Addr],
wait: Duration,
checked: &AtomicUsize,
) -> HashMap<Ipv4Addr, Vec<u16>> {
let deadline = Instant::now() + wait;
let mut liveness = targets
.iter()
.flat_map(|&ip| LIVENESS.iter().map(move |&port| (ip, port)));
let mut services: VecDeque<(Ipv4Addr, u16, Duration)> = VecDeque::new();
let mut set = JoinSet::new();
let mut alive: HashMap<Ipv4Addr, Vec<u16>> = HashMap::new();
loop {
while set.len() < max_in_flight() {
if let Some((ip, port, left)) = services.pop_front() {
set.spawn(async move { (false, check(ip, port, left).await) });
} else if let Some((ip, port)) = liveness.next() {
set.spawn(async move { (true, check(ip, port, wait).await) });
} else {
break;
}
}
let Some(joined) = set.join_next().await else {
break;
};
let Ok((is_liveness, (ip, port, open))) = joined else {
continue;
};
if is_liveness {
checked.fetch_add(1, Ordering::Relaxed);
}
let Some(open) = open else {
continue;
};
if !alive.contains_key(&ip) {
let left = deadline
.saturating_duration_since(Instant::now())
.max(AWAKE_WAIT);
services.extend(SERVICES.iter().map(|&port| (ip, port, left)));
}
let ports = alive.entry(ip).or_default();
if open {
ports.push(port);
}
}
for ports in alive.values_mut() {
ports.sort_unstable();
}
alive
}
pub fn liveness_checks(targets: &[Ipv4Addr]) -> usize {
targets.len() * LIVENESS.len()
}
pub async fn all_ports(ip: Ipv4Addr, wait: Duration) -> Vec<u16> {
let mut set = JoinSet::new();
for &port in LIVENESS.iter().chain(SERVICES) {
set.spawn(check(ip, port, wait));
}
let mut open = Vec::new();
while let Some(joined) = set.join_next().await {
if let Ok((_, port, Some(true))) = joined {
open.push(port);
}
}
open.sort_unstable();
open
}
async fn check(ip: Ipv4Addr, port: u16, wait: Duration) -> (Ipv4Addr, u16, Option<bool>) {
let _permit = sockets().acquire().await.expect("never closed");
let open = match timeout(wait, connect(ip, port)).await {
Ok(Ok(_)) => Some(true),
Ok(Err(e)) if e.kind() == ErrorKind::ConnectionRefused => Some(false),
_ => None,
};
(ip, port, open)
}
#[cfg(not(windows))]
async fn connect(ip: Ipv4Addr, port: u16) -> std::io::Result<TcpStream> {
TcpStream::connect((ip, port)).await
}
#[cfg(windows)]
async fn connect(ip: Ipv4Addr, port: u16) -> std::io::Result<TcpStream> {
let sock = tokio::net::TcpSocket::new_v4()?;
crate::platform::fail_fast_on_refusal(&sock);
sock.connect((ip, port).into()).await
}
fn sockets() -> &'static Semaphore {
static SOCKETS: OnceLock<Semaphore> = OnceLock::new();
SOCKETS.get_or_init(|| Semaphore::new(max_in_flight()))
}
fn max_in_flight() -> usize {
static MAX: OnceLock<usize> = OnceLock::new();
*MAX.get_or_init(|| {
#[cfg(unix)]
let limit = unsafe {
let mut lim: libc::rlimit = std::mem::zeroed();
if libc::getrlimit(libc::RLIMIT_NOFILE, &mut lim) == 0 {
lim.rlim_cur as usize
} else {
256
}
};
#[cfg(windows)]
let limit: usize = 10_240;
limit.saturating_sub(512).clamp(128, 8_192)
})
}
pub fn raise_fd_limit() {
#[cfg(unix)]
unsafe {
let mut lim: libc::rlimit = std::mem::zeroed();
if libc::getrlimit(libc::RLIMIT_NOFILE, &mut lim) == 0 {
lim.rlim_cur = lim.rlim_max.min(10_240);
libc::setrlimit(libc::RLIMIT_NOFILE, &lim);
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn refusals_come_back_at_once() {
let port = std::net::TcpListener::bind("127.0.0.1:0")
.unwrap()
.local_addr()
.unwrap()
.port();
let start = std::time::Instant::now();
let (_, _, open) = check(Ipv4Addr::LOCALHOST, port, Duration::from_secs(5)).await;
assert_eq!(open, Some(false));
assert!(
start.elapsed() < Duration::from_millis(500),
"took {:?}",
start.elapsed()
);
}
}