use ahash::AHashSet;
use ipnet::IpNet;
use std::net::{AddrParseError, IpAddr};
use std::str::FromStr;
#[derive(Clone, Debug)]
pub struct IpRules {
ip_net_list: Vec<IpNet>,
ip_set: AHashSet<IpAddr>,
}
impl IpRules {
pub fn new<T: AsRef<str>>(values: &[T]) -> Self {
let mut ip_net_list = vec![];
let mut ip_set = AHashSet::new();
for item in values {
let item_str = item.as_ref();
if let Ok(value) = IpNet::from_str(item_str) {
ip_net_list.push(value);
} else if let Ok(value) = IpAddr::from_str(item_str) {
ip_set.insert(value);
} else {
}
}
Self {
ip_net_list,
ip_set,
}
}
pub fn is_match(&self, ip: &str) -> Result<bool, AddrParseError> {
let addr = ip.parse::<IpAddr>()?;
Ok(self.is_match_addr(&addr))
}
pub fn is_match_addr(&self, ip_addr: &IpAddr) -> bool {
if self.ip_set.contains(ip_addr) {
return true;
}
self.ip_net_list.iter().any(|net| net.contains(ip_addr))
}
}
#[cfg(test)]
mod tests {
use super::*;
use pretty_assertions::assert_eq;
#[test]
fn test_ip_rules() {
let rules = IpRules::new(&[
"192.168.1.0/24", "10.0.0.1", "2001:db8::/32", "2001:db8:a::1", "not-an-ip", ]);
assert_eq!(rules.ip_net_list.len(), 2);
assert_eq!(rules.ip_set.len(), 2);
let ip_in_net_v4 = "192.168.1.100".parse().unwrap();
let exact_ip_v4 = "10.0.0.1".parse().unwrap();
let outside_ip_v4 = "192.168.2.1".parse().unwrap();
let ip_in_net_v6 = "2001:db8:dead:beef::1".parse().unwrap();
let exact_ip_v6 = "2001:db8:a::1".parse().unwrap();
let outside_ip_v6 = "2001:db9::1".parse().unwrap();
assert!(rules.is_match_addr(&ip_in_net_v4));
assert!(rules.is_match_addr(&exact_ip_v4));
assert!(!rules.is_match_addr(&outside_ip_v4));
assert!(rules.is_match_addr(&ip_in_net_v6));
assert!(rules.is_match_addr(&exact_ip_v6));
assert!(!rules.is_match_addr(&outside_ip_v6));
assert_eq!(rules.is_match("192.168.1.1"), Ok(true));
assert_eq!(rules.is_match("10.0.0.1"), Ok(true));
assert_eq!(rules.is_match("192.168.3.1"), Ok(false));
assert!(rules.is_match("999.999.999.999").is_err());
}
}