1use futures::stream::{self, StreamExt};
24use std::io::ErrorKind;
25use std::net::{IpAddr, Ipv4Addr, Ipv6Addr, SocketAddr};
26use std::sync::{
27 atomic::{AtomicUsize, Ordering},
28 Arc,
29};
30use std::time::{Duration, Instant};
31use tokio::net::UdpSocket;
32use tokio::sync::Semaphore;
33use tokio::time::timeout;
34
35mod probes;
36
37use super::services::guess_service;
38use super::tcp::PortScanProgress;
39use super::types::{PortResult, PortStatus, Protocol};
40use crate::error::Result;
41use probes::{describe_reply, probe_for};
42
43pub const DEFAULT_UDP_PORTS: &[u16] = &[53, 123, 137, 1900, 5353];
46
47const MIN_REPLY_WAIT_MS: u64 = 1_000;
52
53#[derive(Clone)]
55pub struct UdpScanner {
56 semaphore: Arc<Semaphore>,
57 concurrency: usize,
58}
59
60impl UdpScanner {
61 pub fn new(concurrency: usize) -> Self {
62 let concurrency = concurrency.clamp(1, crate::MAX_CONCURRENCY);
63 Self {
64 semaphore: Arc::new(Semaphore::new(concurrency)),
65 concurrency,
66 }
67 }
68
69 pub async fn scan_host_with_progress(
72 &self,
73 target: IpAddr,
74 ports: Vec<u16>,
75 timeout_ms: u64,
76 progress: Option<Arc<dyn Fn(PortScanProgress) + Send + Sync>>,
77 ) -> Result<Vec<PortResult>> {
78 if ports.is_empty() {
79 return Ok(Vec::new());
80 }
81 crate::validate_ports(&ports)?;
82
83 let total = ports.len();
84 let wait = Duration::from_millis(timeout_ms.max(MIN_REPLY_WAIT_MS));
85 let completed = Arc::new(AtomicUsize::new(0));
86 let open_found = Arc::new(AtomicUsize::new(0));
87
88 let results = stream::iter(ports)
89 .map(|port| {
90 let scanner = self.clone();
91 let completed = completed.clone();
92 let open_found = open_found.clone();
93 let progress = progress.clone();
94 async move {
95 let res = scanner.check_port(target, port, wait).await;
96 let done = completed.fetch_add(1, Ordering::SeqCst) + 1;
97 let open_count = if res.open {
98 open_found.fetch_add(1, Ordering::SeqCst) + 1
99 } else {
100 open_found.load(Ordering::SeqCst)
101 };
102 if let Some(cb) = &progress {
103 cb(PortScanProgress {
104 completed: done,
105 total,
106 port,
107 open: res.open,
108 open_found: open_count,
109 });
110 }
111 res
112 }
113 })
114 .buffer_unordered(self.concurrency)
115 .collect::<Vec<PortResult>>()
116 .await;
117
118 Ok(results)
119 }
120
121 async fn check_port(&self, target: IpAddr, port: u16, wait: Duration) -> PortResult {
122 let service = guess_service(port);
123 let result = |status| udp_result(port, status, service.clone());
124 let _permit = match self.semaphore.acquire().await {
125 Ok(permit) => permit,
126 Err(_) => {
127 return result(PortStatus::Error)
128 .with_error("scanner shut down (semaphore closed)".to_string())
129 }
130 };
131
132 let local: SocketAddr = match target {
133 IpAddr::V4(_) => (Ipv4Addr::UNSPECIFIED, 0).into(),
134 IpAddr::V6(_) => (Ipv6Addr::UNSPECIFIED, 0).into(),
135 };
136 let socket = match UdpSocket::bind(local).await {
137 Ok(socket) => socket,
138 Err(e) => return result(PortStatus::Error).with_error(e.to_string()),
139 };
140 if let Err(e) = socket.connect((target, port)).await {
141 return result(PortStatus::Error).with_error(e.to_string());
142 }
143
144 let started = Instant::now();
145 if let Err(e) = socket.send(probe_for(port)).await {
146 return if is_port_unreachable(e.kind()) {
147 result(PortStatus::Closed).with_latency(elapsed_ms(started))
148 } else {
149 result(PortStatus::Error).with_error(e.to_string())
150 };
151 }
152
153 let mut buf = [0u8; 2048];
154 match timeout(wait, socket.recv(&mut buf)).await {
155 Ok(Ok(len)) => {
156 let mut open = result(PortStatus::Open).with_latency(elapsed_ms(started));
157 open.banner = describe_reply(port, &buf[..len]);
158 open
159 }
160 Ok(Err(e)) if is_port_unreachable(e.kind()) => {
161 result(PortStatus::Closed).with_latency(elapsed_ms(started))
162 }
163 Ok(Err(e)) => result(PortStatus::Error).with_error(e.to_string()),
164 Err(_) => result(PortStatus::OpenFiltered),
165 }
166 }
167}
168
169fn udp_result(port: u16, status: PortStatus, service: Option<String>) -> PortResult {
170 let mut result = PortResult::new(port, status, service);
171 result.protocol = Protocol::Udp;
172 result
173}
174
175fn elapsed_ms(started: Instant) -> u64 {
176 started.elapsed().as_millis() as u64
177}
178
179fn is_port_unreachable(kind: ErrorKind) -> bool {
182 matches!(
183 kind,
184 ErrorKind::ConnectionRefused | ErrorKind::ConnectionReset
185 )
186}
187
188#[cfg(test)]
189mod tests {
190 use super::*;
191
192 #[test]
193 fn port_unreachable_is_refused_on_unix_and_reset_on_windows() {
194 assert!(is_port_unreachable(ErrorKind::ConnectionRefused));
195 assert!(is_port_unreachable(ErrorKind::ConnectionReset));
196 assert!(!is_port_unreachable(ErrorKind::TimedOut));
197 }
198
199 #[tokio::test]
200 async fn a_closed_local_port_is_closed_and_an_answering_one_open() {
201 let server = UdpSocket::bind((Ipv4Addr::LOCALHOST, 0)).await.unwrap();
203 let open_port = server.local_addr().unwrap().port();
204 tokio::spawn(async move {
205 let mut buf = [0u8; 64];
206 if let Ok((len, from)) = server.recv_from(&mut buf).await {
207 let _ = server.send_to(&buf[..len.max(1)], from).await;
208 }
209 });
210 let closed_port = {
212 let socket = UdpSocket::bind((Ipv4Addr::LOCALHOST, 0)).await.unwrap();
213 socket.local_addr().unwrap().port()
214 };
215
216 let scanner = UdpScanner::new(4);
217 let results = scanner
218 .scan_host_with_progress(
219 IpAddr::V4(Ipv4Addr::LOCALHOST),
220 vec![open_port, closed_port],
221 500,
222 None,
223 )
224 .await
225 .unwrap();
226 let status_of = |port| results.iter().find(|r| r.port == port).unwrap().status;
227 assert_eq!(status_of(open_port), PortStatus::Open);
228 assert_eq!(status_of(closed_port), PortStatus::Closed);
229 assert!(results.iter().all(|r| r.protocol == Protocol::Udp));
230 }
231}