From 9108af50fe3edb2cd70d8b9d90a3ca864119d0f5 Mon Sep 17 00:00:00 2001 From: JohnsonRan Date: Sat, 8 Aug 2026 14:17:20 +0900 Subject: [PATCH 1/5] fix(proto): validate DNS queries before fast-path handling --- src/main.rs | 92 +++++++++++++++++++++++ src/proto_utils.rs | 180 +++++++++++++++++++++++++++++++-------------- 2 files changed, 218 insertions(+), 54 deletions(-) diff --git a/src/main.rs b/src/main.rs index 93bae8d..6070a1f 100644 --- a/src/main.rs +++ b/src/main.rs @@ -13,6 +13,7 @@ use tracing_subscriber::{EnvFilter, fmt, layer::SubscriberExt, util::SubscriberI use kixdns::config::load_config; use kixdns::engine::{Engine, FastPathResponse, PreParsedData}; use kixdns::matcher::RuntimePipelineConfig; +use kixdns::proto_utils::is_standard_query_header; use kixdns::watcher; #[derive(Parser, Debug)] @@ -573,6 +574,9 @@ async fn run_udp_worker( { // 零拷贝获取 Bytes / Zero-copy obtain Bytes let packet_bytes = buf.split().freeze(); + if !is_standard_query_header(&packet_bytes) { + continue; + } // 每 100 个请求检查一次流控调整 / Check flow control adjustment every 100 requests request_count += 1; @@ -802,6 +806,9 @@ async fn handle_tcp_conn( // 统一 UDP 和 TCP 的行为,避免重复解析 // Unify UDP and TCP behavior to avoid re-parsing let packet_bytes = buf.split().freeze(); + if !is_standard_query_header(&packet_bytes) { + return Ok(()); + } let timeout_dur = Duration::from_millis(timeout_ms); let resp = match engine.handle_packet_fast(&packet_bytes, peer) { @@ -916,3 +923,88 @@ async fn handle_tcp_conn( } } } + +#[cfg(test)] +mod tests { + use super::*; + use hickory_proto::op::{Message, MessageType, OpCode, Query}; + use hickory_proto::rr::{Name, RecordType}; + use std::str::FromStr; + + #[ctor::ctor] + fn init_crypto() { + let _ = rustls::crypto::aws_lc_rs::default_provider().install_default(); + } + + fn static_engine() -> Engine { + use kixdns::config::PipelineConfig; + + let config: PipelineConfig = serde_json::from_value(serde_json::json!({ + "settings": { "default_upstream": "127.0.0.1:9" }, + "pipelines": [{ + "id": "p", + "rules": [{ + "name": "static", + "matchers": [{ "type": "any" }], + "actions": [{ "type": "static_ip_response", "ip": "192.0.2.1" }] + }] + }] + })) + .expect("parse config"); + let runtime = RuntimePipelineConfig::from_config(config).expect("build runtime config"); + Engine::new(runtime, "test".to_string()).expect("initialize engine") + } + + fn dns_query() -> Vec { + let mut message = Message::new(0xCAFE, MessageType::Query, OpCode::Query); + message.metadata.recursion_desired = true; + message.add_query(Query::query( + Name::from_str("data.xiaoheihe.cn").unwrap(), + RecordType::A, + )); + message.to_vec().unwrap() + } + + #[tokio::test] + async fn udp_static_fast_path_rejects_non_query_and_malformed_additional_records() { + let server = Arc::new(UdpSocket::bind("127.0.0.1:0").await.unwrap()); + let server_addr = server.local_addr().unwrap(); + let worker = tokio::spawn(run_udp_worker(0, Arc::clone(&server), static_engine())); + let client = UdpSocket::bind("127.0.0.1:0").await.unwrap(); + + let mut qr_response = dns_query(); + qr_response[2] |= 0x80; + + let mut missing_additional = dns_query(); + missing_additional[11] = 1; + + let mut truncated_edns = dns_query(); + truncated_edns[11] = 1; + truncated_edns.extend_from_slice(&[0, 0, 41, 0x04, 0xD0, 0, 0, 0, 0, 0, 4, 0, 1]); + + let mut invalid_pointer = dns_query(); + invalid_pointer[11] = 1; + invalid_pointer.extend_from_slice(&[0xC0, 0xFF, 0, 16, 0, 1, 0, 0, 0, 0, 0, 0]); + + let mut response_buf = [0u8; 512]; + for (case, packet) in [ + ("QR=1", qr_response), + ("missing Additional RR", missing_additional), + ("truncated EDNS RDATA", truncated_edns), + ("invalid Additional compression pointer", invalid_pointer), + ] { + client.send_to(&packet, server_addr).await.unwrap(); + assert!( + tokio::time::timeout( + Duration::from_millis(100), + client.recv_from(&mut response_buf), + ) + .await + .is_err(), + "{case} must be rejected before static/cache fast paths" + ); + } + + worker.abort(); + } +} diff --git a/src/proto_utils.rs b/src/proto_utils.rs index 3b5d7c4..2c6ce47 100644 --- a/src/proto_utils.rs +++ b/src/proto_utils.rs @@ -9,11 +9,28 @@ pub struct QuickQuery<'a> { pub edns_present: bool, } +/// Check the fixed DNS header fields required by KixDNS's standard query path. +/// This cheap ingress guard must run before cache/static fast paths so response +/// packets cannot be reflected as normal answers. +#[inline] +pub fn is_standard_query_header(packet: &[u8]) -> bool { + if packet.len() < 12 { + return false; + } + + let flags = u16::from_be_bytes([packet[2], packet[3]]); + let qd_count = u16::from_be_bytes([packet[4], packet[5]]); + let an_count = u16::from_be_bytes([packet[6], packet[7]]); + let ns_count = u16::from_be_bytes([packet[8], packet[9]]); + + flags & 0x8000 == 0 && flags & 0x7800 == 0 && qd_count == 1 && an_count == 0 && ns_count == 0 +} + /// 仅解析 DNS 头部和第一个 Query,用于快速缓存查找 / Parse only DNS header and first query for quick cache lookup /// 避免 hickory-proto Message::from_bytes 的全量解析和分配开销 / Avoid full parsing and allocation overhead of hickory-proto Message::from_bytes /// buf: 用于存储归一化(小写)域名的缓冲区,建议至少 256 字节 / buf: buffer for storing normalized (lowercase) domain name, recommend at least 256 bytes pub fn parse_quick<'a>(packet: &[u8], buf: &'a mut [u8]) -> Option> { - if packet.len() < 12 { + if !is_standard_query_header(packet) { return None; } @@ -21,15 +38,8 @@ pub fn parse_quick<'a>(packet: &[u8], buf: &'a mut [u8]) -> Option(packet: &[u8], buf: &'a mut [u8]) -> Option(packet: &[u8], buf: &'a mut [u8]) -> Option 0 && an_count == 0 && ns_count == 0 { - // Fast-path optimization: for standard query messages, AN and NS are expected to be 0. - // We only scan the Additional section in this common case to keep parsing fast, which may miss EDNS - // in non-standard messages (e.g., UPDATE, or responses mistakenly treated as queries). - // 快速路径优化:对于标准查询消息,AN 和 NS 通常为 0。仅在这种常见情况下扫描 Additional 以提升性能, - // 这在非标准消息(例如 UPDATE,或被误当作查询处理的响应)中可能漏检 EDNS。 - let mut ar_pos = pos; - for _ in 0..ar_count { - if ar_pos >= packet.len() { - break; - } - let name_byte = packet[ar_pos]; - let next_pos = if name_byte == 0 { - ar_pos + 1 - } else { - skip_name(packet, ar_pos).unwrap_or(packet.len()) - }; - - if next_pos + 10 > packet.len() { - break; - } - let rr_type = u16::from_be_bytes([packet[next_pos], packet[next_pos + 1]]); - if rr_type == 41 { - // OPT - edns_present = true; - break; - } - let rd_len = u16::from_be_bytes([packet[next_pos + 8], packet[next_pos + 9]]); + let mut ar_pos = pos; + for _ in 0..ar_count { + let owner_start = ar_pos; + let next_pos = skip_name(packet, ar_pos)?; + let fixed_end = next_pos.checked_add(10)?; + if fixed_end > packet.len() { + return None; + } - // Validate arithmetic to prevent overflow and ensure entire record fits in packet - let rd_len_usize = rd_len as usize; - // Check packet length first to avoid underflow in subsequent arithmetic - if packet.len() < 10 + rd_len_usize { - break; - } - if next_pos > packet.len() - 10 - rd_len_usize { - break; + let rr_type = u16::from_be_bytes([packet[next_pos], packet[next_pos + 1]]); + let rd_len = u16::from_be_bytes([packet[next_pos + 8], packet[next_pos + 9]]) as usize; + let rdata_end = fixed_end.checked_add(rd_len)?; + if rdata_end > packet.len() { + return None; + } + + if rr_type == 41 { + // RFC 6891: OPT owner name must be the root and only one OPT is valid. + if edns_present || packet.get(owner_start) != Some(&0) { + return None; } - ar_pos = next_pos + 10 + rd_len_usize; + validate_edns_options(packet, fixed_end, rdata_end)?; + edns_present = true; } + ar_pos = rdata_end; + } + if ar_pos != packet.len() { + return None; } // Return slice of buf as zero-copy reference @@ -213,22 +215,56 @@ pub fn parse_quick<'a>(packet: &[u8], buf: &'a mut [u8]) -> Option Option { let packet_len = packet.len(); + let mut encoded_end = None; + let mut jumps = 0usize; + loop { - if pos >= packet_len { - return None; - } - let len = packet[pos]; + let len = *packet.get(pos)?; if len == 0 { - return Some(pos + 1); + return Some(encoded_end.unwrap_or(pos + 1)); } if (len & 0xC0) == 0xC0 { - if pos + 2 > packet_len { + let pointer_end = pos.checked_add(2)?; + if pointer_end > packet_len { return None; } - return Some(pos + 2); + encoded_end.get_or_insert(pointer_end); + let offset = (((len & 0x3F) as usize) << 8) | packet[pos + 1] as usize; + if offset >= packet_len { + return None; + } + jumps += 1; + if jumps > 16 { + return None; + } + pos = offset; + continue; + } + if len & 0xC0 != 0 || len > 63 { + return None; + } + + let next = pos.checked_add(1 + len as usize)?; + if next > packet_len { + return None; + } + pos = next; + } +} + +fn validate_edns_options(packet: &[u8], mut pos: usize, end: usize) -> Option<()> { + while pos < end { + let header_end = pos.checked_add(4)?; + if header_end > end { + return None; + } + let option_len = u16::from_be_bytes([packet[pos + 2], packet[pos + 3]]) as usize; + pos = header_end.checked_add(option_len)?; + if pos > end { + return None; } - pos += 1 + len as usize; } + (pos == end).then_some(()) } impl QuickQuery<'_> { @@ -809,6 +845,42 @@ mod tests { assert!(parse_quick(&packet, &mut qname).is_some()); } + fn query_with_additional(additional: &[u8]) -> Vec { + let mut packet = query_with_label(b"example"); + packet[10..12].copy_from_slice(&1u16.to_be_bytes()); + packet.extend_from_slice(additional); + packet + } + + #[test] + fn parse_quick_rejects_missing_declared_additional_record() { + let packet = query_with_additional(&[]); + let mut qname = [0u8; 256]; + assert!(parse_quick(&packet, &mut qname).is_none()); + } + + #[test] + fn parse_quick_rejects_truncated_edns_rdata() { + let packet = query_with_additional(&[0, 0, 41, 0x04, 0xD0, 0, 0, 0, 0, 0, 4, 0, 1]); + let mut qname = [0u8; 256]; + assert!(parse_quick(&packet, &mut qname).is_none()); + } + + #[test] + fn parse_quick_rejects_invalid_additional_compression_pointer() { + let packet = query_with_additional(&[0xC0, 0xFF, 0, 16, 0, 1, 0, 0, 0, 0, 0, 0]); + let mut qname = [0u8; 256]; + assert!(parse_quick(&packet, &mut qname).is_none()); + } + + #[test] + fn parse_quick_accepts_valid_empty_edns_record() { + let packet = query_with_additional(&[0, 0, 41, 0x04, 0xD0, 0, 0, 0, 0, 0, 0]); + let mut qname = [0u8; 256]; + let query = parse_quick(&packet, &mut qname).expect("valid EDNS query"); + assert!(query.edns_present); + } + #[test] fn saturating_u64_to_u32_caps_large_values() { assert_eq!(saturating_u64_to_u32(42), 42); From bbb121e9043585b81cc275f0405532f98de28724 Mon Sep 17 00:00:00 2001 From: kix Date: Sat, 8 Aug 2026 15:26:24 +0800 Subject: [PATCH 2/5] merge(three-pr): integrate #31 truncate + #32 validate + #33 SERVFAIL - #32: is_standard_query_header ingress guard (UDP/TCP) + parse_quick full Additional validation (compression pointers, OPT root/unique, EDNS options) - #31: truncate_udp_response at record boundaries with TC bit, integrated into send_udp_response/try_send_udp_response (all 5 UDP send paths) - #33: strict try_build SERVFAIL for UDP/TCP + DoH reuse, write_all TCP framing - fix: permit-exhaustion path uses non-blocking send (backpressure drop) with fast SERVFAIL from pre-parsed data, never blocks the receive loop - tests: 99 lib + 6 bin + 9 doh integration, clippy -D warnings, fmt clean --- src/doh_server.rs | 40 +--- src/engine/execution.rs | 64 +++++++ src/engine/utils.rs | 60 +++++- src/main.rs | 392 ++++++++++++++++++++++++++++++++++++---- src/proto_utils.rs | 223 +++++++++++++++++++++++ 5 files changed, 710 insertions(+), 69 deletions(-) diff --git a/src/doh_server.rs b/src/doh_server.rs index c3b83a8..ebb2f06 100644 --- a/src/doh_server.rs +++ b/src/doh_server.rs @@ -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; @@ -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; @@ -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> { @@ -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] diff --git a/src/engine/execution.rs b/src/engine/execution.rs index 4b97724..0b163bd 100644 --- a/src/engine/execution.rs +++ b/src/engine/execution.rs @@ -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 diff --git a/src/engine/utils.rs b/src/engine/utils.rs index 6c19fd3..a2599ce 100644 --- a/src/engine/utils.rs +++ b/src/engine/utils.rs @@ -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; @@ -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 { + 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] diff --git a/src/main.rs b/src/main.rs index 6070a1f..750bdf8 100644 --- a/src/main.rs +++ b/src/main.rs @@ -11,9 +11,9 @@ use tracing::{debug, error, info, warn}; use tracing_subscriber::{EnvFilter, fmt, layer::SubscriberExt, util::SubscriberInitExt}; use kixdns::config::load_config; -use kixdns::engine::{Engine, FastPathResponse, PreParsedData}; +use kixdns::engine::{Engine, FastPathResponse, PreParsedData, engine_helpers}; use kixdns::matcher::RuntimePipelineConfig; -use kixdns::proto_utils::is_standard_query_header; +use kixdns::proto_utils::{is_standard_query_header, truncate_udp_response}; use kixdns::watcher; #[derive(Parser, Debug)] @@ -527,6 +527,68 @@ fn create_reuseport_udp_socket(addr: SocketAddr) -> anyhow::Result Option { + engine_helpers::try_build_servfail_response_from_wire(request) +} + +/// 发送 UDP 数据报:try_send_to 优先(避免 reactor 开销),WouldBlock 回退异步 +/// Send UDP datagram: try non-blocking first, async fallback on WouldBlock +async fn send_udp_datagram( + socket: &UdpSocket, + data: &[u8], + peer: SocketAddr, +) -> std::io::Result { + match socket.try_send_to(data, peer) { + Ok(sent) => Ok(sent), + Err(error) if error.kind() == std::io::ErrorKind::WouldBlock => { + socket.send_to(data, peer).await + } + Err(error) => Err(error), + } +} + +/// 发送 UDP 响应:超限响应截断(RFC 6891)+ try_send_to 优先,WouldBlock 回退异步 +/// Send UDP response: truncate oversized (RFC 6891), try non-blocking first, async fallback on WouldBlock +async fn send_udp_response(socket: &UdpSocket, request: &[u8], response: &[u8], peer: SocketAddr) { + // 截断超限响应,避免 UDP 报文超过客户端声明上限 / Truncate to the client's advertised UDP limit + let truncated = truncate_udp_response(request, response); + let response = truncated.as_deref().unwrap_or(response); + if let Err(error) = send_udp_datagram(socket, response, peer).await { + debug!(%peer, %error, response_len = response.len(), "failed to send UDP response"); + } +} + +/// 非阻塞发送 UDP 响应(过载保护路径):截断 + 纯 try_send_to,WouldBlock 背压丢弃 +/// Non-blocking UDP send for overload paths: truncate + try_send_to only; drop on WouldBlock +fn try_send_udp_response(socket: &UdpSocket, request: &[u8], response: &[u8], peer: SocketAddr) { + let truncated = truncate_udp_response(request, response); + let response = truncated.as_deref().unwrap_or(response); + if let Err(error) = socket.try_send_to(response, peer) + && error.kind() != std::io::ErrorKind::WouldBlock + { + debug!(%peer, %error, response_len = response.len(), "failed to send UDP response"); + } + // WouldBlock: 发送缓冲满,丢弃(背压)/ send buffer full, drop (backpressure) +} + +/// 异步发送 SERVFAIL(处理失败/超时路径) +async fn send_udp_servfail(socket: &UdpSocket, request: &[u8], peer: SocketAddr) { + let Some(response) = listener_servfail(request) else { + return; + }; + send_udp_response(socket, request, &response, peer).await; +} + +/// 非阻塞发送 SERVFAIL(permit 耗尽路径,避免阻塞接收循环) +fn try_send_udp_servfail(socket: &UdpSocket, request: &[u8], peer: SocketAddr) { + let Some(response) = listener_servfail(request) else { + return; + }; + try_send_udp_response(socket, request, &response, peer); +} + /// 高性能 UDP worker:直接在接收循环中处理请求,避免 spawn 开销 / High-performance UDP worker: process requests directly in receive loop, avoiding spawn overhead async fn run_udp_worker( worker_id: usize, @@ -592,7 +654,7 @@ async fn run_udp_worker( match engine.handle_packet_fast(&packet_bytes, peer) { Ok(Some(FastPathResponse::Direct(bytes))) => { // 已包含正确 TXID,可直接发送 / Already contains correct TXID - let _ = socket.send_to(&bytes, peer).await; + send_udp_response(&socket, &packet_bytes, &bytes, peer).await; } Ok(Some(FastPathResponse::CacheHit { cached, @@ -617,11 +679,7 @@ async fn run_udp_worker( let id_bytes = tx_id.to_be_bytes(); send_buf[0] = id_bytes[0]; send_buf[1] = id_bytes[1]; - if let Err(e) = socket.try_send_to(&send_buf, peer) - && e.kind() == std::io::ErrorKind::WouldBlock - { - let _ = socket.send_to(&send_buf, peer).await; - } + send_udp_response(&socket, &packet_bytes, &send_buf, peer).await; } else { // Slow path: TTL has decayed, must patch all TTLs in-place. // 慢速路径:TTL 已衰减,必须原地 patch 所有 TTL。 @@ -638,7 +696,7 @@ async fn run_udp_worker( send_buf[0] = id_bytes[0]; send_buf[1] = id_bytes[1]; } - let _ = socket.send_to(&send_buf, peer).await; + send_udp_response(&socket, &packet_bytes, &send_buf, peer).await; } } Ok(Some(FastPathResponse::AsyncNeeded { @@ -686,10 +744,11 @@ async fn run_udp_worker( .await { Ok(Ok(resp)) => { - let _ = socket.send_to(&resp, peer).await; + send_udp_response(&socket, &packet_bytes, &resp, peer).await; } Ok(Err(e)) => { - debug!(error = %e, "handle_packet error"); + debug!(error = %e, "handle_packet error, returning SERVFAIL"); + send_udp_servfail(&socket, &packet_bytes, peer).await; } Err(_) => { warn!( @@ -697,9 +756,22 @@ async fn run_udp_worker( upstream_timeout_ms = engine.get_upstream_timeout_ms(), "request timeout after hedge and fallback exhausted" ); + send_udp_servfail(&socket, &packet_bytes, peer).await; } } }); + } else { + // 过载保护:permit 耗尽,用预解析数据快速构造 SERVFAIL 并非阻塞发送(背压丢弃) + // Overload: permit exhausted, fast SERVFAIL from pre-parsed data with non-blocking send + if let Ok(resp) = engine_helpers::build_servfail_response_fast( + tx_id, + &qname, + qtype, + qclass, + packet_bytes[2] & 0x01 != 0, + ) { + try_send_udp_response(&socket, &packet_bytes, &resp, peer); + } } } Ok(None) => { @@ -723,10 +795,11 @@ async fn run_udp_worker( .await { Ok(Ok(resp)) => { - let _ = socket.send_to(&resp, peer).await; + send_udp_response(&socket, &packet_bytes, &resp, peer).await; } Ok(Err(e)) => { - debug!(error = %e, "handle_packet error"); + debug!(error = %e, "handle_packet error, returning SERVFAIL"); + send_udp_servfail(&socket, &packet_bytes, peer).await; } Err(_) => { warn!( @@ -734,13 +807,19 @@ async fn run_udp_worker( upstream_timeout_ms = engine.get_upstream_timeout_ms(), "request timeout" ); + send_udp_servfail(&socket, &packet_bytes, peer).await; } } }); + } else { + // 过载保护:permit 耗尽,非阻塞 SERVFAIL(背压丢弃,不阻塞接收循环) + // Overload: permit exhausted, non-blocking SERVFAIL with backpressure drop + try_send_udp_servfail(&socket, &packet_bytes, peer); } } - Err(_) => { - // 解析错误,忽略 / Parse error, ignore + Err(error) => { + debug!(%error, "UDP fast-path error, returning SERVFAIL for valid query"); + send_udp_servfail(&socket, &packet_bytes, peer).await; } } } @@ -870,14 +949,23 @@ async fn handle_tcp_conn( .await { Ok(Ok(r)) => r, - Ok(Err(_)) => return Ok(()), + Ok(Err(error)) => { + debug!(%error, "TCP request processing error, returning SERVFAIL"); + let Some(response) = listener_servfail(&packet_bytes) else { + return Ok(()); + }; + response + } Err(_) => { warn!( timeout_ms, upstream_timeout_ms = engine.get_upstream_timeout_ms(), "TCP request timeout after hedge and fallback exhausted" ); - return Ok(()); // 关闭连接 / Close connection + let Some(response) = listener_servfail(&packet_bytes) else { + return Ok(()); + }; + response } } } @@ -888,36 +976,43 @@ async fn handle_tcp_conn( .await { Ok(Ok(r)) => r, - Ok(Err(_)) => return Ok(()), + Ok(Err(error)) => { + debug!(%error, "TCP request processing error, returning SERVFAIL"); + let Some(response) = listener_servfail(&packet_bytes) else { + return Ok(()); + }; + response + } Err(_) => { warn!( timeout_ms, upstream_timeout_ms = engine.get_upstream_timeout_ms(), "TCP request timeout" ); - return Ok(()); // 关闭连接 / Close connection + let Some(response) = listener_servfail(&packet_bytes) else { + return Ok(()); + }; + response } } } - Err(_) => { - // 解析错误,关闭连接 / Parse error, close connection - return Ok(()); + Err(error) => { + debug!(%error, "TCP fast-path error, returning SERVFAIL for valid query"); + let Some(response) = listener_servfail(&packet_bytes) else { + return Ok(()); + }; + response } }; if resp.len() <= u16::MAX as usize { - // Single vectored write (scatter-gather): length prefix + body in one - // syscall, avoiding Nagle/delayed-ACK interaction on the wire even - // when TCP_NODELAY fails to take effect. - // - // 单次 vectored write (scatter-gather):长度前缀 + 包体在一次 syscall 中完成, - // 即使 TCP_NODELAY 未生效也能避免 Nagle/delayed-ACK 交互。 - let len_bytes = (resp.len() as u16).to_be_bytes(); - let bufs = [ - std::io::IoSlice::new(&len_bytes), - std::io::IoSlice::new(&resp), - ]; - if stream.write_vectored(&bufs).await.is_err() { + // DNS-over-TCP framing must be written completely; a successful + // write_vectored call may still be partial. + // write_vectored 可能部分写入,必须用 write_all 保证完整帧。 + let mut frame = Vec::with_capacity(2 + resp.len()); + frame.extend_from_slice(&(resp.len() as u16).to_be_bytes()); + frame.extend_from_slice(&resp); + if stream.write_all(&frame).await.is_err() { return Ok(()); } } @@ -927,8 +1022,9 @@ async fn handle_tcp_conn( #[cfg(test)] mod tests { use super::*; - use hickory_proto::op::{Message, MessageType, OpCode, Query}; + use hickory_proto::op::{Message, MessageType, OpCode, Query, ResponseCode}; use hickory_proto::rr::{Name, RecordType}; + use hickory_proto::serialize::binary::BinDecodable; use std::str::FromStr; #[ctor::ctor] @@ -1007,4 +1103,230 @@ mod tests { worker.abort(); } + + fn forwarding_engine() -> Engine { + use kixdns::config::PipelineConfig; + + let config: PipelineConfig = serde_json::from_value(serde_json::json!({ + "settings": { + "default_upstream": "127.0.0.1:9", + "flow_control_enabled": true, + "flow_control_initial_permits": 1, + "flow_control_min_permits": 1, + "flow_control_max_permits": 1 + }, + "pipelines": [{ + "id": "p", + "rules": [{ + "name": "forward", + "matchers": [{ "type": "any" }], + "actions": [{ + "type": "forward", + "upstream": "127.0.0.1:9", + "transport": "udp" + }] + }] + }] + })) + .expect("parse config"); + let runtime = RuntimePipelineConfig::from_config(config).expect("build runtime config"); + Engine::new(runtime, "test".to_string()).expect("initialize engine") + } + + fn doh_timeout_engine(upstream: &str) -> Engine { + use kixdns::config::PipelineConfig; + + let config: PipelineConfig = serde_json::from_value(serde_json::json!({ + "settings": { + "default_upstream": upstream, + "upstream_timeout_ms": 50, + "request_timeout_ms": 50 + }, + "pipelines": [{ + "id": "p", + "rules": [{ + "name": "forward", + "matchers": [{ "type": "any" }], + "actions": [{ + "type": "forward", + "upstream": upstream, + "transport": "doh" + }] + }] + }] + })) + .expect("parse config"); + let runtime = RuntimePipelineConfig::from_config(config).expect("build runtime config"); + Engine::new(runtime, "test".to_string()).expect("initialize engine") + } + + #[tokio::test] + async fn udp_send_datagram_surfaces_socket_errors() { + let socket = UdpSocket::bind("127.0.0.1:0").await.unwrap(); + let ipv6_peer: SocketAddr = "[::1]:53".parse().unwrap(); + assert!(send_udp_datagram(&socket, b"dns", ipv6_peer).await.is_err()); + } + + #[tokio::test] + async fn udp_permit_exhaustion_returns_servfail_but_malformed_packet_is_silent() { + let engine = forwarding_engine(); + engine.permit_manager.set_max_permits(0); + + let server = Arc::new(UdpSocket::bind("127.0.0.1:0").await.unwrap()); + let server_addr = server.local_addr().unwrap(); + let worker = tokio::spawn(run_udp_worker(0, Arc::clone(&server), engine)); + let client = UdpSocket::bind("127.0.0.1:0").await.unwrap(); + + client.send_to(&dns_query(), server_addr).await.unwrap(); + let mut response_buf = [0u8; 512]; + let (response_len, _) = + tokio::time::timeout(Duration::from_secs(1), client.recv_from(&mut response_buf)) + .await + .expect("valid query should receive a response") + .unwrap(); + let response = Message::from_bytes(&response_buf[..response_len]).unwrap(); + assert_eq!(response.metadata.id, 0xCAFE); + assert_eq!(response.metadata.response_code, ResponseCode::ServFail); + assert_eq!(response.queries.len(), 1); + + client.send_to(&[0x12, 0x34], server_addr).await.unwrap(); + assert!( + tokio::time::timeout( + Duration::from_millis(100), + client.recv_from(&mut response_buf), + ) + .await + .is_err(), + "malformed UDP packets must not receive a reflected response" + ); + + worker.abort(); + } + + #[tokio::test] + async fn udp_hanging_doh_returns_servfail() { + let blackhole = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let blackhole_addr = blackhole.local_addr().unwrap(); + let blackhole_task = tokio::spawn(async move { + if let Ok((_stream, _)) = blackhole.accept().await { + tokio::time::sleep(Duration::from_secs(2)).await; + } + }); + + let upstream = format!("https://{blackhole_addr}/dns-query"); + let engine = doh_timeout_engine(&upstream); + let server = Arc::new(UdpSocket::bind("127.0.0.1:0").await.unwrap()); + let server_addr = server.local_addr().unwrap(); + let worker = tokio::spawn(run_udp_worker(0, Arc::clone(&server), engine)); + let client = UdpSocket::bind("127.0.0.1:0").await.unwrap(); + + client.send_to(&dns_query(), server_addr).await.unwrap(); + let mut response_buf = [0u8; 512]; + let (response_len, _) = + tokio::time::timeout(Duration::from_secs(1), client.recv_from(&mut response_buf)) + .await + .expect("timed-out query should receive SERVFAIL") + .unwrap(); + let response = Message::from_bytes(&response_buf[..response_len]).unwrap(); + assert_eq!(response.metadata.id, 0xCAFE); + assert_eq!(response.metadata.response_code, ResponseCode::ServFail); + assert_eq!(response.queries.len(), 1); + + worker.abort(); + blackhole_task.abort(); + } + + #[tokio::test] + async fn tcp_hanging_doh_returns_complete_servfail_frame() { + let blackhole = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let blackhole_addr = blackhole.local_addr().unwrap(); + let blackhole_task = tokio::spawn(async move { + if let Ok((_stream, _)) = blackhole.accept().await { + tokio::time::sleep(Duration::from_secs(2)).await; + } + }); + + let upstream = format!("https://{blackhole_addr}/dns-query"); + let engine = doh_timeout_engine(&upstream); + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let listener_addr = listener.local_addr().unwrap(); + let client_task = tokio::spawn(async move { + let mut client = TcpStream::connect(listener_addr).await.unwrap(); + let query = dns_query(); + let mut frame = Vec::with_capacity(2 + query.len()); + frame.extend_from_slice(&(query.len() as u16).to_be_bytes()); + frame.extend_from_slice(&query); + client.write_all(&frame).await.unwrap(); + + let mut len_buf = [0u8; 2]; + tokio::time::timeout(Duration::from_secs(1), client.read_exact(&mut len_buf)) + .await + .expect("TCP query should receive a length prefix") + .unwrap(); + let response_len = u16::from_be_bytes(len_buf) as usize; + let mut response = vec![0u8; response_len]; + tokio::time::timeout(Duration::from_secs(1), client.read_exact(&mut response)) + .await + .expect("TCP query should receive the complete DNS frame") + .unwrap(); + response + }); + + let (server_stream, peer) = listener.accept().await.unwrap(); + let server_task = tokio::spawn(handle_tcp_conn(server_stream, peer, engine)); + let response_bytes = client_task.await.unwrap(); + let response = Message::from_bytes(&response_bytes).unwrap(); + assert_eq!(response.metadata.id, 0xCAFE); + assert_eq!(response.metadata.response_code, ResponseCode::ServFail); + assert_eq!(response.queries.len(), 1); + + server_task.abort(); + blackhole_task.abort(); + } + + #[tokio::test] + async fn udp_response_truncates_oversized_payload() { + use hickory_proto::rr::{RData, Record, rdata::TXT}; + + let server = UdpSocket::bind("127.0.0.1:0").await.unwrap(); + let server_addr = server.local_addr().unwrap(); + let client = UdpSocket::bind("127.0.0.1:0").await.unwrap(); + + // 客户端无 EDNS → 512 字节限制 / Client without EDNS → 512-byte limit + let request = dns_query(); + // 构造 >512 字节的大 TXT 响应 / Build a large TXT response (>512 bytes) + let name = Name::from_str("data.xiaoheihe.cn.").unwrap(); + let mut big = Message::new(0xCAFE, MessageType::Response, OpCode::Query); + big.metadata.recursion_desired = true; + big.metadata.recursion_available = true; + big.add_query(Query::query(name.clone(), RecordType::TXT)); + for i in 0..20 { + let text = format!("verification-{i:02}-{}", "x".repeat(180)); + big.add_answer(Record::from_rdata( + name.clone(), + 300, + RData::TXT(TXT::new(vec![text])), + )); + } + let big_response = big.to_vec().unwrap(); + assert!(big_response.len() > 512); + + tokio::spawn(async move { + let mut buf = [0u8; 4096]; + let (len, peer) = server.recv_from(&mut buf).await.unwrap(); + send_udp_response(&server, &buf[..len], &big_response, peer).await; + }); + + client.send_to(&request, server_addr).await.unwrap(); + let mut buf = [0u8; 4096]; + let (len, _) = tokio::time::timeout(Duration::from_secs(1), client.recv_from(&mut buf)) + .await + .expect("client should receive a truncated response") + .unwrap(); + let resp = Message::from_bytes(&buf[..len]).unwrap(); + assert!(len <= 512, "response must fit the 512-byte UDP limit"); + assert!(resp.metadata.truncation, "TC bit must be set"); + assert_eq!(resp.queries.len(), 1); + assert_eq!(resp.metadata.id, 0xCAFE); + } } diff --git a/src/proto_utils.rs b/src/proto_utils.rs index 2c6ce47..3d27b32 100644 --- a/src/proto_utils.rs +++ b/src/proto_utils.rs @@ -1,5 +1,9 @@ use std::hash::Hasher; +use bytes::Bytes; +use hickory_proto::op::Message; +use hickory_proto::serialize::binary::{BinDecodable, BinEncodable, BinEncoder}; + /// 快速解析结果,零拷贝实现 / Quick parse result with zero-copy implementation pub struct QuickQuery<'a> { pub tx_id: u16, @@ -676,6 +680,128 @@ pub fn parse_response_quick(packet: &[u8]) -> Option { }) } +/// Re-encode an oversized downstream UDP response so it respects the payload size +/// advertised by the client (RFC 6891). Queries without EDNS are limited to the +/// classic 512-byte DNS/UDP payload. +/// +/// Hickory's bounded encoder emits as many complete records as fit, fixes the +/// section counts, and sets TC=1 when records are omitted. `None` means the +/// original response already fits and can be sent without allocation. +/// +/// Trade-off: oversized responses (>512B, or > the client EDNS payload) are +/// re-parsed and re-encoded on every send. This is intentional — truncation +/// must happen at DNS record boundaries (RFC 1035 §4.2.1) and the correct TC +/// bit + counts can only be produced by a full re-encode. The overwhelmingly +/// common small-response path returns `None` before any parsing, so the cost +/// only applies to responses that would otherwise break the UDP size limit. +/// +/// 权衡:超限响应(>512B 或超过客户端 EDNS 载荷)每次发送都会重新解析并重编码。 +/// 这是有意的——截断必须发生在 DNS 记录边界(RFC 1035 §4.2.1),且正确的 TC 位 +/// 和计数只能由完整重编码产生。绝大多数小响应路径在解析前就返回 None,因此该 +/// 开销只作用于那些会突破 UDP 大小限制的响应。 +pub fn truncate_udp_response(request_packet: &[u8], response_packet: &[u8]) -> Option { + const CLASSIC_DNS_UDP_PAYLOAD: usize = 512; + + // Avoid parsing the request on the overwhelmingly common small-response path. + if response_packet.len() <= CLASSIC_DNS_UDP_PAYLOAD { + return None; + } + + // hickory's Message::max_payload() already clamps to >= 512, so the extra + // max() below is defensive only. + let max_payload = Message::from_bytes(request_packet) + .map(|request| request.max_payload() as usize) + .unwrap_or(CLASSIC_DNS_UDP_PAYLOAD) + .max(CLASSIC_DNS_UDP_PAYLOAD); + + if response_packet.len() <= max_payload { + return None; + } + + if let Ok(mut response) = Message::from_bytes(response_packet) { + // Do not advertise a larger UDP payload in the response than the client + // offered in its request. + if let Some(edns) = response.edns.as_mut() { + edns.set_max_payload(max_payload as u16); + } + + let mut out = Vec::with_capacity(max_payload); + let encoded = { + let mut encoder = BinEncoder::new(&mut out); + encoder.set_max_size(max_payload as u16); + response.emit(&mut encoder) + }; + if encoded.is_ok() && out.len() <= max_payload { + return Some(Bytes::from(out)); + } + } + + // Internal responses should always parse, but never send an oversized UDP + // datagram if a malformed response reaches this boundary. Fall back to a + // minimal TC=1 response containing the original Question section. + Some(build_minimal_truncated_response( + request_packet, + response_packet, + max_payload, + )) +} + +fn build_minimal_truncated_response( + request_packet: &[u8], + response_packet: &[u8], + max_payload: usize, +) -> Bytes { + let mut question_end = 12usize; + let mut copied_questions = false; + + if request_packet.len() >= 12 { + let qd_count = u16::from_be_bytes([request_packet[4], request_packet[5]]); + let mut pos = 12usize; + let mut valid = true; + for _ in 0..qd_count { + let Some(name_end) = skip_name(request_packet, pos) else { + valid = false; + break; + }; + let Some(next) = name_end.checked_add(4) else { + valid = false; + break; + }; + if next > request_packet.len() { + valid = false; + break; + } + pos = next; + } + if valid && pos <= max_payload { + question_end = pos; + copied_questions = true; + } + } + + let header_source = if response_packet.len() >= 12 { + response_packet + } else { + request_packet + }; + let mut out = if header_source.len() >= 12 { + header_source[..12].to_vec() + } else { + vec![0u8; 12] + }; + + // QR=1 and TC=1; preserve the remaining response flags and RCODE. + out[2] |= 0x82; + if copied_questions { + out[4..6].copy_from_slice(&request_packet[4..6]); + out.extend_from_slice(&request_packet[12..question_end]); + } else { + out[4..6].copy_from_slice(&0u16.to_be_bytes()); + } + out[6..12].fill(0); + Bytes::from(out) +} + /// Saturating conversion for DNS TTLs and elapsed seconds. /// DNS TTL fields are u32, while Duration and configuration calculations use u64. #[inline] @@ -845,6 +971,103 @@ mod tests { assert!(parse_quick(&packet, &mut qname).is_some()); } + fn encode_message(message: &Message) -> Vec { + let mut bytes = Vec::new(); + let mut encoder = BinEncoder::new(&mut bytes); + message.emit(&mut encoder).expect("encode DNS message"); + bytes + } + + fn txt_query(max_payload: Option) -> Vec { + use hickory_proto::op::{Edns, MessageType, OpCode, Query}; + use hickory_proto::rr::{Name, RecordType}; + use std::str::FromStr; + + let mut message = Message::new(0x1234, MessageType::Query, OpCode::Query); + message.add_query(Query::query( + Name::from_str("cloudflare.com.").unwrap(), + RecordType::TXT, + )); + if let Some(max_payload) = max_payload { + let mut edns = Edns::new(); + edns.set_max_payload(max_payload); + message.set_edns(edns); + } + encode_message(&message) + } + + fn large_txt_response(answer_count: usize) -> Vec { + use hickory_proto::op::{Edns, MessageType, OpCode, Query}; + use hickory_proto::rr::{Name, RData, Record, RecordType, rdata::TXT}; + use std::str::FromStr; + + let name = Name::from_str("cloudflare.com.").unwrap(); + let mut message = Message::new(0x1234, MessageType::Response, OpCode::Query); + message.metadata.recursion_desired = true; + message.metadata.recursion_available = true; + message.add_query(Query::query(name.clone(), RecordType::TXT)); + for index in 0..answer_count { + let text = format!("verification-{index:02}-{}", "x".repeat(180)); + message.add_answer(Record::from_rdata( + name.clone(), + 300, + RData::TXT(TXT::new(vec![text])), + )); + } + let mut edns = Edns::new(); + edns.set_max_payload(4096); + message.set_edns(edns); + encode_message(&message) + } + + #[test] + fn truncate_udp_response_uses_512_without_edns() { + let request = txt_query(None); + let response = large_txt_response(20); + assert!(response.len() > 512); + + let truncated = truncate_udp_response(&request, &response).expect("must truncate"); + assert!(truncated.len() <= 512); + + let decoded = Message::from_bytes(&truncated).expect("valid truncated response"); + assert!(decoded.metadata.truncation); + assert_eq!(decoded.metadata.id, 0x1234); + assert_eq!(decoded.queries.len(), 1); + } + + #[test] + fn truncate_udp_response_respects_edns_payload_size() { + let request = txt_query(Some(1232)); + let response = large_txt_response(20); + assert!(response.len() > 1232); + + let truncated = truncate_udp_response(&request, &response).expect("must truncate"); + assert!(truncated.len() <= 1232); + + let decoded = Message::from_bytes(&truncated).expect("valid truncated response"); + assert!(decoded.metadata.truncation); + assert_eq!(decoded.queries.len(), 1); + if let Some(edns) = decoded.edns { + assert_eq!(edns.max_payload(), 1232); + } + } + + #[test] + fn truncate_udp_response_keeps_response_that_fits_client_limit() { + let request = txt_query(Some(4096)); + let response = large_txt_response(8); + assert!(response.len() > 512); + assert!(response.len() <= 4096); + + assert!(truncate_udp_response(&request, &response).is_none()); + } + + #[test] + fn truncate_udp_response_skips_small_responses_without_parsing_request() { + let response = vec![0u8; 128]; + assert!(truncate_udp_response(b"not a dns request", &response).is_none()); + } + fn query_with_additional(additional: &[u8]) -> Vec { let mut packet = query_with_label(b"example"); packet[10..12].copy_from_slice(&1u16.to_be_bytes()); From c230f6e29e444905a23930e1fa52468ff045a624 Mon Sep 17 00:00:00 2001 From: kix Date: Sat, 8 Aug 2026 15:33:43 +0800 Subject: [PATCH 3/5] docs(main): document Ok(None) overload-path SERVFAIL parse cost trade-off --- src/main.rs | 8 ++++++-- 1 file changed, 6 insertions(+), 2 deletions(-) diff --git a/src/main.rs b/src/main.rs index 750bdf8..03fe8d7 100644 --- a/src/main.rs +++ b/src/main.rs @@ -812,8 +812,12 @@ async fn run_udp_worker( } }); } else { - // 过载保护:permit 耗尽,非阻塞 SERVFAIL(背压丢弃,不阻塞接收循环) - // Overload: permit exhausted, non-blocking SERVFAIL with backpressure drop + // 过载保护:permit 耗尽。此路径无预解析数据(parse_quick 已失败), + // 需完整解析请求才能构造严格 SERVFAIL 并回显 question。 + // 解析成本有界(≤4096B 报文),发送为非阻塞(背压丢弃),不阻塞接收循环。 + // Overload: no pre-parsed data here (parse_quick failed), so the strict + // SERVFAIL needs a full parse. Cost is bounded by packet size and the + // send is non-blocking (backpressure drop), so the receive loop is safe. try_send_udp_servfail(&socket, &packet_bytes, peer); } } From 38e3cd3ce1f88b4fe21dee0ece860e91dc7e8c92 Mon Sep 17 00:00:00 2001 From: kix Date: Sat, 8 Aug 2026 15:56:26 +0800 Subject: [PATCH 4/5] fix(listener): apply review fixes to SERVFAIL/truncate/validate paths - overload: drop packets that failed parse_quick instead of full re-parse - TCP: use write_all_vectored for zero-alloc complete frame writes - DRY: reuse is_standard_query_header in the strict SERVFAIL builder - document fast-path Additional validation trade-off - inline the listener_servfail one-line wrapper --- Cargo.toml | 2 +- src/engine/utils.rs | 11 +++----- src/main.rs | 64 ++++++++++++++++++++------------------------- src/proto_utils.rs | 8 ++++++ 4 files changed, 42 insertions(+), 43 deletions(-) diff --git a/Cargo.toml b/Cargo.toml index c7c43c6..f48f569 100644 --- a/Cargo.toml +++ b/Cargo.toml @@ -12,7 +12,7 @@ readme = "README.md" [dependencies] tokio = { version = "1", features = ["macros", "rt-multi-thread", "net", "time", "signal", "sync"] } -tokio-util = { version = "0.7", features = ["rt"] } +tokio-util = { version = "0.7", features = ["rt", "io"] } hickory-proto = "0.26" serde = { version = "1", features = ["derive"] } serde_json = "1" diff --git a/src/engine/utils.rs b/src/engine/utils.rs index a2599ce..3a013a5 100644 --- a/src/engine/utils.rs +++ b/src/engine/utils.rs @@ -82,13 +82,10 @@ pub mod engine_helpers { /// 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 { - 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 - { + // Reuse the shared standard-query header guard so listener and fast-path + // validation cannot drift apart. + // 复用共享的标准查询头守卫,避免监听器与快速路径的校验漂移。 + if !crate::proto_utils::is_standard_query_header(request) { return None; } diff --git a/src/main.rs b/src/main.rs index 03fe8d7..d08f52a 100644 --- a/src/main.rs +++ b/src/main.rs @@ -5,7 +5,7 @@ use std::time::Duration; use anyhow::Context; use clap::{Parser, Subcommand}; -use tokio::io::{AsyncReadExt, AsyncWriteExt}; +use tokio::io::AsyncReadExt; use tokio::net::{TcpListener, TcpStream, UdpSocket}; use tracing::{debug, error, info, warn}; use tracing_subscriber::{EnvFilter, fmt, layer::SubscriberExt, util::SubscriberInitExt}; @@ -527,12 +527,6 @@ fn create_reuseport_udp_socket(addr: SocketAddr) -> anyhow::Result Option { - engine_helpers::try_build_servfail_response_from_wire(request) -} - /// 发送 UDP 数据报:try_send_to 优先(避免 reactor 开销),WouldBlock 回退异步 /// Send UDP datagram: try non-blocking first, async fallback on WouldBlock async fn send_udp_datagram( @@ -575,20 +569,12 @@ fn try_send_udp_response(socket: &UdpSocket, request: &[u8], response: &[u8], pe /// 异步发送 SERVFAIL(处理失败/超时路径) async fn send_udp_servfail(socket: &UdpSocket, request: &[u8], peer: SocketAddr) { - let Some(response) = listener_servfail(request) else { + let Some(response) = engine_helpers::try_build_servfail_response_from_wire(request) else { return; }; send_udp_response(socket, request, &response, peer).await; } -/// 非阻塞发送 SERVFAIL(permit 耗尽路径,避免阻塞接收循环) -fn try_send_udp_servfail(socket: &UdpSocket, request: &[u8], peer: SocketAddr) { - let Some(response) = listener_servfail(request) else { - return; - }; - try_send_udp_response(socket, request, &response, peer); -} - /// 高性能 UDP worker:直接在接收循环中处理请求,避免 spawn 开销 / High-performance UDP worker: process requests directly in receive loop, avoiding spawn overhead async fn run_udp_worker( worker_id: usize, @@ -812,13 +798,12 @@ async fn run_udp_worker( } }); } else { - // 过载保护:permit 耗尽。此路径无预解析数据(parse_quick 已失败), - // 需完整解析请求才能构造严格 SERVFAIL 并回显 question。 - // 解析成本有界(≤4096B 报文),发送为非阻塞(背压丢弃),不阻塞接收循环。 - // Overload: no pre-parsed data here (parse_quick failed), so the strict - // SERVFAIL needs a full parse. Cost is bounded by packet size and the - // send is non-blocking (backpressure drop), so the receive loop is safe. - try_send_udp_servfail(&socket, &packet_bytes, peer); + // 过载保护:permit 耗尽且 parse_quick 已失败(大概率畸形/非标准包)。 + // 直接静默丢弃——过载时应做最便宜的拒绝,对已失败包做完整解析 + // 只会加剧过载,且 hickory 大概率同样失败(解析白做)。 + // Overload: permit exhausted and parse_quick already failed (likely + // malformed/non-standard). Drop silently — overload rejection must be + // cheap, and a full parse of an already-failed packet wastes CPU. } } Err(error) => { @@ -955,7 +940,7 @@ async fn handle_tcp_conn( Ok(Ok(r)) => r, Ok(Err(error)) => { debug!(%error, "TCP request processing error, returning SERVFAIL"); - let Some(response) = listener_servfail(&packet_bytes) else { + let Some(response) = engine_helpers::try_build_servfail_response_from_wire(&packet_bytes) else { return Ok(()); }; response @@ -966,7 +951,7 @@ async fn handle_tcp_conn( upstream_timeout_ms = engine.get_upstream_timeout_ms(), "TCP request timeout after hedge and fallback exhausted" ); - let Some(response) = listener_servfail(&packet_bytes) else { + let Some(response) = engine_helpers::try_build_servfail_response_from_wire(&packet_bytes) else { return Ok(()); }; response @@ -982,7 +967,7 @@ async fn handle_tcp_conn( Ok(Ok(r)) => r, Ok(Err(error)) => { debug!(%error, "TCP request processing error, returning SERVFAIL"); - let Some(response) = listener_servfail(&packet_bytes) else { + let Some(response) = engine_helpers::try_build_servfail_response_from_wire(&packet_bytes) else { return Ok(()); }; response @@ -993,7 +978,7 @@ async fn handle_tcp_conn( upstream_timeout_ms = engine.get_upstream_timeout_ms(), "TCP request timeout" ); - let Some(response) = listener_servfail(&packet_bytes) else { + let Some(response) = engine_helpers::try_build_servfail_response_from_wire(&packet_bytes) else { return Ok(()); }; response @@ -1002,7 +987,7 @@ async fn handle_tcp_conn( } Err(error) => { debug!(%error, "TCP fast-path error, returning SERVFAIL for valid query"); - let Some(response) = listener_servfail(&packet_bytes) else { + let Some(response) = engine_helpers::try_build_servfail_response_from_wire(&packet_bytes) else { return Ok(()); }; response @@ -1010,13 +995,21 @@ async fn handle_tcp_conn( }; if resp.len() <= u16::MAX as usize { - // DNS-over-TCP framing must be written completely; a successful - // write_vectored call may still be partial. - // write_vectored 可能部分写入,必须用 write_all 保证完整帧。 - let mut frame = Vec::with_capacity(2 + resp.len()); - frame.extend_from_slice(&(resp.len() as u16).to_be_bytes()); - frame.extend_from_slice(&resp); - if stream.write_all(&frame).await.is_err() { + // DNS-over-TCP framing must be written completely; write_vectored may be + // partial, so use write_all_vectored: scatter-gather single syscall with + // no frame allocation (write_all would copy prefix+body into a new Vec). + // + // 长度前缀 + 包体用 write_all_vectored 保证完整写入,保持单 syscall 零分配; + // write_all 需要拼帧(一次堆分配 + 拷贝)。 + let len_bytes = (resp.len() as u16).to_be_bytes(); + let mut bufs = [ + std::io::IoSlice::new(&len_bytes), + std::io::IoSlice::new(&resp), + ]; + if tokio_util::io::write_all_vectored(&mut stream, &mut bufs) + .await + .is_err() + { return Ok(()); } } @@ -1026,6 +1019,7 @@ async fn handle_tcp_conn( #[cfg(test)] mod tests { use super::*; + use tokio::io::AsyncWriteExt; use hickory_proto::op::{Message, MessageType, OpCode, Query, ResponseCode}; use hickory_proto::rr::{Name, RecordType}; use hickory_proto::serialize::binary::BinDecodable; diff --git a/src/proto_utils.rs b/src/proto_utils.rs index 3d27b32..3ec3c4b 100644 --- a/src/proto_utils.rs +++ b/src/proto_utils.rs @@ -165,6 +165,14 @@ pub fn parse_quick<'a>(packet: &[u8], buf: &'a mut [u8]) -> Option Date: Sat, 8 Aug 2026 16:26:59 +0800 Subject: [PATCH 5/5] fix(proto): trim truncated UDP responses to record boundary (RFC 1035 4.2.1) hickory's bounded encoder rolls back its write offset but not the backing buffer when a record does not fit, so truncated output could carry partial bytes from the failed record past the (corrected) header counts. Re-walk the emitted sections with valid_message_len and trim to the last complete record boundary. Verified red-green: without the trim, a >512B TXT response emitted 13 bytes of trailing garbage. test(listener): cover overload drop of parse_quick-rejected packets - trailing-data packet passes the header gate but fails parse_quick, so it reaches the overload branch and must stay silent there ci: run binary listener tests in CI test job (were never executed) docs: document inbound validation/SERVFAIL/truncation behavior --- .github/workflows/build.yml | 3 ++ README.md | 6 +++ README.zh-CN.md | 6 +++ src/main.rs | 37 ++++++++++++++--- src/proto_utils.rs | 79 +++++++++++++++++++++++++++++++++++++ 5 files changed, 125 insertions(+), 6 deletions(-) diff --git a/.github/workflows/build.yml b/.github/workflows/build.yml index c177132..81f0f0f 100644 --- a/.github/workflows/build.yml +++ b/.github/workflows/build.yml @@ -232,6 +232,9 @@ jobs: - name: Run unit tests run: cargo test --lib + - name: Run listener integration tests + run: cargo test --bin kixdns + - name: Run DoH integration tests run: cargo test --test doh_integration diff --git a/README.md b/README.md index 9f3314e..fcc3f92 100644 --- a/README.md +++ b/README.md @@ -153,6 +153,12 @@ sudo systemctl enable --now kixdns UDP and TCP listeners are always created from settings.bind_udp and settings.bind_tcp. +Inbound listener behavior: + +- Only standard single-question QUERY messages are processed. Malformed packets, non-QUERY opcodes, and DNS response packets are silently dropped. +- A listener processing error, an overall request timeout, or an exhausted flow-control permit returns SERVFAIL to the client for a valid query. +- UDP responses larger than the client's advertised size (the EDNS payload size, or 512 bytes without EDNS) are truncated at record boundaries with the TC bit set, so the client can retry over TCP. + Inbound DoH is disabled unless settings.bind_doh is set. When it is set, settings.doh_tls_cert and settings.doh_tls_key are required and must point to PEM files. The path defaults to /dns-query and can be changed with settings.doh_path. The inbound DoH handler: diff --git a/README.zh-CN.md b/README.zh-CN.md index 978e800..fdc6afc 100644 --- a/README.zh-CN.md +++ b/README.zh-CN.md @@ -152,6 +152,12 @@ sudo systemctl enable --now kixdns UDP 和 TCP 监听器分别由 settings.bind_udp 和 settings.bind_tcp 创建。 +入站监听器行为: + +- 只处理标准单 question 的 QUERY 消息。畸形包、非 QUERY opcode 和 DNS 响应包会被静默丢弃。 +- 监听器处理错误、整体请求超时或流控 permit 耗尽时,对合法查询返回 SERVFAIL。 +- 大于客户端声明大小(EDNS 载荷大小;无 EDNS 时为 512 字节)的 UDP 响应会在记录边界截断并置 TC 位,客户端可据此改用 TCP 重试。 + 只有设置 settings.bind_doh 后才会启用入站 DoH。启用时必须同时设置 settings.doh_tls_cert 和 settings.doh_tls_key,且路径指向 PEM 文件。路径默认为 /dns-query,可通过 settings.doh_path 修改。 入站 DoH 处理器: diff --git a/src/main.rs b/src/main.rs index d08f52a..a3e0213 100644 --- a/src/main.rs +++ b/src/main.rs @@ -940,7 +940,9 @@ async fn handle_tcp_conn( Ok(Ok(r)) => r, Ok(Err(error)) => { debug!(%error, "TCP request processing error, returning SERVFAIL"); - let Some(response) = engine_helpers::try_build_servfail_response_from_wire(&packet_bytes) else { + let Some(response) = + engine_helpers::try_build_servfail_response_from_wire(&packet_bytes) + else { return Ok(()); }; response @@ -951,7 +953,9 @@ async fn handle_tcp_conn( upstream_timeout_ms = engine.get_upstream_timeout_ms(), "TCP request timeout after hedge and fallback exhausted" ); - let Some(response) = engine_helpers::try_build_servfail_response_from_wire(&packet_bytes) else { + let Some(response) = + engine_helpers::try_build_servfail_response_from_wire(&packet_bytes) + else { return Ok(()); }; response @@ -967,7 +971,9 @@ async fn handle_tcp_conn( Ok(Ok(r)) => r, Ok(Err(error)) => { debug!(%error, "TCP request processing error, returning SERVFAIL"); - let Some(response) = engine_helpers::try_build_servfail_response_from_wire(&packet_bytes) else { + let Some(response) = + engine_helpers::try_build_servfail_response_from_wire(&packet_bytes) + else { return Ok(()); }; response @@ -978,7 +984,9 @@ async fn handle_tcp_conn( upstream_timeout_ms = engine.get_upstream_timeout_ms(), "TCP request timeout" ); - let Some(response) = engine_helpers::try_build_servfail_response_from_wire(&packet_bytes) else { + let Some(response) = + engine_helpers::try_build_servfail_response_from_wire(&packet_bytes) + else { return Ok(()); }; response @@ -987,7 +995,9 @@ async fn handle_tcp_conn( } Err(error) => { debug!(%error, "TCP fast-path error, returning SERVFAIL for valid query"); - let Some(response) = engine_helpers::try_build_servfail_response_from_wire(&packet_bytes) else { + let Some(response) = + engine_helpers::try_build_servfail_response_from_wire(&packet_bytes) + else { return Ok(()); }; response @@ -1019,11 +1029,11 @@ async fn handle_tcp_conn( #[cfg(test)] mod tests { use super::*; - use tokio::io::AsyncWriteExt; use hickory_proto::op::{Message, MessageType, OpCode, Query, ResponseCode}; use hickory_proto::rr::{Name, RecordType}; use hickory_proto::serialize::binary::BinDecodable; use std::str::FromStr; + use tokio::io::AsyncWriteExt; #[ctor::ctor] fn init_crypto() { @@ -1198,6 +1208,21 @@ mod tests { "malformed UDP packets must not receive a reflected response" ); + // A packet that passes the header gate but fails parse_quick (trailing + // data) still reaches the overload branch; it must stay silent there. + let mut trailing = dns_query(); + trailing.extend_from_slice(&[0x00]); + client.send_to(&trailing, server_addr).await.unwrap(); + assert!( + tokio::time::timeout( + Duration::from_millis(100), + client.recv_from(&mut response_buf), + ) + .await + .is_err(), + "overload branch must drop packets that failed parse_quick" + ); + worker.abort(); } diff --git a/src/proto_utils.rs b/src/proto_utils.rs index 3ec3c4b..4233e88 100644 --- a/src/proto_utils.rs +++ b/src/proto_utils.rs @@ -264,6 +264,42 @@ fn skip_name(packet: &[u8], mut pos: usize) -> Option { } } +/// Byte length of the last complete record, or `None` on malformed structure. +/// +/// hickory's bounded encoder rolls back its write offset but not the backing +/// buffer when a record does not fit, so truncated output can carry partial +/// bytes from the failed record past the (corrected) header counts. Re-walking +/// the emitted sections yields the exact valid boundary so the response stays +/// within RFC 1035 §4.2.1 record boundaries. +fn valid_message_len(packet: &[u8]) -> Option { + if packet.len() < 12 { + return None; + } + + let qd_count = u16::from_be_bytes([packet[4], packet[5]]) as usize; + let an_count = u16::from_be_bytes([packet[6], packet[7]]) as usize; + let ns_count = u16::from_be_bytes([packet[8], packet[9]]) as usize; + let ar_count = u16::from_be_bytes([packet[10], packet[11]]) as usize; + + let mut pos = 12usize; + // Questions: NAME + QTYPE + QCLASS. + for _ in 0..qd_count { + pos = skip_name(packet, pos)?; + pos = pos.checked_add(4)?; + } + // Resource records: NAME + TYPE/CLASS/TTL/RDLENGTH + RDATA. + for _ in 0..an_count + ns_count + ar_count { + pos = skip_name(packet, pos)?; + let fixed_end = pos.checked_add(10)?; + if fixed_end > packet.len() { + return None; + } + let rd_len = u16::from_be_bytes([packet[pos + 8], packet[pos + 9]]) as usize; + pos = fixed_end.checked_add(rd_len)?; + } + (pos <= packet.len()).then_some(pos) +} + fn validate_edns_options(packet: &[u8], mut pos: usize, end: usize) -> Option<()> { while pos < end { let header_end = pos.checked_add(4)?; @@ -740,6 +776,13 @@ pub fn truncate_udp_response(request_packet: &[u8], response_packet: &[u8]) -> O response.emit(&mut encoder) }; if encoded.is_ok() && out.len() <= max_payload { + // hickory's rollback does not shrink the backing buffer, so the + // output can carry partial bytes from the record that did not fit. + // Trim to the last complete record boundary from the corrected + // header counts; fall back to the raw output if the walk fails. + if let Some(valid_len) = valid_message_len(&out) { + out.truncate(valid_len); + } return Some(Bytes::from(out)); } } @@ -1043,6 +1086,42 @@ mod tests { assert_eq!(decoded.queries.len(), 1); } + #[test] + fn truncate_udp_response_has_no_trailing_bytes() { + use hickory_proto::op::{MessageType, OpCode, Query}; + use hickory_proto::rr::{Name, RData, Record, RecordType, rdata::TXT}; + use std::str::FromStr; + + let request = txt_query(None); + // No EDNS in the response: the failed answer's partial bytes are not + // overwritten by a later OPT record, so hickory's bounded encoder + // leaves them as trailing garbage past the last complete record. + let name = Name::from_str("cloudflare.com.").unwrap(); + let mut response = Message::new(0x1234, MessageType::Response, OpCode::Query); + response.metadata.recursion_desired = true; + response.metadata.recursion_available = true; + response.add_query(Query::query(name.clone(), RecordType::TXT)); + for index in 0..20 { + let text = format!("verification-{index:02}-{}", "x".repeat(180)); + response.add_answer(Record::from_rdata( + name.clone(), + 300, + RData::TXT(TXT::new(vec![text])), + )); + } + let response = encode_message(&response); + assert!(response.len() > 512); + + let truncated = truncate_udp_response(&request, &response).expect("must truncate"); + // Truncation must respect RFC 1035 §4.2.1 record boundaries with no + // trailing data beyond the last complete record. + assert_eq!( + valid_message_len(&truncated), + Some(truncated.len()), + "truncated response must end exactly at a record boundary" + ); + } + #[test] fn truncate_udp_response_respects_edns_payload_size() { let request = txt_query(Some(1232));