Skip to content
Closed
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
40 changes: 7 additions & 33 deletions src/doh_server.rs
Original file line number Diff line number Diff line change
Expand Up @@ -10,7 +10,9 @@ use std::time::Duration;

use anyhow::Context;
use bytes::Bytes;
use hickory_proto::op::{Message, MessageType, ResponseCode};
#[cfg(test)]
use hickory_proto::op::{Message, ResponseCode};
#[cfg(test)]
use hickory_proto::serialize::binary::BinDecodable;
use http_body_util::{BodyExt, Full};
use hyper::server::conn::http1;
Expand All @@ -22,7 +24,7 @@ use tokio::net::TcpListener;
use tokio_rustls::TlsAcceptor;
use tracing::{info, warn};

use crate::engine::{Engine, FastPathResponse, PreParsedData};
use crate::engine::{Engine, FastPathResponse, PreParsedData, engine_helpers};
use crate::proto_utils;

const MAX_DNS_MESSAGE: usize = 64 * 1024;
Expand Down Expand Up @@ -232,35 +234,7 @@ async fn process_dns_wire(packet: &[u8], peer: SocketAddr, engine: &Engine) -> B

/// 构造空 DNS 响应(SERVFAIL) / Build empty SERVFAIL response
fn empty_dns_response(request: &[u8]) -> Bytes {
fn header_only(request: &[u8]) -> Bytes {
if request.len() < 2 {
return Bytes::new();
}
let mut response = vec![0u8; 12];
response[0] = request[0];
response[1] = request[1];
response[2] = 0x80;
response[3] = 0x02;
Bytes::from(response)
}

let Ok(request_message) = Message::from_bytes(request) else {
return header_only(request);
};

let mut response = Message::new(
request_message.metadata.id,
MessageType::Response,
request_message.metadata.op_code,
);
response.metadata.response_code = ResponseCode::ServFail;
response.metadata.recursion_desired = request_message.metadata.recursion_desired;
response.add_queries(request_message.queries.iter().cloned());

response
.to_vec()
.map(Bytes::from)
.unwrap_or_else(|_| header_only(request))
engine_helpers::build_servfail_response_from_wire(request)
}

