use std::io::ErrorKind;
use std::net::{IpAddr, Ipv6Addr, TcpStream, ToSocketAddrs};
use std::time::{Duration, Instant};
pub const CONNECT_TIMEOUT: Duration = Duration::from_secs(2);
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum Probe {
Open(Duration),
Closed,
NoAnswer,
}
fn with_retries<F: FnMut() -> Probe>(retries: u32, mut probe: F) -> Probe {
let mut outcome = Probe::NoAnswer;
for _ in 0..=retries {
outcome = probe();
if outcome != Probe::NoAnswer {
break;
}
}
outcome
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct PortHit {
pub port: u16,
pub latency: Duration,
}
fn host_port(host: &str, port: u16) -> String {
if host.parse::<Ipv6Addr>().is_ok() {
format!("[{}]:{}", host, port)
} else {
format!("{}:{}", host, port)
}
}
pub fn resolve_host(host: &str) -> Option<IpAddr> {
host_port(host, 80)
.to_socket_addrs()
.ok()?
.next()
.map(|addr| addr.ip())
}
pub fn is_resolvable(host: &str) -> bool {
resolve_host(host).is_some()
}
pub fn scan_port(host: String, port: u16, timeout: Option<Duration>) -> Option<PortHit> {
scan_port_with_retries(host, port, timeout, 0)
}
pub fn scan_port_with_retries(
host: String,
port: u16,
timeout: Option<Duration>,
retries: u32,
) -> Option<PortHit> {
let timeout = timeout.unwrap_or(CONNECT_TIMEOUT);
match with_retries(retries, || probe_port(&host, port, timeout)) {
Probe::Open(latency) => Some(PortHit { port, latency }),
Probe::Closed | Probe::NoAnswer => None,
}
}
fn probe_port(host: &str, port: u16, timeout: Duration) -> Probe {
crate::rate::gate();
let Some(socket_addr) = host_port(host, port)
.to_socket_addrs()
.ok()
.and_then(|mut addrs| addrs.next())
else {
return Probe::Closed;
};
let start = Instant::now();
match TcpStream::connect_timeout(&socket_addr, timeout) {
Ok(_) => Probe::Open(start.elapsed()),
Err(e)
if matches!(
e.kind(),
ErrorKind::ConnectionRefused | ErrorKind::ConnectionReset
) =>
{
Probe::Closed
}
Err(_) => Probe::NoAnswer,
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub struct UdpHit {
pub port: u16,
pub latency: Duration,
pub open: bool,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
enum UdpProbe {
Open,
Closed,
OpenFiltered,
}
pub fn scan_udp_port(
host: String,
port: u16,
timeout: Option<Duration>,
retries: u32,
) -> Option<UdpHit> {
let timeout = timeout.unwrap_or(CONNECT_TIMEOUT);
let start = Instant::now();
let mut outcome = UdpProbe::OpenFiltered;
for _ in 0..=retries {
outcome = probe_udp(&host, port, timeout);
if outcome != UdpProbe::OpenFiltered {
break;
}
}
match outcome {
UdpProbe::Open => Some(UdpHit {
port,
latency: start.elapsed(),
open: true,
}),
UdpProbe::OpenFiltered => Some(UdpHit {
port,
latency: start.elapsed(),
open: false,
}),
UdpProbe::Closed => None,
}
}
fn probe_udp(host: &str, port: u16, timeout: Duration) -> UdpProbe {
use std::net::UdpSocket;
crate::rate::gate();
let Some(target) = host_port(host, port)
.to_socket_addrs()
.ok()
.and_then(|mut addrs| addrs.next())
else {
return UdpProbe::Closed;
};
let bind_addr = if target.is_ipv4() {
"0.0.0.0:0"
} else {
"[::]:0"
};
let Ok(socket) = UdpSocket::bind(bind_addr) else {
return UdpProbe::OpenFiltered;
};
if socket.connect(target).is_err() {
return UdpProbe::OpenFiltered;
}
if socket.set_read_timeout(Some(timeout)).is_err() {
return UdpProbe::OpenFiltered;
}
if socket.send(&udp_probe_payload(port)).is_err() {
return UdpProbe::OpenFiltered;
}
let mut buf = [0u8; 512];
classify_udp(socket.recv(&mut buf))
}
fn classify_udp(result: std::io::Result<usize>) -> UdpProbe {
match result {
Ok(_) => UdpProbe::Open,
Err(e) if e.kind() == ErrorKind::ConnectionRefused => UdpProbe::Closed,
Err(_) => UdpProbe::OpenFiltered,
}
}
fn udp_probe_payload(port: u16) -> Vec<u8> {
match port {
53 => vec![
0x12, 0x34, 0x01, 0x00, 0x00, 0x01, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x02, 0x00, 0x01, ],
123 => {
let mut pkt = vec![0u8; 48];
pkt[0] = 0x1b;
pkt
}
_ => Vec::new(),
}
}
#[cfg(test)]
mod tests {
use super::*;
const TEST_TIMEOUT: Option<Duration> = Some(Duration::from_millis(100));
#[test]
fn test_scan_port_success() {
let result = scan_port("127.0.0.1".to_string(), 0, TEST_TIMEOUT);
assert!(result.is_none());
}
#[test]
fn test_scan_port_failure() {
let result = scan_port("127.0.0.1".to_string(), 1, TEST_TIMEOUT); assert!(result.is_none());
}
#[test]
fn test_is_resolvable_numeric_ip() {
assert!(is_resolvable("127.0.0.1"));
}
#[test]
fn test_is_resolvable_invalid_host() {
assert!(!is_resolvable(""));
}
#[test]
fn test_host_port_ipv6_is_bracketed() {
assert_eq!(host_port("::1", 80), "[::1]:80");
assert_eq!(host_port("2001:db8::1", 443), "[2001:db8::1]:443");
}
#[test]
fn test_host_port_ipv4_and_hostname_plain() {
assert_eq!(host_port("127.0.0.1", 80), "127.0.0.1:80");
assert_eq!(host_port("example.com", 8080), "example.com:8080");
}
#[test]
fn retries_stop_immediately_on_a_definitive_open() {
let mut calls = 0;
let outcome = with_retries(5, || {
calls += 1;
Probe::Open(Duration::from_millis(1))
});
assert_eq!(outcome, Probe::Open(Duration::from_millis(1)));
assert_eq!(calls, 1, "an open port must not be retried");
}
#[test]
fn retries_stop_immediately_on_a_definitive_closed() {
let mut calls = 0;
let outcome = with_retries(5, || {
calls += 1;
Probe::Closed
});
assert_eq!(outcome, Probe::Closed);
assert_eq!(calls, 1, "a refused/reset port must not be retried");
}
#[test]
fn no_answer_is_retried_exactly_retries_plus_one_times() {
let mut calls = 0;
let outcome = with_retries(3, || {
calls += 1;
Probe::NoAnswer
});
assert_eq!(outcome, Probe::NoAnswer);
assert_eq!(calls, 4);
}
#[test]
fn no_answer_then_open_succeeds_within_the_retry_budget() {
let mut calls = 0;
let outcome = with_retries(3, || {
calls += 1;
if calls < 3 {
Probe::NoAnswer
} else {
Probe::Open(Duration::from_millis(2))
}
});
assert_eq!(outcome, Probe::Open(Duration::from_millis(2)));
assert_eq!(calls, 3, "should stop on the first definitive answer");
}
#[test]
fn zero_retries_is_a_single_attempt() {
let mut calls = 0;
let _ = with_retries(0, || {
calls += 1;
Probe::NoAnswer
});
assert_eq!(calls, 1);
}
#[test]
fn udp_classification_maps_recv_outcomes() {
use std::io::{Error, ErrorKind};
assert_eq!(classify_udp(Ok(12)), UdpProbe::Open);
assert_eq!(
classify_udp(Err(Error::from(ErrorKind::ConnectionRefused))),
UdpProbe::Closed
);
assert_eq!(
classify_udp(Err(Error::from(ErrorKind::WouldBlock))),
UdpProbe::OpenFiltered
);
assert_eq!(
classify_udp(Err(Error::from(ErrorKind::TimedOut))),
UdpProbe::OpenFiltered
);
}
#[test]
fn dns_payload_is_a_wellformed_query_header() {
let dns = udp_probe_payload(53);
assert!(!dns.is_empty());
assert_eq!(&dns[4..6], &[0x00, 0x01]);
assert_eq!(&dns[dns.len() - 5..], &[0x00, 0x00, 0x02, 0x00, 0x01]);
}
#[test]
fn ntp_payload_is_a_48_byte_client_request() {
let ntp = udp_probe_payload(123);
assert_eq!(ntp.len(), 48);
assert_eq!(ntp[0], 0x1b);
}
#[test]
fn unknown_ports_get_an_empty_payload() {
assert!(udp_probe_payload(4444).is_empty());
}
#[test]
fn udp_scan_reports_open_when_a_local_server_replies() {
use std::net::UdpSocket;
use std::thread;
let server = UdpSocket::bind("127.0.0.1:0").expect("bind server");
let port = server.local_addr().unwrap().port();
let handle = thread::spawn(move || {
let mut buf = [0u8; 512];
if let Ok((n, peer)) = server.recv_from(&mut buf) {
let _ = server.send_to(&buf[..n], peer);
}
});
let hit = scan_udp_port("127.0.0.1".to_string(), port, TEST_TIMEOUT, 0)
.expect("a replying port is open|filtered or open, never closed");
assert!(hit.open, "a replying UDP port must be reported open");
let _ = handle.join();
}
}