diff --git a/.unreleased/telio_dns_cleanup b/.unreleased/telio_dns_cleanup new file mode 100644 index 000000000..e69de29bb diff --git a/crates/telio-core/src/device.rs b/crates/telio-core/src/device.rs index 6cb6a7bfc..0a48b24bc 100644 --- a/crates/telio-core/src/device.rs +++ b/crates/telio-core/src/device.rs @@ -2332,7 +2332,7 @@ impl Runtime { if is_meshnet_exit_node { if let Some(dns) = &self.entities.dns.lock().await.resolver { - self.reconfigure_dns_peer(dns, &dns.get_default_dns_servers()) + self.reconfigure_dns_peer(dns, &dns.get_exit_node_dns_servers()) .await?; } } else { diff --git a/crates/telio-dns/src/dns.rs b/crates/telio-dns/src/dns.rs index 924170913..cdb41ab01 100644 --- a/crates/telio-dns/src/dns.rs +++ b/crates/telio-dns/src/dns.rs @@ -6,7 +6,7 @@ use crate::{ use async_trait::async_trait; use ipnet::IpNet; use neptun::noise::Tunn; -use std::net::{IpAddr, Ipv4Addr, Ipv6Addr}; +use std::net::IpAddr; use std::{net::SocketAddr, sync::Arc}; use telio_crypto::{PublicKey, SecretKey}; use telio_wg::uapi::Peer; @@ -15,7 +15,10 @@ use tokio::net::UdpSocket; use tokio::sync::{Mutex, RwLock}; use x25519_dalek::{PublicKey as PublicKeyDalek, StaticSecret}; -use telio_model::features::{FeatureDns, TtlValue}; +use telio_model::{ + constants::{DNS_EXIT_IPV4, DNS_EXIT_IPV6, DNS_VIRTUAL_IPV4, DNS_VIRTUAL_IPV6}, + features::{FeatureDns, TtlValue}, +}; //debug tools use telio_utils::{telio_log_debug, telio_log_error}; @@ -38,10 +41,10 @@ pub trait DnsResolver { fn get_peer(&self, allowed_ips: Vec) -> Peer; /// Get default allowed IPs of this DNS server. fn get_default_dns_allowed_ips(&self) -> Vec; - /// Get DNS virtual peer addresses. + /// Get allowed IPs when connected to exit node. fn get_exit_connected_dns_allowed_ips(&self) -> Vec; - /// Get default DNS server IP addresses. - fn get_default_dns_servers(&self) -> Vec; + /// Get DNS server IP addresses when connected to exit node. + fn get_exit_node_dns_servers(&self) -> Vec; /// Change DNS peer's public key async fn set_peer_public_key(&self, key: PublicKey); } @@ -156,25 +159,22 @@ impl DnsResolver for LocalDnsResolver { fn get_default_dns_allowed_ips(&self) -> Vec { vec![ - IpAddr::V4(Ipv4Addr::new(100, 64, 0, 2)).into(), - IpAddr::V4(Ipv4Addr::new(100, 64, 0, 3)).into(), - IpAddr::V6(Ipv6Addr::new(0xfd74, 0x656c, 0x696f, 0, 0, 0, 0, 2)).into(), - IpAddr::V6(Ipv6Addr::new(0xfd74, 0x656c, 0x696f, 0, 0, 0, 0, 3)).into(), + IpAddr::V4(DNS_VIRTUAL_IPV4).into(), + IpAddr::V4(DNS_EXIT_IPV4).into(), + IpAddr::V6(DNS_VIRTUAL_IPV6).into(), + IpAddr::V6(DNS_EXIT_IPV6).into(), ] } fn get_exit_connected_dns_allowed_ips(&self) -> Vec { vec![ - IpAddr::V4(Ipv4Addr::new(100, 64, 0, 2)).into(), - IpAddr::V6(Ipv6Addr::new(0xfd74, 0x656c, 0x696f, 0, 0, 0, 0, 2)).into(), + IpAddr::V4(DNS_VIRTUAL_IPV4).into(), + IpAddr::V6(DNS_VIRTUAL_IPV6).into(), ] } - fn get_default_dns_servers(&self) -> Vec { - vec![ - IpAddr::V4(Ipv4Addr::new(100, 64, 0, 3)), - IpAddr::V6(Ipv6Addr::new(0xfd74, 0x656c, 0x696f, 0, 0, 0, 0, 3)), - ] + fn get_exit_node_dns_servers(&self) -> Vec { + vec![IpAddr::V4(DNS_EXIT_IPV4), IpAddr::V6(DNS_EXIT_IPV6)] } async fn set_peer_public_key(&self, pubkey: PublicKey) { @@ -257,7 +257,7 @@ mod tests { "100.64.0.3".parse::().unwrap(), "fd74:656c:696f::3".parse::().unwrap(), ], - resolver.get_default_dns_servers() + resolver.get_exit_node_dns_servers() ); } } diff --git a/crates/telio-dns/src/error.rs b/crates/telio-dns/src/error.rs index aa7c823bc..17de73565 100644 --- a/crates/telio-dns/src/error.rs +++ b/crates/telio-dns/src/error.rs @@ -1,4 +1,4 @@ -use crate::forwarder::ForwardError; +use crate::udp_forwarder::ForwardError; use crate::zone::NordZoneError; use std::{io, net::AddrParseError}; use thiserror::Error; diff --git a/crates/telio-dns/src/lib.rs b/crates/telio-dns/src/lib.rs index 5992e4903..3dcc6513c 100644 --- a/crates/telio-dns/src/lib.rs +++ b/crates/telio-dns/src/lib.rs @@ -5,13 +5,12 @@ //! Easily create and run in process dns resolver. mod dns; -// TODO: LLT-7053 remove after integrating forwarder -#[allow(dead_code)] -mod forwarder; mod nameserver; mod packet_decoder; mod packet_encoder; mod resolver; +mod udp_forwarder; +mod upstream; mod zone; pub mod bind_tun; @@ -35,3 +34,6 @@ pub mod fuzz { pub use super::packet_decoder::{find_nord_query, parse_dns_query_packet}; pub use super::packet_encoder::fuzz_build_response; } + +/// DNS port number +pub const DNS_PORT: u16 = 53; diff --git a/crates/telio-dns/src/nameserver.rs b/crates/telio-dns/src/nameserver.rs index af4ab9762..296d2eb19 100644 --- a/crates/telio-dns/src/nameserver.rs +++ b/crates/telio-dns/src/nameserver.rs @@ -1,9 +1,11 @@ use crate::error::Result as DnsResult; +use crate::DNS_PORT; use crate::{ - forwarder::UdpForwarder, packet_decoder::{find_nord_query, normalize_qname, parse_dns_query_packet, DnsParseError}, packet_encoder::{DnsBuildError, DnsResponseBuilder}, resolver::Resolver, + udp_forwarder::UdpForwarder, + upstream::UpstreamState, zone::{AuthoritativeZone, ClonableZones, ForwardZone, NordZone, Records, NORD_ZONE}, }; use async_trait::async_trait; @@ -44,7 +46,6 @@ const UDP_HEADER: usize = 8; const TCP_MIN_HEADER: usize = 20; const MAX_CONCURRENT_QUERIES: usize = 256; const IDLE_TIME: Duration = Duration::from_secs(1); -const DNS_PORT: u16 = 53; #[derive(Debug, Error)] enum PacketError { @@ -157,6 +158,7 @@ pub struct LocalNameServer { zones: Arc, task_handle: Option>, forwarder: Option, + upstreams: Arc>, } impl LocalNameServer { @@ -166,8 +168,9 @@ impl LocalNameServer { forward_ips: &[IpAddr], use_raw_forwarder: bool, ) -> DnsResult>> { + let upstreams = Arc::new(Mutex::new(UpstreamState::default())); let raw_forwarder: Option = if use_raw_forwarder { - Some(UdpForwarder::new().await?) + Some(UdpForwarder::new(upstreams.clone()).await?) } else { None }; @@ -176,6 +179,7 @@ impl LocalNameServer { zones: Arc::new(ClonableZones::new()), task_handle: None, forwarder: raw_forwarder, + upstreams, })); ns.forward(forward_ips).await?; Ok(ns) @@ -908,11 +912,8 @@ impl NameServer for Arc> { } async fn forward_to_addrs(&self, to: &[SocketAddr]) -> DnsResult<()> { - let ns = self.read().await; - if let Some(forwarder) = &ns.forwarder { - forwarder.set_upstreams(to.to_vec()).await; - } - + let upstreams = self.read().await.upstreams.clone(); + upstreams.lock().await.set(to.to_vec()); Ok(()) } @@ -1070,6 +1071,41 @@ mod tests { assert!(ns.forwarder.is_some()); } + #[tokio::test] + async fn forward_updates_shared_upstreams() { + let nameserver = LocalNameServer::new(&[IpAddr::V4(Ipv4Addr::new(8, 8, 8, 8))], false) + .await + .unwrap(); + { + let ns = nameserver.read().await; + let state = ns.upstreams.lock().await; + assert_eq!( + state.addrs(), + [SocketAddr::new( + IpAddr::V4(Ipv4Addr::new(8, 8, 8, 8)), + DNS_PORT + )] + ); + assert_eq!(state.generation(), 1); + } + + nameserver + .forward(&[IpAddr::V4(Ipv4Addr::new(1, 1, 1, 1))]) + .await + .unwrap(); + + let ns = nameserver.read().await; + let state = ns.upstreams.lock().await; + assert_eq!( + state.addrs(), + [SocketAddr::new( + IpAddr::V4(Ipv4Addr::new(1, 1, 1, 1)), + DNS_PORT + )] + ); + assert_eq!(state.generation(), 2); + } + #[tokio::test] async fn nameserver_skips_forwarder_by_default() { let nameserver = LocalNameServer::new(&[IpAddr::V4(Ipv4Addr::new(8, 8, 8, 8))], false) @@ -1195,7 +1231,7 @@ mod tests { // Tests PacketError::InvalidUdpChecksum #[test] fn packet_error_invalid_udp_checksum() { - let mut udp_seg = build_udp_segment(12345, 53, &[0; 4]); + let mut udp_seg = build_udp_segment(12345, DNS_PORT, &[0; 4]); // Corrupt UDP checksum (bytes 6-7) udp_seg[6] ^= 0xFF; let packet = build_ipv4_packet(IpNextHeaderProtocols::Udp, &udp_seg); @@ -1235,7 +1271,7 @@ mod tests { 0x01, b'a', 0x00, // QTYPE and QCLASS intentionally missing → hickory fails to decode ]; - let udp_seg = build_udp_segment(12345, 53, dns_payload); + let udp_seg = build_udp_segment(12345, DNS_PORT, dns_payload); let packet = build_ipv4_packet(IpNextHeaderProtocols::Udp, &udp_seg); let ns = test_nameserver().await; let mut response = vec![0u8; MAX_PACKET]; @@ -1260,7 +1296,7 @@ mod tests { }, payload: PayloadRequestInfo::Udp { source_port: 12345, - destination_port: 53, + destination_port: DNS_PORT, dns_request: None, }, }; @@ -1282,7 +1318,7 @@ mod tests { }, payload: PayloadRequestInfo::Udp { source_port: 12345, - destination_port: 53, + destination_port: DNS_PORT, dns_request: None, }, }; @@ -1308,7 +1344,7 @@ mod tests { }, payload: PayloadRequestInfo::Udp { source_port: 12345, - destination_port: 53, + destination_port: DNS_PORT, dns_request: None, }, }; diff --git a/crates/telio-dns/src/forwarder.rs b/crates/telio-dns/src/udp_forwarder.rs similarity index 75% rename from crates/telio-dns/src/forwarder.rs rename to crates/telio-dns/src/udp_forwarder.rs index bfb027f8f..0459da868 100644 --- a/crates/telio-dns/src/forwarder.rs +++ b/crates/telio-dns/src/udp_forwarder.rs @@ -4,12 +4,16 @@ //! Multiple queries can be in flight concurrently. //! //! Upstream resolvers are tried in order they are set. -//! If no response is received after the set timeout, the next +//! If no response is received within `QUERY_TIMEOUT`, the next //! upstream resolver is tried. -use crate::{bind_tun::bind_to_tun, packet_encoder::DNS_HEADER_OFFSET}; +use crate::{ + bind_tun::bind_to_tun, + packet_encoder::DNS_HEADER_OFFSET, + upstream::{AttemptState, UpstreamState}, +}; use rand::RngExt; -use std::{collections::HashMap, io, net::SocketAddr, sync::Arc, time::Duration}; +use std::{collections::HashMap, io, sync::Arc, time::Duration}; use telio_utils::{sleep_until, telio_log_debug, telio_log_warn, Instant}; use thiserror::Error; use tokio::{ @@ -21,8 +25,8 @@ use tokio::{ const FORWARDER_BUFFER_SIZE: usize = 4096; /// Channel size for forward messages const CHANNEL_SIZE: usize = 256; -/// Default timeout for upstream DNS queries -const DEFAULT_QUERY_TIMEOUT: Duration = Duration::from_secs(2); +/// Timeout for one upstream DNS query attempt +const QUERY_TIMEOUT: Duration = Duration::from_secs(2); /// Maximum number of DNS IDs const DNS_ID_SPACE: u32 = 65536; @@ -64,19 +68,10 @@ pub enum ForwardError { /// DNS query forwarder bound to the tunnel interface /// /// Forwards raw DNS queries to upstream resolvers. -/// Next upstream resolver is tried if no response is received after `timeout`. +/// Next upstream resolver is tried if no response is received within `QUERY_TIMEOUT`. #[derive(Clone, Debug)] pub(crate) struct UdpForwarder { tx: mpsc::Sender, - upstreams: Arc>, - timeout: Arc>, -} - -/// Upstream resolver list with a generation counter to detect changes -#[derive(Clone, Debug, Default)] -struct UpstreamState { - addrs: Vec, - generation: u64, } /// Internal message for forwarding a DNS query @@ -93,10 +88,8 @@ struct PendingQuery { original_id: u16, /// Raw DNS packet bytes query_bytes: Vec, - /// Current upstream index for request - upstream_index: usize, - /// Generation of the upstream list when this query was last dispatched - upstream_generation: u64, + /// Failover cursor over the upstream list + attempt: AttemptState, /// Channel to send the response back to the caller respond_to: oneshot::Sender, ForwardError>>, /// Instant when request times out @@ -144,22 +137,16 @@ fn allocate_id(pending: &HashMap, next_id: &mut u16) -> Optio impl UdpForwarder { /// Create a new DNS forwarder - pub(crate) async fn new() -> Result { + pub(crate) async fn new(upstreams: Arc>) -> Result { let socket = Arc::new(UdpSocket::bind("0.0.0.0:0").await?); bind_to_tun(&socket)?; let (tx, rx) = mpsc::channel(CHANNEL_SIZE); - let upstreams = Arc::new(Mutex::new(UpstreamState::default())); - let timeout = Arc::new(Mutex::new(DEFAULT_QUERY_TIMEOUT)); - tokio::spawn(Self::run(socket, rx, upstreams.clone(), timeout.clone())); + tokio::spawn(Self::run(socket, rx, upstreams)); - Ok(UdpForwarder { - tx, - upstreams, - timeout, - }) + Ok(UdpForwarder { tx }) } /// Forward a DNS query to the configured upstream resolver @@ -178,26 +165,11 @@ impl UdpForwarder { recv.await.map_err(|_| ForwardError::ChannelClosed)? } - /// Update the list of upstream resolvers - pub(crate) async fn set_upstreams(&self, addrs: Vec) { - let mut state = self.upstreams.lock().await; - state.generation = state.generation.wrapping_add(1); - state.addrs = addrs; - } - - // TODO: LLT-7054: remove if not needed after complete regression testing - #[cfg(test)] - /// Update the timeout for upstream DNS queries - pub(crate) async fn set_timeout(&self, timeout: Duration) { - *self.timeout.lock().await = timeout; - } - /// Main async loop that forwards incoming requests async fn run( socket: Arc, mut rx: mpsc::Receiver, upstreams: Arc>, - timeout: Arc>, ) { telio_log_debug!("Forwarder starting"); let mut recv_buf = vec![0u8; FORWARDER_BUFFER_SIZE]; @@ -216,7 +188,6 @@ impl UdpForwarder { &socket, forward_msg, &upstreams, - &timeout, &mut pending, &mut next_id, ).await; @@ -234,7 +205,7 @@ impl UdpForwarder { // validate the src IP let is_known_upstream = { let state = upstreams.lock().await; - state.addrs.iter().any(|u| u.ip() == src.ip()) + state.addrs().iter().any(|u| u.ip() == src.ip()) }; if !is_known_upstream { telio_log_warn!("Received DNS response from unknown source: {src}, ignoring"); @@ -265,7 +236,6 @@ impl UdpForwarder { Self::handle_timeouts( &socket, &upstreams, - &timeout, &mut pending, ).await; } @@ -289,7 +259,6 @@ impl UdpForwarder { socket: &UdpSocket, msg: ForwardQuery, upstreams: &Arc>, - timeout: &Arc>, pending: &mut HashMap, next_id: &mut u16, ) { @@ -302,10 +271,11 @@ impl UdpForwarder { } }; - let (upstream_addr, generation) = { + let (upstream_addr, attempt) = { let state = upstreams.lock().await; - match state.addrs.first().cloned() { - Some(addr) => (addr, state.generation), + let mut attempt = AttemptState::new(state.generation()); + match attempt.advance(&state) { + Some(addr) => (addr, attempt), None => { send_channel_response!(msg.respond_to, Err(ForwardError::NoUpstreams)); return; @@ -318,16 +288,14 @@ impl UdpForwarder { return; } - let query_timeout = *timeout.lock().await; pending.insert( internal_id, PendingQuery { original_id, query_bytes: rewritten_bytes, - upstream_index: 0, - upstream_generation: generation, + attempt, respond_to: msg.respond_to, - deadline: Instant::now() + query_timeout, + deadline: Instant::now() + QUERY_TIMEOUT, }, ); } @@ -382,7 +350,6 @@ impl UdpForwarder { async fn handle_timeouts( socket: &UdpSocket, upstreams: &Arc>, - timeout: &Arc>, pending: &mut HashMap, ) { let now = Instant::now(); @@ -398,11 +365,6 @@ impl UdpForwarder { return; } - let current_state = { - let locked = upstreams.lock().await; - locked.clone() - }; - for (internal_id, is_closed) in expired_ids { let mut entry = match pending.remove(&internal_id) { Some(e) => e, @@ -414,23 +376,24 @@ impl UdpForwarder { continue; } - let next_index = if entry.upstream_generation == current_state.generation { - entry.upstream_index + 1 - } else { - telio_log_debug!( - "Upstreams changed for request: {internal_id}, restarting from index 0" - ); - 0 + let next_upstream = { + let current_state = upstreams.lock().await; + + if entry.attempt.generation() != current_state.generation() { + telio_log_debug!( + "Upstreams changed for request: {internal_id}, restarting from index 0" + ); + } + + entry.attempt.advance(¤t_state) }; - match current_state.addrs.get(next_index) { - Some(&next_upstream) => { + match next_upstream { + Some(next_upstream) => { telio_log_debug!( "Upstream timed out for request: {internal_id}, trying next: {next_upstream}" ); - entry.upstream_index = next_index; - entry.upstream_generation = current_state.generation; - entry.deadline = Instant::now() + *timeout.lock().await; + entry.deadline = Instant::now() + QUERY_TIMEOUT; if let Err(e) = socket.send_to(&entry.query_bytes, next_upstream).await { send_channel_response!(entry.respond_to, Err(ForwardError::SendFailed(e))); @@ -451,6 +414,7 @@ impl UdpForwarder { #[cfg(test)] mod tests { use super::*; + use std::net::SocketAddr; use tokio::task::JoinHandle; const TEST_PACKET_ID: u16 = 0x1234; @@ -507,12 +471,19 @@ mod tests { (addr, handle) } + /// Forwarder over a fresh shared upstream state preloaded with `addrs` + async fn new_forwarder(addrs: Vec) -> (UdpForwarder, Arc>) { + let upstreams = Arc::new(Mutex::new(UpstreamState::default())); + upstreams.lock().await.set(addrs); + let forwarder = UdpForwarder::new(upstreams.clone()).await.unwrap(); + (forwarder, upstreams) + } + fn dummy_pending() -> PendingQuery { PendingQuery { original_id: 0, query_bytes: vec![], - upstream_index: 0, - upstream_generation: 0, + attempt: AttemptState::new(0), respond_to: oneshot::channel().0, deadline: Instant::now(), } @@ -575,33 +546,9 @@ mod tests { assert_eq!(allocate_id(&pending, &mut next_id), None); } - #[tokio::test] - async fn set_upstreams_stores_resolvers() { - let forwarder = UdpForwarder::new().await.unwrap(); - let addrs = vec!["8.8.8.8:53".parse().unwrap(), "1.1.1.1:53".parse().unwrap()]; - - forwarder.set_upstreams(addrs.clone()).await; - - let state = forwarder.upstreams.lock().await; - assert_eq!(state.addrs, addrs); - } - - #[tokio::test] - async fn set_upstreams_replaces_existing() { - let forwarder = UdpForwarder::new().await.unwrap(); - let first_addrs = vec!["8.8.8.8:53".parse().unwrap()]; - let second_addrs = vec!["1.1.1.1:53".parse().unwrap(), "8.8.4.4:53".parse().unwrap()]; - - forwarder.set_upstreams(first_addrs).await; - forwarder.set_upstreams(second_addrs.clone()).await; - - let state = forwarder.upstreams.lock().await; - assert_eq!(state.addrs, second_addrs); - } - #[tokio::test] async fn query_returns_no_upstreams_when_empty() { - let forwarder = UdpForwarder::new().await.unwrap(); + let (forwarder, _upstreams) = new_forwarder(vec![]).await; let request = make_dns_packet(TEST_PACKET_ID, TEST_DNS_PAYLOAD); let result = forwarder.query(&request).await; @@ -612,40 +559,12 @@ mod tests { } } - #[tokio::test] - async fn cloned_forwarders_share_upstream_state() { - let forwarder1 = UdpForwarder::new().await.unwrap(); - let forwarder2 = forwarder1.clone(); - - let addrs = vec!["8.8.8.8:53".parse().unwrap()]; - forwarder1.set_upstreams(addrs.clone()).await; - - let state = forwarder2.upstreams.lock().await; - assert_eq!(state.addrs, addrs); - } - - #[tokio::test] - async fn set_timeout_updates_value() { - let forwarder = UdpForwarder::new().await.unwrap(); - let target_timeout = Duration::from_secs(10); - - forwarder.set_timeout(target_timeout).await; - - let timeout = *forwarder.timeout.lock().await; - assert_eq!(timeout, target_timeout); - } - #[tokio::test] async fn forward_query_timeout() { let (blackhole_addr1, _bh1) = spawn_stub(StubBehavior::BlackHole).await; let (blackhole_addr2, _bh2) = spawn_stub(StubBehavior::BlackHole).await; - let target_timeout = Duration::from_millis(50); - let forwarder = UdpForwarder::new().await.unwrap(); - forwarder - .set_upstreams(vec![blackhole_addr1, blackhole_addr2]) - .await; - forwarder.set_timeout(target_timeout).await; + let (forwarder, _upstreams) = new_forwarder(vec![blackhole_addr1, blackhole_addr2]).await; let request = make_dns_packet(TEST_PACKET_ID, TEST_DNS_PAYLOAD); let result = forwarder.query(&request).await; @@ -663,8 +582,7 @@ mod tests { async fn forward_query_returns_response_from_upstream() { let (addr, _handle) = spawn_stub(StubBehavior::Echo).await; - let forwarder = UdpForwarder::new().await.unwrap(); - forwarder.set_upstreams(vec![addr]).await; + let (forwarder, _upstreams) = new_forwarder(vec![addr]).await; let request = make_dns_packet(TEST_PACKET_ID, TEST_DNS_PAYLOAD); let result = forwarder.query(&request).await.unwrap(); @@ -678,11 +596,7 @@ mod tests { let (blackhole_addr, _bh_handle) = spawn_stub(StubBehavior::BlackHole).await; let (reply_addr, _reply_handle) = spawn_stub(StubBehavior::Echo).await; - let forwarder = UdpForwarder::new().await.unwrap(); - forwarder - .set_upstreams(vec![blackhole_addr, reply_addr]) - .await; - forwarder.set_timeout(Duration::from_millis(50)).await; + let (forwarder, _upstreams) = new_forwarder(vec![blackhole_addr, reply_addr]).await; let request = make_dns_packet(TEST_PACKET_ID, TEST_DNS_PAYLOAD); let result = forwarder.query(&request).await.unwrap(); @@ -696,8 +610,7 @@ mod tests { let large_payload = vec![0xAB; FORWARDER_BUFFER_SIZE - 3]; let (addr, _handle) = spawn_stub(StubBehavior::Echo).await; - let forwarder = UdpForwarder::new().await.unwrap(); - forwarder.set_upstreams(vec![addr]).await; + let (forwarder, _upstreams) = new_forwarder(vec![addr]).await; let request = make_dns_packet(TEST_PACKET_ID, &large_payload); let result = forwarder.query(&request).await.unwrap(); @@ -711,14 +624,13 @@ mod tests { let (first_addr, _h1) = spawn_stub(StubBehavior::Echo).await; let (second_addr, _h2) = spawn_stub(StubBehavior::Echo).await; - let forwarder = UdpForwarder::new().await.unwrap(); - forwarder.set_upstreams(vec![first_addr]).await; + let (forwarder, upstreams) = new_forwarder(vec![first_addr]).await; let request1 = make_dns_packet(0x1111, TEST_DNS_PAYLOAD); let r1 = forwarder.query(&request1).await.unwrap(); assert_eq!(get_dns_id(&r1).unwrap(), 0x1111); - forwarder.set_upstreams(vec![second_addr]).await; + upstreams.lock().await.set(vec![second_addr]); let request2 = make_dns_packet(0x2222, TEST_DNS_PAYLOAD); let r2 = forwarder.query(&request2).await.unwrap(); @@ -730,9 +642,7 @@ mod tests { let query_count = 20; let (stub_addr, _stub_handle) = spawn_multi_stub(query_count).await; - let forwarder = UdpForwarder::new().await.unwrap(); - forwarder.set_upstreams(vec![stub_addr]).await; - forwarder.set_timeout(Duration::from_secs(5)).await; + let (forwarder, _upstreams) = new_forwarder(vec![stub_addr]).await; let mut handles = Vec::new(); for i in 0..query_count { @@ -771,9 +681,7 @@ mod tests { } }); - let forwarder = UdpForwarder::new().await.unwrap(); - forwarder.set_upstreams(vec![stub_addr]).await; - forwarder.set_timeout(Duration::from_secs(5)).await; + let (forwarder, _upstreams) = new_forwarder(vec![stub_addr]).await; let mut handles = Vec::new(); for i in 0..3u16 { @@ -800,8 +708,7 @@ mod tests { async fn id_rewrite_preserves_rest_of_packet() { let (addr, _handle) = spawn_stub(StubBehavior::Echo).await; - let forwarder = UdpForwarder::new().await.unwrap(); - forwarder.set_upstreams(vec![addr]).await; + let (forwarder, _upstreams) = new_forwarder(vec![addr]).await; let payload: Vec = (0..200u8).collect(); let request = make_dns_packet(0xBEEF, &payload); @@ -813,10 +720,7 @@ mod tests { #[tokio::test] async fn packet_too_short_returns_error() { - let forwarder = UdpForwarder::new().await.unwrap(); - forwarder - .set_upstreams(vec!["127.0.0.1:53".parse().unwrap()]) - .await; + let (forwarder, _upstreams) = new_forwarder(vec!["127.0.0.1:53".parse().unwrap()]).await; let result = forwarder.query(&[0x12]).await; match result { @@ -825,38 +729,12 @@ mod tests { } } - #[tokio::test] - async fn set_upstreams_increments_generation() { - let forwarder = UdpForwarder::new().await.unwrap(); - - let state = forwarder.upstreams.lock().await; - assert_eq!(state.generation, 0); - drop(state); - - forwarder - .set_upstreams(vec!["8.8.8.8:53".parse().unwrap()]) - .await; - - let state = forwarder.upstreams.lock().await; - assert_eq!(state.generation, 1); - drop(state); - - forwarder - .set_upstreams(vec!["1.1.1.1:53".parse().unwrap()]) - .await; - - let state = forwarder.upstreams.lock().await; - assert_eq!(state.generation, 2); - } - #[tokio::test] async fn upstream_change_during_pending_query_retries_from_new_list() { let (blackhole_addr, _bh) = spawn_stub(StubBehavior::BlackHole).await; let (reply_addr, _reply) = spawn_stub(StubBehavior::Echo).await; - let forwarder = UdpForwarder::new().await.unwrap(); - forwarder.set_upstreams(vec![blackhole_addr]).await; - forwarder.set_timeout(Duration::from_millis(50)).await; + let (forwarder, upstreams) = new_forwarder(vec![blackhole_addr]).await; let f = forwarder.clone(); let query_handle = tokio::spawn(async move { @@ -865,7 +743,7 @@ mod tests { }); tokio::time::sleep(Duration::from_millis(10)).await; - forwarder.set_upstreams(vec![reply_addr]).await; + upstreams.lock().await.set(vec![reply_addr]); let result = query_handle.await.unwrap().unwrap(); assert_eq!(get_dns_id(&result).unwrap(), TEST_PACKET_ID); @@ -876,9 +754,7 @@ mod tests { async fn upstream_change_to_empty_during_pending_query_times_out() { let (blackhole_addr, _bh) = spawn_stub(StubBehavior::BlackHole).await; - let forwarder = UdpForwarder::new().await.unwrap(); - forwarder.set_upstreams(vec![blackhole_addr]).await; - forwarder.set_timeout(Duration::from_millis(50)).await; + let (forwarder, upstreams) = new_forwarder(vec![blackhole_addr]).await; let f = forwarder.clone(); let query_handle = tokio::spawn(async move { @@ -887,7 +763,7 @@ mod tests { }); tokio::time::sleep(Duration::from_millis(10)).await; - forwarder.set_upstreams(vec![]).await; + upstreams.lock().await.set(vec![]); match query_handle.await.unwrap() { Err(ForwardError::Timeout) => {} @@ -900,9 +776,7 @@ mod tests { let (blackhole_addr, _bh) = spawn_stub(StubBehavior::BlackHole).await; let (reply_addr, _reply) = spawn_stub(StubBehavior::Echo).await; - let forwarder = UdpForwarder::new().await.unwrap(); - forwarder.set_upstreams(vec![blackhole_addr]).await; - forwarder.set_timeout(Duration::from_millis(100)).await; + let (forwarder, upstreams) = new_forwarder(vec![blackhole_addr]).await; let f = forwarder.clone(); let query_handle = tokio::spawn(async move { @@ -911,7 +785,7 @@ mod tests { }); tokio::time::sleep(Duration::from_millis(50)).await; - forwarder.set_upstreams(vec![reply_addr]).await; + upstreams.lock().await.set(vec![reply_addr]); let result = query_handle.await.unwrap().unwrap(); assert_eq!(get_dns_id(&result).unwrap(), TEST_PACKET_ID); diff --git a/crates/telio-dns/src/upstream.rs b/crates/telio-dns/src/upstream.rs new file mode 100644 index 000000000..97ca14e3f --- /dev/null +++ b/crates/telio-dns/src/upstream.rs @@ -0,0 +1,120 @@ +//! Upstream resolver list shared by the DNS forwarders + +use std::net::SocketAddr; + +/// Upstream resolver list with a generation counter to detect changes +#[derive(Clone, Debug, Default)] +pub(crate) struct UpstreamState { + addrs: Vec, + generation: u64, +} + +impl UpstreamState { + /// Upstream resolvers in the order they are tried + pub(crate) fn addrs(&self) -> &[SocketAddr] { + &self.addrs + } + + /// Generation of the current list + pub(crate) fn generation(&self) -> u64 { + self.generation + } + + /// Replace the upstream list and bump the generation + pub(crate) fn set(&mut self, addrs: Vec) { + self.generation = self.generation.wrapping_add(1); + self.addrs = addrs; + } +} + +/// Cursor over the upstream list for one query attempt +#[derive(Clone, Debug)] +pub(crate) struct AttemptState { + index: usize, + generation: u64, +} + +impl AttemptState { + pub(crate) fn new(generation: u64) -> Self { + Self { + index: 0, + generation, + } + } + + /// Generation of the upstream list this cursor is walking + pub(crate) fn generation(&self) -> u64 { + self.generation + } + + /// Pick the next upstream to try, restarting if the upstream list changed + pub(crate) fn advance(&mut self, current_upstream: &UpstreamState) -> Option { + if self.generation != current_upstream.generation { + self.index = 0; + self.generation = current_upstream.generation; + } + let addr = current_upstream.addrs.get(self.index).copied(); + self.index = self.index.saturating_add(1); + addr + } +} + +#[cfg(test)] +mod tests { + use super::*; + + fn upstream_state(addrs: Vec, generation: u64) -> UpstreamState { + UpstreamState { addrs, generation } + } + + #[test] + fn set_replaces_addrs_and_increments_generation() { + let a1: SocketAddr = "10.0.0.1:53".parse().unwrap(); + let a2: SocketAddr = "10.0.0.2:53".parse().unwrap(); + let mut state = UpstreamState::default(); + assert_eq!(state.generation(), 0); + assert!(state.addrs().is_empty()); + + state.set(vec![a1]); + assert_eq!(state.addrs(), [a1]); + assert_eq!(state.generation(), 1); + + state.set(vec![a2, a1]); + assert_eq!(state.addrs(), [a2, a1]); + assert_eq!(state.generation(), 2); + } + + #[test] + fn advance_walks_list_in_order() { + let a1: SocketAddr = "10.0.0.1:53".parse().unwrap(); + let a2: SocketAddr = "10.0.0.2:53".parse().unwrap(); + let state = upstream_state(vec![a1, a2], 0); + let mut attempt = AttemptState::new(0); + + assert_eq!(attempt.advance(&state), Some(a1)); + assert_eq!(attempt.advance(&state), Some(a2)); + assert_eq!(attempt.advance(&state), None, "exhausted"); + } + + #[test] + fn advance_restarts_on_generation_change() { + let a1: SocketAddr = "10.0.0.1:53".parse().unwrap(); + let b1: SocketAddr = "10.0.0.9:53".parse().unwrap(); + let mut attempt = AttemptState::new(0); + + let state = upstream_state(vec![a1], 0); + assert_eq!(attempt.advance(&state), Some(a1)); + assert_eq!(attempt.advance(&state), None); + + let state = upstream_state(vec![b1], 1); + assert_eq!(attempt.advance(&state), Some(b1)); + assert_eq!(attempt.advance(&state), None); + } + + #[test] + fn advance_empty_list_returns_none() { + let state = upstream_state(vec![], 0); + let mut attempt = AttemptState::new(0); + assert_eq!(attempt.advance(&state), None); + } +} diff --git a/crates/telio-dns/src/zone.rs b/crates/telio-dns/src/zone.rs index 821a5f5b0..69159e3c1 100644 --- a/crates/telio-dns/src/zone.rs +++ b/crates/telio-dns/src/zone.rs @@ -1,5 +1,6 @@ use crate::{ forward::ForwardAuthority, packet_decoder::normalize_qname, packet_encoder::ResponseKind, + DNS_PORT, }; use async_trait::async_trait; use hickory_server::{ @@ -343,7 +344,7 @@ impl ForwardZone { ZoneType::Forward, ForwardConfig { options: Some(options), - name_servers: NameServerConfigGroup::from_ips_clear(ips, 53, true), + name_servers: NameServerConfigGroup::from_ips_clear(ips, DNS_PORT, true), }, ) .await?; diff --git a/crates/telio-dns/tests/nameserver.rs b/crates/telio-dns/tests/nameserver.rs index 22a7c6b58..6bed443f2 100644 --- a/crates/telio-dns/tests/nameserver.rs +++ b/crates/telio-dns/tests/nameserver.rs @@ -18,7 +18,7 @@ use std::{ str::FromStr, sync::Arc, }; -use telio_dns::{LocalNameServer, NameServer, Records}; +use telio_dns::{LocalNameServer, NameServer, Records, DNS_PORT}; use telio_model::features::TtlValue; use tokio::task::JoinHandle; use tokio::time::sleep; @@ -242,7 +242,7 @@ impl WGClient { .expect("Failed to parse tcp response"); assert_eq!(tcp_response.get_flags(), TcpFlags::RST); // Server should reply from 53 - assert_eq!(tcp_response.get_source(), 53); + assert_eq!(tcp_response.get_source(), DNS_PORT); None } else { let udp_response = UdpPacket::new(ip_response.payload()) @@ -260,7 +260,7 @@ impl WGClient { let tcp_response = TcpPacket::new(ip_response.payload()) .expect("Failed to parse tcp response"); assert_eq!(tcp_response.get_flags(), TcpFlags::RST); - assert_eq!(tcp_response.get_source(), 53); + assert_eq!(tcp_response.get_source(), DNS_PORT); None } else { let udp_response = UdpPacket::new(ip_response.payload()) @@ -339,7 +339,7 @@ impl WGClient { let mut tcp_packet = MutableTcpPacket::new(&mut buffer[IPV4_HEADER..length]) .expect("Failed to create MutableTcpPacket"); tcp_packet.set_source(100); - tcp_packet.set_destination(53); + tcp_packet.set_destination(DNS_PORT); tcp_packet.set_sequence(42); tcp_packet.set_payload(dns_query); tcp_packet.set_checksum(0); @@ -361,9 +361,10 @@ impl WGClient { let mut udp_packet = MutableUdpPacket::new(&mut buffer[IPV4_HEADER..length]) .expect("Failed to create MutableUdpPacket"); udp_packet.set_source(100); - udp_packet.set_destination(53); if matches!(test_type, DnsTestType::BadUdpPortIpv4) { udp_packet.set_destination(54); + } else { + udp_packet.set_destination(DNS_PORT); } udp_packet.set_length((UDP_HEADER + dns_query.len()) as u16); udp_packet.set_payload(dns_query); @@ -420,7 +421,7 @@ impl WGClient { let mut tcp_packet = MutableTcpPacket::new(&mut buffer[IPV6_HEADER..total_length]) .expect("Failed to create MutableTcpPacket"); tcp_packet.set_source(100); - tcp_packet.set_destination(53); + tcp_packet.set_destination(DNS_PORT); tcp_packet.set_payload(dns_query); tcp_packet.set_checksum(0); tcp_packet.set_checksum(pnet_packet::tcp::ipv6_checksum( @@ -450,7 +451,7 @@ impl WGClient { if matches!(test_type, DnsTestType::BadUdpPortIpv6) { udp_response.set_destination(54); } else { - udp_response.set_destination(53); + udp_response.set_destination(DNS_PORT); } udp_response.set_length(length as u16); udp_response.set_payload(dns_query); diff --git a/crates/telio-model/src/constants.rs b/crates/telio-model/src/constants.rs index 7c14aa97e..42e81c4a0 100644 --- a/crates/telio-model/src/constants.rs +++ b/crates/telio-model/src/constants.rs @@ -24,3 +24,11 @@ pub const IPV6_STARCAST_ADDRESS: Ipv6Addr = Ipv6Addr::new(0xfd74, 0x656c, 0x696f pub const IPV4_STARCAST_NETWORK: ConstIpv4Net = ConstIpv4Net::new(IPV4_STARCAST_ADDRESS, 32); /// Ipv6 starcast's virtual peer network pub const IPV6_STARCAST_NETWORK: ConstIpv6Net = ConstIpv6Net::new(IPV6_STARCAST_ADDRESS, 128); +/// IPv4 DNS virtual peer address +pub const DNS_VIRTUAL_IPV4: Ipv4Addr = Ipv4Addr::new(100, 64, 0, 2); +/// IPv6 DNS virtual peer address +pub const DNS_VIRTUAL_IPV6: Ipv6Addr = Ipv6Addr::new(0xfd74, 0x656c, 0x696f, 0, 0, 0, 0, 0x2); +/// IPv4 DNS peer address on exit node +pub const DNS_EXIT_IPV4: Ipv4Addr = Ipv4Addr::new(100, 64, 0, 3); +/// IPv6 DNS peer address on exit node +pub const DNS_EXIT_IPV6: Ipv6Addr = Ipv6Addr::new(0xfd74, 0x656c, 0x696f, 0, 0, 0, 0, 0x3);