remove address hijack

This commit is contained in:
Luna Yao
2026-04-27 23:45:03 +02:00
parent 87f2905360
commit 7d536d3353
9 changed files with 91 additions and 743 deletions
+1 -19
View File
@@ -4,7 +4,6 @@ use dashmap::DashMap;
use hmac::{Hmac, Mac};
use sha2::Sha256;
use socket2::Protocol;
use std::sync::RwLock;
use std::{
collections::{HashMap, hash_map::DefaultHasher},
hash::Hasher,
@@ -22,10 +21,7 @@ use super::{
stun::{StunInfoCollector, StunInfoCollectorTrait},
};
#[cfg(feature = "magic-dns")]
use crate::dns::{
config::{DnsConfigLoaderExt, DnsExportConfig, DnsGlobalCtxExt},
server::DnsServer,
};
use crate::dns::config::{DnsConfigLoaderExt, DnsExportConfig, DnsGlobalCtxExt};
use crate::{
common::{
config::ProxyNetworkConfig, shrink_dashmap, stats_manager::StatsManager,
@@ -217,9 +213,6 @@ pub struct GlobalCtx {
hostname: Mutex<String>,
#[cfg(feature = "magic-dns")]
dns_server: RwLock<Option<Arc<DnsServer>>>,
stun_info_collection: Mutex<Arc<dyn StunInfoCollectorTrait>>,
running_listeners: Mutex<Vec<url::Url>>,
@@ -318,9 +311,6 @@ impl GlobalCtx {
stun_info_collector.clone(),
)))),
#[cfg(feature = "magic-dns")]
dns_server: RwLock::new(None),
hostname: Mutex::new(hostname),
stun_info_collection: Mutex::new(stun_info_collector),
@@ -723,14 +713,6 @@ impl GlobalCtx {
#[cfg(feature = "magic-dns")]
impl DnsGlobalCtxExt for GlobalCtx {
fn dns_server(&self) -> Option<Arc<DnsServer>> {
self.dns_server.read().unwrap().clone()
}
fn set_dns_server(&self, dns: Option<Arc<DnsServer>>) {
*self.dns_server.write().unwrap() = dns;
}
fn dns_self_zone(&self) -> crate::dns::config::zone::ZoneConfig {
let dns = self.config.get_dns();
let mut hostname = dns.name.to_string();
+3 -27
View File
@@ -1,15 +1,12 @@
use crate::dns::config::policy::DnsPolicyConfig;
use crate::dns::config::zone::ZoneConfig;
use crate::dns::config::{DNS_DEFAULT_ADDRESS, DNS_DEFAULT_DOMAIN};
use crate::dns::server::DnsServer;
use crate::dns::config::{DNS_DEFAULT_ADDRESSES, DNS_DEFAULT_DOMAIN};
use crate::dns::utils::addr::NameServerAddrGroup;
use crate::proto::dns::GetExportConfigResponse;
use derivative::Derivative;
use hickory_net::xfer::Protocol;
use hickory_proto::rr::LowerName;
use serde::{Deserialize, Deserializer, Serialize};
use serde::{Deserialize, Serialize};
use std::collections::HashMap;
use std::sync::Arc;
#[derive(Derivative, Debug, Clone, Deserialize, Serialize, PartialEq)]
#[derivative(Default)]
@@ -22,30 +19,11 @@ pub struct DnsConfig {
pub name: LowerName,
#[derivative(Default(value = "DNS_DEFAULT_DOMAIN.clone()"))]
pub domain: LowerName,
#[derivative(Default(value = "vec![DNS_DEFAULT_ADDRESS].into()"))]
#[serde(deserialize_with = "DnsConfig::deserialize_addresses")]
#[derivative(Default(value = "DNS_DEFAULT_ADDRESSES.clone()"))]
pub addresses: NameServerAddrGroup,
pub listeners: NameServerAddrGroup,
}
impl DnsConfig {
pub fn deserialize_addresses<'de, D>(deserializer: D) -> Result<NameServerAddrGroup, D::Error>
where
D: Deserializer<'de>,
{
let addresses = NameServerAddrGroup::deserialize(deserializer)?;
for address in &addresses {
if address.protocol != Protocol::Udp {
return Err(serde::de::Error::custom(format!(
"unsupported address protocol: {}, only udp is supported",
address.protocol
)));
}
}
Ok(addresses)
}
}
#[auto_impl::auto_impl(Box, &)]
pub trait DnsConfigLoaderExt {
fn get_dns(&self) -> DnsConfig;
@@ -55,8 +33,6 @@ pub trait DnsConfigLoaderExt {
pub type DnsExportConfig = GetExportConfigResponse;
pub trait DnsGlobalCtxExt {
fn dns_server(&self) -> Option<Arc<DnsServer>>; // TODO: remove this
fn set_dns_server(&self, dns: Option<Arc<DnsServer>>); // TODO: remove this
fn dns_self_zone(&self) -> ZoneConfig;
fn dns_export_config(&self) -> DnsExportConfig;
fn dns_iter_zones(&self) -> impl Iterator<Item = ZoneConfig>;
+4 -7
View File
@@ -1,7 +1,6 @@
use crate::dns::utils::addr::NameServerAddr;
use hickory_net::xfer::Protocol;
use crate::dns::utils::addr::NameServerAddrGroup;
use hickory_proto::rr::LowerName;
use std::net::{Ipv4Addr, SocketAddr, SocketAddrV4};
use std::net::IpAddr;
use std::str::FromStr;
use std::sync::LazyLock;
use std::time::Duration;
@@ -14,10 +13,8 @@ pub mod zone;
pub static DNS_DEFAULT_DOMAIN: LazyLock<LowerName> =
LazyLock::new(|| LowerName::from_str("et.net.").unwrap());
pub const DNS_DEFAULT_ADDRESS: NameServerAddr = NameServerAddr {
protocol: Protocol::Udp,
addr: SocketAddr::V4(SocketAddrV4::new(Ipv4Addr::new(100, 100, 100, 101), 53)),
};
pub static DNS_DEFAULT_ADDRESSES: LazyLock<NameServerAddrGroup> =
LazyLock::new(|| IpAddr::from_str("100.100.100.101").unwrap().into());
pub static DNS_SERVER_RPC_ADDR: LazyLock<Url> =
LazyLock::new(|| Url::parse("tcp://127.0.0.1:49813").unwrap());
+2 -80
View File
@@ -1,19 +1,17 @@
use crate::common::global_ctx::{ArcGlobalCtx, GlobalCtxEvent};
use crate::dns::config::{
DNS_NODE_RR_INTERVAL, DNS_SERVER_ELECTION_INTERVAL, DNS_SERVER_RPC_ADDR, DnsGlobalCtxExt,
DNS_NODE_RR_INTERVAL, DNS_SERVER_ELECTION_INTERVAL, DNS_SERVER_RPC_ADDR,
};
use crate::dns::peer_mgr::DnsPeerMgr;
use crate::dns::server::DnsServer;
#[cfg(feature = "tun")]
use crate::instance::instance::ArcNicCtx;
use crate::peers::NicPacketFilter;
use crate::peers::peer_manager::PeerManager;
use crate::proto::dns::{DnsNodeMgrRpcClientFactory, HeartbeatRequest};
use crate::proto::rpc_impl::standalone::{StandAloneClient, StandAloneServer};
use crate::proto::rpc_types::controller::BaseController;
use crate::tunnel::tcp::{TcpTunnelConnector, TcpTunnelListener};
use crate::utils::task::CancellableTask;
use guarden::guard;
use std::io;
use std::sync::Arc;
use tokio::sync::{Notify, broadcast};
@@ -75,27 +73,8 @@ impl DnsNodeRuntime {
#[cfg(feature = "tun")]
self.nic_ctx.clone(),
));
server.register(&rpc);
self.global_ctx.set_dns_server(Some(server.clone()));
let guard = guard! {
[
global_ctx = self.global_ctx.clone(),
peer_mgr = self.peer_mgr.clone(),
id = server.id()
]
global_ctx.set_dns_server(None);
async move { let _ = peer_mgr.remove_nic_packet_process_pipeline(id).await; }
};
tokio::join!(
self.peer_mgr
.add_nic_packet_process_pipeline(Box::new(server.clone())),
server.run(token.child_token())
);
guard.trigger().await;
server.run(token.child_token()).await;
tracing::warn!("DnsServer exited, will retry election");
}
@@ -572,61 +551,4 @@ mod tests {
let node = build_test_node().await;
node.stop().await.unwrap();
}
#[tokio::test]
#[serial_test::serial(dns_node_rpc_addr)]
async fn election_wins_and_sets_dns_server() {
let node = build_test_node().await;
let runtime = node.runtime.clone();
wait_for_condition(
async || runtime.global_ctx.dns_server().is_some(),
Duration::from_secs(2),
)
.await;
let global_ctx = runtime.global_ctx.clone();
node.stop().await.unwrap();
wait_for_condition(
async || global_ctx.dns_server().is_none(),
Duration::from_secs(2),
)
.await;
}
#[tokio::test]
#[serial_test::serial(dns_node_rpc_addr)]
async fn election_loses_when_rpc_addr_is_occupied() {
let holder = occupy_dns_rpc_addr().await;
let node = build_test_node().await;
let runtime = node.runtime.clone();
sleep(Duration::from_millis(300)).await;
assert!(runtime.global_ctx.dns_server().is_none());
node.stop().await.unwrap();
drop(holder);
}
#[tokio::test]
#[serial_test::serial(dns_node_rpc_addr)]
async fn election_retries_after_losing_then_wins() {
let holder = occupy_dns_rpc_addr().await;
let node = build_test_node().await;
let runtime = node.runtime.clone();
sleep(Duration::from_millis(300)).await;
assert!(runtime.global_ctx.dns_server().is_none());
drop(holder);
sleep(DNS_NODE_RR_INTERVAL).await;
wait_for_condition(
async || runtime.global_ctx.dns_server().is_some(),
Duration::from_secs(2),
)
.await;
node.stop().await.unwrap();
}
}
+7 -548
View File
@@ -2,17 +2,13 @@ use crate::common::global_ctx::ArcGlobalCtx;
use crate::dns::node_mgr::DnsNodeMgr;
use crate::dns::system;
use crate::dns::utils::addr::NameServerAddr;
use crate::dns::utils::response::ResponseHandle;
use crate::peer_center::instance::PeerCenterPeerManagerTrait;
use crate::peers::NicPacketFilter;
use crate::peers::peer_manager::PeerManager;
use crate::proto::dns::DnsNodeMgrRpcServer;
use crate::proto::rpc_impl::standalone::StandAloneServer;
use crate::tunnel::packet_def::ZCPacket;
use crate::tunnel::tcp::TcpTunnelListener;
use derivative::Derivative;
use guarden::guarded;
use hickory_net::runtime::{Time, TokioTime};
use hickory_net::runtime::Time;
use hickory_net::xfer::Protocol;
use hickory_server::{
Server,
@@ -20,13 +16,8 @@ use hickory_server::{
zone_handler::Catalog,
};
use parking_lot::RwLock;
use pnet::packet::icmp::{IcmpTypes, MutableIcmpPacket};
use pnet::packet::ip::IpNextHeaderProtocols;
use pnet::packet::ipv4::{Ipv4Packet, MutableIpv4Packet};
use pnet::packet::udp::{MutableUdpPacket, UdpPacket};
use pnet::packet::{MutablePacket, Packet, icmp, ipv4, udp};
use std::collections::HashSet;
use std::net::{IpAddr, Ipv4Addr, SocketAddr, SocketAddrV4};
use std::net::SocketAddr;
use std::{sync::Arc, time::Duration};
use tokio_util::sync::CancellationToken;
use tracing::{Instrument, instrument};
@@ -315,183 +306,10 @@ impl DnsServer {
}
}
// region NIC packet filter
const NIC_PIPELINE_NAME: &str = "magic_dns_server";
#[async_trait::async_trait]
impl NicPacketFilter for DnsServer {
async fn try_process_packet_from_nic(&self, zc_packet: &mut ZCPacket) -> bool {
self.handle_ip_packet(zc_packet).await.is_some()
}
fn id(&self) -> String {
NIC_PIPELINE_NAME.to_string()
}
}
impl DnsServer {
fn is_hijacked_ip(&self, ip: &IpAddr) -> bool {
self.addresses.read().iter().any(|a| a.addr.ip() == *ip)
}
fn is_hijacked_addr(&self, addr: SocketAddr) -> bool {
self.addresses.read().contains(&addr.into())
}
/// Replace the content of an incoming UDP DNS request and ICMP echo request packet with reply data,
/// and swap source and destination IP addresses to send it back.
async fn handle_ip_packet(&self, zc_packet: &mut ZCPacket) -> Option<()> {
let (ip_header_length, ip_protocol, src_ip, dst_ip) = {
let ip_packet = Ipv4Packet::new(zc_packet.payload())?;
if ip_packet.get_version() != 4 {
return None;
}
(
ip_packet.get_header_length() as usize * 4,
ip_packet.get_next_level_protocol(),
ip_packet.get_source(),
ip_packet.get_destination(),
)
};
if !self.is_hijacked_ip(&dst_ip.into()) {
return None;
}
match ip_protocol {
IpNextHeaderProtocols::Udp => {
self.handle_udp_packet(zc_packet, ip_header_length, src_ip, dst_ip)
.await?;
}
IpNextHeaderProtocols::Icmp => {
self.handle_icmp_packet(zc_packet, ip_header_length)?;
}
_ => {
return None;
}
}
// Swap source and destination IP addresses for the reply.
let mut ip_packet = MutableIpv4Packet::new(zc_packet.mut_payload())?;
ip_packet.set_source(dst_ip);
ip_packet.set_destination(src_ip);
ip_packet.set_checksum(ipv4::checksum(&ip_packet.to_immutable()));
// Route the response back to ourselves so it goes through the tun device.
zc_packet.mut_peer_manager_header().unwrap().to_peer_id = self.peer_mgr.my_peer_id().into();
Some(())
}
/// Extract the DNS request message from a UDP packet and send it to the catalog.
/// Replace the content of the UDP packet with the response message.
async fn handle_udp_packet(
&self,
zc_packet: &mut ZCPacket,
ip_header_length: usize,
src_ip: Ipv4Addr,
dst_ip: Ipv4Addr,
) -> Option<()> {
let (src_port, dst_port, request, request_length) = {
let udp_packet = UdpPacket::new(&zc_packet.payload()[ip_header_length..])?;
let src_port = udp_packet.get_source();
let dst_port = udp_packet.get_destination();
let request_payload = udp_packet.payload();
(
src_port,
dst_port,
Request::from_bytes(
request_payload.to_vec(),
SocketAddr::from(SocketAddrV4::new(src_ip, src_port)),
Protocol::Udp,
)
.ok()?,
request_payload.len(),
)
};
if !self.is_hijacked_addr(SocketAddr::new(dst_ip.into(), dst_port)) {
return None;
}
let response_payload = {
let response = ResponseHandle::new(512);
self.catalog
.handle_request::<_, TokioTime>(&request, response.clone())
.await;
response.into_inner()?
};
let response_length = response_payload.len();
let delta_length = response_length as isize - request_length as isize;
// Resize the packet buffer to accommodate the response.
let inner_length = (zc_packet.buf_len() as isize + delta_length) as usize;
if zc_packet.mut_inner().capacity() < inner_length {
let header_length = inner_length - response_length;
zc_packet.mut_inner().truncate(header_length);
}
zc_packet.mut_inner().resize(inner_length, 0);
let mut ip_packet = MutableIpv4Packet::new(zc_packet.mut_payload())?;
let ip_length = (ip_packet.get_total_length() as isize + delta_length) as u16;
ip_packet.set_total_length(ip_length);
let mut udp_packet = MutableUdpPacket::new(ip_packet.payload_mut())?;
let udp_length = (udp_packet.get_length() as isize + delta_length) as u16;
udp_packet.set_length(udp_length);
udp_packet.set_source(dst_port);
udp_packet.set_destination(src_port);
udp_packet.payload_mut().copy_from_slice(&response_payload);
udp_packet.set_checksum(udp::ipv4_checksum(
&udp_packet.to_immutable(),
&dst_ip,
&src_ip,
));
Some(())
}
/// Handle ICMP echo request by turning it into an echo reply.
fn handle_icmp_packet(&self, zc_packet: &mut ZCPacket, ip_header_length: usize) -> Option<()> {
let mut icmp_packet =
MutableIcmpPacket::new(&mut zc_packet.mut_payload()[ip_header_length..])?;
if icmp_packet.get_icmp_type() != IcmpTypes::EchoRequest {
return None;
}
icmp_packet.set_icmp_type(IcmpTypes::EchoReply);
icmp_packet.set_checksum(icmp::checksum(&icmp_packet.to_immutable()));
Some(())
}
}
// endregion
#[cfg(test)]
mod tests {
use super::*;
use crate::dns::tests::{
dns_snapshot_with as snapshot_with, heartbeat_with_snapshot, zone_data_a as valid_zone_data,
};
use crate::peers::tests::create_mock_peer_manager;
use crate::proto::dns::DnsNodeMgrRpc;
use crate::proto::rpc_types::controller::BaseController;
use hickory_net::client::{Client, ClientHandle};
use hickory_net::runtime::TokioRuntimeProvider;
use hickory_net::udp::UdpClientStream;
@@ -501,17 +319,15 @@ mod tests {
use hickory_server::store::in_memory::InMemoryZoneHandler;
use hickory_server::zone_handler::ZoneType;
use hickory_server::zone_handler::{AxfrPolicy, Catalog};
use pnet::packet::icmp::{IcmpPacket, IcmpTypes, MutableIcmpPacket};
use pnet::packet::ip::IpNextHeaderProtocols;
use pnet::packet::ipv4::{Ipv4Packet, MutableIpv4Packet};
use pnet::packet::udp::{MutableUdpPacket, UdpPacket};
use pnet::packet::{MutablePacket, Packet, icmp, ipv4, udp};
use std::net::{Ipv4Addr, SocketAddr};
use pnet::packet::icmp::{IcmpTypes, MutableIcmpPacket};
use pnet::packet::ipv4::MutableIpv4Packet;
use pnet::packet::udp::MutableUdpPacket;
use pnet::packet::{MutablePacket, icmp, ipv4, udp};
use std::net::Ipv4Addr;
use std::str::FromStr;
use std::time::Duration;
use tokio::net::UdpSocket;
use tokio::time::{sleep, timeout};
use uuid::Uuid;
/// Build a `Catalog` containing a single A record: `test.example.com -> 1.2.3.4`.
fn build_test_catalog() -> Catalog {
@@ -643,188 +459,6 @@ mod tests {
// ─── Tests ───────────────────────────────────────────────────────────
#[tokio::test]
async fn should_match_hijacked_ip_and_addr_when_address_is_registered() {
let server = create_test_server().await;
let addr: SocketAddr = "10.0.0.53:53".parse().unwrap();
assert!(!server.is_hijacked_ip(&addr.ip()));
assert!(!server.is_hijacked_addr(addr));
server.addresses.write().insert(addr.into());
assert!(server.is_hijacked_ip(&addr.ip()));
assert!(server.is_hijacked_addr(addr));
// Different port on same IP — ip matches, but addr does not.
let other_addr: SocketAddr = "10.0.0.53:5353".parse().unwrap();
assert!(server.is_hijacked_ip(&other_addr.ip()));
assert!(!server.is_hijacked_addr(other_addr));
}
#[tokio::test]
async fn should_reply_icmp_echo_and_swap_endpoints_when_packet_is_hijacked() {
let server = create_test_server().await;
let dst_ip: Ipv4Addr = "10.0.0.53".parse().unwrap();
let src_ip: Ipv4Addr = "10.0.0.1".parse().unwrap();
// Register the dst IP as hijacked.
server
.addresses
.write()
.insert(SocketAddr::new(dst_ip.into(), 53).into());
let icmp_payload = build_icmp_echo_request();
let ip_bytes =
build_ipv4_packet(src_ip, dst_ip, IpNextHeaderProtocols::Icmp, &icmp_payload);
let mut zc = ZCPacket::new_with_payload(&ip_bytes);
zc.fill_peer_manager_hdr(1, 2, crate::tunnel::packet_def::PacketType::Data as u8);
let result = server.handle_ip_packet(&mut zc).await;
assert!(
result.is_some(),
"handle_ip_packet should succeed for echo request"
);
// Verify ICMP type is now EchoReply.
let ip = Ipv4Packet::new(zc.payload()).unwrap();
let icmp = IcmpPacket::new(ip.payload()).unwrap();
assert_eq!(icmp.get_icmp_type(), IcmpTypes::EchoReply);
// Verify IP addresses are swapped.
assert_eq!(ip.get_source(), dst_ip);
assert_eq!(ip.get_destination(), src_ip);
// Verify route-to-self rewrite in peer manager header.
let hdr = zc.peer_manager_header().unwrap();
assert_eq!(hdr.to_peer_id.get(), server.peer_mgr.my_peer_id() as u32);
}
#[tokio::test]
async fn should_ignore_icmp_when_type_is_not_echo_request() {
let server = create_test_server().await;
let dst_ip: Ipv4Addr = "10.0.0.53".parse().unwrap();
let src_ip: Ipv4Addr = "10.0.0.1".parse().unwrap();
server
.addresses
.write()
.insert(SocketAddr::new(dst_ip.into(), 53).into());
// Build an ICMP Destination Unreachable (not echo request).
let mut icmp_buf = vec![0u8; 8];
{
let mut pkt = MutableIcmpPacket::new(&mut icmp_buf).unwrap();
pkt.set_icmp_type(IcmpTypes::DestinationUnreachable);
pkt.set_checksum(icmp::checksum(&pkt.to_immutable()));
}
let ip_bytes = build_ipv4_packet(src_ip, dst_ip, IpNextHeaderProtocols::Icmp, &icmp_buf);
let mut zc = ZCPacket::new_with_payload(&ip_bytes);
zc.fill_peer_manager_hdr(1, 2, crate::tunnel::packet_def::PacketType::Data as u8);
let result = server.handle_ip_packet(&mut zc).await;
assert!(result.is_none(), "non-echo ICMP should be ignored");
}
#[tokio::test]
async fn should_ignore_packet_when_destination_ip_is_not_hijacked() {
let server = create_test_server().await;
// Do NOT register any hijacked addresses.
let icmp_payload = build_icmp_echo_request();
let ip_bytes = build_ipv4_packet(
"10.0.0.1".parse().unwrap(),
"10.0.0.99".parse().unwrap(),
IpNextHeaderProtocols::Icmp,
&icmp_payload,
);
let mut zc = ZCPacket::new_with_payload(&ip_bytes);
zc.fill_peer_manager_hdr(1, 2, crate::tunnel::packet_def::PacketType::Data as u8);
let result = server.handle_ip_packet(&mut zc).await;
assert!(
result.is_none(),
"packet to non-hijacked IP should be ignored"
);
}
#[tokio::test]
async fn should_rewrite_udp_dns_packet_when_query_targets_hijacked_dns_addr() {
let server = create_test_server().await;
let dst_ip: Ipv4Addr = "10.0.0.53".parse().unwrap();
let src_ip: Ipv4Addr = "10.0.0.1".parse().unwrap();
let dns_port: u16 = 53;
let client_port: u16 = 12345;
// Register dst as hijacked.
server
.addresses
.write()
.insert(SocketAddr::new(dst_ip.into(), dns_port).into());
// Load a catalog with test.example.com -> 1.2.3.4.
server.catalog.replace(build_test_catalog()).await;
// Build DNS query.
let dns_bytes = build_dns_query_bytes("test.example.com.");
let udp_bytes = build_udp_packet(client_port, dns_port, &dns_bytes, src_ip, dst_ip);
let ip_bytes = build_ipv4_packet(src_ip, dst_ip, IpNextHeaderProtocols::Udp, &udp_bytes);
let mut zc = ZCPacket::new_with_payload(&ip_bytes);
zc.fill_peer_manager_hdr(1, 2, crate::tunnel::packet_def::PacketType::Data as u8);
let result = server.handle_ip_packet(&mut zc).await;
assert!(result.is_some(), "DNS query should be handled");
// Parse the response IP packet => UDP => DNS message.
let ip = Ipv4Packet::new(zc.payload()).unwrap();
assert_eq!(
ip.get_source(),
dst_ip,
"reply source should be the DNS server IP"
);
assert_eq!(
ip.get_destination(),
src_ip,
"reply dest should be the client IP"
);
let udp_reply = UdpPacket::new(ip.payload()).unwrap();
assert_eq!(udp_reply.get_source(), dns_port);
assert_eq!(udp_reply.get_destination(), client_port);
assert_eq!(
udp_reply.get_length() as usize,
8 + udp_reply.payload().len(),
"UDP length should match payload size"
);
assert_eq!(
udp_reply.get_checksum(),
udp::ipv4_checksum(&udp_reply, &dst_ip, &src_ip),
"UDP checksum should be recomputed for swapped src/dst IP"
);
assert_eq!(
ip.get_total_length() as usize,
20 + udp_reply.packet().len(),
"IP total length should match rewritten packet"
);
let dns_reply = Message::from_vec(udp_reply.payload()).unwrap();
assert_eq!(dns_reply.id, 0x1234);
assert!(
!dns_reply.answers.is_empty(),
"DNS reply should contain answers"
);
let answer = &dns_reply.answers[0];
if let RData::A(a) = answer.data {
assert_eq!(a.0, Ipv4Addr::new(1, 2, 3, 4));
} else {
panic!("expected A record in answer, got {:?}", answer.data);
}
}
/// Full end-to-end test: start a real DNS UDP listener via `ServerFuture`,
/// send a query with a `hickory_client`, and verify the response.
#[tokio::test]
@@ -878,117 +512,6 @@ mod tests {
shutdown_token.cancel();
}
#[tokio::test]
async fn should_process_icmp_packet_and_expose_pipeline_id() {
let server = create_test_server().await;
let src_ip: Ipv4Addr = "10.0.0.1".parse().unwrap();
let dst_ip: Ipv4Addr = "10.0.0.53".parse().unwrap();
server
.addresses
.write()
.insert(SocketAddr::new(dst_ip.into(), 53).into());
let icmp_payload = build_icmp_echo_request();
let ip_bytes =
build_ipv4_packet(src_ip, dst_ip, IpNextHeaderProtocols::Icmp, &icmp_payload);
let mut zc = ZCPacket::new_with_payload(&ip_bytes);
zc.fill_peer_manager_hdr(1, 2, crate::tunnel::packet_def::PacketType::Data as u8);
assert!(server.try_process_packet_from_nic(&mut zc).await);
assert_eq!(server.id(), NIC_PIPELINE_NAME);
}
#[tokio::test]
async fn should_ignore_packet_when_ipv4_header_has_non_ipv4_version() {
let server = create_test_server().await;
let src_ip: Ipv4Addr = "10.0.0.1".parse().unwrap();
let dst_ip: Ipv4Addr = "10.0.0.53".parse().unwrap();
server
.addresses
.write()
.insert(SocketAddr::new(dst_ip.into(), 53).into());
let icmp_payload = build_icmp_echo_request();
let mut ip_bytes =
build_ipv4_packet(src_ip, dst_ip, IpNextHeaderProtocols::Icmp, &icmp_payload);
ip_bytes[0] = (6 << 4) | 5; // fake IPv6 version in IPv4 header
let mut zc = ZCPacket::new_with_payload(&ip_bytes);
zc.fill_peer_manager_hdr(1, 2, crate::tunnel::packet_def::PacketType::Data as u8);
assert!(server.handle_ip_packet(&mut zc).await.is_none());
}
#[tokio::test]
async fn should_ignore_packet_when_protocol_is_unsupported() {
let server = create_test_server().await;
let src_ip: Ipv4Addr = "10.0.0.1".parse().unwrap();
let dst_ip: Ipv4Addr = "10.0.0.53".parse().unwrap();
server
.addresses
.write()
.insert(SocketAddr::new(dst_ip.into(), 53).into());
let tcp_like_payload = vec![0u8; 20];
let ip_bytes = build_ipv4_packet(
src_ip,
dst_ip,
IpNextHeaderProtocols::Tcp,
&tcp_like_payload,
);
let mut zc = ZCPacket::new_with_payload(&ip_bytes);
zc.fill_peer_manager_hdr(1, 2, crate::tunnel::packet_def::PacketType::Data as u8);
assert!(server.handle_ip_packet(&mut zc).await.is_none());
}
#[tokio::test]
async fn should_ignore_udp_dns_packet_when_destination_port_is_not_hijacked() {
let server = create_test_server().await;
let dst_ip: Ipv4Addr = "10.0.0.53".parse().unwrap();
let src_ip: Ipv4Addr = "10.0.0.1".parse().unwrap();
server
.addresses
.write()
.insert(SocketAddr::new(dst_ip.into(), 53).into());
server.catalog.replace(build_test_catalog()).await;
let dns_bytes = build_dns_query_bytes("test.example.com.");
let udp_bytes = build_udp_packet(12345, 5353, &dns_bytes, src_ip, dst_ip);
let ip_bytes = build_ipv4_packet(src_ip, dst_ip, IpNextHeaderProtocols::Udp, &udp_bytes);
let mut zc = ZCPacket::new_with_payload(&ip_bytes);
zc.fill_peer_manager_hdr(1, 2, crate::tunnel::packet_def::PacketType::Data as u8);
assert!(server.handle_ip_packet(&mut zc).await.is_none());
}
#[tokio::test]
async fn should_ignore_udp_packet_when_dns_payload_is_invalid() {
let server = create_test_server().await;
let dst_ip: Ipv4Addr = "10.0.0.53".parse().unwrap();
let src_ip: Ipv4Addr = "10.0.0.1".parse().unwrap();
server
.addresses
.write()
.insert(SocketAddr::new(dst_ip.into(), 53).into());
let invalid_dns = vec![0xde, 0xad, 0xbe];
let udp_bytes = build_udp_packet(12345, 53, &invalid_dns, src_ip, dst_ip);
let ip_bytes = build_ipv4_packet(src_ip, dst_ip, IpNextHeaderProtocols::Udp, &udp_bytes);
let mut zc = ZCPacket::new_with_payload(&ip_bytes);
zc.fill_peer_manager_hdr(1, 2, crate::tunnel::packet_def::PacketType::Data as u8);
assert!(server.handle_ip_packet(&mut zc).await.is_none());
}
#[tokio::test]
async fn should_update_public_addresses_when_reload_addresses_is_called() {
let server = create_test_server().await;
@@ -1060,68 +583,4 @@ mod tests {
let _ = runtime.stop(None).await;
}
}
#[tokio::test]
async fn should_apply_snapshot_updates_and_clear_state_on_shutdown() {
let server = create_test_server().await;
let token = CancellationToken::new();
let run_server = server.clone();
let run_token = token.clone();
let run_task = tokio::spawn(async move {
run_server.run(run_token).await;
});
let node_id = Uuid::new_v4();
let snapshot = snapshot_with(
vec![valid_zone_data("run-loop.test", "7.7.7.7")],
vec!["udp://10.0.0.53:53"],
vec![],
);
DnsNodeMgrRpc::heartbeat(
&*server.mgr,
BaseController::default(),
heartbeat_with_snapshot(node_id, snapshot),
)
.await
.unwrap();
wait_until(|| {
server
.addresses()
.contains(&"10.0.0.53:53".parse::<SocketAddr>().unwrap())
})
.await;
// Verify catalog hot-reload by issuing a hijacked DNS packet query.
let src_ip: Ipv4Addr = "10.0.0.1".parse().unwrap();
let dst_ip: Ipv4Addr = "10.0.0.53".parse().unwrap();
let dns_bytes = build_dns_query_bytes("run-loop.test.");
let udp_bytes = build_udp_packet(12000, 53, &dns_bytes, src_ip, dst_ip);
let ip_bytes = build_ipv4_packet(src_ip, dst_ip, IpNextHeaderProtocols::Udp, &udp_bytes);
let mut zc = ZCPacket::new_with_payload(&ip_bytes);
zc.fill_peer_manager_hdr(1, 2, crate::tunnel::packet_def::PacketType::Data as u8);
let mut handled = false;
for _ in 0..80 {
if server.handle_ip_packet(&mut zc).await.is_some() {
handled = true;
break;
}
sleep(Duration::from_millis(50)).await;
}
assert!(handled, "catalog should be hot-reloaded before timeout");
token.cancel();
let _ = run_task.await;
assert!(
server.addresses().is_empty(),
"run() should clear addresses on exit"
);
assert!(
server.listeners.read().is_empty(),
"run() should clear listeners on exit"
);
}
}
+2 -2
View File
@@ -175,7 +175,7 @@ mod tests {
use std::{str::FromStr as _, sync::Arc, time::Duration};
use crate::dns::{
config::DNS_DEFAULT_ADDRESS,
config::DNS_DEFAULT_ADDRESSES,
tests::{prepare_env, start_dns_node},
};
use crate::instance::proxy_cidrs_monitor::ProxyCidrsMonitor;
@@ -193,7 +193,7 @@ mod tests {
let dns_node = start_dns_node(peer_mgr, virtual_nic);
println!("dev_name: {}", tun_name);
let fake_ip = match DNS_DEFAULT_ADDRESS.addr.ip() {
let fake_ip = match DNS_DEFAULT_ADDRESSES[0].addr.ip() {
IpAddr::V4(ip) => ip,
IpAddr::V6(ip) => panic!("unexpected ipv6 default dns address in test: {ip}"),
};
+61 -25
View File
@@ -4,10 +4,10 @@ use anyhow::{Error, anyhow};
use hickory_net::xfer::Protocol;
use hickory_resolver::config::{ConnectionConfig, NameServerConfig, ProtocolConfig};
use serde::de::IntoDeserializer;
use serde::{Deserialize, de};
use serde::{Deserialize, Deserializer, de};
use serde_with::{DeserializeFromStr, SerializeDisplay};
use std::fmt::{Display, Formatter};
use std::net::{IpAddr, SocketAddr};
use std::net::{IpAddr, Ipv6Addr, SocketAddr};
use std::str::FromStr;
use url::Url;
@@ -39,21 +39,6 @@ impl From<(IpAddr, &ConnectionConfig)> for NameServerAddr {
}
}
impl From<SocketAddr> for NameServerAddr {
fn from(value: SocketAddr) -> Self {
Self {
protocol: Protocol::Udp,
addr: value,
}
}
}
impl From<IpAddr> for NameServerAddr {
fn from(value: IpAddr) -> Self {
SocketAddr::new(value, 53).into()
}
}
impl From<NameServerAddr> for Url {
fn from(value: NameServerAddr) -> Self {
Url::parse(&format!("{}://{}", value.protocol, value.addr)).unwrap()
@@ -102,14 +87,6 @@ impl TryFrom<&proto::common::Url> for NameServerAddr {
impl FromStr for NameServerAddr {
type Err = Error;
fn from_str(s: &str) -> Result<Self, Self::Err> {
macro_rules! try_parse {
($($t:ty),+) => {
$( if let Ok(v) = s.parse::<$t>() { return Ok(v.into()); } )+
};
}
try_parse!(IpAddr, SocketAddr);
(&Url::parse(s)?).try_into()
}
}
@@ -131,3 +108,62 @@ impl From<&NameServerConfig> for NameServerAddrGroup {
.collect()
}
}
impl From<SocketAddr> for NameServerAddrGroup {
fn from(value: SocketAddr) -> Self {
vec![
NameServerAddr {
protocol: Protocol::Udp,
addr: value,
},
NameServerAddr {
protocol: Protocol::Tcp,
addr: value,
},
]
.into()
}
}
impl From<IpAddr> for NameServerAddrGroup {
fn from(value: IpAddr) -> Self {
SocketAddr::new(value, 53).into()
}
}
impl From<u16> for NameServerAddrGroup {
fn from(value: u16) -> Self {
SocketAddr::new(Ipv6Addr::UNSPECIFIED.into(), value).into()
}
}
impl<'de> Deserialize<'de> for NameServerAddrGroup {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: Deserializer<'de>,
{
#[derive(Deserialize)]
#[serde(untagged)]
enum Candidate {
NameServerAddr(NameServerAddr),
U16(u16),
IpAddr(IpAddr),
SocketAddr(SocketAddr),
}
let items = Vec::<Candidate>::deserialize(deserializer)?;
let items = items
.into_iter()
.flat_map(|item| -> NameServerAddrGroup {
match item {
Candidate::NameServerAddr(addr) => vec![addr].into(),
Candidate::U16(port) => port.into(),
Candidate::IpAddr(ip) => ip.into(),
Candidate::SocketAddr(addr) => addr.into(),
}
})
.collect();
Ok(items)
}
}
+9 -22
View File
@@ -1,9 +1,8 @@
use std::collections::BTreeSet;
use std::net::IpAddr;
use std::sync::{Arc, Weak};
use crate::common::global_ctx::{ArcGlobalCtx, GlobalCtxEvent};
use crate::peers::peer_manager::PeerManager;
use std::collections::BTreeSet;
use std::sync::{Arc, Weak};
use std::time::Instant;
use tokio_util::task::AbortOnDropHandle;
/// ProxyCidrsMonitor monitors changes in proxy CIDRs from peer routes
@@ -44,17 +43,6 @@ impl ProxyCidrsMonitor {
proxy_cidrs.insert(vpn_cfg.client_cidr);
}
#[cfg(feature = "magic-dns")]
{
use crate::dns::config::DnsGlobalCtxExt;
if let Some(dns) = global_ctx.dns_server() {
proxy_cidrs.extend(dns.addresses().into_iter().filter_map(|a| match a.ip() {
IpAddr::V4(ip) => Some(cidr::Ipv4Cidr::new_host(ip)),
_ => None,
}))
}
}
proxy_cidrs
};
@@ -72,7 +60,7 @@ impl ProxyCidrsMonitor {
pub fn start(self) -> AbortOnDropHandle<()> {
AbortOnDropHandle::new(tokio::spawn(async move {
let mut cur_proxy_cidrs = BTreeSet::new();
// let mut last_update = None::<Instant>;
let mut last_update = None::<Instant>;
loop {
tokio::time::sleep(std::time::Duration::from_secs(1)).await;
@@ -82,13 +70,12 @@ impl ProxyCidrsMonitor {
break;
};
// TODO: same logic for DNS
// Check if route info has been updated
// let last_update_time = peer_mgr.get_route_peer_info_last_update_time().await;
// if last_update == Some(last_update_time) {
// continue;
// }
// last_update = Some(last_update_time);
let last_update_time = peer_mgr.get_route_peer_info_last_update_time().await;
if last_update == Some(last_update_time) {
continue;
}
last_update = Some(last_update_time);
let (new_proxy_cidrs, added, removed) =
Self::diff_proxy_cidrs(peer_mgr.as_ref(), &self.global_ctx, &cur_proxy_cidrs)
+2 -13
View File
@@ -2,7 +2,7 @@ use delegate::delegate;
use derivative::Derivative;
use derive_more::{Deref, DerefMut, From, IntoIterator};
use prost::Message;
use serde::{Deserialize, Serialize};
use serde::Serialize;
use sha2::{Digest, Sha256};
/// Generates a stable digest strictly within the lifecycle of the current process.
@@ -36,18 +36,7 @@ where
}
#[derive(
Derivative,
Debug,
Clone,
PartialEq,
Eq,
Hash,
From,
Deref,
DerefMut,
Serialize,
Deserialize,
IntoIterator,
Derivative, Debug, Clone, PartialEq, Eq, Hash, From, Deref, DerefMut, Serialize, IntoIterator,
)]
#[derivative(Default(bound = ""))]
#[serde(transparent)]