Skip to main content

netscli_core/
discover.rs

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        // Clamp both ends. `.max(1)` alone left the upper bound to
42        // `Semaphore::new`, which asserts `permits <= usize::MAX >> 3` --
43        // so a caller passing `usize::MAX` got a panic rather than an
44        // error, from a constructor that returns `Self` and cannot report
45        // one. Absurd input, but this is public API and a process abort is
46        // the wrong failure mode for it.
47        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    /// Ping-sweep a subnet, then resolve names for the hosts that answered.
61    ///
62    /// Enforces the /16 cap itself rather than trusting the caller. This is
63    /// public API re-exported at the crate root, and it used to collect every
64    /// address of whatever `Ipv4Net` it was given straight into a `Vec` --
65    /// `0.0.0.0/0` is 4,294,967,294 entries, roughly 73 GB, allocated before
66    /// a single packet is sent. `Ops` checked, the engine did not.
67    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        // 0) Start the mDNS browse now, so it overlaps everything below.
77        // See discover/mdns_fusion.rs for why this exists and what it costs.
78        #[cfg(feature = "mdns")]
79        let mdns_browse = mdns_fusion::spawn_browse();
80
81        // 1) Ping first (fast), to avoid reverse-DNS work on dead hosts.
82        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                    // Capture `found` *after* both atomic updates settle.
96                    // Using a single load for both branches avoids the
97                    // previous inconsistency where an unreachable IP could
98                    // be reported with a `found` count that was updated by
99                    // a concurrent task between this task's two loads.
100                    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        // 2) Load ARP/neighbor table once and reuse it.
128        //
129        // On a blocking thread: this shells out to `arp` on Windows and
130        // macOS, and running it inline parked a runtime worker on the child
131        // process for every discover and sweep.
132        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        // 3) Optionally reverse-DNS alive hosts using a single resolver.
142        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        // 4) Fuse in mDNS names for hosts the reverse lookup left unnamed.
189        // Fill-blanks-only, so nothing that resolves today changes. Placed
190        // after `hostname_map` so ARP-only neighbours are named too -- a
191        // device that never answered a probe is the one least likely to have
192        // a PTR record.
193        #[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            // One path for both builds: without the mdns feature the map is
205            // simply empty, so the fallback never fires.
206            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        // Anything the OS has an ARP entry for is on this link, whether or
230        // not it answered us. Plenty of devices do not: consumer IoT
231        // routinely ignores ICMP, and Windows drops echo requests by
232        // default. Discarding them meant discovery reported a fraction of
233        // the network -- measured at 13 of 25 known devices on one ordinary
234        // LAN -- while the table needed to find them was already loaded, and
235        // used only to decorate the hosts that had replied.
236        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                // Only within the range that was asked for. The neighbour
243                // table spans every interface, so it holds addresses from
244                // other subnets entirely.
245                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
261/// An address in `subnet` that a device can hold: not the network address
262/// and not the broadcast address, except on /31 and /32 where every address
263/// is usable. The Windows neighbour table lists x.x.x.255
264/// (ff:ff:ff:ff:ff:ff), and discovery reported it as a host.
265fn 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}