use std::net::IpAddr;
#[cfg(feature = "host")]
use std::net::SocketAddr;
use serde::Deserialize;
use crate::Decision;
use crate::grant::PolicyMode;
#[derive(Debug, Clone, Default, Deserialize, PartialEq)]
pub struct NetworkRule {
#[serde(default)]
pub host: Option<String>,
#[serde(default)]
pub ports: Option<Vec<u16>>,
#[serde(default)]
pub cidr: Option<String>,
#[serde(default, rename = "except-ports")]
pub except_ports: Option<Vec<u16>>,
}
pub fn cidr_contains(cidr: &str, ip: IpAddr) -> bool {
cidr.parse::<cidr::IpCidr>().is_ok_and(|c| c.contains(&ip))
}
pub fn host_matches(pattern: &str, host: &str) -> bool {
if pattern == "*" {
return true;
}
if let Some(suffix) = pattern.strip_prefix("*.") {
return host.eq_ignore_ascii_case(suffix)
|| host
.to_ascii_lowercase()
.ends_with(&format!(".{}", suffix.to_ascii_lowercase()));
}
host.eq_ignore_ascii_case(pattern)
}
#[derive(Debug, Clone, Copy)]
pub struct NetworkCheck<'a> {
pub host: &'a str,
pub port: u16,
pub resolved_ips: &'a [IpAddr],
}
impl<'a> NetworkCheck<'a> {
pub fn new(host: &'a str, port: u16) -> Self {
Self {
host,
port,
resolved_ips: &[],
}
}
#[allow(dead_code)] pub fn with_resolved(host: &'a str, port: u16, resolved_ips: &'a [IpAddr]) -> Self {
Self {
host,
port,
resolved_ips,
}
}
}
pub fn rule_matches(rule: &NetworkRule, check: &NetworkCheck) -> bool {
if let Some(cidr_spec) = rule.cidr.as_deref() {
let ip_literal_match = check
.host
.parse::<IpAddr>()
.is_ok_and(|ip| cidr_contains(cidr_spec, ip));
let resolved_match = check
.resolved_ips
.iter()
.any(|ip| cidr_contains(cidr_spec, *ip));
if !ip_literal_match && !resolved_match {
return false;
}
if let Some(except) = &rule.except_ports
&& except.contains(&check.port)
{
return false;
}
} else if let Some(want_host) = rule.host.as_deref() {
if !host_matches(want_host, check.host) {
return false;
}
} else {
return false;
}
if let Some(want_ports) = rule.ports.as_deref()
&& !want_ports.contains(&check.port)
{
return false;
}
true
}
#[allow(dead_code)] pub fn decide(
mode: PolicyMode,
allow: &[NetworkRule],
deny: &[NetworkRule],
check: &NetworkCheck,
) -> Decision {
let on_match = match mode {
PolicyMode::Deny => return Decision::Deny,
PolicyMode::Open => return Decision::Allow,
PolicyMode::Ask => Decision::Ask,
PolicyMode::Allowlist => Decision::Allow,
};
if deny.iter().any(|r| rule_matches(r, check)) {
return Decision::Deny;
}
if allow.iter().any(|r| rule_matches(r, check)) {
on_match
} else {
Decision::Deny
}
}
#[cfg(feature = "host")]
#[allow(dead_code)] pub async fn resolve_host(host: &str, port: u16) -> Vec<SocketAddr> {
let target = format!("{host}:{port}");
match tokio::net::lookup_host(&target).await {
Ok(addrs) => addrs.collect(),
Err(_) => Vec::new(),
}
}
#[allow(dead_code)] pub fn any_deny_cidr_matches(deny_rules: &[NetworkRule], ip: IpAddr, port: u16) -> bool {
let ips = [ip];
let check = NetworkCheck::with_resolved("", port, &ips);
deny_rules
.iter()
.any(|rule| rule.cidr.is_some() && rule_matches(rule, &check))
}
#[cfg(feature = "host")]
#[allow(dead_code)] pub async fn first_cidr_deny_hit(
deny_rules: &[NetworkRule],
host: &str,
port: u16,
) -> Option<SocketAddr> {
if deny_rules.iter().all(|r| r.cidr.is_none()) {
return None;
}
for addr in resolve_host(host, port).await {
if any_deny_cidr_matches(deny_rules, addr.ip(), addr.port()) {
return Some(addr);
}
}
None
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn cidr_basic_ipv4() {
assert!(cidr_contains("10.0.0.0/8", "10.1.2.3".parse().unwrap()));
assert!(!cidr_contains("10.0.0.0/8", "11.0.0.0".parse().unwrap()));
assert!(cidr_contains("0.0.0.0/0", "8.8.8.8".parse().unwrap()));
}
#[test]
fn cidr_basic_ipv6() {
assert!(cidr_contains("fc00::/7", "fc00::1".parse().unwrap()));
assert!(!cidr_contains("fc00::/7", "2001:db8::1".parse().unwrap()));
}
#[test]
fn cidr_malformed_is_no_match() {
assert!(!cidr_contains("not-a-cidr", "10.0.0.1".parse().unwrap()));
assert!(!cidr_contains("10.0.0.0/99", "10.0.0.1".parse().unwrap()));
}
#[test]
fn cidr_family_mismatch_is_no_match() {
assert!(!cidr_contains("10.0.0.0/8", "::1".parse().unwrap()));
assert!(!cidr_contains("fc00::/7", "10.0.0.1".parse().unwrap()));
}
#[test]
fn host_exact_case_insensitive() {
assert!(host_matches("api.example.com", "api.example.com"));
assert!(host_matches("api.example.com", "API.EXAMPLE.COM"));
assert!(!host_matches("api.example.com", "api2.example.com"));
}
#[test]
fn host_wildcard_matches_apex_and_subdomains() {
assert!(host_matches("*.example.com", "example.com"));
assert!(host_matches("*.example.com", "api.example.com"));
assert!(host_matches("*.example.com", "a.b.example.com"));
}
#[test]
fn host_wildcard_rejects_unrelated_and_confusable_suffixes() {
assert!(!host_matches("*.example.com", "api.other.com"));
assert!(!host_matches("*.example.com", "notexample.com"));
assert!(!host_matches("*.example.com", "example.com.evil.com"));
}
fn rule(
host: Option<&str>,
ports: Option<Vec<u16>>,
cidr: Option<&str>,
except_ports: Option<Vec<u16>>,
) -> NetworkRule {
NetworkRule {
host: host.map(String::from),
ports,
cidr: cidr.map(String::from),
except_ports,
}
}
#[test]
fn rule_requires_host_or_cidr() {
let r = rule(None, None, None, None);
assert!(!rule_matches(&r, &NetworkCheck::new("example.com", 443)));
}
#[test]
fn host_rule_with_port_narrowing() {
let r = rule(Some("*.example.com"), Some(vec![443]), None, None);
assert!(rule_matches(&r, &NetworkCheck::new("api.example.com", 443)));
assert!(!rule_matches(
&r,
&NetworkCheck::new("api.example.com", 8443)
));
assert!(!rule_matches(&r, &NetworkCheck::new("api.other.com", 443)));
}
#[test]
fn cidr_rule_matches_ip_literal_host() {
let r = rule(None, None, Some("10.0.0.0/8"), None);
assert!(rule_matches(&r, &NetworkCheck::new("10.1.2.3", 80)));
assert!(!rule_matches(&r, &NetworkCheck::new("11.1.2.3", 80)));
}
#[test]
fn cidr_rule_matches_resolved_ip_when_host_is_name() {
let r = rule(None, None, Some("10.0.0.0/8"), None);
let ips = ["10.1.2.3".parse().unwrap()];
let check = NetworkCheck::with_resolved("internal.example.com", 443, &ips);
assert!(rule_matches(&r, &check));
}
#[test]
fn cidr_except_ports_carves_exceptions() {
let r = rule(None, None, Some("127.0.0.0/8"), Some(vec![3000]));
assert!(rule_matches(&r, &NetworkCheck::new("127.0.0.1", 80)));
assert!(!rule_matches(&r, &NetworkCheck::new("127.0.0.1", 3000)));
}
#[test]
fn decide_allowlist_deny_beats_allow() {
let allow = vec![rule(Some("*.example.com"), None, None, None)];
let deny = vec![rule(Some("admin.example.com"), None, None, None)];
let good = NetworkCheck::new("api.example.com", 443);
let bad = NetworkCheck::new("admin.example.com", 443);
assert_eq!(
decide(PolicyMode::Allowlist, &allow, &deny, &good),
Decision::Allow
);
assert_eq!(
decide(PolicyMode::Allowlist, &allow, &deny, &bad),
Decision::Deny
);
}
#[test]
fn decide_ask_is_bounded_by_allow_ceiling() {
let allow = vec![rule(Some("*.example.com"), None, None, None)];
let deny = vec![rule(Some("admin.example.com"), None, None, None)];
let in_ceiling = NetworkCheck::new("api.example.com", 443);
let out_of_ceiling = NetworkCheck::new("evil.com", 443);
let denied = NetworkCheck::new("admin.example.com", 443);
assert_eq!(
decide(PolicyMode::Ask, &allow, &deny, &in_ceiling),
Decision::Ask
);
assert_eq!(
decide(PolicyMode::Ask, &allow, &deny, &out_of_ceiling),
Decision::Deny
);
assert_eq!(
decide(PolicyMode::Ask, &allow, &deny, &denied),
Decision::Deny
);
}
#[test]
fn decide_modes_short_circuit() {
let allow = vec![rule(Some("example.com"), None, None, None)];
let deny = vec![];
let check = NetworkCheck::new("example.com", 443);
assert_eq!(
decide(PolicyMode::Open, &allow, &deny, &check),
Decision::Allow
);
assert_eq!(
decide(PolicyMode::Deny, &allow, &deny, &check),
Decision::Deny
);
}
#[test]
fn any_deny_cidr_matches_ip() {
let deny = vec![rule(None, None, Some("10.0.0.0/8"), None)];
assert!(any_deny_cidr_matches(
&deny,
"10.1.2.3".parse().unwrap(),
443
));
assert!(!any_deny_cidr_matches(
&deny,
"8.8.8.8".parse().unwrap(),
443
));
}
#[test]
fn any_deny_cidr_respects_except_ports() {
let deny = vec![rule(None, None, Some("127.0.0.0/8"), Some(vec![3000]))];
assert!(any_deny_cidr_matches(
&deny,
"127.0.0.1".parse().unwrap(),
80
));
assert!(!any_deny_cidr_matches(
&deny,
"127.0.0.1".parse().unwrap(),
3000
));
}
#[test]
fn any_deny_cidr_ignores_host_only_rules() {
let deny = vec![rule(Some("example.com"), None, None, None)];
assert!(!any_deny_cidr_matches(
&deny,
"10.1.2.3".parse().unwrap(),
443
));
}
#[test]
fn host_matches_star_wildcard() {
assert!(host_matches("*", "example.com"));
assert!(host_matches("*", "foo.bar.example.com"));
assert!(host_matches("*", "localhost"));
assert!(host_matches("*", "127.0.0.1"));
}
}