fn error_response(status: StatusCode) -> Response<Full<Bytes>> {
Expand Down Expand Up @@ -385,8 +359,8 @@ mod tests {
assert_eq!(resp.len(), 12);
// Byte 2: QR=1, Opcode=0, AA=0, TC=0, RD=0 → 0x80
assert_eq!(resp[2], 0x80, "QR should be 1");
// Byte 3: RA=0, Z=0, RCODE=2 (SERVFAIL) → 0x02
assert_eq!(resp[3], 0x02, "RCODE should be SERVFAIL (2)");
// Byte 3: RA=1, Z=0, RCODE=2 (SERVFAIL) → 0x82
assert_eq!(resp[3], 0x82, "RA and SERVFAIL should be set");
}

#[test]
Expand Down
64 changes: 64 additions & 0 deletions src/engine/execution.rs
Original file line number Diff line number Diff line change
Expand Up @@ -1191,6 +1191,70 @@ mod tests {
assert_eq!(msg.queries.len(), 1, "Should have one query");
}

#[test]
fn test_engine_helpers_build_servfail_response_from_wire() {
let mut req = Message::new(0xBEEF, MessageType::Query, OpCode::Query);
req.metadata.recursion_desired = true;
req.add_query(Query::query(
Name::from_str("data.xiaoheihe.cn").unwrap(),
RecordType::AAAA,
));
let request = req.to_vec().unwrap();

let response = engine_helpers::build_servfail_response_from_wire(&request);
let decoded = Message::from_bytes(&response).unwrap();

assert_eq!(decoded.metadata.id, 0xBEEF);
assert_eq!(decoded.metadata.response_code, ResponseCode::ServFail);
assert!(decoded.metadata.recursion_desired);
assert!(decoded.metadata.recursion_available);
assert_eq!(decoded.queries.len(), 1);
assert_eq!(decoded.queries[0].name().to_utf8(), "data.xiaoheihe.cn.");
assert_eq!(decoded.queries[0].query_type(), RecordType::AAAA);
}

#[test]
fn test_engine_helpers_build_servfail_response_from_malformed_wire() {
let response = engine_helpers::build_servfail_response_from_wire(&[0x12, 0x34, 0x01, 0x10]);

assert_eq!(response.len(), 12);
assert_eq!(&response[..2], &[0x12, 0x34]);
assert_ne!(response[2] & 0x80, 0, "QR must be set");
assert_ne!(response[3] & 0x80, 0, "RA must be set");
assert_eq!(response[3] & 0x0F, 2, "RCODE must be SERVFAIL");
assert_ne!(response[2] & 0x01, 0, "RD must be preserved");
assert_ne!(response[3] & 0x10, 0, "CD must be preserved");
}

#[test]
fn test_listener_servfail_rejects_non_queries_and_malformed_packets() {
assert!(engine_helpers::try_build_servfail_response_from_wire(&[0x12, 0x34]).is_none());

let mut no_question = vec![0u8; 12];
no_question[..2].copy_from_slice(&0x1234u16.to_be_bytes());
assert!(
engine_helpers::try_build_servfail_response_from_wire(&no_question).is_none(),
"QDCOUNT=0 must not receive a response"
);

let mut response_packet = vec![0u8; 12];
response_packet[2] = 0x80;
response_packet[5] = 1;
response_packet.extend_from_slice(&[0, 0, 1, 0, 1]);
assert!(
engine_helpers::try_build_servfail_response_from_wire(&response_packet).is_none(),
"QR=1 packets must not receive a response"
);

let mut invalid_name = vec![0u8; 12];
invalid_name[5] = 1;
invalid_name.extend_from_slice(&[0x3f, b'a', 0, 0, 1, 0, 1]);
assert!(
engine_helpers::try_build_servfail_response_from_wire(&invalid_name).is_none(),
"malformed names must not receive a response"
);
}

#[test]
fn test_engine_helpers_build_refused_response() {
// Arrange: Create test request
Expand Down
60 changes: 59 additions & 1 deletion src/engine/utils.rs
Original file line number Diff line number Diff line change
Expand Up @@ -4,7 +4,7 @@ use bytes::Bytes;
use dashmap::DashSet;
use hickory_proto::op::{Message, MessageType, OpCode, Query, ResponseCode};
use hickory_proto::rr::{DNSClass, Name, RecordType};
use hickory_proto::serialize::binary::{BinEncodable, BinEncoder};
use hickory_proto::serialize::binary::{BinDecodable, BinEncodable, BinEncoder};
use rustc_hash::FxHashSet;
use std::str::FromStr;
use std::sync::Arc;
Expand Down Expand Up @@ -77,6 +77,64 @@ pub mod engine_helpers {
build_response(req, ResponseCode::ServFail, Vec::new())
}

/// Build a SERVFAIL response for a validated, standard DNS query.
///
/// Listener code must use this strict variant so malformed packets and DNS
/// responses cannot trigger reflection traffic.
pub fn try_build_servfail_response_from_wire(request: &[u8]) -> Option<Bytes> {
if request.len() < 12
|| request[2] & 0x80 != 0
|| request[2] & 0x78 != 0
|| u16::from_be_bytes([request[4], request[5]]) != 1
|| u16::from_be_bytes([request[6], request[7]]) != 0
|| u16::from_be_bytes([request[8], request[9]]) != 0
{
return None;
}

let request_message = Message::from_bytes(request).ok()?;
if request_message.metadata.message_type != MessageType::Query
|| request_message.metadata.op_code != OpCode::Query
|| request_message.queries.len() != 1
{
return None;
}

build_servfail_response(&request_message).ok()
}

/// Build a SERVFAIL response directly from a DNS wire request.
///
/// Valid requests preserve TXID, opcode, RD, and the Question section. The
/// header-only fallback is intentionally allocation-bounded and is used only
/// by request/response transports such as DoH, never by a UDP listener.
pub fn build_servfail_response_from_wire(request: &[u8]) -> Bytes {
fn header_only(request: &[u8]) -> Bytes {
if request.len() < 2 {
return Bytes::new();
}

let mut response = vec![0u8; 12];
response[0] = request[0];
response[1] = request[1];
if request.len() >= 4 {
// Preserve opcode, RD, and CD while setting QR=1, RA=1, SERVFAIL.
response[2] = 0x80 | (request[2] & 0x79);
response[3] = 0x82 | (request[3] & 0x10);
} else {
response[2] = 0x80;
response[3] = 0x82;
}
Bytes::from(response)
}

let Ok(request_message) = Message::from_bytes(request) else {
return header_only(request);
};

build_servfail_response(&request_message).unwrap_or_else(|_| header_only(request))
}

/// 快速构建错误响应(ServFail),避免解析完整请求
/// Fast build ServFail response without parsing full request
#[inline]
Expand Down
Loading