diff --git a/config/default.toml b/config/default.toml index 352a8ec..b0b9d31 100644 --- a/config/default.toml +++ b/config/default.toml @@ -1,6 +1,9 @@ [server] listen = "0.0.0.0:8800" +[dns] +listen = "0.0.0.0:5353" + [upstream] resolvers = [ { name = "cloudflare", url = "https://cloudflare-dns.com/dns-query" }, diff --git a/src/config.rs b/src/config.rs index e955574..c26c298 100644 --- a/src/config.rs +++ b/src/config.rs @@ -8,10 +8,13 @@ use crate::upstream::doh_client::Strategy; pub struct AppConfig { /// Server configuration. pub server: ServerConfig, + /// DNS listener configuration. + #[serde(default)] + pub dns: DnsListenerConfig, /// Upstream resolver configuration. pub upstream: UpstreamConfig, /// Cache configuration. - #[allow(dead_code)] // Used once M3 caching is implemented + #[allow(dead_code)] // Used once M4 caching is implemented pub cache: CacheConfig, } @@ -22,6 +25,21 @@ pub struct ServerConfig { pub listen: String, } +/// DNS listener configuration for UDP/TCP. +#[derive(Debug, Clone, Deserialize)] +pub struct DnsListenerConfig { + /// Address to listen on for DNS queries (e.g., "0.0.0.0:5353"). + pub listen: String, +} + +impl Default for DnsListenerConfig { + fn default() -> Self { + Self { + listen: "0.0.0.0:5353".to_string(), + } + } +} + /// Upstream DoH resolver settings. #[derive(Debug, Clone, Deserialize)] pub struct UpstreamConfig { diff --git a/src/dns/listener.rs b/src/dns/listener.rs new file mode 100644 index 0000000..51b17ba --- /dev/null +++ b/src/dns/listener.rs @@ -0,0 +1,206 @@ +use std::sync::atomic::{AtomicUsize, Ordering}; +use std::sync::Arc; + +use tokio::net::{TcpListener, UdpSocket}; +use tracing::{debug, error, info, warn}; + +use crate::dns::query::parse_query; +use crate::dns::response::build_servfail; +use crate::upstream::doh_client::{DohClient, Strategy}; + +/// Errors that can occur in the DNS listener. +#[derive(Debug, thiserror::Error)] +pub enum ListenerError { + /// Failed to bind the UDP socket. + #[error("failed to bind UDP socket: {0}")] + UdpBind(#[source] std::io::Error), + /// Failed to bind the TCP listener. + #[error("failed to bind TCP listener: {0}")] + TcpBind(#[source] std::io::Error), +} + +/// A DNS listener that accepts UDP and TCP queries and forwards them +/// through the DoH pipeline. +#[derive(Clone)] +pub struct DnsListener { + doh_client: DohClient, + strategy: Strategy, + rr_counter: Arc, +} + +impl DnsListener { + /// Create a new DNS listener with the given DoH client and strategy. + pub fn new(doh_client: DohClient, strategy: Strategy) -> Self { + Self { + doh_client, + strategy, + rr_counter: Arc::new(AtomicUsize::new(0)), + } + } + + /// Start the UDP and TCP listeners on the given address. + pub async fn run(&self, addr: &str) -> Result<(), ListenerError> { + let udp_listener = self.clone(); + let tcp_listener = self.clone(); + let udp_addr = addr.to_string(); + let tcp_addr = addr.to_string(); + + let udp_handle = tokio::spawn(async move { + if let Err(e) = udp_listener.run_udp(&udp_addr).await { + error!(error = %e, "UDP listener failed"); + } + }); + + let tcp_handle = tokio::spawn(async move { + if let Err(e) = tcp_listener.run_tcp(&tcp_addr).await { + error!(error = %e, "TCP listener failed"); + } + }); + + // Wait for either to complete (they shouldn't unless they error) + tokio::select! { + _ = udp_handle => {}, + _ = tcp_handle => {}, + } + + Ok(()) + } + + /// Run the UDP listener loop. + async fn run_udp(&self, addr: &str) -> Result<(), ListenerError> { + let socket = UdpSocket::bind(addr) + .await + .map_err(ListenerError::UdpBind)?; + info!(addr = %addr, "UDP DNS listener started"); + + let mut buf = vec![0u8; 512]; // Standard DNS UDP packet size + + loop { + let (len, src) = match socket.recv_from(&mut buf).await { + Ok(v) => v, + Err(e) => { + warn!(error = %e, "UDP recv error"); + continue; + } + }; + + let query_bytes = &buf[..len]; + debug!(from = %src, bytes = len, "received UDP DNS query"); + + let response = self.process_query(query_bytes).await; + + if let Err(e) = socket.send_to(&response, src).await { + warn!(error = %e, "UDP send error"); + } + } + } + + /// Run the TCP listener loop. + async fn run_tcp(&self, addr: &str) -> Result<(), ListenerError> { + let listener = TcpListener::bind(addr) + .await + .map_err(ListenerError::TcpBind)?; + info!(addr = %addr, "TCP DNS listener started"); + + loop { + let (stream, src) = match listener.accept().await { + Ok(v) => v, + Err(e) => { + warn!(error = %e, "TCP accept error"); + continue; + } + }; + + let this = self.clone(); + tokio::spawn(async move { + if let Err(e) = this.handle_tcp(stream, src).await { + debug!(error = %e, from = %src, "TCP connection error"); + } + }); + } + } + + /// Handle a single TCP DNS connection. + /// + /// DNS over TCP uses a 2-byte length prefix (RFC 1035 ยง4.2.2). + async fn handle_tcp( + &self, + mut stream: tokio::net::TcpStream, + src: std::net::SocketAddr, + ) -> anyhow::Result<()> { + use tokio::io::{AsyncReadExt, AsyncWriteExt}; + + // Read 2-byte length prefix + let mut len_buf = [0u8; 2]; + stream.read_exact(&mut len_buf).await?; + let query_len = u16::from_be_bytes(len_buf) as usize; + + if query_len > 65535 { + return Err(anyhow::anyhow!("TCP query too large: {query_len} bytes")); + } + + // Read the query + let mut query_buf = vec![0u8; query_len]; + stream.read_exact(&mut query_buf).await?; + + debug!(from = %src, bytes = query_len, "received TCP DNS query"); + + let response = self.process_query(&query_buf).await; + + // Write 2-byte length prefix + response + let resp_len = (response.len() as u16).to_be_bytes(); + stream.write_all(&resp_len).await?; + stream.write_all(&response).await?; + + Ok(()) + } + + /// Process a DNS query: parse, forward through DoH, return wire-format response. + async fn process_query(&self, query_bytes: &[u8]) -> Vec { + let query = match parse_query(query_bytes) { + Ok(q) => q, + Err(e) => { + warn!(error = %e, "invalid DNS query"); + // Return a minimal SERVFAIL for unparseable queries + return build_minimal_servfail(); + } + }; + + let qname = query + .queries() + .first() + .map(|q| q.name().to_string()) + .unwrap_or_default(); + let qtype = query + .queries() + .first() + .map(|q| q.query_type().to_string()) + .unwrap_or_default(); + + let idx = self.rr_counter.fetch_add(1, Ordering::Relaxed); + + match self + .doh_client + .forward(query_bytes, self.strategy, idx) + .await + { + Ok(response) => { + info!(query = %qname, qtype = %qtype, bytes = response.len(), "DNS query forwarded"); + response + } + Err(e) => { + error!(error = %e, query = %qname, "upstream forward failed"); + build_servfail(&query).unwrap_or_else(|_| build_minimal_servfail()) + } + } + } +} + +/// Build a minimal SERVFAIL response (12-byte header, no questions). +fn build_minimal_servfail() -> Vec { + let mut buf = vec![0u8; 12]; + buf[2] = 0x81; // QR=1, Opcode=0, AA=0, TC=0, RD=1 + buf[3] = 0x80; // RA=1, Z=0, RCODE=0 (no error in this byte, RCODE in lower 4 bits) + buf[3] |= 0x02; // RCODE = SERVFAIL (2) + buf +} diff --git a/src/dns/mod.rs b/src/dns/mod.rs index 60c0187..1c9d9e8 100644 --- a/src/dns/mod.rs +++ b/src/dns/mod.rs @@ -1,4 +1,5 @@ -//! DNS wire-format parsing and response building. +//! DNS wire-format parsing, response building, and UDP/TCP listener. +pub mod listener; pub mod query; pub mod response; diff --git a/src/main.rs b/src/main.rs index e7d65a6..a81576a 100644 --- a/src/main.rs +++ b/src/main.rs @@ -9,6 +9,7 @@ use tracing::info; use tracing_subscriber::EnvFilter; use doh_forwarder::config::AppConfig; +use doh_forwarder::dns::listener::DnsListener; use doh_forwarder::routes::{dns_query_get, dns_query_post, AppState}; use doh_forwarder::upstream::doh_client::DohClient; @@ -44,6 +45,16 @@ async fn main() -> anyhow::Result<()> { let strategy = config.upstream.strategy()?; let doh_client = DohClient::new(config.upstream.resolvers.clone()); + // Start DNS listener (UDP/TCP) in background + let dns_listener = DnsListener::new(doh_client.clone(), strategy); + let dns_addr = config.dns.listen.clone(); + tokio::spawn(async move { + if let Err(e) = dns_listener.run(&dns_addr).await { + tracing::error!(error = %e, "DNS listener failed"); + } + }); + info!(addr = %config.dns.listen, "DNS listener starting"); + let state = AppState { doh_client, rr_counter: Arc::new(std::sync::atomic::AtomicUsize::new(0)),