From a1fe803c4a294128cd887543bc18317ef6843a95 Mon Sep 17 00:00:00 2001 From: Kevin Schaller Date: Wed, 29 Jul 2026 03:04:20 +0200 Subject: [PATCH 1/2] feat: add DNS-over-TCP support alongside UDP listener - Implemented TcpListener bound to the configured listen address & port - Added RFC 1035 TCP DNS frame handler (2-byte length prefix) - Refactored query resolution pipeline to process both UDP and TCP queries through unified cache and DoH upstream logic --- src/handler.rs | 244 ++++++++++++++++++++----------------------------- src/run.rs | 32 ++++++- 2 files changed, 131 insertions(+), 145 deletions(-) diff --git a/src/handler.rs b/src/handler.rs index cb71281..f03946b 100644 --- a/src/handler.rs +++ b/src/handler.rs @@ -4,182 +4,140 @@ use bytes::Bytes; use dns_message_parser::question::Question; use dns_message_parser::Dns; use futures::lock::Mutex; -use std::future::Future; +use log::{debug, error, info}; use std::net::SocketAddr; use std::sync::Arc; -use std::time::Duration; -use tokio::net::UdpSocket; +use tokio::io::{AsyncReadExt, AsyncWriteExt}; use tokio::time::timeout as create_timeout; -async fn send_response( - dns_response: &mut Dns, - id: u16, - addr: SocketAddr, - sender: Arc, -) -> DohResult<()> { - dns_response.id = id; - let bytes = dns_response.encode()?; - sender.send_to(bytes.as_ref(), addr).await?; - Ok(()) -} +pub async fn resolve_request(msg: Bytes, context: &Context) -> DohResult { + let mut dns_request = Dns::decode(msg)?; + if dns_request.is_response() { + return Err(DohError::DnsNotRequest(dns_request)); + } -enum CacheReturn<'a> { - Found(DohResult<()>), - NotFound(Option<(&'a Mutex>, Question)>), -} + let id = dns_request.id; -#[allow(clippy::needless_lifetimes)] -async fn get_response_from_cache<'a>( - context: &'a Context, - dns_request: &Dns, - addr: &SocketAddr, -) -> CacheReturn<'a> { + // 1. Check cache if let Some(cache) = &context.cache { - let questions = &dns_request.questions; - if questions.len() == 1 { - let question = &questions[0]; + if dns_request.questions.len() == 1 { + let question = &dns_request.questions[0]; let mut guard_cache = cache.lock().await; let entry = if context.cache_fallback { guard_cache.get_expired(question) } else { guard_cache.get(question) }; - if let Some(dns_response) = entry { - let id = dns_request.id; - let sender = context.sender.clone(); - let addr = *addr; + let mut resp = dns_response.clone(); + resp.id = id; debug!("Question is found in cache"); - let result = send_response(dns_response, id, addr, sender).await; - CacheReturn::Found(result) - } else { - debug!("Question is not found in cache"); - CacheReturn::NotFound(Some((cache, question.clone()))) + return Ok(resp.encode()?.into()); } - } else { - debug!("The amount of questions is not equal 1"); - CacheReturn::NotFound(None) } - } else { - debug!("Cache is disable"); - CacheReturn::NotFound(None) } -} -async fn get_response( - context: &Context, - cache_question: &Option<(&Mutex>, Question)>, - response: ( - impl Future)>>, - u32, - ), - id: u16, - addr: &SocketAddr, -) -> Option> { - let (response_future, connection_id) = response; - let timeout = context.timeout; - match create_timeout(timeout, response_future).await { - Ok(Ok((mut dns_response, duration))) => { - let addr = *addr; - let sender = context.sender.clone(); - let result = send_response(&mut dns_response, id, addr, sender).await; - if let Some(duration) = duration { - if let Some((cache, question)) = cache_question { - let mut guard_cache = cache.lock().await; - debug!( - "Add records in cache: {}, {}, {:?}", - question, dns_response, duration - ); - guard_cache.put(question.clone(), dns_response, duration); + // 2. Fetch from remote DoH + let mut guard_remote = context.remote_session.lock().await; + let remote_res = guard_remote.start_request(&mut dns_request).await; + drop(guard_remote); + + if let Ok((response_future, connection_id)) = remote_res { + match create_timeout(context.timeout, response_future).await { + Ok(Ok((mut dns_response, duration))) => { + if let Some(duration) = duration { + if let Some(cache) = &context.cache { + if dns_request.questions.len() == 1 { + let question = &dns_request.questions[0]; + let mut guard_cache = cache.lock().await; + debug!( + "Add records in cache: {}, {}, {:?}", + question, dns_response, duration + ); + guard_cache.put(question.clone(), dns_response.clone(), duration); + } + } } + dns_response.id = id; + return Ok(dns_response.encode()?.into()); + } + Ok(Err(e)) => { + error!("Could not retrieve DNS response from server: {}", e); + } + Err(e) => { + error!("Timeout: {}", e); } - return Some(result); - } - Ok(Err(e)) => { - error!("Could not retrieve DNS response from server: {}", e); - } - Err(e) => { - error!("Timeout: {}", e); - } - } - let mut guard_remote_session = context.remote_session.lock().await; - guard_remote_session.disconnect(connection_id); - None -} - -async fn get_response_from_remote( - context: &Context, - cache_question: &Option<(&Mutex>, Question)>, - dns_request: &mut Dns, - addr: &SocketAddr, -) -> Option> { - let mut guard_remote_session = context.remote_session.lock().await; - let result = guard_remote_session.start_request(dns_request).await; - drop(guard_remote_session); - match result { - Ok(response) => { - let id = dns_request.id; - get_response(context, cache_question, response, id, addr).await - } - Err(e) => { - info!("Could not contact DNS server: {}", e); - None } + let mut guard_remote = context.remote_session.lock().await; + guard_remote.disconnect(connection_id); } -} -#[allow(clippy::needless_lifetimes)] -async fn get_response_from_cache_fallback<'a>( - context: &'a Context, - cache_question: Option<(&Mutex>, Question)>, - dns_request: &Dns, - addr: SocketAddr, -) -> Option> { + // 3. Fallback cache if context.cache_fallback { - if let Some((cache, question)) = &cache_question { - let mut guard_cache = cache.lock().await; - if let Some(dns_response) = guard_cache.get_expired_fallback(question) { - let id = dns_request.id; - let sender = context.sender.clone(); - debug!("Question is found in cache fallback"); - let result = send_response(dns_response, id, addr, sender).await; - Some(result) - } else { - debug!("Question is not found in cache fallback"); - None + if let Some(cache) = &context.cache { + if dns_request.questions.len() == 1 { + let question = &dns_request.questions[0]; + let mut guard_cache = cache.lock().await; + if let Some(dns_response) = guard_cache.get_expired_fallback(question) { + let mut resp = dns_response.clone(); + resp.id = id; + debug!("Question is found in cache fallback"); + return Ok(resp.encode()?.into()); + } } - } else { - debug!("Question cannot be cached"); - None } - } else { - debug!("Cache fallback is disable"); - None } + + Err(DohError::CouldNotGetResponse(dns_request)) } pub async fn request_handler(msg: Bytes, addr: SocketAddr, context: &Context) -> DohResult<()> { - let mut dns_request = Dns::decode(msg)?; - if dns_request.is_response() { - return Err(DohError::DnsNotRequest(dns_request)); - } + let resp_bytes = resolve_request(msg, context).await?; + context.sender.send_to(resp_bytes.as_ref(), addr).await?; + Ok(()) +} - let cache = get_response_from_cache(context, &dns_request, &addr).await; - let cache_question = match cache { - CacheReturn::Found(result) => return result, - CacheReturn::NotFound(cache_question) => cache_question, - }; +pub async fn tcp_connection_handler( + mut stream: tokio::net::TcpStream, + context: Arc, +) -> DohResult<()> { + let peer_addr = stream.peer_addr()?; + debug!("New TCP DNS connection from {}", peer_addr); - let remote = get_response_from_remote(context, &cache_question, &mut dns_request, &addr).await; - if let Some(result) = remote { - return result; - } + loop { + let mut len_buf = [0u8; 2]; + if stream.read_exact(&mut len_buf).await.is_err() { + break; + } + let packet_len = u16::from_be_bytes(len_buf) as usize; + if packet_len == 0 { + break; + } - let fallback = - get_response_from_cache_fallback(context, cache_question, &dns_request, addr).await; - if let Some(result) = fallback { - return result; - } + let mut msg_buf = vec![0u8; packet_len]; + if stream.read_exact(&mut msg_buf).await.is_err() { + break; + } - Err(DohError::CouldNotGetResponse(dns_request)) + let msg = Bytes::from(msg_buf); + match resolve_request(msg, context.as_ref()).await { + Ok(resp_bytes) => { + let resp_len = (resp_bytes.len() as u16).to_be_bytes(); + if stream.write_all(&resp_len).await.is_err() { + break; + } + if stream.write_all(resp_bytes.as_ref()).await.is_err() { + break; + } + if stream.flush().await.is_err() { + break; + } + } + Err(e) => { + debug!("TCP query handling failed for {}: {}", peer_addr, e); + break; + } + } + } + Ok(()) } diff --git a/src/run.rs b/src/run.rs index 1d92251..1b8ceea 100644 --- a/src/run.rs +++ b/src/run.rs @@ -1,17 +1,45 @@ use crate::config::Config; use crate::error::Result as DohResult; -use crate::handler::request_handler; +use crate::handler::{request_handler, tcp_connection_handler}; use bytes::Bytes; use dns_message_parser::MAXIMUM_DNS_PACKET_SIZE; +use log::{debug, error, info}; use std::sync::Arc; use tokio::spawn; /// Run the `doh-client` with a specific configuration. pub async fn run(config: Config) -> DohResult<()> { let (recv, context) = config.into().await?; - + let local_addr = recv.local_addr()?; let context = Arc::new(context); + match tokio::net::TcpListener::bind(local_addr).await { + Ok(tcp_listener) => { + info!("Listening for DNS requests on UDP and TCP address {}", local_addr); + let tcp_context = context.clone(); + spawn(async move { + loop { + match tcp_listener.accept().await { + Ok((stream, _)) => { + let c = tcp_context.clone(); + spawn(async move { + if let Err(e) = tcp_connection_handler(stream, c).await { + debug!("TCP connection closed: {}", e); + } + }); + } + Err(e) => { + error!("TCP accept error: {}", e); + } + } + } + }); + } + Err(e) => { + error!("Failed to bind TCP listener on {}: {}", local_addr, e); + } + } + let mut buffer: [u8; MAXIMUM_DNS_PACKET_SIZE] = [0; MAXIMUM_DNS_PACKET_SIZE]; loop { let (n, addr) = recv.recv_from(&mut buffer[..]).await?; From 8b75cf1287b9190e8ec2542a4cf472768407f439 Mon Sep 17 00:00:00 2001 From: Kevin Schaller Date: Wed, 29 Jul 2026 03:29:30 +0200 Subject: [PATCH 2/2] fix: return SERVFAIL DNS response when upstream DoH query fails - Prevents premature TCP socket teardown when upstream lookup times out or fails - Ensures compliance with DNS over TCP standards by sending proper RCode::ServFail response --- src/handler.rs | 14 ++++++++------ 1 file changed, 8 insertions(+), 6 deletions(-) diff --git a/src/handler.rs b/src/handler.rs index f03946b..4caefc3 100644 --- a/src/handler.rs +++ b/src/handler.rs @@ -1,10 +1,8 @@ use crate::context::Context; -use crate::{Cache, DohError, DohResult}; +use crate::{DohError, DohResult}; use bytes::Bytes; -use dns_message_parser::question::Question; -use dns_message_parser::Dns; -use futures::lock::Mutex; -use log::{debug, error, info}; +use dns_message_parser::{Dns, RCode}; +use log::{debug, error}; use std::net::SocketAddr; use std::sync::Arc; use tokio::io::{AsyncReadExt, AsyncWriteExt}; @@ -88,7 +86,11 @@ pub async fn resolve_request(msg: Bytes, context: &Context) -> DohResult } } - Err(DohError::CouldNotGetResponse(dns_request)) + // Upstream failed: return SERVFAIL DNS response to client instead of dropping TCP connection + let mut servfail_resp = dns_request; + servfail_resp.flags.qr = true; + servfail_resp.flags.rcode = RCode::ServFail; + Ok(servfail_resp.encode()?.into()) } pub async fn request_handler(msg: Bytes, addr: SocketAddr, context: &Context) -> DohResult<()> {