feat: implement M3 — UDP/TCP DNS listener

- Add DnsListener that binds UDP and TCP on configurable port
- Add [dns] config section with listen address (default 0.0.0.0:5353)
- DNS over TCP uses 2-byte length prefix (RFC 1035 §4.2.2)
- Forward queries through existing DoH pipeline
- Start DNS listener as background task alongside DoH server
This commit is contained in:
2026-07-06 17:57:04 +03:30
parent 72198bb868
commit a9584ab940
5 changed files with 241 additions and 2 deletions

View File

@@ -1,6 +1,9 @@
[server] [server]
listen = "0.0.0.0:8800" listen = "0.0.0.0:8800"
[dns]
listen = "0.0.0.0:5353"
[upstream] [upstream]
resolvers = [ resolvers = [
{ name = "cloudflare", url = "https://cloudflare-dns.com/dns-query" }, { name = "cloudflare", url = "https://cloudflare-dns.com/dns-query" },

View File

@@ -8,10 +8,13 @@ use crate::upstream::doh_client::Strategy;
pub struct AppConfig { pub struct AppConfig {
/// Server configuration. /// Server configuration.
pub server: ServerConfig, pub server: ServerConfig,
/// DNS listener configuration.
#[serde(default)]
pub dns: DnsListenerConfig,
/// Upstream resolver configuration. /// Upstream resolver configuration.
pub upstream: UpstreamConfig, pub upstream: UpstreamConfig,
/// Cache configuration. /// Cache configuration.
#[allow(dead_code)] // Used once M3 caching is implemented #[allow(dead_code)] // Used once M4 caching is implemented
pub cache: CacheConfig, pub cache: CacheConfig,
} }
@@ -22,6 +25,21 @@ pub struct ServerConfig {
pub listen: String, 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. /// Upstream DoH resolver settings.
#[derive(Debug, Clone, Deserialize)] #[derive(Debug, Clone, Deserialize)]
pub struct UpstreamConfig { pub struct UpstreamConfig {

206
src/dns/listener.rs Normal file
View File

@@ -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<AtomicUsize>,
}
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<u8> {
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<u8> {
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
}

View File

@@ -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 query;
pub mod response; pub mod response;

View File

@@ -9,6 +9,7 @@ use tracing::info;
use tracing_subscriber::EnvFilter; use tracing_subscriber::EnvFilter;
use doh_forwarder::config::AppConfig; 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::routes::{dns_query_get, dns_query_post, AppState};
use doh_forwarder::upstream::doh_client::DohClient; use doh_forwarder::upstream::doh_client::DohClient;
@@ -44,6 +45,16 @@ async fn main() -> anyhow::Result<()> {
let strategy = config.upstream.strategy()?; let strategy = config.upstream.strategy()?;
let doh_client = DohClient::new(config.upstream.resolvers.clone()); 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 { let state = AppState {
doh_client, doh_client,
rr_counter: Arc::new(std::sync::atomic::AtomicUsize::new(0)), rr_counter: Arc::new(std::sync::atomic::AtomicUsize::new(0)),