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 93bae8d..b5e6665 100644 --- a/src/main.rs +++ b/src/main.rs @@ -11,7 +11,7 @@ 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::watcher; @@ -526,6 +526,37 @@ fn create_reuseport_udp_socket(addr: SocketAddr) -> anyhow::Result Option { + engine_helpers::try_build_servfail_response_from_wire(request) +} + +async fn send_udp_datagram( + socket: &UdpSocket, + response: &[u8], + peer: SocketAddr, +) -> std::io::Result { + match socket.try_send_to(response, peer) { + Ok(sent) => Ok(sent), + Err(error) if error.kind() == std::io::ErrorKind::WouldBlock => { + socket.send_to(response, peer).await + } + Err(error) => Err(error), + } +} + +async fn send_udp_response(socket: &UdpSocket, _request: &[u8], response: &[u8], peer: SocketAddr) { + if let Err(error) = send_udp_datagram(socket, response, peer).await { + debug!(%peer, %error, response_len = response.len(), "failed to send UDP response"); + } +} + +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; +} + /// 高性能 UDP worker:直接在接收循环中处理请求,避免 spawn 开销 / High-performance UDP worker: process requests directly in receive loop, avoiding spawn overhead async fn run_udp_worker( worker_id: usize, @@ -588,7 +619,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, @@ -613,11 +644,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。 @@ -634,7 +661,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 { @@ -682,10 +709,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!( @@ -693,9 +721,12 @@ 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 { + send_udp_servfail(&socket, &packet_bytes, peer).await; } } Ok(None) => { @@ -719,10 +750,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!( @@ -730,13 +762,17 @@ async fn run_udp_worker( upstream_timeout_ms = engine.get_upstream_timeout_ms(), "request timeout" ); + send_udp_servfail(&socket, &packet_bytes, peer).await; } } }); + } else { + send_udp_servfail(&socket, &packet_bytes, peer).await; } } - 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; } } } @@ -863,14 +899,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 } } } @@ -881,38 +926,248 @@ 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. + 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(()); } } } } + +#[cfg(test)] +mod tests { + use super::*; + 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] + fn init_crypto() { + let _ = rustls::crypto::aws_lc_rs::default_provider().install_default(); + } + + 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") + } + + 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_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(); + } +}