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:
@@ -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 {
|
||||
|
||||
206
src/dns/listener.rs
Normal file
206
src/dns/listener.rs
Normal 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
|
||||
}
|
||||
@@ -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;
|
||||
|
||||
11
src/main.rs
11
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)),
|
||||
|
||||
Reference in New Issue
Block a user