1use futures::stream::{self, StreamExt};
2use std::net::{IpAddr, SocketAddr};
3use std::sync::{
4 atomic::{AtomicUsize, Ordering},
5 Arc,
6};
7use std::time::Instant;
8use tokio::net::TcpStream;
9use tokio::sync::Semaphore;
10
11use super::probes::{as_text, ask, first_banner_line, probe_http, probe_tls, read_greeting, Quiet};
12use super::services::{classify_connect_error, guess_service, is_http_port, is_tls_port};
13use super::types::{PortResult, PortStatus};
14use crate::error::Result;
15
16pub struct PortScanner {
24 semaphore: Arc<Semaphore>,
25 pub(super) concurrency: usize,
26}
27
28#[derive(Debug, Clone)]
29pub struct PortScanProgress {
30 pub completed: usize,
31 pub total: usize,
32 pub port: u16,
33 pub open: bool,
34 pub open_found: usize,
35}
36
37impl PortScanner {
38 pub fn new(concurrency: usize) -> Self {
39 let concurrency = concurrency.clamp(1, crate::MAX_CONCURRENCY);
40 Self {
41 semaphore: Arc::new(Semaphore::new(concurrency)),
42 concurrency,
43 }
44 }
45
46 pub async fn scan_host(
47 &self,
48 target: IpAddr,
49 ports: Vec<u16>,
50 timeout_ms: u64,
51 ) -> Result<Vec<PortResult>> {
52 self.scan_host_with_progress(target, ports, timeout_ms, None)
53 .await
54 }
55
56 pub async fn scan_host_with_progress(
66 &self,
67 target: IpAddr,
68 ports: Vec<u16>,
69 timeout_ms: u64,
70 progress: Option<Arc<dyn Fn(PortScanProgress) + Send + Sync>>,
71 ) -> Result<Vec<PortResult>> {
72 if ports.is_empty() {
73 return Ok(Vec::new());
74 }
75 crate::validate_ports(&ports)?;
76
77 let total = ports.len();
78 let completed = Arc::new(AtomicUsize::new(0));
79 let open_found = Arc::new(AtomicUsize::new(0));
80
81 let results = stream::iter(ports)
82 .map(|port| {
83 let scanner = self.clone();
84 let completed = completed.clone();
85 let open_found = open_found.clone();
86 let progress = progress.clone();
87 async move {
88 let res = scanner.check_port(target, port, timeout_ms).await;
89 let done = completed.fetch_add(1, Ordering::SeqCst) + 1;
90 let open_count = if res.open {
91 open_found.fetch_add(1, Ordering::SeqCst) + 1
92 } else {
93 open_found.load(Ordering::SeqCst)
94 };
95
96 if let Some(cb) = &progress {
97 cb(PortScanProgress {
98 completed: done,
99 total,
100 port,
101 open: res.open,
102 open_found: open_count,
103 });
104 }
105
106 res
107 }
108 })
109 .buffer_unordered(self.concurrency)
110 .collect::<Vec<PortResult>>()
111 .await;
112
113 Ok(results)
114 }
115
116 async fn check_port(&self, target: IpAddr, port: u16, timeout_ms: u64) -> PortResult {
117 let _permit = match self.semaphore.acquire().await {
118 Ok(p) => p,
119 Err(_) => {
120 return PortResult::new(port, PortStatus::Error, None)
123 .with_error("scanner shut down (semaphore closed)".to_string());
124 }
125 };
126
127 let addr = SocketAddr::new(target, port);
128 let service = Self::guess_service(port);
129 let started = Instant::now();
130 match crate::common::tcp_connect(addr, timeout_ms).await {
131 Ok(stream) => {
132 let latency_ms = started.elapsed().as_millis() as u64;
133 let mut result = PortResult::new(port, PortStatus::Open, service.clone())
134 .with_latency(latency_ms);
135 self.enrich_open_port(target, port, stream, timeout_ms, &mut result)
136 .await;
137 if result.product.is_none() {
138 if let Some((product, version)) = super::version::identify(&result) {
139 result.product = Some(product);
140 result.version = version;
141 }
142 }
143 result
144 }
145 Err(e) => match classify_connect_error(e.kind()) {
146 PortStatus::Closed => PortResult::new(port, PortStatus::Closed, service)
147 .with_latency(started.elapsed().as_millis() as u64),
148 PortStatus::Filtered => PortResult::new(port, PortStatus::Filtered, service),
149 PortStatus::Error | PortStatus::Open | PortStatus::OpenFiltered => {
150 PortResult::new(port, PortStatus::Error, service)
151 .with_latency(started.elapsed().as_millis() as u64)
152 .with_error(e.to_string())
153 }
154 },
155 }
156 }
157
158 async fn enrich_open_port(
159 &self,
160 target: IpAddr,
161 port: u16,
162 stream: TcpStream,
163 timeout_ms: u64,
164 result: &mut PortResult,
165 ) {
166 let service = result.service.as_deref();
167 if is_tls_port(port, service) {
168 if let Some((tls, http, banner, raw)) =
169 probe_tls(target, port, stream, timeout_ms, service).await
170 {
171 result.tls = Some(tls);
172 result.http = http;
173 result.banner = banner;
174 result.raw = raw;
175 }
176 return;
177 }
178
179 if is_http_port(port, service) {
180 let mut stream = stream;
181 if let Some((http, banner, raw)) =
182 probe_http(&mut stream, &target.to_string(), timeout_ms).await
183 {
184 result.http = Some(http);
185 result.banner = banner;
186 result.raw = raw;
187 }
188 return;
189 }
190
191 let mut stream = stream;
192 if let Some(bytes) = read_greeting(&mut stream, timeout_ms).await {
193 let raw = as_text(&bytes);
197 if let Some((product, version)) = super::version::from_mysql_greeting(&bytes) {
198 result.product = Some(product);
201 result.version = version;
202 } else {
203 result.banner = Some(first_banner_line(&raw));
204 }
205 result.raw = Some(raw);
206 return;
207 }
208
209 if let Some(quiet) = Quiet::on(port, service) {
212 if let Some(bytes) = ask(&mut stream, quiet, timeout_ms).await {
213 let reply = as_text(&bytes);
214 let found = match quiet {
215 Quiet::Redis => super::version::from_redis_info(&reply),
216 Quiet::Memcached => super::version::from_memcached(&reply),
217 };
218 if let Some((product, version)) = found {
219 result.product = Some(product);
220 result.version = version;
221 }
222 result.raw = Some(reply);
223 }
224 }
225 }
226
227 pub(super) fn guess_service(port: u16) -> Option<String> {
228 guess_service(port)
229 }
230}
231
232impl Clone for PortScanner {
233 fn clone(&self) -> Self {
234 Self {
235 semaphore: self.semaphore.clone(),
236 concurrency: self.concurrency,
237 }
238 }
239}