Skip to content
Merged
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
130 changes: 101 additions & 29 deletions src/scanner/socket_iterator.rs
Original file line number Diff line number Diff line change
@@ -1,55 +1,78 @@
use itertools::{iproduct, Product};
use std::iter::FusedIterator;
use std::net::{IpAddr, SocketAddr};
use std::slice;

/// An iterator that receives a slice of IPs and ports and returns a
/// [`SocketAddr`] for each IP/port pair until all combinations are exhausted.
///
/// The goal of this iterator is to walk every IP and port combination
/// *without* allocating a big intermediate buffer — the alternative would be
/// materialising a `Vec<SocketAddr>` of every pair up front.
///
/// # Ordering
///
/// The IP/port order is intentionally reversed inside `product_it`: we want
/// `iproduct!` to iterate *all IPs for one port* before advancing to the next
/// port ("hold the port, go through all the IPs, then advance the port").
///
/// # Example
///
/// `SocketIterator` lives in a private module, so this example is not run as
/// a doctest; `goes_through_every_ip_port_combination` checks the same order.
///
/// ```ignore
/// # use std::net::IpAddr;
/// let ips = [
/// "127.0.0.1".parse::<IpAddr>().unwrap(),
/// "192.168.0.1".parse::<IpAddr>().unwrap(),
/// ];
/// let ports = [80u16, 443];
///
/// let mut it = SocketIterator::new(&ips, &ports);
/// assert_eq!(it.next(), Some("127.0.0.1:80".parse().unwrap()));
/// assert_eq!(it.next(), Some("192.168.0.1:80".parse().unwrap()));
/// assert_eq!(it.next(), Some("127.0.0.1:443".parse().unwrap()));
/// assert_eq!(it.next(), Some("192.168.0.1:443".parse().unwrap()));
/// assert_eq!(it.next(), None);
/// ```
#[derive(Clone)]
pub struct SocketIterator<'s> {
// product_it is a cartesian product iterator over
// the slices of ports and IP addresses.
//
// The IP/port order is intentionally reversed here since we want
// the itertools::iproduct! macro below to generate the pairs with
// all the IPs for one port before moving on to the next one
// ("hold the port, go through all the IPs, then advance the port...").
// See also the comments in the iterator implementation for an example.
product_it:
Product<Box<std::slice::Iter<'s, u16>>, Box<std::slice::Iter<'s, std::net::IpAddr>>>,
// No boxing: the iterator owns the concrete slice iterators directly.
// This keeps `SocketIterator` allocation-free and makes `Clone` free too.
product_it: Product<slice::Iter<'s, u16>, slice::Iter<'s, IpAddr>>,
}

/// An iterator that receives a slice of IPs and ports and returns a Socket
/// for each IP and port pair until all of these combinations are exhausted.
/// The goal of this iterator is to go over every IP and port combination
/// without generating a big memory footprint. The alternative would be
/// generating a vector containing all these combinations.
impl<'s> SocketIterator<'s> {
pub fn new(ips: &'s [IpAddr], ports: &'s [u16]) -> Self {
let ports_it = Box::new(ports.iter());
let ips_it = Box::new(ips.iter());
Self {
product_it: iproduct!(ports_it, ips_it),
// `iproduct!` calls `.into_iter()` on each argument; `slice::Iter`
// is already an iterator, so this is a no-op (no allocation).
product_it: iproduct!(ports.iter(), ips.iter()),
}
}
}

#[allow(clippy::doc_link_with_quotes)]
impl Iterator for SocketIterator<'_> {
type Item = SocketAddr;

/// Returns a socket based on the combination of one of the provided
/// IPs and ports or None when these combinations are exhausted. Every
/// IP will have the same port until a port is incremented.
///
/// let it = SocketIterator::new(&["127.0.0.1", "192.168.0.1"], &[80, 443]);
/// it.next(); // 127.0.0.1:80
/// it.next(); // 192.168.0.1:80
/// it.next(); // 127.0.0.1:443
/// it.next(); // 192.168.0.1:443
/// it.next(); // None
/// Returns the next socket, or `None` when all combinations are exhausted.
/// Every IP is paired with the same port until the port advances.
fn next(&mut self) -> Option<Self::Item> {
self.product_it
.next()
.map(|(port, ip)| SocketAddr::new(*ip, *port))
}

// Delegate the size hint so callers like `collect()` can pre-size buffers.
fn size_hint(&self) -> (usize, Option<usize>) {
self.product_it.size_hint()
}
}

// `Product` and `slice::Iter` are both fused, so we are too.
impl FusedIterator for SocketIterator<'_> {}

#[cfg(test)]
mod tests {
use super::SocketIterator;
Expand All @@ -72,4 +95,53 @@ mod tests {
assert_eq!(Some(SocketAddr::new(addrs[1], ports[2])), it.next());
assert_eq!(None, it.next());
}

#[test]
fn size_hint_is_exact() {
let addrs = ["127.0.0.1".parse::<IpAddr>().unwrap()];
let ports: Vec<u16> = vec![22, 80, 443];
let mut it = SocketIterator::new(&addrs, &ports);

assert_eq!(it.size_hint(), (3, Some(3)));
it.next();
assert_eq!(it.size_hint(), (2, Some(2)));
}

#[test]
fn clone_resumes_independently() {
let addrs = [
"127.0.0.1".parse::<IpAddr>().unwrap(),
"192.168.0.1".parse::<IpAddr>().unwrap(),
];
let ports: Vec<u16> = vec![22, 80];

let mut a = SocketIterator::new(&addrs, &ports);
assert_eq!(a.next(), Some(SocketAddr::new(addrs[0], ports[0])));
let b = a.clone(); // b starts where a is now

// 2 IPs x 2 ports = 4 sockets, one of which was consumed before cloning.
let rest = [
SocketAddr::new(addrs[1], ports[0]),
SocketAddr::new(addrs[0], ports[1]),
SocketAddr::new(addrs[1], ports[1]),
];
assert_eq!(a.collect::<Vec<_>>(), rest);
assert_eq!(b.collect::<Vec<_>>(), rest);
}

#[test]
fn empty_inputs_yield_nothing() {
let addrs = ["127.0.0.1".parse::<IpAddr>().unwrap()];
let ports = [22u16];

for mut it in [
SocketIterator::new(&[], &ports),
SocketIterator::new(&addrs, &[]),
] {
assert_eq!(it.size_hint(), (0, Some(0)));
assert_eq!(it.next(), None);
// Stays exhausted, as promised by the `FusedIterator` impl.
assert_eq!(it.next(), None);
}
}
}
Loading