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