1use crate::arp::NetworkManager;
2use crate::error::Result;
3use crate::oui::lookup_vendor;
4use crate::ping::PingScanner;
5use futures::stream::{self, StreamExt};
6use ipnet::Ipv4Net;
7use std::collections::HashMap;
8use std::net::IpAddr;
9use std::sync::{
10 atomic::{AtomicUsize, Ordering},
11 Arc,
12};
13
14#[cfg(feature = "mdns")]
15mod mdns_fusion;
16mod types;
17
18pub use types::{DiscoverPhase, DiscoverProgress, FoundBy, Host, NameSource};
19
20pub struct DiscoverEngine {
21 ping_scanner: PingScanner,
22 concurrency: usize,
23 ping_timeout_ms: u64,
24 dns_timeout_ms: u64,
25}
26
27impl DiscoverEngine {
28 pub fn new(concurrency: usize) -> Self {
29 Self::new_with_timeouts(
30 concurrency,
31 crate::DEFAULT_PING_TIMEOUT_MS,
32 crate::DEFAULT_DNS_TIMEOUT_MS,
33 )
34 }
35
36 pub fn new_with_timeouts(
37 concurrency: usize,
38 ping_timeout_ms: u64,
39 dns_timeout_ms: u64,
40 ) -> Self {
41 let concurrency = concurrency.clamp(1, crate::MAX_CONCURRENCY);
48 Self {
49 ping_scanner: PingScanner::new(concurrency),
50 concurrency,
51 ping_timeout_ms: ping_timeout_ms.max(1),
52 dns_timeout_ms: dns_timeout_ms.max(1),
53 }
54 }
55
56 pub async fn scan_subnet(&self, subnet: Ipv4Net, resolve: bool) -> Result<Vec<Host>> {
57 self.scan_subnet_with_progress(subnet, resolve, None).await
58 }
59
60 pub async fn scan_subnet_with_progress(
68 &self,
69 subnet: Ipv4Net,
70 resolve: bool,
71 progress: Option<Arc<dyn Fn(DiscoverProgress) + Send + Sync>>,
72 ) -> Result<Vec<Host>> {
73 crate::ops::validation::ensure_subnet_limit(&subnet, &subnet.to_string())?;
74 let ips: Vec<IpAddr> = subnet.hosts().map(IpAddr::V4).collect();
75
76 #[cfg(feature = "mdns")]
79 let mdns_browse = mdns_fusion::spawn_browse();
80
81 let total = ips.len();
83 let completed = Arc::new(AtomicUsize::new(0));
84 let found = Arc::new(AtomicUsize::new(0));
85 let ping_results = stream::iter(ips)
86 .map(|ip| {
87 let scanner = self.ping_scanner.clone();
88 let timeout_ms = self.ping_timeout_ms;
89 let completed = completed.clone();
90 let found = found.clone();
91 let progress = progress.clone();
92 async move {
93 let res = scanner.ping(ip, timeout_ms).await;
94 let done = completed.fetch_add(1, Ordering::SeqCst) + 1;
95 if res.alive {
101 found.fetch_add(1, Ordering::SeqCst);
102 }
103 let found_snapshot = found.load(Ordering::SeqCst);
104
105 if let Some(cb) = &progress {
106 if res.alive || done == total || done.is_multiple_of(10) {
107 cb(DiscoverProgress {
108 phase: DiscoverPhase::Ping,
109 completed: done,
110 total,
111 found: found_snapshot,
112 ip,
113 });
114 }
115 }
116
117 res
118 }
119 })
120 .buffer_unordered(self.concurrency)
121 .collect::<Vec<crate::ping::PingResult>>()
122 .await;
123
124 let alive: Vec<crate::ping::PingResult> =
125 ping_results.into_iter().filter(|r| r.alive).collect();
126
127 let arp_map: HashMap<IpAddr, crate::arp::ArpEntry> =
133 tokio::task::spawn_blocking(NetworkManager::get_arp_table)
134 .await
135 .unwrap_or_else(|_| Ok(Vec::new()))
136 .unwrap_or_default()
137 .into_iter()
138 .map(|e| (e.ip, e))
139 .collect();
140
141 let hostname_map: HashMap<IpAddr, Option<String>> = if resolve {
143 let ips = alive.iter().map(|r| r.ip).collect::<Vec<_>>();
144 let concurrency = self.concurrency.min(32);
145 let total = ips.len();
146 let completed = Arc::new(AtomicUsize::new(0));
147 let resolved = Arc::new(AtomicUsize::new(0));
148 stream::iter(ips)
149 .map(|ip| {
150 let dns_timeout_ms = self.dns_timeout_ms;
151 let completed = completed.clone();
152 let resolved = resolved.clone();
153 let progress = progress.clone();
154 async move {
155 let name =
156 crate::dns::reverse_lookup_best_effort_timeout(ip, dns_timeout_ms)
157 .await;
158 let done = completed.fetch_add(1, Ordering::SeqCst) + 1;
159 if name.is_some() {
160 resolved.fetch_add(1, Ordering::SeqCst);
161 }
162 let resolved_count = resolved.load(Ordering::SeqCst);
163
164 if let Some(cb) = &progress {
165 if name.is_some() || done == total || done.is_multiple_of(5) {
166 cb(DiscoverProgress {
167 phase: DiscoverPhase::Resolve,
168 completed: done,
169 total,
170 found: resolved_count,
171 ip,
172 });
173 }
174 }
175
176 (ip, name)
177 }
178 })
179 .buffer_unordered(concurrency)
180 .collect::<Vec<(IpAddr, Option<String>)>>()
181 .await
182 .into_iter()
183 .collect()
184 } else {
185 HashMap::new()
186 };
187
188 #[cfg(feature = "mdns")]
194 let mdns_names = mdns_fusion::collect_names(mdns_browse).await;
195 #[cfg(not(feature = "mdns"))]
196 let mdns_names: HashMap<IpAddr, String> = HashMap::new();
197
198 let build = |ip: IpAddr, rtt_ms: Option<u64>, found_by: FoundBy| {
199 let mac_entry = arp_map.get(&ip);
200 let mac_str = mac_entry.map(|e| e.mac.to_string());
201 let vendor = mac_entry
202 .and_then(|e| e.vendor.clone())
203 .or_else(|| mac_str.as_deref().and_then(lookup_vendor));
204 let (hostname, hostname_source) = match hostname_map.get(&ip).cloned().unwrap_or(None) {
207 Some(name) => (Some(name), Some(NameSource::Reverse)),
208 None => match mdns_names.get(&ip) {
209 Some(name) => (Some(name.clone()), Some(NameSource::Mdns)),
210 None => (None, None),
211 },
212 };
213 Host {
214 ip,
215 hostname,
216 mac: mac_str,
217 vendor,
218 rtt_ms,
219 found_by,
220 hostname_source,
221 }
222 };
223
224 let mut hosts: Vec<Host> = alive
225 .iter()
226 .map(|r| build(r.ip, r.rtt_ms, FoundBy::Probe))
227 .collect();
228
229 let answered: std::collections::HashSet<IpAddr> = alive.iter().map(|r| r.ip).collect();
237 let mut neighbors: Vec<IpAddr> = arp_map
238 .keys()
239 .copied()
240 .filter(|ip| !answered.contains(ip))
241 .filter(|ip| match ip {
242 IpAddr::V4(v4) => is_host_address(&subnet, *v4),
246 IpAddr::V6(_) => false,
247 })
248 .filter(|ip| arp_map.get(ip).is_none_or(|e| e.mac.bytes()[0] & 1 == 0))
249 .collect();
250 neighbors.sort_unstable();
251 hosts.extend(
252 neighbors
253 .into_iter()
254 .map(|ip| build(ip, None, FoundBy::Neighbor)),
255 );
256
257 Ok(hosts)
258 }
259}
260
261fn is_host_address(subnet: &Ipv4Net, ip: std::net::Ipv4Addr) -> bool {
266 subnet.contains(&ip)
267 && (subnet.prefix_len() >= 31 || (ip != subnet.network() && ip != subnet.broadcast()))
268}
269
270#[cfg(test)]
271mod host_address_tests {
272 use super::is_host_address;
273
274 #[test]
275 fn network_and_broadcast_are_not_hosts() {
276 let net = "192.168.1.0/24".parse().unwrap();
277 assert!(!is_host_address(&net, "192.168.1.0".parse().unwrap()));
278 assert!(!is_host_address(&net, "192.168.1.255".parse().unwrap()));
279 assert!(is_host_address(&net, "192.168.1.1".parse().unwrap()));
280 assert!(is_host_address(&net, "192.168.1.254".parse().unwrap()));
281 assert!(!is_host_address(&net, "192.168.2.1".parse().unwrap()));
282 }
283
284 #[test]
285 fn every_address_counts_on_a_31_or_32() {
286 let p2p = "10.0.0.0/31".parse().unwrap();
287 assert!(is_host_address(&p2p, "10.0.0.0".parse().unwrap()));
288 assert!(is_host_address(&p2p, "10.0.0.1".parse().unwrap()));
289 let single = "10.0.0.5/32".parse().unwrap();
290 assert!(is_host_address(&single, "10.0.0.5".parse().unwrap()));
291 }
292}