use ipnet::{IpBitAnd, IpBitOr, IpNet, Ipv4Net, Ipv6Net};
use std::net::{IpAddr, Ipv4Addr, Ipv6Addr};
use rand::Rng;
pub fn gen_ip<R: Rng>(rng: &mut R, network: IpNet) -> IpAddr {
match network {
IpNet::V4(network) => IpAddr::V4(gen_ipv4(rng, network)),
IpNet::V6(network) => IpAddr::V6(gen_ipv6(rng, network)),
}
}
pub fn gen_ipv4<R: Rng>(rng: &mut R, network: Ipv4Net) -> Ipv4Addr {
let random: Ipv4Addr = rng.gen::<u32>().into();
network.network().bitor(random.bitand(network.hostmask()))
}
pub fn gen_ipv6<R: Rng>(rng: &mut R, network: Ipv6Net) -> Ipv6Addr {
let random: Ipv6Addr = rng.gen::<u128>().into();
network.network().bitor(random.bitand(network.hostmask()))
}
pub fn probe_targets<R: Rng>(rng: &mut R, nets: &[IpNet]) -> Vec<IpAddr> {
nets.iter().map(|&net| gen_ip(rng, net)).collect()
}
pub fn net_of(nets: &[IpNet], addr: IpAddr) -> Option<IpNet> {
nets.iter().copied().find(|net| net.contains(&addr))
}
pub fn host_net(addr: IpAddr) -> IpNet {
let prefix_len = match addr {
IpAddr::V4(_) => 32,
IpAddr::V6(_) => 128,
};
IpNet::new(addr, prefix_len).expect("host prefix length is always valid")
}
#[cfg(test)]
mod tests {
use std::collections::HashSet;
use rand::SeedableRng;
use super::gen_ip;
const N_ADDR: usize = 1000;
#[test]
fn rand_ipv4() {
let mut rng = rand::rngs::StdRng::seed_from_u64(42);
let net = "127.0.0.0/8".parse().unwrap();
let addrs: HashSet<_> = (0..N_ADDR).map(|_| gen_ip(&mut rng, net)).collect();
assert_eq!(addrs.len(), N_ADDR);
for addr in addrs {
assert!(net.contains(&addr), "{net} should contain {addr}");
}
}
#[test]
fn rand_ipv6() {
let mut rng = rand::rngs::StdRng::seed_from_u64(42);
let net = "2001:db8::/32".parse().unwrap();
let addrs: HashSet<_> = (0..N_ADDR).map(|_| gen_ip(&mut rng, net)).collect();
assert_eq!(addrs.len(), N_ADDR);
for addr in addrs {
assert!(net.contains(&addr), "{net} should contain {addr}");
}
}
#[test]
fn probe_targets_one_per_net() {
use super::probe_targets;
let mut rng = rand::rngs::StdRng::seed_from_u64(42);
let nets = [
"127.0.0.0/30".parse().unwrap(),
"127.0.1.0/30".parse().unwrap(),
"10.0.0.0/8".parse().unwrap(),
];
let targets = probe_targets(&mut rng, &nets);
assert_eq!(targets.len(), nets.len());
for (target, net) in targets.iter().zip(nets.iter()) {
assert!(net.contains(target), "{net} should contain {target}");
}
}
#[test]
fn net_classification() {
use super::{host_net, net_of};
let nets: Vec<_> = ["127.0.0.0/30", "127.0.1.0/30"]
.iter()
.map(|s| s.parse().unwrap())
.collect();
assert_eq!(net_of(&nets, "127.0.0.1".parse().unwrap()), Some(nets[0]));
assert_eq!(net_of(&nets, "127.0.1.1".parse().unwrap()), Some(nets[1]));
assert_eq!(net_of(&nets, "10.0.0.1".parse().unwrap()), None);
assert_eq!(
host_net("10.0.0.1".parse().unwrap()),
"10.0.0.1/32".parse().unwrap()
);
assert_eq!(host_net("::1".parse().unwrap()), "::1/128".parse().unwrap());
}
}