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