Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
254 changes: 107 additions & 147 deletions src/handler.rs
Original file line number Diff line number Diff line change
@@ -1,185 +1,145 @@
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 std::future::Future;
use dns_message_parser::{Dns, RCode};
use log::{debug, error};
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<UdpSocket>,
) -> 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<Bytes> {
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<Cache<Question, Dns>>, 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<Cache<Question, Dns>>, Question)>,
response: (
impl Future<Output = DohResult<(Dns, Option<Duration>)>>,
u32,
),
id: u16,
addr: &SocketAddr,
) -> Option<DohResult<()>> {
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<Cache<Question, Dns>>, Question)>,
dns_request: &mut Dns,
addr: &SocketAddr,
) -> Option<DohResult<()>> {
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<Cache<Question, Dns>>, Question)>,
dns_request: &Dns,
addr: SocketAddr,
) -> Option<DohResult<()>> {
// 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
}

// 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<()> {
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<Context>,
) -> 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(())
}
32 changes: 30 additions & 2 deletions src/run.rs
Original file line number Diff line number Diff line change
@@ -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?;
Expand Down