diff --git a/src/dns/inspect.rs b/src/dns/inspect.rs index 56eb177b..60c18528 100644 --- a/src/dns/inspect.rs +++ b/src/dns/inspect.rs @@ -522,10 +522,12 @@ async fn lookup_udp_record( ) -> Result, QueryError> { let id = dns_query_id(); let raw = wire::build_query(id, host, query_type.dns_type).map_err(QueryError::other)?; + let matcher = wire::ResponseMatcher::new(id, host, query_type.dns_type, DNS_CLASS_IN); let udp_timeout = udp_dns_timeout(timeout.remaining().map_err(QueryError::other)?); - let mut response = crate::dns::transport::query_udp_on_socket(socket, &raw, udp_timeout) - .await - .map_err(QueryError::other)?; + let mut response = + crate::dns::transport::query_udp_on_socket(socket, &raw, &matcher, udp_timeout) + .await + .map_err(QueryError::other)?; let raw_records = match wire::parse_response(&response, id, host, query_type.dns_type, DNS_CLASS_IN) { Ok(records) => records, diff --git a/src/dns/resolver.rs b/src/dns/resolver.rs index d090e422..122a3c67 100644 --- a/src/dns/resolver.rs +++ b/src/dns/resolver.rs @@ -101,8 +101,9 @@ pub(crate) async fn query_udp_type( ) -> Result, ResolverError> { let id = dns_query_id(); let raw = wire::build_query(id, host, dns_type).map_err(resolver_error)?; + let matcher = wire::ResponseMatcher::new(id, host, dns_type, DNS_CLASS_IN); let timeout = udp_dns_timeout(budget.remaining().map_err(resolver_error)?); - let response = crate::dns::transport::query_udp(*server_addr, &raw, timeout) + let response = crate::dns::transport::query_udp(*server_addr, &raw, &matcher, timeout) .await .map_err(resolver_error)?; match wire_records_from_response(&response, id, host, dns_type) { diff --git a/src/dns/transport.rs b/src/dns/transport.rs index e0ee828d..8a8fcfc7 100644 --- a/src/dns/transport.rs +++ b/src/dns/transport.rs @@ -10,6 +10,7 @@ use tokio::net::{TcpStream, UdpSocket}; use tokio_rustls::TlsConnector; use crate::dns::util::udp_dns_timeout; +use crate::dns::wire::ResponseMatcher; use crate::duration::TimeoutBudget; use crate::error::FetchError; @@ -27,10 +28,11 @@ impl std::error::Error for DnsTransportError {} pub(crate) async fn query_udp( server_addr: SocketAddr, query: &[u8], + matcher: &ResponseMatcher, timeout: Duration, ) -> Result, DnsTransportError> { let socket = udp_socket(server_addr).await?; - query_udp_on_socket(&socket, query, timeout).await + query_udp_on_socket(&socket, query, matcher, timeout).await } pub(crate) async fn udp_socket(server_addr: SocketAddr) -> Result { @@ -48,18 +50,27 @@ pub(crate) async fn udp_socket(server_addr: SocketAddr) -> Result Result, DnsTransportError> { + let deadline = tokio::time::Instant::now() + timeout; + let expected_source = socket.peer_addr().map_err(transport_error)?; socket.send(query).await.map_err(transport_error)?; let mut buf = vec![0u8; 4096]; - let n = match tokio::time::timeout(timeout, socket.recv(&mut buf)).await { - Ok(Ok(n)) => n, - Ok(Err(err)) => return Err(transport_error(err)), - Err(_) => return Err(DnsTransportError("DNS lookup timed out".to_string())), - }; - buf.truncate(n); - Ok(buf) + loop { + let (n, source) = match tokio::time::timeout_at(deadline, socket.recv_from(&mut buf)).await + { + Ok(Ok(received)) => received, + Ok(Err(err)) => return Err(transport_error(err)), + Err(_) => return Err(DnsTransportError("DNS lookup timed out".to_string())), + }; + if source != expected_source || !matcher.matches(&buf[..n]) { + continue; + } + buf.truncate(n); + return Ok(buf); + } } pub(crate) async fn query_tcp( @@ -258,6 +269,110 @@ fn transport_error(err: impl ToString) -> DnsTransportError { #[cfg(test)] mod tests { use super::*; + use crate::dns::wire::{self, CLASS_IN, TYPE_A, TYPE_AAAA}; + + fn response_for(query: &[u8], flags: u16) -> Vec { + let mut response = query.to_vec(); + response[2..4].copy_from_slice(&flags.to_be_bytes()); + response + } + + #[tokio::test] + async fn udp_discards_stale_response_from_timed_out_query() { + let server = UdpSocket::bind("127.0.0.1:0").await.unwrap(); + let server_addr = server.local_addr().unwrap(); + let client = udp_socket(server_addr).await.unwrap(); + let first = wire::build_query(0x1001, "example.com", TYPE_A).unwrap(); + let second = wire::build_query(0x1002, "example.com", TYPE_AAAA).unwrap(); + + let server_task = tokio::spawn(async move { + let mut buf = [0u8; 512]; + let (first_len, peer) = server.recv_from(&mut buf).await.unwrap(); + let first_response = response_for(&buf[..first_len], 0x8180); + let (second_len, second_peer) = server.recv_from(&mut buf).await.unwrap(); + assert_eq!(peer, second_peer); + let second_response = response_for(&buf[..second_len], 0x8180); + server.send_to(&first_response, peer).await.unwrap(); + server.send_to(&second_response, peer).await.unwrap(); + }); + + let first_matcher = ResponseMatcher::new(0x1001, "example.com", TYPE_A, CLASS_IN); + let err = query_udp_on_socket(&client, &first, &first_matcher, Duration::from_millis(20)) + .await + .unwrap_err(); + assert_eq!(err.to_string(), "DNS lookup timed out"); + + let second_matcher = ResponseMatcher::new(0x1002, "example.com", TYPE_AAAA, CLASS_IN); + let response = + query_udp_on_socket(&client, &second, &second_matcher, Duration::from_secs(1)) + .await + .unwrap(); + assert_eq!(u16::from_be_bytes([response[0], response[1]]), 0x1002); + server_task.await.unwrap(); + } + + #[tokio::test] + async fn udp_discards_mismatched_response_fields() { + let server = UdpSocket::bind("127.0.0.1:0").await.unwrap(); + let server_addr = server.local_addr().unwrap(); + let query = wire::build_query(0x2001, "example.com", TYPE_A).unwrap(); + let matcher = ResponseMatcher::new(0x2001, "example.com", TYPE_A, CLASS_IN); + let server_query = query.clone(); + + let server_task = tokio::spawn(async move { + let mut buf = [0u8; 512]; + let (_, peer) = server.recv_from(&mut buf).await.unwrap(); + let mut wrong_class = response_for(&server_query, 0x8180); + wrong_class[27..29].copy_from_slice(&2u16.to_be_bytes()); + let packets = vec![ + response_for(&server_query, 0x0100), + response_for(&server_query, 0x8800), + response_for( + &wire::build_query(0x2001, "other.example", TYPE_A).unwrap(), + 0x8180, + ), + response_for( + &wire::build_query(0x2001, "example.com", TYPE_AAAA).unwrap(), + 0x8180, + ), + wrong_class, + response_for(&server_query, 0x8180), + ]; + for packet in packets { + server.send_to(&packet, peer).await.unwrap(); + } + }); + + let response = query_udp(server_addr, &query, &matcher, Duration::from_secs(1)) + .await + .unwrap(); + assert_eq!(response, response_for(&query, 0x8180)); + server_task.await.unwrap(); + } + + #[tokio::test] + async fn udp_discards_response_from_wrong_source() { + let server = UdpSocket::bind("127.0.0.1:0").await.unwrap(); + let attacker = UdpSocket::bind("127.0.0.1:0").await.unwrap(); + let server_addr = server.local_addr().unwrap(); + let query = wire::build_query(0x3001, "example.com", TYPE_A).unwrap(); + let matcher = ResponseMatcher::new(0x3001, "example.com", TYPE_A, CLASS_IN); + + let server_task = tokio::spawn(async move { + let mut buf = [0u8; 512]; + let (len, peer) = server.recv_from(&mut buf).await.unwrap(); + let spoof = response_for(&buf[..len], 0x8183); + attacker.send_to(&spoof, peer).await.unwrap(); + let valid = response_for(&buf[..len], 0x8180); + server.send_to(&valid, peer).await.unwrap(); + }); + + let response = query_udp(server_addr, &query, &matcher, Duration::from_secs(1)) + .await + .unwrap(); + assert_eq!(u16::from_be_bytes([response[2], response[3]]), 0x8180); + server_task.await.unwrap(); + } #[tokio::test] async fn write_read_framed_query_round_trips() { diff --git a/src/dns/wire.rs b/src/dns/wire.rs index 766e7c42..eacf3cb1 100644 --- a/src/dns/wire.rs +++ b/src/dns/wire.rs @@ -249,6 +249,45 @@ pub(crate) fn build_query(id: u16, host: &str, dns_type: u16) -> Result, Ok(raw) } +pub(crate) struct ResponseMatcher { + id: u16, + name: CanonicalName, + typ: u16, + class: u16, +} + +impl ResponseMatcher { + pub(crate) fn new(id: u16, name: &str, typ: u16, class: u16) -> Self { + Self { + id, + name: CanonicalName::from_text(name), + typ, + class, + } + } + + pub(crate) fn matches(&self, raw: &[u8]) -> bool { + if raw.len() < 12 || !read_u16(raw, 0).is_ok_and(|id| id == self.id) { + return false; + } + let Ok(flags) = read_u16(raw, 2) else { + return false; + }; + if flags & FLAG_RESPONSE == 0 + || flags & FLAG_OPCODE != 0 + || !read_u16(raw, 4).is_ok_and(|count| count == 1) + { + return false; + } + let Ok(question_name) = read_parsed_name_bounded(raw, 12, raw.len()) else { + return false; + }; + question_name.canonical == self.name + && read_u16(raw, question_name.next).is_ok_and(|typ| typ == self.typ) + && read_u16(raw, question_name.next + 2).is_ok_and(|class| class == self.class) + } +} + pub(crate) fn parse_response<'a>( raw: &'a [u8], expected_id: u16, diff --git a/tests/network.rs b/tests/network.rs index 41915e31..d866d996 100644 --- a/tests/network.rs +++ b/tests/network.rs @@ -168,7 +168,7 @@ fn ech_dns_discovery_failure_is_reported_and_auto_falls_back() { ]); assert_exit(&required, 1); assert!( - required.stderr.contains("mismatched DNS response ID"), + required.stderr.contains("ServerFailure"), "{}", required.stderr ); diff --git a/tests/support/dns.rs b/tests/support/dns.rs index 14c00560..2961181f 100644 --- a/tests/support/dns.rs +++ b/tests/support/dns.rs @@ -308,9 +308,9 @@ pub(crate) fn start_udp_dns_server_with_failing_https(host: &'static str, ip: Ip dns_response(&buf[..n], question_end, None) }; if name == host && qtype == TYPE_HTTPS { - // Return a response with a mismatched transaction ID so the - // resolver reports a protocol failure instead of no records. - response[0] ^= 1; + // Return a matching SERVFAIL response so the resolver reports + // the DNS error instead of treating the response as no records. + response[3] = 0x82; } let _ = socket.send_to(&response, peer); } diff --git a/tests/websocket.rs b/tests/websocket.rs index f69f4ba0..c0a20c56 100644 --- a/tests/websocket.rs +++ b/tests/websocket.rs @@ -1306,7 +1306,7 @@ fn websocket_ech_discovery_uses_http_error_policy() { ]); assert_exit(&required, 1); assert!( - required.stderr.contains("mismatched DNS response ID"), + required.stderr.contains("ServerFailure"), "{}", required.stderr );