use std::{
net::{IpAddr, Ipv4Addr, Ipv6Addr},
str::FromStr,
};
use thiserror::Error;
#[derive(Debug, Error, PartialEq)]
pub enum IpError {
#[error("Invalid range: start address {0} is greater than end address {1}")]
InvalidRange(IpAddr, IpAddr),
#[error("Invalid CIDR prefix: {0}")]
InvalidPrefix(u8),
#[error("Network error: {0}")]
NetworkError(String),
#[error("Failed to parse IP address: {0}")]
AddrParse(#[from] std::net::AddrParseError),
#[error("Invalid IP range format: {0}")]
InvalidFormat(String),
#[error("Invalid prefix number format: {0}")]
PrefixParse(#[from] std::num::ParseIntError),
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub struct Ipv4Range {
pub start_addr: Ipv4Addr,
pub end_addr: Ipv4Addr,
}
impl Ipv4Range {
pub fn new(start: Ipv4Addr, end: Ipv4Addr) -> Result<Self, IpError> {
if u32::from(start) <= u32::from(end) {
Ok(Self {
start_addr: start,
end_addr: end,
})
} else {
Err(IpError::InvalidRange(IpAddr::V4(start), IpAddr::V4(end)))
}
}
pub fn to_iter(&self) -> impl Iterator<Item = IpAddr> {
let start: u32 = self.start_addr.into();
let end: u32 = self.end_addr.into();
(start..=end).map(|ip| IpAddr::V4(Ipv4Addr::from(ip)))
}
pub fn contains(&self, ip: &Ipv4Addr) -> bool {
let start: u32 = self.start_addr.into();
let end: u32 = self.end_addr.into();
let ip_u32: u32 = (*ip).into();
ip_u32 >= start && ip_u32 <= end
}
pub fn len(&self) -> u64 {
let s_u32: u64 = u32::from(self.start_addr) as u64;
let e_u32: u64 = u32::from(self.end_addr) as u64;
(e_u32 - s_u32) + 1
}
pub fn is_empty(&self) -> bool {
self.len() == 0
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub struct Ipv6Range {
pub start_addr: Ipv6Addr,
pub end_addr: Ipv6Addr,
}
impl Ipv6Range {
pub fn new(start: Ipv6Addr, end: Ipv6Addr) -> Result<Self, IpError> {
if u128::from(start) <= u128::from(end) {
Ok(Self {
start_addr: start,
end_addr: end,
})
} else {
Err(IpError::InvalidRange(IpAddr::V6(start), IpAddr::V6(end)))
}
}
pub fn to_iter(&self) -> impl Iterator<Item = IpAddr> {
let start: u128 = self.start_addr.into();
let end: u128 = self.end_addr.into();
(start..=end).map(|ip| IpAddr::V6(Ipv6Addr::from(ip)))
}
pub fn contains(&self, ip: &Ipv6Addr) -> bool {
let start: u128 = self.start_addr.into();
let end: u128 = self.end_addr.into();
let ip_u128: u128 = (*ip).into();
ip_u128 >= start && ip_u128 <= end
}
pub fn len(&self) -> u128 {
let s_u128: u128 = u128::from(self.start_addr);
let e_u128: u128 = u128::from(self.end_addr);
(e_u128 - s_u128) + 1
}
pub fn is_empty(&self) -> bool {
self.len() == 0
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)]
pub enum IpRange {
V4(Ipv4Range),
V6(Ipv6Range),
}
impl IpRange {
pub fn start_addr(&self) -> IpAddr {
match self {
IpRange::V4(r) => IpAddr::V4(r.start_addr),
IpRange::V6(r) => IpAddr::V6(r.start_addr),
}
}
pub fn end_addr(&self) -> IpAddr {
match self {
IpRange::V4(r) => IpAddr::V4(r.end_addr),
IpRange::V6(r) => IpAddr::V6(r.end_addr),
}
}
pub fn contains(&self, ip: &IpAddr) -> bool {
match (self, ip) {
(IpRange::V4(r), IpAddr::V4(ip)) => r.contains(ip),
(IpRange::V6(r), IpAddr::V6(ip)) => r.contains(ip),
_ => false,
}
}
pub fn len(&self) -> u128 {
match self {
IpRange::V4(r) => r.len() as u128,
IpRange::V6(r) => r.len(),
}
}
pub fn is_empty(&self) -> bool {
self.len() == 0
}
}
impl FromStr for IpRange {
type Err = IpError;
fn from_str(s: &str) -> Result<Self, Self::Err> {
let s = s.trim();
if let Some(pos) = s.find('/') {
let ip = s[..pos].parse::<IpAddr>()?;
let prefix = s[pos + 1..].parse::<u8>()?;
return cidr_range(ip, prefix);
}
if let Some(pos) = s.find('-') {
let start_str = s[..pos].trim();
let end_str = s[pos + 1..].trim();
if let Ok(start) = start_str.parse::<Ipv4Addr>() {
let end = end_str.parse::<Ipv4Addr>()?;
return Ok(IpRange::V4(Ipv4Range::new(start, end)?));
} else if let Ok(start) = start_str.parse::<Ipv6Addr>() {
let end = end_str.parse::<Ipv6Addr>()?;
return Ok(IpRange::V6(Ipv6Range::new(start, end)?));
}
return Err(IpError::InvalidFormat(s.to_string()));
}
let ip = s.parse::<IpAddr>()?;
match ip {
IpAddr::V4(v4) => Ok(IpRange::V4(Ipv4Range::new(v4, v4).unwrap())),
IpAddr::V6(v6) => Ok(IpRange::V6(Ipv6Range::new(v6, v6).unwrap())),
}
}
}
pub fn cidr_range(ip: IpAddr, prefix: u8) -> Result<IpRange, IpError> {
match ip {
IpAddr::V4(v4) => {
if prefix > 32 {
return Err(IpError::InvalidPrefix(prefix));
}
let ip_u32 = u32::from(v4);
let mask = if prefix == 0 {
0
} else {
!u32::MAX.checked_shr(prefix as u32).unwrap_or(0)
};
let network = ip_u32 & mask;
let broadcast = ip_u32 | !mask;
Ok(IpRange::V4(
Ipv4Range::new(Ipv4Addr::from(network), Ipv4Addr::from(broadcast)).unwrap(),
))
}
IpAddr::V6(v6) => {
if prefix > 128 {
return Err(IpError::InvalidPrefix(prefix));
}
let ip_u128 = u128::from(v6);
let mask = if prefix == 0 {
0
} else {
!u128::MAX.checked_shr(prefix as u32).unwrap_or(0)
};
let network = ip_u128 & mask;
let broadcast = ip_u128 | !mask;
Ok(IpRange::V6(
Ipv6Range::new(Ipv6Addr::from(network), Ipv6Addr::from(broadcast)).unwrap(),
))
}
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn new_valid_v4() {
let start = Ipv4Addr::new(192, 168, 1, 1);
let end = Ipv4Addr::new(192, 168, 1, 10);
let range = Ipv4Range::new(start, end).unwrap();
assert_eq!(range.start_addr, start);
assert_eq!(range.end_addr, end);
}
#[test]
fn len_calculations_v4() {
let cases = vec![
(Ipv4Addr::new(10, 0, 0, 0), Ipv4Addr::new(10, 0, 0, 0), 1),
(Ipv4Addr::new(10, 0, 0, 0), Ipv4Addr::new(10, 0, 0, 255), 256),
(Ipv4Addr::new(0, 0, 0, 0), Ipv4Addr::new(0, 0, 0, 10), 11),
];
for (start, end, expected_len) in cases {
let range = Ipv4Range::new(start, end).unwrap();
assert_eq!(range.len(), expected_len);
}
}
#[test]
fn contains_logic_v4() {
let range = Ipv4Range::new(Ipv4Addr::new(172, 16, 0, 10), Ipv4Addr::new(172, 16, 0, 20)).unwrap();
assert!(range.contains(&Ipv4Addr::new(172, 16, 0, 10)));
assert!(range.contains(&Ipv4Addr::new(172, 16, 0, 15)));
assert!(range.contains(&Ipv4Addr::new(172, 16, 0, 20)));
assert!(!range.contains(&Ipv4Addr::new(172, 16, 0, 9)));
assert!(!range.contains(&Ipv4Addr::new(172, 16, 0, 21)));
}
#[test]
fn iteration_values_v4() {
let range = Ipv4Range::new(Ipv4Addr::new(1, 1, 1, 1), Ipv4Addr::new(1, 1, 1, 3)).unwrap();
let ips: Vec<IpAddr> = range.to_iter().collect();
assert_eq!(ips, vec![
IpAddr::V4(Ipv4Addr::new(1, 1, 1, 1)),
IpAddr::V4(Ipv4Addr::new(1, 1, 1, 2)),
IpAddr::V4(Ipv4Addr::new(1, 1, 1, 3)),
]);
}
#[test]
fn max_u32_range_boundaries() {
let start = Ipv4Addr::new(255, 255, 255, 254);
let end = Ipv4Addr::new(255, 255, 255, 255);
let range = Ipv4Range::new(start, end).unwrap();
assert_eq!(range.len(), 2);
}
#[test]
fn ipv6_range_basics() {
let start = Ipv6Addr::from(100);
let end = Ipv6Addr::from(200);
let range = Ipv6Range::new(start, end).unwrap();
assert_eq!(range.len(), 101);
assert!(range.contains(&Ipv6Addr::from(150)));
assert!(!range.contains(&Ipv6Addr::from(201)));
}
#[test]
fn ipv6_large_len() {
let range = cidr_range(IpAddr::V6(Ipv6Addr::UNSPECIFIED), 64).unwrap();
assert_eq!(range.len(), 1u128 << 64);
}
#[test]
fn iteration_ipv6_small() {
let range = Ipv6Range::new(Ipv6Addr::from(1), Ipv6Addr::from(3)).unwrap();
let ips: Vec<_> = range.to_iter().collect();
assert_eq!(ips.len(), 3);
assert_eq!(ips[0], IpAddr::V6(Ipv6Addr::from(1)));
}
#[test]
fn from_str_comprehensive() {
assert_eq!("10.0.0.0/24".parse::<IpRange>().unwrap().len(), 256);
assert_eq!("192.168.1.0/24".parse::<IpRange>().unwrap().len(), 256);
assert_eq!("::1/120".parse::<IpRange>().unwrap().len(), 256);
assert_eq!("1.1.1.1-1.1.1.5".parse::<IpRange>().unwrap().len(), 5);
assert_eq!("8.8.8.8".parse::<IpRange>().unwrap().len(), 1);
}
#[test]
fn invalid_range_order() {
let v4_err = Ipv4Range::new(Ipv4Addr::new(1, 1, 1, 2), Ipv4Addr::new(1, 1, 1, 1));
assert!(matches!(v4_err, Err(IpError::InvalidRange(_, _))));
let v6_err = Ipv6Range::new(Ipv6Addr::from(2), Ipv6Addr::from(1));
assert!(matches!(v6_err, Err(IpError::InvalidRange(_, _))));
}
#[test]
fn error_formatting() {
let prefix_err = IpError::InvalidPrefix(40);
assert_eq!(format!("{prefix_err}"), "Invalid CIDR prefix: 40");
let range_err = IpError::InvalidRange(
IpAddr::V4(Ipv4Addr::new(1, 1, 1, 2)),
IpAddr::V4(Ipv4Addr::new(1, 1, 1, 1)),
);
assert!(format!("{range_err}").contains("is greater than"));
}
}
#[cfg(test)]
mod property_tests {
use super::*;
use proptest::prelude::*;
fn any_ipv4() -> impl Strategy<Value = Ipv4Addr> {
proptest::prelude::any::<u32>().prop_map(Ipv4Addr::from)
}
fn any_ipv6() -> impl Strategy<Value = Ipv6Addr> {
proptest::prelude::any::<u128>().prop_map(Ipv6Addr::from)
}
fn any_ipv4_range() -> impl Strategy<Value = Ipv4Range> {
(any_ipv4(), 0..5000u32).prop_map(|(start, len)| {
let start_u32 = u32::from(start);
let end_u32 = start_u32.saturating_add(len);
Ipv4Range::new(start, Ipv4Addr::from(end_u32)).unwrap()
})
}
fn any_ipv6_range() -> impl Strategy<Value = Ipv6Range> {
(any_ipv6(), 0..5000u128).prop_map(|(start, len)| {
let start_u128 = u128::from(start);
let end_u128 = start_u128.saturating_add(len);
Ipv6Range::new(start, Ipv6Addr::from(end_u128)).unwrap()
})
}
proptest::proptest! {
#[test]
fn ipv4_range_invariant(a in any_ipv4(), b in any_ipv4()) {
let start = std::cmp::min(a, b);
let end = std::cmp::max(a, b);
let range = Ipv4Range::new(start, end).unwrap();
prop_assert!(range.contains(&start));
prop_assert!(range.contains(&end));
prop_assert_eq!(range.len(), (u32::from(end) - u32::from(start)) as u64 + 1);
}
#[test]
fn ipv6_range_invariant(a in any_ipv6(), b in any_ipv6()) {
let start = std::cmp::min(a, b);
let end = std::cmp::max(a, b);
let range = Ipv6Range::new(start, end).unwrap();
prop_assert!(range.contains(&start));
prop_assert!(range.contains(&end));
prop_assert_eq!(range.len(), (u128::from(end) - u128::from(start)) + 1);
}
#[test]
fn ipv4_iterator_consistency(range in any_ipv4_range()) {
prop_assert_eq!(range.to_iter().count() as u64, range.len());
}
#[test]
fn ipv6_iterator_consistency(range in any_ipv6_range()) {
prop_assert_eq!(range.to_iter().count() as u128, range.len());
}
#[test]
fn cidr_v4_roundtrip(v4 in any_ipv4(), prefix in 1..=32u8) {
let range = cidr_range(IpAddr::V4(v4), prefix).unwrap();
prop_assert_eq!(range.len() as u128, 1u128 << (32 - prefix));
}
#[test]
fn cidr_v6_roundtrip(v6 in any_ipv6(), prefix in 1..=128u8) {
let range = cidr_range(IpAddr::V6(v6), prefix).unwrap();
prop_assert_eq!(range.len(), 1u128 << (128 - prefix));
}
}
}