use std::collections::HashSet;
use std::net::{IpAddr, SocketAddr};
use std::sync::Arc;
use std::time::Duration;
use bytes::Bytes;
use futures::StreamExt;
use hickory_net::proto::op::{DnsRequest, Message, Query, ResponseCode};
use hickory_net::proto::rr::rdata::{A, AAAA, CNAME};
use hickory_net::proto::rr::{RData, Record, RecordType};
use hickory_net::proto::serialize::binary::{BinDecodable, BinEncodable};
use hickory_net::xfer::DnsHandle;
use tokio::sync::{OnceCell, watch};
use super::client::{Client, build_direct_client, build_tcp_client, build_udp_client};
use super::common::config::NormalizedDnsConfig;
use super::common::filter::{is_private_ipv4, is_private_ipv6};
use super::common::transport::Transport;
use super::nameserver::{read_host_dns_servers, resolve_nameservers};
use crate::policy::{Action, DomainName, NetworkPolicy};
use crate::shared::{ResolvedHostnameFamily, SharedState};
use crate::stack::GatewayIps;
const RESOLVED_HOSTNAME_MIN_TTL_SECS: u32 = 1;
const HOST_ALIAS_TTL_SECS: u32 = 60;
pub(crate) type DnsForwarderHandle = watch::Receiver<Option<Arc<DnsForwarder>>>;
pub(crate) struct DnsForwarder {
configured: Vec<ConfiguredUpstream>,
gateway_ips: Arc<HashSet<IpAddr>>,
network_policy: Arc<NetworkPolicy>,
platform_policy: Option<Arc<NetworkPolicy>>,
shared: Arc<SharedState>,
gateway: GatewayIps,
config: Arc<NormalizedDnsConfig>,
}
struct ConfiguredUpstream {
addr: SocketAddr,
udp: Client,
tcp: OnceCell<Client>,
}
enum UpstreamChoice {
Configured,
Direct(Client),
PolicyDenied,
ServFail,
}
#[derive(Debug, PartialEq, Eq)]
enum UpstreamDecision {
Configured,
Direct(SocketAddr),
PolicyDenied,
}
impl DnsForwarder {
pub(crate) async fn forward(
&self,
raw_query: &[u8],
original_dst: Option<IpAddr>,
transport: Transport,
sni: Option<&str>,
) -> Option<Bytes> {
let query_msg = Message::from_bytes(raw_query).ok()?;
let guest_id = query_msg.metadata.id;
let question = match single_question(&query_msg) {
Ok(question) => question,
Err(rcode) => return build_status_response(&query_msg, rcode),
};
let query_type = question.query_type();
let domain = question.name().to_string();
let domain = domain.trim_end_matches('.').to_owned();
if decide_dns_action(&self.network_policy, &domain, transport).is_deny() {
tracing::debug!(domain = %domain, "DNS query denied by network policy");
return build_status_response(&query_msg, ResponseCode::NXDomain);
}
if let Some(family) = inactive_query_family(query_type, self.gateway) {
tracing::debug!(
domain = %domain,
?family,
"DNS query family is inactive for this sandbox",
);
self.shared.clear_resolved_hostname(&domain, family);
return build_status_response(&query_msg, ResponseCode::NoError);
}
if is_host_alias_query(&domain)
&& let Some(response) =
synthesize_host_alias_response(&query_msg, self.gateway, query_type)
{
return Some(response);
}
let response = match self.select_upstream(original_dst, transport, sni).await {
UpstreamChoice::PolicyDenied => {
tracing::debug!(
domain = %domain,
?original_dst,
"DNS resolver denied by network policy"
);
return build_status_response(&query_msg, ResponseCode::NXDomain);
}
UpstreamChoice::ServFail => None,
UpstreamChoice::Direct(client) => self.send_query(&client, &query_msg, &domain).await,
UpstreamChoice::Configured => {
self.forward_to_configured(&query_msg, &domain, transport)
.await
}
};
let Some(mut response_msg) = response else {
return build_status_response(&query_msg, ResponseCode::ServFail);
};
if self.config.rebind_protection {
for record in &response_msg.answers {
let private_addr = match &record.data {
RData::A(a) => {
let addr = IpAddr::V4((*a).into());
is_private_ipv4((*a).into()).then_some(addr)
}
RData::AAAA(aaaa) => {
let addr = IpAddr::V6((*aaaa).into());
is_private_ipv6((*aaaa).into()).then_some(addr)
}
_ => None,
};
if private_addr.is_some_and(|addr| {
!policies_allow_rebind_address(
&self.network_policy,
self.platform_policy.as_deref(),
&self.shared,
addr,
)
}) {
tracing::debug!(
domain = %domain,
"DNS rebind protection: response contains private IP"
);
return build_status_response(&query_msg, ResponseCode::NXDomain);
}
}
}
if let Some(family) = family_for_query_type(query_type) {
if let Some((addrs, ttl)) = extract_addrs_and_ttl(&response_msg, family, &domain) {
self.shared
.cache_resolved_hostname(&domain, family, addrs, ttl);
} else {
self.shared.clear_resolved_hostname(&domain, family);
}
}
response_msg.metadata.id = guest_id;
let response_bytes = response_msg.to_bytes().ok()?;
if transport == Transport::Udp {
let max_size = query_msg.max_payload() as usize;
if response_bytes.len() > max_size {
tracing::debug!(
domain = %domain,
response_size = response_bytes.len(),
advertised = max_size,
"DNS response exceeds guest UDP buffer; setting TC=1"
);
return build_truncated_response(&query_msg).map(Bytes::from);
}
}
Some(Bytes::from(response_bytes))
}
async fn select_upstream(
&self,
original_dst: Option<IpAddr>,
transport: Transport,
sni: Option<&str>,
) -> UpstreamChoice {
match decide_upstream_with_platform(
&self.gateway_ips,
&self.network_policy,
self.platform_policy.as_deref(),
&self.shared,
original_dst,
transport,
) {
UpstreamDecision::Configured => UpstreamChoice::Configured,
UpstreamDecision::PolicyDenied => UpstreamChoice::PolicyDenied,
UpstreamDecision::Direct(addr) => {
match build_direct_client(addr, transport, sni, self.config.query_timeout).await {
Some(client) => UpstreamChoice::Direct(client),
None => UpstreamChoice::ServFail,
}
}
}
}
async fn forward_to_configured(
&self,
query_msg: &Message,
domain: &str,
transport: Transport,
) -> Option<Message> {
let total = self.configured.len();
for (index, upstream) in self.configured.iter().enumerate() {
let Some(client) = self.client_for(upstream, transport).await else {
continue;
};
if let Some(response) = self.send_query(&client, query_msg, domain).await {
return Some(response);
}
if index + 1 < total {
tracing::debug!(
domain = %domain,
upstream = %upstream.addr,
"upstream DNS unusable, trying next configured nameserver",
);
}
}
None
}
async fn send_query(
&self,
client: &Client,
query_msg: &Message,
domain: &str,
) -> Option<Message> {
let mut send = client.send(DnsRequest::from(query_msg.clone()));
match send.next().await {
Some(Ok(resp)) => Some(resp.into()),
Some(Err(e)) => {
tracing::warn!(domain = %domain, error = %e, "upstream DNS send failed");
None
}
None => {
tracing::warn!(domain = %domain, "upstream DNS closed stream without a response");
None
}
}
}
async fn client_for(
&self,
upstream: &ConfiguredUpstream,
transport: Transport,
) -> Option<Client> {
match transport {
Transport::Udp => Some(upstream.udp.clone()),
Transport::Tcp | Transport::Dot => {
let timeout = self.config.query_timeout;
let addr = upstream.addr;
upstream
.tcp
.get_or_try_init(
|| async move { build_tcp_client(addr, timeout).await.ok_or(()) },
)
.await
.ok()
.cloned()
}
}
}
#[allow(clippy::too_many_arguments)]
pub(super) fn spawn(
handle: &tokio::runtime::Handle,
config: Arc<NormalizedDnsConfig>,
gateway_ips: Arc<HashSet<IpAddr>>,
network_policy: Arc<NetworkPolicy>,
platform_policy: Option<Arc<NetworkPolicy>>,
shared: Arc<SharedState>,
gateway: GatewayIps,
) -> DnsForwarderHandle {
let (forwarder_tx, forwarder_rx) = watch::channel(None);
handle.spawn(async move {
let Some(forwarder) = Self::build(
config,
gateway_ips,
network_policy,
platform_policy,
shared,
gateway,
)
.await
else {
return;
};
let _ = forwarder_tx.send(Some(forwarder));
});
forwarder_rx
}
async fn build(
config: Arc<NormalizedDnsConfig>,
gateway_ips: Arc<HashSet<IpAddr>>,
network_policy: Arc<NetworkPolicy>,
platform_policy: Option<Arc<NetworkPolicy>>,
shared: Arc<SharedState>,
gateway: GatewayIps,
) -> Option<Arc<Self>> {
let upstreams = if !config.nameservers.is_empty() {
match resolve_nameservers(&config.nameservers).await {
Ok(s) if !s.is_empty() => s,
Ok(_) => {
tracing::error!("no configured nameservers resolved to an address");
return None;
}
Err(e) => {
tracing::error!(error = %e, "failed to resolve configured nameservers");
return None;
}
}
} else {
match read_host_dns_servers().await {
Ok(s) if !s.is_empty() => s,
Ok(_) => {
tracing::error!("no upstream DNS servers discovered from host");
return None;
}
Err(e) => {
tracing::error!(error = %e, "failed to read host DNS configuration");
return None;
}
}
};
let mut configured = Vec::with_capacity(upstreams.len());
for addr in upstreams {
let Some(udp) = build_udp_client(addr, config.query_timeout).await else {
tracing::warn!(upstream = %addr, "skipping upstream: failed to build UDP client");
continue;
};
configured.push(ConfiguredUpstream {
addr,
udp,
tcp: OnceCell::new(),
});
}
if configured.is_empty() {
tracing::error!("no upstream DNS client could be built");
return None;
}
Some(Arc::new(Self {
configured,
gateway_ips,
network_policy,
platform_policy,
shared,
gateway,
config,
}))
}
pub(crate) async fn wait(mut handle: DnsForwarderHandle) -> Option<Arc<Self>> {
if let Some(f) = handle.borrow().clone() {
return Some(f);
}
handle.changed().await.ok()?;
handle.borrow().clone()
}
#[cfg(test)]
pub(crate) async fn for_proxy_test(shared: Arc<SharedState>, gateway: GatewayIps) -> Arc<Self> {
let config = Arc::new(NormalizedDnsConfig::from_config(
crate::config::DnsConfig::default(),
));
let upstream = SocketAddr::from(([127, 0, 0, 1], 9));
let udp = build_udp_client(upstream, config.query_timeout)
.await
.expect("test UDP client should initialize");
let gateway_ips = Arc::new(
gateway
.ipv4
.map(IpAddr::V4)
.into_iter()
.chain(gateway.ipv6.map(IpAddr::V6))
.collect(),
);
Arc::new(Self {
configured: vec![ConfiguredUpstream {
addr: upstream,
udp,
tcp: OnceCell::new(),
}],
gateway_ips,
network_policy: Arc::new(NetworkPolicy::allow_all()),
platform_policy: None,
shared,
gateway,
config,
})
}
}
#[cfg(test)]
fn policy_allows_rebind_address(
policy: &NetworkPolicy,
shared: &SharedState,
addr: IpAddr,
) -> bool {
policies_allow_rebind_address(policy, None, shared, addr)
}
fn policies_allow_rebind_address(
policy: &NetworkPolicy,
platform_policy: Option<&NetworkPolicy>,
shared: &SharedState,
addr: IpAddr,
) -> bool {
[crate::policy::Protocol::Tcp, crate::policy::Protocol::Udp]
.into_iter()
.any(|protocol| {
let platform_allows = platform_policy.is_none_or(|platform| {
platform
.evaluate_egress_ip(addr, protocol, shared)
.is_allow()
});
platform_allows
&& policy.evaluate_explicit_egress_ip(addr, protocol, shared) == Some(Action::Allow)
})
}
#[cfg(test)]
fn decide_upstream(
gateway_ips: &HashSet<IpAddr>,
policy: &NetworkPolicy,
shared: &SharedState,
original_dst: Option<IpAddr>,
transport: Transport,
) -> UpstreamDecision {
decide_upstream_with_platform(gateway_ips, policy, None, shared, original_dst, transport)
}
fn decide_upstream_with_platform(
gateway_ips: &HashSet<IpAddr>,
policy: &NetworkPolicy,
platform_policy: Option<&NetworkPolicy>,
shared: &SharedState,
original_dst: Option<IpAddr>,
transport: Transport,
) -> UpstreamDecision {
let Some(dst) = original_dst else {
return UpstreamDecision::Configured;
};
if gateway_ips.contains(&dst) {
return UpstreamDecision::Configured;
}
let policy_dst = SocketAddr::new(dst, transport.upstream_port());
if platform_policy.is_some_and(|platform| {
platform
.evaluate_egress(policy_dst, transport.policy_protocol(), shared)
.is_deny()
}) || policy
.evaluate_egress(policy_dst, transport.policy_protocol(), shared)
.is_deny()
{
return UpstreamDecision::PolicyDenied;
}
UpstreamDecision::Direct(policy_dst)
}
fn decide_dns_action(policy: &NetworkPolicy, domain: &str, transport: Transport) -> Action {
match domain.parse::<DomainName>() {
Ok(canonical) => policy.evaluate_dns_query(
&canonical,
transport.policy_protocol(),
transport.upstream_port(),
),
Err(_) => policy.evaluate_dns_query_without_name(
transport.policy_protocol(),
transport.upstream_port(),
),
}
}
fn build_status_response(query: &Message, rcode: ResponseCode) -> Option<Bytes> {
let mut response = Message::response(query.metadata.id, query.metadata.op_code);
response.metadata.recursion_desired = query.metadata.recursion_desired;
response.metadata.response_code = rcode;
response.metadata.recursion_available = true;
if let Some(q) = query.queries.first() {
response.add_query(q.clone());
}
response.to_bytes().ok().map(Bytes::from)
}
fn single_question(query: &Message) -> Result<&Query, ResponseCode> {
if query.queries.len() == 1 {
return Ok(&query.queries[0]);
}
Err(ResponseCode::FormErr)
}
fn family_for_query_type(query_type: RecordType) -> Option<ResolvedHostnameFamily> {
match query_type {
RecordType::A => Some(ResolvedHostnameFamily::Ipv4),
RecordType::AAAA => Some(ResolvedHostnameFamily::Ipv6),
_ => None,
}
}
fn inactive_query_family(
query_type: RecordType,
gateway: GatewayIps,
) -> Option<ResolvedHostnameFamily> {
match query_type {
RecordType::A if gateway.ipv4.is_none() => Some(ResolvedHostnameFamily::Ipv4),
RecordType::AAAA if gateway.ipv6.is_none() => Some(ResolvedHostnameFamily::Ipv6),
_ => None,
}
}
fn extract_addrs_and_ttl(
response: &Message,
family: ResolvedHostnameFamily,
query_name: &str,
) -> Option<(Vec<IpAddr>, Duration)> {
if response.metadata.response_code != ResponseCode::NoError {
return None;
}
let mut eligible_names = HashSet::from([normalize_dns_name(query_name)]);
let mut ttl: Option<Duration> = None;
let mut changed = true;
while changed {
changed = false;
for record in &response.answers {
let owner = normalize_dns_name(&record.name.to_string());
if !eligible_names.contains(&owner) {
continue;
}
if let RData::CNAME(CNAME(canonical)) = &record.data {
let record_ttl = dns_record_ttl(record.ttl);
ttl = Some(ttl.map_or(record_ttl, |current| current.min(record_ttl)));
changed |= eligible_names.insert(normalize_dns_name(&canonical.to_string()));
}
}
}
let mut addrs = Vec::new();
for record in &response.answers {
if !eligible_names.contains(&normalize_dns_name(&record.name.to_string())) {
continue;
}
let addr = match (family, &record.data) {
(ResolvedHostnameFamily::Ipv4, RData::A(a)) => IpAddr::V4((*a).into()),
(ResolvedHostnameFamily::Ipv6, RData::AAAA(aaaa)) => IpAddr::V6((*aaaa).into()),
_ => continue,
};
addrs.push(addr);
let record_ttl = dns_record_ttl(record.ttl);
ttl = Some(ttl.map_or(record_ttl, |current| current.min(record_ttl)));
}
if addrs.is_empty() {
None
} else {
ttl.map(|ttl| (addrs, ttl))
}
}
fn dns_record_ttl(ttl: u32) -> Duration {
Duration::from_secs(u64::from(ttl.max(RESOLVED_HOSTNAME_MIN_TTL_SECS)))
}
fn normalize_dns_name(name: &str) -> String {
name.trim_end_matches('.').to_ascii_lowercase()
}
fn is_host_alias_query(query_name: &str) -> bool {
query_name
.trim_end_matches('.')
.eq_ignore_ascii_case(crate::HOST_ALIAS)
}
fn synthesize_host_alias_response(
query: &Message,
gateway: GatewayIps,
qtype: RecordType,
) -> Option<Bytes> {
let question = query.queries.first()?;
let name = question.name().clone();
let rdata = match qtype {
RecordType::A => RData::A(A::from(gateway.ipv4?)),
RecordType::AAAA => RData::AAAA(AAAA::from(gateway.ipv6?)),
_ => return None,
};
let mut response = Message::response(query.metadata.id, query.metadata.op_code);
response.metadata.recursion_desired = query.metadata.recursion_desired;
response.metadata.response_code = ResponseCode::NoError;
response.metadata.recursion_available = true;
response.metadata.authoritative = true;
response.add_query(question.clone());
response.add_answer(Record::from_rdata(name, HOST_ALIAS_TTL_SECS, rdata));
response.to_bytes().ok().map(Bytes::from)
}
fn build_truncated_response(query: &Message) -> Option<Vec<u8>> {
let mut response = Message::response(query.metadata.id, query.metadata.op_code);
response.metadata.recursion_desired = query.metadata.recursion_desired;
response.metadata.response_code = ResponseCode::NoError;
response.metadata.recursion_available = true;
response.metadata.truncation = true;
if let Some(q) = query.queries.first() {
response.add_query(q.clone());
}
response.to_bytes().ok()
}
#[cfg(test)]
mod tests {
use super::*;
use crate::policy::{Action, Destination, NetworkProfile, Protocol, Rule};
use hickory_net::proto::op::{Edns, MessageType, OpCode, Query};
use hickory_net::proto::rr::{DNSClass, Name, RecordType};
use std::net::Ipv4Addr;
use std::sync::atomic::{AtomicUsize, Ordering};
fn make_query(name: &str, qtype: RecordType) -> Message {
let mut msg = Message::new(0x4242, MessageType::Query, OpCode::Query);
msg.metadata.recursion_desired = true;
let parsed = Name::from_ascii(name).expect("valid dns name");
let mut q = Query::new();
q.set_name(parsed);
q.set_query_type(qtype);
q.set_query_class(DNSClass::IN);
msg.add_query(q);
msg
}
async fn blackhole_udp() -> SocketAddr {
let sock = tokio::net::UdpSocket::bind("127.0.0.1:0").await.unwrap();
let addr = sock.local_addr().unwrap();
tokio::spawn(async move {
let mut buf = [0u8; 4096];
loop {
let _ = sock.recv_from(&mut buf).await;
}
});
addr
}
async fn responding_udp(answer_ip: Ipv4Addr) -> (SocketAddr, Arc<AtomicUsize>) {
let sock = tokio::net::UdpSocket::bind("127.0.0.1:0").await.unwrap();
let addr = sock.local_addr().unwrap();
let hits = Arc::new(AtomicUsize::new(0));
let seen = Arc::clone(&hits);
tokio::spawn(async move {
let mut buf = [0u8; 4096];
loop {
let Ok((len, from)) = sock.recv_from(&mut buf).await else {
continue;
};
seen.fetch_add(1, Ordering::SeqCst);
let Ok(query) = Message::from_bytes(&buf[..len]) else {
continue;
};
let mut resp =
Message::new(query.metadata.id, MessageType::Response, OpCode::Query);
resp.metadata.recursion_desired = query.metadata.recursion_desired;
resp.metadata.recursion_available = true;
if let Some(q) = query.queries.first() {
resp.add_query(q.clone());
resp.answers.push(Record::from_rdata(
q.name().clone(),
60,
RData::A(A::from(answer_ip)),
));
}
if let Ok(bytes) = resp.to_bytes() {
let _ = sock.send_to(&bytes, from).await;
}
}
});
(addr, hits)
}
async fn forwarder_over(upstreams: &[SocketAddr]) -> Arc<DnsForwarder> {
let config = Arc::new(NormalizedDnsConfig {
rebind_protection: false,
nameservers: Vec::new(),
query_timeout: Duration::from_millis(300),
});
let mut configured = Vec::new();
for addr in upstreams {
configured.push(ConfiguredUpstream {
addr: *addr,
udp: build_udp_client(*addr, config.query_timeout)
.await
.expect("udp client"),
tcp: OnceCell::new(),
});
}
let gateway_ip: IpAddr = "10.0.0.1".parse().unwrap();
Arc::new(DnsForwarder {
configured,
gateway_ips: Arc::new(HashSet::from([gateway_ip])),
network_policy: Arc::new(NetworkPolicy::from_profiles([NetworkProfile::Public])),
platform_policy: None,
shared: Arc::new(SharedState::new(4)),
gateway: GatewayIps {
ipv4: Some("10.0.0.1".parse().unwrap()),
ipv6: None,
},
config,
})
}
async fn resolve_via_gateway(forwarder: &DnsForwarder) -> Option<Ipv4Addr> {
let query = make_query("example.com.", RecordType::A);
let raw = query.to_bytes().expect("encode query");
let gateway: IpAddr = "10.0.0.1".parse().unwrap();
let bytes = forwarder
.forward(&raw, Some(gateway), Transport::Udp, None)
.await?;
let msg = Message::from_bytes(&bytes).expect("parse response");
if msg.metadata.response_code != ResponseCode::NoError {
return None;
}
msg.answers.iter().find_map(|r| match &r.data {
RData::A(a) => Some(Ipv4Addr::from(*a)),
_ => None,
})
}
#[tokio::test]
async fn configured_upstream_falls_over_to_the_next_on_timeout() {
let dead = blackhole_udp().await;
let (live, live_hits) = responding_udp(Ipv4Addr::new(93, 184, 216, 34)).await;
let forwarder = forwarder_over(&[dead, live]).await;
assert_eq!(
resolve_via_gateway(&forwarder).await,
Some(Ipv4Addr::new(93, 184, 216, 34)),
"a stalled first upstream must fall over to the next one"
);
assert_eq!(
live_hits.load(Ordering::SeqCst),
1,
"the working upstream should have been queried exactly once"
);
}
#[tokio::test]
async fn first_usable_upstream_answers_without_consulting_the_rest() {
let (first, first_hits) = responding_udp(Ipv4Addr::new(198, 51, 100, 7)).await;
let (second, second_hits) = responding_udp(Ipv4Addr::new(198, 51, 100, 8)).await;
let forwarder = forwarder_over(&[first, second]).await;
assert_eq!(
resolve_via_gateway(&forwarder).await,
Some(Ipv4Addr::new(198, 51, 100, 7)),
"the answer must come from the first upstream"
);
assert_eq!(
first_hits.load(Ordering::SeqCst),
1,
"a query answered by the first upstream must not be re-sent"
);
assert_eq!(
second_hits.load(Ordering::SeqCst),
0,
"later upstreams must not be consulted once one answers"
);
}
#[tokio::test]
async fn configured_upstreams_fall_over_past_several_dead_servers() {
let first = blackhole_udp().await;
let second = blackhole_udp().await;
let (live, _) = responding_udp(Ipv4Addr::new(203, 0, 113, 5)).await;
let forwarder = forwarder_over(&[first, second, live]).await;
assert_eq!(
resolve_via_gateway(&forwarder).await,
Some(Ipv4Addr::new(203, 0, 113, 5))
);
}
#[tokio::test]
async fn all_upstreams_unusable_yields_servfail() {
let first = blackhole_udp().await;
let second = blackhole_udp().await;
let forwarder = forwarder_over(&[first, second]).await;
let query = make_query("example.com.", RecordType::A);
let raw = query.to_bytes().expect("encode query");
let gateway: IpAddr = "10.0.0.1".parse().unwrap();
let bytes = forwarder
.forward(&raw, Some(gateway), Transport::Udp, None)
.await
.expect("a synthesized response");
let msg = Message::from_bytes(&bytes).expect("parse response");
assert_eq!(msg.metadata.response_code, ResponseCode::ServFail);
assert_eq!(msg.metadata.id, 0x4242, "guest transaction id is preserved");
}
fn make_response(query: &Message) -> Message {
let mut response = Message::response(query.metadata.id, query.metadata.op_code);
response.metadata.response_code = ResponseCode::NoError;
response.metadata.recursion_available = true;
response.add_query(query.queries[0].clone());
response
}
#[test]
fn rebind_filter_allows_private_answers_for_private_profile() {
let policy = NetworkPolicy::from_profiles([NetworkProfile::Private]);
let shared = SharedState::new(4);
assert!(policy_allows_rebind_address(
&policy,
&shared,
"10.20.30.40".parse().unwrap()
));
}
#[test]
fn platform_public_floor_rejects_tenant_allowed_private_dns_answer() {
let tenant = NetworkPolicy::from_profiles([NetworkProfile::Private]);
let platform = NetworkPolicy::from_profiles([NetworkProfile::Public]);
let shared = SharedState::new(4);
assert!(!policies_allow_rebind_address(
&tenant,
Some(&platform),
&shared,
"10.0.0.7".parse().unwrap(),
));
}
#[test]
fn rebind_filter_rejects_private_answers_for_public_profile() {
let policy = NetworkPolicy::from_profiles([NetworkProfile::Public]);
let shared = SharedState::new(4);
assert!(!policy_allows_rebind_address(
&policy,
&shared,
"10.20.30.40".parse().unwrap()
));
}
#[test]
fn rebind_filter_rejects_unspecified_answers_for_public_profile() {
let policy = NetworkPolicy::from_profiles([NetworkProfile::Public]);
let shared = SharedState::new(4);
for addr in ["0.0.0.0", "::"] {
assert!(
!policy_allows_rebind_address(&policy, &shared, addr.parse().unwrap()),
"expected {addr} to remain blocked by rebind protection"
);
}
}
#[test]
fn rebind_filter_remains_enabled_for_allow_all_default() {
let policy = NetworkPolicy::allow_all();
let shared = SharedState::new(4);
assert!(!policy_allows_rebind_address(
&policy,
&shared,
"10.20.30.40".parse().unwrap()
));
}
#[test]
fn rebind_filter_does_not_treat_explicit_any_as_private_intent() {
let policy = NetworkPolicy {
default_egress: Action::Deny,
default_ingress: Action::Allow,
rules: vec![Rule::allow_egress(Destination::Any)],
};
let shared = SharedState::new(4);
assert!(!policy_allows_rebind_address(
&policy,
&shared,
"10.20.30.40".parse().unwrap()
));
}
#[test]
fn rebind_filter_honors_ordered_deny_before_private_allow() {
let mut policy = NetworkPolicy::from_profiles([NetworkProfile::Private]);
policy.rules.insert(
0,
Rule {
direction: crate::policy::Direction::Egress,
destination: Destination::Cidr("10.20.30.40/32".parse().unwrap()),
protocols: Vec::new(),
ports: Vec::new(),
action: Action::Deny,
},
);
let shared = SharedState::new(4);
assert!(!policy_allows_rebind_address(
&policy,
&shared,
"10.20.30.40".parse().unwrap()
));
}
#[test]
fn build_status_response_preserves_header_and_question() {
let query = make_query("slack.com.", RecordType::AAAA);
let bytes = build_status_response(&query, ResponseCode::Refused).expect("built");
let msg = Message::from_bytes(&bytes).expect("parse response");
assert_eq!(msg.metadata.id, 0x4242);
assert_eq!(msg.metadata.response_code, ResponseCode::Refused);
assert_eq!(msg.metadata.message_type, MessageType::Response);
assert_eq!(msg.metadata.op_code, OpCode::Query);
assert!(msg.metadata.recursion_desired);
assert!(msg.metadata.recursion_available);
assert_eq!(msg.queries.len(), 1);
assert_eq!(msg.queries[0].query_type(), RecordType::AAAA);
assert_eq!(msg.answers.len(), 0);
}
#[test]
fn build_status_response_servfail_variant() {
let query = make_query("example.com.", RecordType::A);
let bytes = build_status_response(&query, ResponseCode::ServFail).expect("built");
let msg = Message::from_bytes(&bytes).expect("parse response");
assert_eq!(msg.metadata.response_code, ResponseCode::ServFail);
assert_eq!(msg.answers.len(), 0);
}
#[test]
fn single_question_rejects_multi_question_packets() {
let mut query = make_query("example.com.", RecordType::A);
let mut extra = Query::new();
extra.set_name(Name::from_ascii("other.example.").unwrap());
extra.set_query_type(RecordType::AAAA);
extra.set_query_class(DNSClass::IN);
query.add_query(extra);
assert!(matches!(
single_question(&query),
Err(ResponseCode::FormErr)
));
let bytes = build_status_response(&query, ResponseCode::FormErr).expect("built");
let msg = Message::from_bytes(&bytes).expect("parse response");
assert_eq!(msg.metadata.response_code, ResponseCode::FormErr);
assert_eq!(msg.queries.len(), 1);
assert_eq!(msg.queries[0].query_type(), RecordType::A);
}
#[test]
fn build_status_response_nxdomain_variant() {
let query = make_query("example.com.", RecordType::A);
let bytes = build_status_response(&query, ResponseCode::NXDomain).expect("built");
let msg = Message::from_bytes(&bytes).expect("parse response");
assert_eq!(msg.metadata.response_code, ResponseCode::NXDomain);
assert_eq!(msg.answers.len(), 0);
assert_eq!(msg.queries.len(), 1);
}
#[test]
fn build_status_response_noerror_variant_is_nodata() {
let query = make_query("example.com.", RecordType::AAAA);
let bytes = build_status_response(&query, ResponseCode::NoError).expect("built");
let msg = Message::from_bytes(&bytes).expect("parse response");
assert_eq!(msg.metadata.response_code, ResponseCode::NoError);
assert_eq!(msg.answers.len(), 0);
assert_eq!(msg.queries.len(), 1);
assert_eq!(msg.queries[0].query_type(), RecordType::AAAA);
}
#[test]
fn build_truncated_response_sets_tc_and_keeps_question() {
let query = make_query("example.com.", RecordType::TXT);
let bytes = build_truncated_response(&query).expect("built");
let msg = Message::from_bytes(&bytes).expect("parse response");
assert_eq!(msg.metadata.id, 0x4242);
assert_eq!(msg.metadata.message_type, MessageType::Response);
assert_eq!(msg.metadata.response_code, ResponseCode::NoError);
assert!(msg.metadata.truncation, "TC bit should be set");
assert_eq!(msg.queries.len(), 1);
assert_eq!(msg.queries[0].query_type(), RecordType::TXT);
assert!(msg.answers.is_empty());
}
#[test]
fn edns_opt_round_trips_through_wire() {
let mut query = make_query("example.com.", RecordType::A);
let mut edns = Edns::new();
edns.set_max_payload(4096);
edns.set_dnssec_ok(true);
edns.set_version(0);
query.edns = Some(edns);
let bytes = query.to_bytes().expect("serialize");
let parsed = Message::from_bytes(&bytes).expect("parse");
let opt = parsed.edns.as_ref().expect("OPT preserved");
assert_eq!(opt.max_payload(), 4096);
assert!(opt.flags().dnssec_ok, "DO bit preserved");
assert_eq!(parsed.max_payload(), 4096);
}
#[test]
fn max_payload_defaults_to_512_without_opt() {
let query = make_query("example.com.", RecordType::A);
assert!(query.edns.is_none());
assert_eq!(query.max_payload(), 512);
}
#[test]
fn inactive_query_family_detects_missing_ipv6_gateway() {
let gateway = GatewayIps {
ipv4: Some(std::net::Ipv4Addr::new(172, 16, 0, 1)),
ipv6: None,
};
assert_eq!(
inactive_query_family(RecordType::AAAA, gateway),
Some(ResolvedHostnameFamily::Ipv6)
);
assert_eq!(inactive_query_family(RecordType::A, gateway), None);
}
#[test]
fn inactive_query_family_detects_missing_ipv4_gateway() {
let gateway = GatewayIps {
ipv4: None,
ipv6: Some("fd42:6d73:62::1".parse().unwrap()),
};
assert_eq!(
inactive_query_family(RecordType::A, gateway),
Some(ResolvedHostnameFamily::Ipv4)
);
assert_eq!(inactive_query_family(RecordType::AAAA, gateway), None);
}
#[test]
fn inactive_query_family_ignores_non_address_queries() {
let gateway = GatewayIps {
ipv4: None,
ipv6: None,
};
assert_eq!(inactive_query_family(RecordType::MX, gateway), None);
}
#[test]
fn extract_addrs_and_ttl_ignores_unrelated_answers() {
let query = make_query("example.com.", RecordType::A);
let mut response = make_response(&query);
response.add_answer(Record::from_rdata(
Name::from_ascii("example.com.").unwrap(),
30,
RData::A(A::from(std::net::Ipv4Addr::new(93, 184, 216, 34))),
));
response.add_answer(Record::from_rdata(
Name::from_ascii("unrelated.example.").unwrap(),
10,
RData::A(A::from(std::net::Ipv4Addr::new(198, 51, 100, 7))),
));
let (addrs, ttl) =
extract_addrs_and_ttl(&response, ResolvedHostnameFamily::Ipv4, "example.com").unwrap();
assert_eq!(
addrs,
vec![IpAddr::V4(std::net::Ipv4Addr::new(93, 184, 216, 34))]
);
assert_eq!(ttl, Duration::from_secs(30));
}
#[test]
fn extract_addrs_and_ttl_follows_cname_chain() {
let query = make_query("example.com.", RecordType::A);
let mut response = make_response(&query);
response.add_answer(Record::from_rdata(
Name::from_ascii("example.com.").unwrap(),
20,
RData::CNAME(CNAME(Name::from_ascii("cdn.example.net.").unwrap())),
));
response.add_answer(Record::from_rdata(
Name::from_ascii("cdn.example.net.").unwrap(),
40,
RData::A(A::from(std::net::Ipv4Addr::new(203, 0, 113, 10))),
));
response.add_answer(Record::from_rdata(
Name::from_ascii("other.example.net.").unwrap(),
1,
RData::A(A::from(std::net::Ipv4Addr::new(203, 0, 113, 11))),
));
let (addrs, ttl) =
extract_addrs_and_ttl(&response, ResolvedHostnameFamily::Ipv4, "example.com").unwrap();
assert_eq!(
addrs,
vec![IpAddr::V4(std::net::Ipv4Addr::new(203, 0, 113, 10))]
);
assert_eq!(ttl, Duration::from_secs(20));
}
#[test]
fn extract_addrs_and_ttl_ignores_error_responses() {
let query = make_query("example.com.", RecordType::A);
let mut response = make_response(&query);
response.metadata.response_code = ResponseCode::NXDomain;
response.add_answer(Record::from_rdata(
Name::from_ascii("example.com.").unwrap(),
30,
RData::A(A::from(std::net::Ipv4Addr::new(93, 184, 216, 34))),
));
assert!(
extract_addrs_and_ttl(&response, ResolvedHostnameFamily::Ipv4, "example.com").is_none()
);
}
fn gateway_set() -> HashSet<IpAddr> {
HashSet::from([
IpAddr::V4(std::net::Ipv4Addr::new(10, 0, 0, 1)),
IpAddr::V6(std::net::Ipv6Addr::LOCALHOST),
])
}
#[test]
fn decide_upstream_configured_when_dst_is_gateway_v4() {
let gw = gateway_set();
let shared = SharedState::new(4);
let policy = NetworkPolicy::allow_all();
let dst = Some(IpAddr::V4(std::net::Ipv4Addr::new(10, 0, 0, 1)));
assert_eq!(
decide_upstream(&gw, &policy, &shared, dst, Transport::Udp),
UpstreamDecision::Configured
);
}
#[test]
fn platform_public_floor_denies_private_direct_resolver() {
let gateways = gateway_set();
let shared = SharedState::new(4);
let tenant = NetworkPolicy::allow_all();
let platform = NetworkPolicy::from_profiles([NetworkProfile::Public]);
let dst = Some(IpAddr::V4("10.0.0.53".parse().unwrap()));
assert_eq!(
decide_upstream_with_platform(
&gateways,
&tenant,
Some(&platform),
&shared,
dst,
Transport::Udp,
),
UpstreamDecision::PolicyDenied
);
}
#[test]
fn decide_upstream_configured_when_dst_is_gateway_v6() {
let gw = gateway_set();
let shared = SharedState::new(4);
let policy = NetworkPolicy::allow_all();
let dst = Some(IpAddr::V6(std::net::Ipv6Addr::LOCALHOST));
assert_eq!(
decide_upstream(&gw, &policy, &shared, dst, Transport::Tcp),
UpstreamDecision::Configured
);
}
#[test]
fn decide_upstream_configured_when_dst_unknown() {
let gw = gateway_set();
let shared = SharedState::new(4);
let policy = NetworkPolicy::allow_all();
assert_eq!(
decide_upstream(&gw, &policy, &shared, None, Transport::Udp),
UpstreamDecision::Configured
);
}
#[test]
fn decide_upstream_direct_when_dst_external_and_policy_allows() {
let gw = gateway_set();
let shared = SharedState::new(4);
let policy = NetworkPolicy::allow_all();
let dst = Some(IpAddr::V4(std::net::Ipv4Addr::new(1, 1, 1, 1)));
assert_eq!(
decide_upstream(&gw, &policy, &shared, dst, Transport::Udp),
UpstreamDecision::Direct(SocketAddr::from(([1, 1, 1, 1], 53)))
);
}
#[test]
fn decide_upstream_policy_denied_when_policy_denies_resolver() {
let gw = gateway_set();
let shared = SharedState::new(4);
let policy = NetworkPolicy::default();
let dst = Some(IpAddr::V4(std::net::Ipv4Addr::new(192, 168, 1, 53)));
assert_eq!(
decide_upstream(&gw, &policy, &shared, dst, Transport::Udp),
UpstreamDecision::PolicyDenied
);
}
#[test]
fn decide_upstream_policy_denied_when_policy_denies_all() {
let gw = gateway_set();
let shared = SharedState::new(4);
let policy = NetworkPolicy::none();
let dst = Some(IpAddr::V4(std::net::Ipv4Addr::new(1, 1, 1, 1)));
assert_eq!(
decide_upstream(&gw, &policy, &shared, dst, Transport::Tcp),
UpstreamDecision::PolicyDenied
);
let gw_dst = Some(IpAddr::V4(std::net::Ipv4Addr::new(10, 0, 0, 1)));
assert_eq!(
decide_upstream(&gw, &policy, &shared, gw_dst, Transport::Tcp),
UpstreamDecision::Configured
);
}
#[test]
fn decide_upstream_uses_correct_transport_protocol() {
use crate::policy::{Action, Destination, Direction, Rule};
let gw = gateway_set();
let shared = SharedState::new(4);
let dst_ip = std::net::Ipv4Addr::new(8, 8, 8, 8);
let policy = NetworkPolicy {
default_egress: Action::Allow,
default_ingress: Action::Allow,
rules: vec![Rule {
direction: Direction::Egress,
destination: Destination::Cidr("8.8.8.8/32".parse().unwrap()),
protocols: vec![Protocol::Tcp],
ports: vec![],
action: Action::Deny,
}],
};
let dst = Some(IpAddr::V4(dst_ip));
assert_eq!(
decide_upstream(&gw, &policy, &shared, dst, Transport::Udp),
UpstreamDecision::Direct(SocketAddr::from(([8, 8, 8, 8], 53)))
);
assert_eq!(
decide_upstream(&gw, &policy, &shared, dst, Transport::Tcp),
UpstreamDecision::PolicyDenied
);
}
#[test]
fn decide_upstream_dot_configured_when_dst_is_gateway() {
let gw = gateway_set();
let shared = SharedState::new(4);
let policy = NetworkPolicy::allow_all();
let dst = Some(IpAddr::V4(std::net::Ipv4Addr::new(10, 0, 0, 1)));
assert_eq!(
decide_upstream(&gw, &policy, &shared, dst, Transport::Dot),
UpstreamDecision::Configured
);
}
#[test]
fn decide_upstream_dot_direct_targets_port_853() {
let gw = gateway_set();
let shared = SharedState::new(4);
let policy = NetworkPolicy::allow_all();
let dst = Some(IpAddr::V4(std::net::Ipv4Addr::new(1, 1, 1, 1)));
assert_eq!(
decide_upstream(&gw, &policy, &shared, dst, Transport::Dot),
UpstreamDecision::Direct(SocketAddr::from(([1, 1, 1, 1], 853))),
);
}
#[test]
fn decide_upstream_dot_policy_denied_when_policy_denies_853() {
use crate::policy::{Action, Destination, Direction, Rule};
let gw = gateway_set();
let shared = SharedState::new(4);
let policy = NetworkPolicy {
default_egress: Action::Allow,
default_ingress: Action::Allow,
rules: vec![Rule {
direction: Direction::Egress,
destination: Destination::Cidr("1.1.1.1/32".parse().unwrap()),
protocols: vec![Protocol::Tcp],
ports: vec![],
action: Action::Deny,
}],
};
let dst = Some(IpAddr::V4(std::net::Ipv4Addr::new(1, 1, 1, 1)));
assert_eq!(
decide_upstream(&gw, &policy, &shared, dst, Transport::Dot),
UpstreamDecision::PolicyDenied
);
}
#[test]
fn decide_dns_action_allows_under_default_allow() {
let policy = NetworkPolicy::allow_all();
assert_eq!(
decide_dns_action(&policy, "example.com", Transport::Udp),
Action::Allow
);
}
#[test]
fn decide_dns_action_denies_under_deny_by_default() {
let policy = NetworkPolicy::none();
assert_eq!(
decide_dns_action(&policy, "example.com", Transport::Udp),
Action::Deny
);
assert_eq!(
decide_dns_action(&policy, "example.com", Transport::Tcp),
Action::Deny
);
assert_eq!(
decide_dns_action(&policy, "example.com", Transport::Dot),
Action::Deny
);
}
#[test]
fn decide_dns_action_any_rule_grants_dns_when_protocol_and_port_match() {
use crate::policy::{Destination, Direction, PortRange, Rule};
let policy = NetworkPolicy {
default_egress: Action::Deny,
default_ingress: Action::Allow,
rules: vec![Rule {
direction: Direction::Egress,
destination: Destination::Any,
protocols: vec![Protocol::Udp],
ports: vec![PortRange::single(53)],
action: Action::Allow,
}],
};
assert_eq!(
decide_dns_action(&policy, "example.com", Transport::Udp),
Action::Allow
);
assert_eq!(
decide_dns_action(&policy, "example.com", Transport::Tcp),
Action::Deny
);
}
#[test]
fn decide_dns_action_dot_uses_tcp_and_port_853() {
use crate::policy::{Destination, Direction, PortRange, Rule};
let policy_853 = NetworkPolicy {
default_egress: Action::Deny,
default_ingress: Action::Allow,
rules: vec![Rule {
direction: Direction::Egress,
destination: Destination::Any,
protocols: vec![Protocol::Tcp],
ports: vec![PortRange::single(853)],
action: Action::Allow,
}],
};
assert_eq!(
decide_dns_action(&policy_853, "example.com", Transport::Dot),
Action::Allow
);
let policy_53 = NetworkPolicy {
default_egress: Action::Deny,
default_ingress: Action::Allow,
rules: vec![Rule {
direction: Direction::Egress,
destination: Destination::Any,
protocols: vec![Protocol::Tcp],
ports: vec![PortRange::single(53)],
action: Action::Allow,
}],
};
assert_eq!(
decide_dns_action(&policy_53, "example.com", Transport::Dot),
Action::Deny
);
}
#[test]
fn decide_dns_action_unparseable_name_takes_nameless_path() {
let policy = NetworkPolicy::allow_all()
.deny_domain("evil.com")
.expect("valid name");
assert_eq!(
decide_dns_action(&policy, "", Transport::Udp),
Action::Allow
);
let deny = NetworkPolicy::none();
assert_eq!(decide_dns_action(&deny, "", Transport::Udp), Action::Deny);
}
#[test]
fn decide_dns_action_domain_rule_denies_specific_name() {
let policy = NetworkPolicy::allow_all()
.deny_domain("evil.com")
.expect("valid name");
assert_eq!(
decide_dns_action(&policy, "evil.com", Transport::Udp),
Action::Deny
);
assert_eq!(
decide_dns_action(&policy, "good.com", Transport::Udp),
Action::Allow
);
}
}