This commit is contained in:
Luna Yao
2026-04-17 18:26:21 +02:00
parent 38cb4a22fd
commit ba7fc1098b
7 changed files with 95 additions and 99 deletions
+24 -18
View File
@@ -5,15 +5,6 @@ use std::{
sync::{Arc, Mutex}, sync::{Arc, Mutex},
}; };
use super::env_parser;
use crate::{
common::stun::StunInfoCollector,
proto::{
acl::Acl,
common::{CompressionAlgoPb, PortForwardConfigPb, SecureModeConfig, SocketType},
},
tunnel::generate_digest_from_str,
};
use anyhow::Context; use anyhow::Context;
use base64::{Engine as _, prelude::BASE64_STANDARD}; use base64::{Engine as _, prelude::BASE64_STANDARD};
use clap::ValueEnum; use clap::ValueEnum;
@@ -22,6 +13,18 @@ use serde::{Deserialize, Serialize};
use strum::{Display, EnumString, VariantArray}; use strum::{Display, EnumString, VariantArray};
use tokio::io::AsyncReadExt as _; use tokio::io::AsyncReadExt as _;
use crate::{
common::stun::StunInfoCollector,
instance::dns_server::DEFAULT_ET_DNS_ZONE,
proto::{
acl::Acl,
common::{CompressionAlgoPb, PortForwardConfigPb, SecureModeConfig, SocketType},
},
tunnel::generate_digest_from_str,
};
use super::env_parser;
pub type Flags = crate::proto::common::FlagsInConfig; pub type Flags = crate::proto::common::FlagsInConfig;
pub fn gen_default_flags() -> Flags { pub fn gen_default_flags() -> Flags {
@@ -116,10 +119,12 @@ impl Default for EncryptionAlgorithm {
} }
} }
cfg_if! { cfg_select! {
if #[cfg(feature = "magic-dns")] { feature = "magic-dns" => {
use crate::dns::config::{DnsConfig, DnsConfigLoaderExt}; use crate::dns::config::{DnsConfig, DnsConfigLoaderExt};
} else { }
_ => {
#[auto_impl::auto_impl(Box, &)] #[auto_impl::auto_impl(Box, &)]
pub trait DnsConfigLoaderExt {} pub trait DnsConfigLoaderExt {}
} }
@@ -536,16 +541,17 @@ impl TomlConfigLoader {
} }
impl DnsConfigLoaderExt for TomlConfigLoader { impl DnsConfigLoaderExt for TomlConfigLoader {
cfg_if! { cfg_select! {
if #[cfg(feature = "magic-dns")] { feature = "magic-dns" => {
fn get_dns(&self) -> DnsConfig { fn get_dns(&self) -> DnsConfig {
self.config.lock().unwrap().dns.clone().unwrap_or_default() self.config.lock().unwrap().dns.clone().unwrap_or_default()
} }
fn set_dns(&self, config: Option<DnsConfig>) {
self.config.lock().unwrap().dns = config;
}
}
fn set_dns(&self, dns: Option<DnsConfig>) { _ => {}
self.config.lock().unwrap().dns = dns;
}
}
} }
} }
+7 -10
View File
@@ -166,7 +166,6 @@ impl DnsNodeMgrRpc for DnsNodeMgr {
#[cfg(test)] #[cfg(test)]
mod tests { mod tests {
use super::*; use super::*;
use crate::common::log;
use crate::dns::tests::{ use crate::dns::tests::{
dns_snapshot_with as snapshot_with, heartbeat_with_snapshot, new_request, dns_snapshot_with as snapshot_with, heartbeat_with_snapshot, new_request,
zone_data_a_with_forwarders as valid_zone_data, zone_data_a_with_forwarders as valid_zone_data,
@@ -214,8 +213,6 @@ mod tests {
#[tokio::test] #[tokio::test]
async fn catalog_lookup_returns_record_after_snapshot_heartbeat() -> anyhow::Result<()> { async fn catalog_lookup_returns_record_after_snapshot_heartbeat() -> anyhow::Result<()> {
log::tests::init();
let mgr = DnsNodeMgr::new(); let mgr = DnsNodeMgr::new();
let id = Uuid::new_v4(); let id = Uuid::new_v4();
let snapshot = snapshot_with( let snapshot = snapshot_with(
@@ -291,7 +288,7 @@ mod tests {
let full = send_heartbeat(&mgr, heartbeat_with_snapshot(id, snapshot)).await; let full = send_heartbeat(&mgr, heartbeat_with_snapshot(id, snapshot)).await;
assert!(!full.resync); assert!(!full.resync);
let same = send_heartbeat(&mgr, heartbeat_digest_only(id, digest.clone())).await; let same = send_heartbeat(&mgr, heartbeat_digest_only(id, digest.into())).await;
assert!(!same.resync); assert!(!same.resync);
let different = send_heartbeat(&mgr, heartbeat_digest_only(id, vec![9, 9, 9])).await; let different = send_heartbeat(&mgr, heartbeat_digest_only(id, vec![9, 9, 9])).await;
@@ -381,7 +378,7 @@ mod tests {
.insert( .insert(
Uuid::new_v4(), Uuid::new_v4(),
DnsNodeInfo { DnsNodeInfo {
digest: vec![1], digest: [1; 32],
zones: vec![zone_a].into(), zones: vec![zone_a].into(),
addresses: [ns("udp://10.100.0.1:53"), ns("udp://10.100.0.2:53")] addresses: [ns("udp://10.100.0.1:53"), ns("udp://10.100.0.2:53")]
.into_iter() .into_iter()
@@ -394,7 +391,7 @@ mod tests {
.insert( .insert(
Uuid::new_v4(), Uuid::new_v4(),
DnsNodeInfo { DnsNodeInfo {
digest: vec![2], digest: [2; 32],
zones: vec![zone_b].into(), zones: vec![zone_b].into(),
addresses: [ns("udp://10.100.0.2:53"), ns("udp://10.100.0.3:53")] addresses: [ns("udp://10.100.0.2:53"), ns("udp://10.100.0.3:53")]
.into_iter() .into_iter()
@@ -438,7 +435,7 @@ mod tests {
.insert( .insert(
Uuid::new_v4(), Uuid::new_v4(),
DnsNodeInfo { DnsNodeInfo {
digest: vec![1], digest: [1; 32],
zones: vec![zone].into(), zones: vec![zone].into(),
addresses: [ns("udp://10.0.0.10:53")].into_iter().collect(), addresses: [ns("udp://10.0.0.10:53")].into_iter().collect(),
listeners: [ns("tcp://10.0.0.11:53")].into_iter().collect(), listeners: [ns("tcp://10.0.0.11:53")].into_iter().collect(),
@@ -531,7 +528,7 @@ mod tests {
let _ = send_heartbeat(&mgr, heartbeat_with_snapshot(node_a, snap_a)).await; let _ = send_heartbeat(&mgr, heartbeat_with_snapshot(node_a, snap_a)).await;
let a_same = send_heartbeat(&mgr, heartbeat_digest_only(node_a, digest_a)).await; let a_same = send_heartbeat(&mgr, heartbeat_digest_only(node_a, digest_a.into())).await;
assert!(!a_same.resync); assert!(!a_same.resync);
let b_unknown = send_heartbeat(&mgr, heartbeat_digest_only(node_b, vec![1, 2, 3])).await; let b_unknown = send_heartbeat(&mgr, heartbeat_digest_only(node_b, vec![1, 2, 3])).await;
@@ -550,12 +547,12 @@ mod tests {
let digest = snapshot.digest(); let digest = snapshot.digest();
let _ = send_heartbeat(&mgr, heartbeat_with_snapshot(id, snapshot)).await; let _ = send_heartbeat(&mgr, heartbeat_with_snapshot(id, snapshot)).await;
let before_expiry = send_heartbeat(&mgr, heartbeat_digest_only(id, digest.clone())).await; let before_expiry = send_heartbeat(&mgr, heartbeat_digest_only(id, digest.to_vec())).await;
assert!(!before_expiry.resync); assert!(!before_expiry.resync);
sleep(DNS_NODE_TTI + Duration::from_millis(300)).await; sleep(DNS_NODE_TTI + Duration::from_millis(300)).await;
let after_expiry = send_heartbeat(&mgr, heartbeat_digest_only(id, digest)).await; let after_expiry = send_heartbeat(&mgr, heartbeat_digest_only(id, digest.into())).await;
assert!(after_expiry.resync); assert!(after_expiry.resync);
} }
} }
+14 -10
View File
@@ -272,7 +272,7 @@ mod tests {
.insert( .insert(
999_999, 999_999,
DnsPeerInfo { DnsPeerInfo {
digest: vec![1, 2, 3], digest: [9; 32],
zones: vec![valid_zone_data("peer-cache.test", "10.20.30.40")], zones: vec![valid_zone_data("peer-cache.test", "10.20.30.40")],
}, },
) )
@@ -351,7 +351,7 @@ mod tests {
.insert( .insert(
11, 11,
DnsPeerInfo { DnsPeerInfo {
digest: vec![11], digest: [11; 32],
zones: vec![valid_zone_data("peer-a.test", "10.20.30.41")], zones: vec![valid_zone_data("peer-a.test", "10.20.30.41")],
}, },
) )
@@ -360,7 +360,7 @@ mod tests {
.insert( .insert(
12, 12,
DnsPeerInfo { DnsPeerInfo {
digest: vec![12], digest: [12; 32],
zones: vec![valid_zone_data("peer-b.test", "10.20.30.42")], zones: vec![valid_zone_data("peer-b.test", "10.20.30.42")],
}, },
) )
@@ -390,7 +390,7 @@ mod tests {
.insert( .insert(
13, 13,
DnsPeerInfo { DnsPeerInfo {
digest: vec![13], digest: [13; 32],
zones: vec![], zones: vec![],
}, },
) )
@@ -489,7 +489,9 @@ mod tests {
.insert( .insert(
remote_id, remote_id,
DnsPeerInfo { DnsPeerInfo {
digest: remote_route_dns, digest: remote_route_dns
.try_into()
.expect("route dns digest should be 32 bytes"),
zones: vec![valid_zone_data("cached-same.test", "10.0.1.9")], zones: vec![valid_zone_data("cached-same.test", "10.0.1.9")],
}, },
) )
@@ -617,7 +619,7 @@ mod tests {
.insert( .insert(
fail_id, fail_id,
DnsPeerInfo { DnsPeerInfo {
digest: vec![1], digest: [1; 32],
zones: vec![valid_zone_data("cached-fail.test", "10.2.1.20")], zones: vec![valid_zone_data("cached-fail.test", "10.2.1.20")],
}, },
) )
@@ -627,7 +629,7 @@ mod tests {
.insert( .insert(
keep_id, keep_id,
DnsPeerInfo { DnsPeerInfo {
digest: vec![2], digest: [2; 32],
zones: vec![valid_zone_data("cached-keep.test", "10.2.1.21")], zones: vec![valid_zone_data("cached-keep.test", "10.2.1.21")],
}, },
) )
@@ -700,7 +702,7 @@ mod tests {
.insert( .insert(
changed_peer.my_peer_id(), changed_peer.my_peer_id(),
DnsPeerInfo { DnsPeerInfo {
digest: vec![0], digest: [0; 32],
zones: vec![valid_zone_data("stale-changed.test", "10.2.2.20")], zones: vec![valid_zone_data("stale-changed.test", "10.2.2.20")],
}, },
) )
@@ -710,7 +712,9 @@ mod tests {
.insert( .insert(
unchanged_id, unchanged_id,
DnsPeerInfo { DnsPeerInfo {
digest: unchanged_digest, digest: unchanged_digest
.try_into()
.expect("route dns digest should be 32 bytes"),
zones: vec![valid_zone_data("cached-unchanged.test", "10.2.2.21")], zones: vec![valid_zone_data("cached-unchanged.test", "10.2.2.21")],
}, },
) )
@@ -753,7 +757,7 @@ mod tests {
.insert( .insert(
cached_peer_id, cached_peer_id,
DnsPeerInfo { DnsPeerInfo {
digest: vec![6, 6, 6], digest: [6; 32],
zones: vec![valid_zone_data("cached-expire.test", "10.3.0.2")], zones: vec![valid_zone_data("cached-expire.test", "10.3.0.2")],
}, },
) )
+12 -17
View File
@@ -134,8 +134,7 @@ impl DnsServer {
if let Some(nic_ctx) = nic_ctx if let Some(nic_ctx) = nic_ctx
.as_ref() .as_ref()
.and_then(|nic_ctx| nic_ctx.downcast_ref::<NicCtx>()) .and_then(|nic_ctx| nic_ctx.downcast_ref::<NicCtx>())
{ && let Some(system) = nic_ctx
if let Some(system) = nic_ctx
.ifname() .ifname()
.await .await
.map(|ifname| system::get(&ifname)) .map(|ifname| system::get(&ifname))
@@ -160,7 +159,6 @@ impl DnsServer {
})?; })?;
} }
} }
}
*self.addresses.write() = addresses; *self.addresses.write() = addresses;
@@ -181,11 +179,11 @@ impl DnsServer {
} }
tracing::info!(?listeners, "reloading"); tracing::info!(?listeners, "reloading");
if let Some(runtime) = runtime.as_ref() { if let Some(runtime) = runtime.as_ref()
if let Some(Err(error)) = runtime.stop(None).await { && let Some(Err(error)) = runtime.stop(None).await
{
tracing::error!(?error, "failed to stop old DNS server runtime"); tracing::error!(?error, "failed to stop old DNS server runtime");
} }
}
let runtime = runtime.get_or_insert_default(); let runtime = runtime.get_or_insert_default();
@@ -213,7 +211,7 @@ impl DnsServer {
.unwrap_or_else(|e| tracing::error!("DNS server exited with error: {:?}", e)); .unwrap_or_else(|e| tracing::error!("DNS server exited with error: {:?}", e));
} }
.instrument(tracing::info_span!("DNS server backend runtime")) .instrument(tracing::info_span!("DNS server backend runtime"))
}); })?;
*self.listeners.write() = listeners; *self.listeners.write() = listeners;
@@ -238,12 +236,12 @@ impl DnsServer {
let reload_addresses = async { let reload_addresses = async {
loop { loop {
dirty.addresses.wait().await; dirty.addresses.wait().await;
if dirty.addresses.reset() { if dirty.addresses.reset()
if let Err(error) = self.reload_addresses(self.mgr.iter_addresses()).await { && let Err(error) = self.reload_addresses(self.mgr.iter_addresses()).await
{
tracing::error!(?error, "failed to reload addresses"); tracing::error!(?error, "failed to reload addresses");
dirty.addresses.mark(); dirty.addresses.mark();
} }
}
tokio::time::sleep(Duration::from_secs(1)).await; tokio::time::sleep(Duration::from_secs(1)).await;
} }
}; };
@@ -251,15 +249,14 @@ impl DnsServer {
let reload_listeners = async { let reload_listeners = async {
loop { loop {
dirty.listeners.wait().await; dirty.listeners.wait().await;
if dirty.listeners.reset() { if dirty.listeners.reset()
if let Err(error) = self && let Err(error) = self
.reload_listeners(self.mgr.iter_listeners(), &mut runtime) .reload_listeners(self.mgr.iter_listeners(), &mut runtime)
.await .await
{ {
tracing::error!(?error, "failed to reload listeners"); tracing::error!(?error, "failed to reload listeners");
dirty.listeners.mark(); dirty.listeners.mark();
} }
}
tokio::time::sleep(Duration::from_secs(1)).await; tokio::time::sleep(Duration::from_secs(1)).await;
} }
}; };
@@ -284,8 +281,7 @@ impl DnsServer {
.await .await
.as_ref() .as_ref()
.and_then(|nic_ctx| nic_ctx.downcast_ref::<NicCtx>()) .and_then(|nic_ctx| nic_ctx.downcast_ref::<NicCtx>())
{ && let Some(system) = nic_ctx
if let Some(system) = nic_ctx
.ifname() .ifname()
.await .await
.and_then(|ifname| system::get(&ifname).ok()) .and_then(|ifname| system::get(&ifname).ok())
@@ -293,7 +289,6 @@ impl DnsServer {
{ {
let _ = system.clean(); let _ = system.clean();
} }
}
if let Some(runtime) = runtime.take() { if let Some(runtime) = runtime.take() {
let _ = runtime.stop(None).await; let _ = runtime.stop(None).await;
@@ -1057,7 +1052,7 @@ mod tests {
assert!(!response.answers().is_empty()); assert!(!response.answers().is_empty());
if let Some(runtime) = runtime.take() { if let Some(runtime) = runtime.take() {
let _ = runtime.stop().await; let _ = runtime.stop(None).await;
} }
} }
+3 -8
View File
@@ -167,10 +167,8 @@ impl SystemConfigurator for WindowsDNSManager {
#[cfg(all(test, target_os = "windows", feature = "magic-dns", feature = "tun"))] #[cfg(all(test, target_os = "windows", feature = "magic-dns", feature = "tun"))]
mod tests { mod tests {
use std::net::IpAddr;
use crate::common::log;
use cidr::Ipv4Inet; use cidr::Ipv4Inet;
use std::net::IpAddr;
#[tokio::test] #[tokio::test]
async fn test_windows_set_primary_server() { async fn test_windows_set_primary_server() {
@@ -184,8 +182,6 @@ mod tests {
use crate::instance::virtual_nic::NicCtx; use crate::instance::virtual_nic::NicCtx;
use crate::peers::peer_manager::PeerManager; use crate::peers::peer_manager::PeerManager;
log::tests::init();
let tun_ip = Ipv4Inet::from_str("10.144.144.10/24").unwrap(); let tun_ip = Ipv4Inet::from_str("10.144.144.10/24").unwrap();
let (peer_mgr, virtual_nic): (Arc<PeerManager>, NicCtx) = let (peer_mgr, virtual_nic): (Arc<PeerManager>, NicCtx) =
prepare_env("test1", tun_ip).await; prepare_env("test1", tun_ip).await;
@@ -212,16 +208,15 @@ mod tests {
.arg("1") .arg("1")
.arg("-w") .arg("-w")
.arg("100") .arg("100")
.arg(&fake_ip.to_string()) .arg(fake_ip.to_string())
.output() .output()
.await .await
&& o.status.success()
{ {
if o.status.success() {
ping_ready = true; ping_ready = true;
break; break;
} }
} }
}
if !ping_ready { if !ping_ready {
tracing::warn!( tracing::warn!(
"dns test endpoint {} did not respond to ping in time; continue with dns checks", "dns test endpoint {} did not respond to ping in time; continue with dns checks",
+6 -7
View File
@@ -76,7 +76,7 @@ pub fn start_dns_node(peer_mgr: Arc<PeerManager>, virtual_nic: NicCtx) -> DnsNod
let nic_ctx: ArcNicCtx = Arc::new(tokio::sync::Mutex::new(Some(Box::new(virtual_nic)))); let nic_ctx: ArcNicCtx = Arc::new(tokio::sync::Mutex::new(Some(Box::new(virtual_nic))));
let dns_node = DnsNode::new(peer_mgr, global_ctx, nic_ctx); let dns_node = DnsNode::new(peer_mgr, global_ctx, nic_ctx);
dns_node.start(); dns_node.start().expect("failed to start dns node");
dns_node dns_node
} }
@@ -85,7 +85,7 @@ pub fn start_dns_node_without_nic(peer_mgr: Arc<PeerManager>) -> DnsNode {
let nic_ctx: ArcNicCtx = Arc::new(tokio::sync::Mutex::new(None)); let nic_ctx: ArcNicCtx = Arc::new(tokio::sync::Mutex::new(None));
let dns_node = DnsNode::new(peer_mgr, global_ctx, nic_ctx); let dns_node = DnsNode::new(peer_mgr, global_ctx, nic_ctx);
dns_node.start(); dns_node.start().expect("failed to start dns node");
dns_node dns_node
} }
@@ -178,13 +178,12 @@ pub async fn check_dns_record_at(server_addr: SocketAddr, domain: &str, expected
let attempt_err = match query_result { let attempt_err = match query_result {
Ok(Ok(response)) => { Ok(Ok(response)) => {
if response.answers().len() == 1 { if response.answers().len() == 1
if let Some(resp) = response.answers().first() { && let Some(resp) = response.answers().first()
if resp.clone().into_parts().rdata.into_a().unwrap().0 == expected { && resp.clone().into_parts().rdata.into_a().unwrap().0 == expected
{
return; return;
} }
}
}
format!("unexpected response: {:?}", response.answers()) format!("unexpected response: {:?}", response.answers())
} }
Ok(Err(e)) => { Ok(Err(e)) => {
+2 -2
View File
@@ -535,14 +535,14 @@ pub fn bind<B: Bindable>(
B::finalize(socket) B::finalize(socket)
} }
// endregion
pub fn reserve_buf(buf: &mut BytesMut, min_size: usize, max_size: usize) { pub fn reserve_buf(buf: &mut BytesMut, min_size: usize, max_size: usize) {
if buf.capacity() < min_size { if buf.capacity() < min_size {
buf.reserve(max_size); buf.reserve(max_size);
} }
} }
// endregion
pub mod tests { pub mod tests {
use atomic_shim::AtomicU64; use atomic_shim::AtomicU64;
use std::{sync::Arc, time::Instant}; use std::{sync::Arc, time::Instant};