From 7d536d33530f7bdbbbee816f0850d4ff95061053 Mon Sep 17 00:00:00 2001 From: Luna Yao <40349250+ZnqbuZ@users.noreply.github.com> Date: Mon, 27 Apr 2026 23:45:03 +0200 Subject: [PATCH] remove address hijack --- easytier/src/common/global_ctx.rs | 20 +- easytier/src/dns/config/dns.rs | 30 +- easytier/src/dns/config/mod.rs | 11 +- easytier/src/dns/node.rs | 82 +-- easytier/src/dns/server.rs | 555 +------------------ easytier/src/dns/system/windows.rs | 4 +- easytier/src/dns/utils/addr.rs | 86 ++- easytier/src/instance/proxy_cidrs_monitor.rs | 31 +- easytier/src/proto/utils.rs | 15 +- 9 files changed, 91 insertions(+), 743 deletions(-) diff --git a/easytier/src/common/global_ctx.rs b/easytier/src/common/global_ctx.rs index c65443dd..e64673cc 100644 --- a/easytier/src/common/global_ctx.rs +++ b/easytier/src/common/global_ctx.rs @@ -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, - #[cfg(feature = "magic-dns")] - dns_server: RwLock>>, - stun_info_collection: Mutex>, running_listeners: Mutex>, @@ -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> { - self.dns_server.read().unwrap().clone() - } - - fn set_dns_server(&self, dns: Option>) { - *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(); diff --git a/easytier/src/dns/config/dns.rs b/easytier/src/dns/config/dns.rs index 477e93ca..5b9333fc 100644 --- a/easytier/src/dns/config/dns.rs +++ b/easytier/src/dns/config/dns.rs @@ -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 - 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>; // TODO: remove this - fn set_dns_server(&self, dns: Option>); // TODO: remove this fn dns_self_zone(&self) -> ZoneConfig; fn dns_export_config(&self) -> DnsExportConfig; fn dns_iter_zones(&self) -> impl Iterator; diff --git a/easytier/src/dns/config/mod.rs b/easytier/src/dns/config/mod.rs index 920f1f5e..63b7736f 100644 --- a/easytier/src/dns/config/mod.rs +++ b/easytier/src/dns/config/mod.rs @@ -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 = 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 = + LazyLock::new(|| IpAddr::from_str("100.100.100.101").unwrap().into()); pub static DNS_SERVER_RPC_ADDR: LazyLock = LazyLock::new(|| Url::parse("tcp://127.0.0.1:49813").unwrap()); diff --git a/easytier/src/dns/node.rs b/easytier/src/dns/node.rs index 64ac5de9..e24d2fc8 100644 --- a/easytier/src/dns/node.rs +++ b/easytier/src/dns/node.rs @@ -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(); - } } diff --git a/easytier/src/dns/server.rs b/easytier/src/dns/server.rs index 3f52e592..4082baa2 100644 --- a/easytier/src/dns/server.rs +++ b/easytier/src/dns/server.rs @@ -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::().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" - ); - } } diff --git a/easytier/src/dns/system/windows.rs b/easytier/src/dns/system/windows.rs index 731b1230..6c10d5b3 100644 --- a/easytier/src/dns/system/windows.rs +++ b/easytier/src/dns/system/windows.rs @@ -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}"), }; diff --git a/easytier/src/dns/utils/addr.rs b/easytier/src/dns/utils/addr.rs index 77f744b0..c8d3fdff 100644 --- a/easytier/src/dns/utils/addr.rs +++ b/easytier/src/dns/utils/addr.rs @@ -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 for NameServerAddr { - fn from(value: SocketAddr) -> Self { - Self { - protocol: Protocol::Udp, - addr: value, - } - } -} - -impl From for NameServerAddr { - fn from(value: IpAddr) -> Self { - SocketAddr::new(value, 53).into() - } -} - impl From 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 { - 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 for NameServerAddrGroup { + fn from(value: SocketAddr) -> Self { + vec![ + NameServerAddr { + protocol: Protocol::Udp, + addr: value, + }, + NameServerAddr { + protocol: Protocol::Tcp, + addr: value, + }, + ] + .into() + } +} + +impl From for NameServerAddrGroup { + fn from(value: IpAddr) -> Self { + SocketAddr::new(value, 53).into() + } +} + +impl From for NameServerAddrGroup { + fn from(value: u16) -> Self { + SocketAddr::new(Ipv6Addr::UNSPECIFIED.into(), value).into() + } +} + +impl<'de> Deserialize<'de> for NameServerAddrGroup { + fn deserialize(deserializer: D) -> Result + where + D: Deserializer<'de>, + { + #[derive(Deserialize)] + #[serde(untagged)] + enum Candidate { + NameServerAddr(NameServerAddr), + U16(u16), + IpAddr(IpAddr), + SocketAddr(SocketAddr), + } + + let items = Vec::::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) + } +} diff --git a/easytier/src/instance/proxy_cidrs_monitor.rs b/easytier/src/instance/proxy_cidrs_monitor.rs index 3e3314b8..e206da45 100644 --- a/easytier/src/instance/proxy_cidrs_monitor.rs +++ b/easytier/src/instance/proxy_cidrs_monitor.rs @@ -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::; + let mut last_update = None::; 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) diff --git a/easytier/src/proto/utils.rs b/easytier/src/proto/utils.rs index c9ab016e..b82aa800 100644 --- a/easytier/src/proto/utils.rs +++ b/easytier/src/proto/utils.rs @@ -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)]