Files
Easytier/easytier/src/dns/server.rs
T
2026-04-06 11:54:50 +02:00

875 lines
30 KiB
Rust

use crate::common::config::ConfigLoader;
use crate::common::global_ctx::ArcGlobalCtx;
use crate::dns::node_mgr::DnsNodeMgr;
use crate::dns::system;
use crate::dns::utils::addr::NameServerAddr;
use crate::instance::instance::{ArcNicCtx, NicCtx};
use crate::peer_center::instance::PeerCenterPeerManagerTrait;
use crate::peers::peer_manager::PeerManager;
use crate::peers::NicPacketFilter;
use crate::proto::dns::DnsNodeMgrRpcServer;
use crate::proto::rpc_impl::standalone::StandAloneServer;
use crate::tunnel::common::bind_socket;
use crate::tunnel::packet_def::ZCPacket;
use crate::tunnel::tcp::TcpTunnelListener;
use crate::utils::AsyncRuntime;
use derivative::Derivative;
use hickory_proto::rr::Record;
use hickory_proto::serialize::binary::{BinDecodable, BinEncoder};
use hickory_proto::xfer::Protocol;
use hickory_server::authority::MessageRequest;
use hickory_server::{
authority::{Catalog, MessageResponse},
server::{Request, RequestHandler, ResponseHandler, ResponseInfo},
ServerFuture,
};
use parking_lot::{Mutex, 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::{icmp, ipv4, udp, MutablePacket, Packet};
use std::collections::HashSet;
use std::io;
use std::net::{IpAddr, Ipv4Addr, SocketAddr, SocketAddrV4};
use std::{sync::Arc, time::Duration};
use tokio_util::sync::CancellationToken;
use tracing::{instrument, Instrument};
#[derive(Clone)]
pub struct DynamicCatalog {
inner: Arc<tokio::sync::RwLock<Catalog>>,
}
impl DynamicCatalog {
pub fn new() -> Self {
Self {
inner: Arc::new(tokio::sync::RwLock::new(Catalog::new())),
}
}
pub async fn replace(&self, new: Catalog) {
*self.inner.write().await = new;
}
}
#[async_trait::async_trait]
impl RequestHandler for DynamicCatalog {
async fn handle_request<R: ResponseHandler>(
&self,
request: &Request,
response_handle: R,
) -> ResponseInfo {
self.inner
.read()
.await
.handle_request(request, response_handle)
.await
}
}
// ResponseWrapper for serializing DNS responses into a byte buffer.
// Used by the address hijacking NIC packet filter to produce DNS replies in-place.
#[derive(Debug, Clone)]
struct ResponseHandle {
inner: Arc<Mutex<Vec<u8>>>,
}
impl ResponseHandle {
pub fn new(capacity: usize) -> Self {
Self {
inner: Arc::new(Mutex::new(Vec::with_capacity(capacity))),
}
}
pub fn into_inner(self) -> Option<Vec<u8>> {
Arc::into_inner(self.inner).map(Mutex::into_inner)
}
}
trait RecordIter<'r>: Iterator<Item = &'r Record> + Send + 'r {}
impl<'r, T> RecordIter<'r> for T where T: Iterator<Item = &'r Record> + Send + 'r {}
#[async_trait::async_trait]
impl ResponseHandler for ResponseHandle {
async fn send_response<'r>(
&mut self,
response: MessageResponse<
'_,
'r,
impl RecordIter<'r>,
impl RecordIter<'r>,
impl RecordIter<'r>,
impl RecordIter<'r>,
>,
) -> io::Result<ResponseInfo> {
let max_size = if let Some(edns) = response.get_edns() {
edns.max_payload()
} else {
hickory_proto::udp::MAX_RECEIVE_BUFFER_SIZE as u16
};
let mut inner = self.inner.lock();
let mut encoder = BinEncoder::new(inner.as_mut());
encoder.set_max_size(max_size);
response
.destructive_emit(&mut encoder)
.map_err(io::Error::other)
}
}
#[derive(Derivative)]
#[derivative(Debug)]
pub struct DnsServer {
mgr: Arc<DnsNodeMgr>,
#[cfg(feature = "tun")]
nic_ctx: ArcNicCtx, // TODO: REMOVE THIS
peer_mgr: Arc<PeerManager>,
global_ctx: ArcGlobalCtx,
#[derivative(Debug = "ignore")]
catalog: DynamicCatalog,
listeners: Arc<RwLock<HashSet<NameServerAddr>>>,
addresses: Arc<RwLock<HashSet<NameServerAddr>>>,
}
const DNS_SERVER_LISTENER_TCP_TIMEOUT: Duration = Duration::from_secs(5);
impl DnsServer {
pub fn new(
peer_mgr: Arc<PeerManager>,
global_ctx: ArcGlobalCtx,
#[cfg(feature = "tun")] nic_ctx: ArcNicCtx, // TODO: REMOVE THIS
) -> Self {
Self {
mgr: Arc::new(DnsNodeMgr::new()),
nic_ctx,
peer_mgr,
global_ctx,
catalog: DynamicCatalog::new(),
listeners: Default::default(),
addresses: Default::default(),
}
}
pub fn register(&self, rpc: &StandAloneServer<TcpTunnelListener>) {
rpc.registry()
.register(DnsNodeMgrRpcServer::new_arc(self.mgr.clone()), "");
}
pub fn addresses(&self) -> HashSet<SocketAddr> {
self.addresses.read().iter().map(|a| a.addr).collect()
}
#[instrument(skip_all)]
async fn reload_listeners(
&self,
listeners: impl IntoIterator<Item = NameServerAddr>,
runtime: &mut Option<AsyncRuntime>,
) -> anyhow::Result<()> {
let listeners = listeners.into_iter().collect();
if &*self.listeners.read() == &listeners {
tracing::info!("listeners unchanged, no need to reload");
return Ok(());
}
tracing::info!(?listeners, "reloading");
if let Some(runtime) = runtime.as_ref() {
if let Some(Err(e)) = runtime.stop().await {
tracing::error!("failed to stop old DNS server runtime: {}", e);
}
}
let runtime = runtime.get_or_insert_default();
let mut server = ServerFuture::new(self.catalog.clone());
for listener in listeners {
let addr = listener.addr;
tracing::info!(?addr, "binding listener");
if let Err(error) = match listener.protocol {
Protocol::Udp => bind_socket(addr, None).map(|s| server.register_socket(s)),
Protocol::Tcp => bind_socket(addr, None)
.map(|s| server.register_listener(s, DNS_SERVER_LISTENER_TCP_TIMEOUT)),
_ => unimplemented!(),
} {
tracing::error!(?addr, ?error, "failed to bind listener");
}
}
runtime.start(Some(server.shutdown_token().clone()), |_| {
async move {
server
.block_until_done()
.await
.unwrap_or_else(|e| tracing::error!("DNS server exited with error: {:?}", e));
}
.instrument(tracing::info_span!("DNS server backend runtime"))
});
Ok(())
}
#[instrument(skip_all)]
async fn reload_addresses(
&self,
addresses: impl IntoIterator<Item = NameServerAddr>,
) -> anyhow::Result<()> {
let addresses = addresses.into_iter().collect();
if &*self.addresses.read() == &addresses {
tracing::info!("addresses unchanged, no need to reload");
return Ok(());
}
tracing::info!(?addresses, "reloading");
#[cfg(feature = "tun")]
{
let nic_ctx = self.nic_ctx.lock().await;
if let Some(nic_ctx) = nic_ctx
.as_ref()
.and_then(|nic_ctx| nic_ctx.downcast_ref::<NicCtx>())
{
if let Some(system) = nic_ctx
.ifname()
.await
.map(|ifname| system::get(&ifname))
.transpose()?
.flatten()
{
let config = self.global_ctx.config.get_dns();
let domain = vec![config.domain.to_string()];
system.set_dns(&system::SystemConfig {
nameservers: addresses
.iter()
.filter_map(|a| {
(a.protocol == Protocol::Udp && a.addr.port() == 53)
.then_some(a.addr.ip().to_string())
})
.collect(),
search_domains: domain.clone(),
match_domains: domain
.into_iter()
.chain(config.zones.iter().map(|z| z.origin.to_string()))
.collect(),
})?;
}
}
}
*self.addresses.write() = addresses;
Ok(())
}
#[instrument(skip_all, name = "DnsServer main loop")]
pub async fn run(&self, token: CancellationToken) {
let dirty = &self.mgr.dirty;
let mut runtime = None;
let reload_catalog = async {
loop {
dirty.catalog.notified().await;
if dirty.catalog.reset() {
self.catalog.replace(self.mgr.catalog()).await;
}
tokio::time::sleep(Duration::from_secs(1)).await;
}
};
let reload_addresses = async {
loop {
dirty.addresses.notified().await;
if dirty.addresses.reset() {
if let Err(e) = self.reload_addresses(self.mgr.iter_addresses()).await {
tracing::error!("failed to reload addresses: {:?}", e);
dirty.addresses.mark();
}
}
tokio::time::sleep(Duration::from_secs(1)).await;
}
};
let reload_listeners = async {
loop {
dirty.listeners.notified().await;
if dirty.listeners.reset() {
if let Err(e) = self
.reload_listeners(self.mgr.iter_listeners(), &mut runtime)
.await
{
tracing::error!("failed to reload listeners: {:?}", e);
dirty.listeners.mark();
}
}
tokio::time::sleep(Duration::from_secs(1)).await;
}
};
tokio::select!(
_ = token.cancelled() => {
tracing::info!("DnsServer received shutdown signal, exiting server loop");
}
_ = reload_catalog => {},
_ = reload_addresses => {},
_ = reload_listeners => {},
);
self.addresses.write().clear();
if let Some(runtime) = runtime.take() {
let _ = runtime.stop().await;
}
}
}
impl Drop for DnsServer {
fn drop(&mut self) {
tracing::info!("DnsServer is dropped");
self.addresses.write().clear();
}
}
// 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::new(
MessageRequest::from_bytes(request_payload).ok()?,
SocketAddr::from(SocketAddrV4::new(src_ip, src_port)),
Protocol::Udp,
),
request_payload.len(),
)
};
if !self.is_hijacked_addr(SocketAddr::new(dst_ip.into(), dst_port)) {
return None;
}
tracing::warn!("HIJACKING PACKET");
let response_payload = {
let response = ResponseHandle::new(512);
self.catalog
.handle_request(&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::peers::tests::create_mock_peer_manager;
use hickory_client::client::{Client, ClientHandle};
use hickory_proto::op::{Message, MessageType, OpCode, Query};
use hickory_proto::rr::{rdata, DNSClass, Name, RData, Record, RecordType};
use hickory_proto::runtime::TokioRuntimeProvider;
use hickory_proto::serialize::binary::BinEncodable;
use hickory_proto::udp::UdpClientStream;
use hickory_server::authority::Catalog;
use hickory_server::authority::ZoneType;
use hickory_server::store::in_memory::InMemoryAuthority;
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::{icmp, ipv4, udp, MutablePacket, Packet};
use std::net::{Ipv4Addr, SocketAddr};
use std::str::FromStr;
use std::time::Duration;
/// Build a `Catalog` containing a single A record: `test.example.com -> 1.2.3.4`.
fn build_test_catalog() -> Catalog {
let origin = Name::from_str("example.com.").unwrap();
let mut authority = InMemoryAuthority::empty(origin.clone(), ZoneType::Primary, false);
let record = Record::from_rdata(
Name::from_str("test.example.com.").unwrap(),
60,
RData::A(rdata::a::A(Ipv4Addr::new(1, 2, 3, 4))),
);
let rr_key =
hickory_proto::rr::RrKey::new(record.name().clone().into(), record.record_type());
let mut rr_set =
hickory_proto::rr::RecordSet::new(record.name().clone(), record.record_type(), 0);
rr_set.insert(record, 0);
authority.records_get_mut().insert(rr_key, Arc::new(rr_set));
let mut catalog = Catalog::new();
catalog.upsert(
origin.into(),
vec![Arc::new(authority) as Arc<dyn hickory_server::authority::AuthorityObject>],
);
catalog
}
/// Create a test `DnsServer` with `create_mock_peer_manager()`.
async fn create_test_server() -> Arc<DnsServer> {
let peer_mgr = create_mock_peer_manager().await;
let global_ctx = peer_mgr.get_global_ctx();
Arc::new(DnsServer::new(
peer_mgr,
global_ctx,
#[cfg(feature = "tun")]
ArcNicCtx::default(),
))
}
/// Build a raw IPv4 packet (as `Vec<u8>`) carrying the given L4 payload bytes.
/// `protocol` selects ICMP / UDP etc.
fn build_ipv4_packet(
src: Ipv4Addr,
dst: Ipv4Addr,
protocol: pnet::packet::ip::IpNextHeaderProtocol,
l4_payload: &[u8],
) -> Vec<u8> {
let ip_header_len = 20usize;
let total_len = ip_header_len + l4_payload.len();
let mut buf = vec![0u8; total_len];
{
let mut ip = MutableIpv4Packet::new(&mut buf).unwrap();
ip.set_version(4);
ip.set_header_length(5); // 20 bytes
ip.set_total_length(total_len as u16);
ip.set_ttl(64);
ip.set_next_level_protocol(protocol);
ip.set_source(src);
ip.set_destination(dst);
ip.payload_mut().copy_from_slice(l4_payload);
ip.set_checksum(ipv4::checksum(&ip.to_immutable()));
}
buf
}
/// Build ICMP Echo Request payload (8 bytes minimum).
fn build_icmp_echo_request() -> Vec<u8> {
let mut buf = vec![0u8; 8];
{
let mut icmp_pkt = MutableIcmpPacket::new(&mut buf).unwrap();
icmp_pkt.set_icmp_type(IcmpTypes::EchoRequest);
icmp_pkt.set_icmp_code(icmp::IcmpCode::new(0));
icmp_pkt.set_checksum(icmp::checksum(&icmp_pkt.to_immutable()));
}
buf
}
/// Build a minimal DNS query message for `name` and encode it to bytes.
fn build_dns_query_bytes(name: &str) -> Vec<u8> {
let mut msg = Message::new();
msg.set_id(0x1234);
msg.set_message_type(MessageType::Query);
msg.set_op_code(OpCode::Query);
msg.set_recursion_desired(true);
let mut query = Query::new();
query.set_name(Name::from_str(name).unwrap());
query.set_query_type(RecordType::A);
query.set_query_class(DNSClass::IN);
msg.add_query(query);
msg.to_bytes().unwrap().to_vec()
}
/// Build a UDP packet carrying `payload`, with given src/dst ports.
fn build_udp_packet(
src_port: u16,
dst_port: u16,
payload: &[u8],
src_ip: Ipv4Addr,
dst_ip: Ipv4Addr,
) -> Vec<u8> {
let udp_len = 8 + payload.len();
let mut buf = vec![0u8; udp_len];
{
let mut udp_pkt = MutableUdpPacket::new(&mut buf).unwrap();
udp_pkt.set_source(src_port);
udp_pkt.set_destination(dst_port);
udp_pkt.set_length(udp_len as u16);
udp_pkt.payload_mut().copy_from_slice(payload);
udp_pkt.set_checksum(udp::ipv4_checksum(
&udp_pkt.to_immutable(),
&src_ip,
&dst_ip,
));
}
buf
}
// ─── Tests ───────────────────────────────────────────────────────────
#[tokio::test]
async fn test_dynamic_catalog_replace() {
let catalog = DynamicCatalog::new();
let new_catalog = build_test_catalog();
catalog.replace(new_catalog).await;
// After replacement the catalog should resolve test.example.com
// This is implicitly verified by the UDP DNS test below; here we
// just make sure `replace` does not panic and completes.
}
#[tokio::test]
async fn test_is_hijacked_ip_and_addr() {
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 test_handle_icmp_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();
// 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);
}
#[tokio::test]
async fn test_handle_icmp_non_echo_ignored() {
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 test_non_hijacked_ip_ignored() {
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 test_handle_udp_dns_packet() {
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);
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]
async fn test_full_udp_dns_query() {
use hickory_server::ServerFuture;
use tokio::net::UdpSocket;
use tokio::time::timeout;
// Build a catalog with test.example.com -> 1.2.3.4.
let catalog = build_test_catalog();
// Bind to a random port.
let socket = UdpSocket::bind("127.0.0.1:0").await.unwrap();
let addr = socket.local_addr().unwrap();
let mut server = ServerFuture::new(catalog);
server.register_socket(socket);
let shutdown_token = server.shutdown_token().clone();
tokio::spawn(async move {
server.block_until_done().await.ok();
});
// Send a real DNS query using hickory_client.
let conn = UdpClientStream::builder(addr, TokioRuntimeProvider::default()).build();
let (mut client, bg) = timeout(Duration::from_secs(2), Client::connect(conn))
.await
.expect("client connect timeout")
.expect("client connect failed");
tokio::spawn(async move {
bg.await.ok();
});
let response = timeout(
Duration::from_secs(2),
client.query(
Name::from_str("test.example.com.").unwrap(),
DNSClass::IN,
RecordType::A,
),
)
.await
.expect("query timeout")
.expect("query failed");
assert!(!response.answers().is_empty(), "should get answers");
let a_record = &response.answers()[0];
if let RData::A(a) = a_record.data() {
assert_eq!(a.0, Ipv4Addr::new(1, 2, 3, 4));
} else {
panic!("expected A record, got {:?}", a_record.data());
}
// Shutdown the server.
shutdown_token.cancel();
}
}