From 5dee96297adf7917d11c7b898b12c3e34fa34be8 Mon Sep 17 00:00:00 2001 From: Ryan Fowler Date: Wed, 5 Aug 2026 10:15:18 -0400 Subject: [PATCH] fix: bound DNS RDATA decoding --- src/dns/doh.rs | 114 +++++++--------- src/dns/inspect/rdata.rs | 111 +++++---------- src/dns/wire.rs | 283 ++++++++++++++++++++++++++++++++++++++- 3 files changed, 354 insertions(+), 154 deletions(-) diff --git a/src/dns/doh.rs b/src/dns/doh.rs index 23b9b083..7f7905fa 100644 --- a/src/dns/doh.rs +++ b/src/dns/doh.rs @@ -366,79 +366,38 @@ fn doh_records_from_wire_response( } fn wire_record_data(packet: &[u8], record: &wire::ResourceRecord<'_>) -> Result { - let offset = record.data_offset; - let len = record.data.len(); - let rdata = record.data; - let value = match (record.typ, len) { - (DNS_TYPE_A, 4) => IpAddr::from([rdata[0], rdata[1], rdata[2], rdata[3]]).to_string(), - (DNS_TYPE_AAAA, 16) => { - let mut octets = [0u8; 16]; - octets.copy_from_slice(rdata); - IpAddr::from(octets).to_string() - } - (wire::TYPE_CNAME | wire::TYPE_NS, _) => { - wire::read_name(packet, offset) - .map_err(|err| DnsError(err.to_string()))? - .0 - } - (wire::TYPE_TXT, _) => parse_txt_rdata(rdata), - (wire::TYPE_MX, 3..) => { - let pref = wire::read_u16(packet, offset).map_err(|err| DnsError(err.to_string()))?; - let name = wire::read_name(packet, offset + 2) - .map_err(|err| DnsError(err.to_string()))? - .0; - format!("{pref} {name}") - } - (wire::TYPE_SOA, _) => parse_soa_rdata(packet, offset)?, - (wire::TYPE_SRV, 7..) => { - let priority = - wire::read_u16(packet, offset).map_err(|err| DnsError(err.to_string()))?; - let weight = - wire::read_u16(packet, offset + 2).map_err(|err| DnsError(err.to_string()))?; - let port = - wire::read_u16(packet, offset + 4).map_err(|err| DnsError(err.to_string()))?; - let target = wire::read_name(packet, offset + 6) - .map_err(|err| DnsError(err.to_string()))? - .0; - format!("{priority} {weight} {port} {target}") - } - _ => generic_rdata(rdata), + let decoded = wire::decode_rdata(packet, record.typ, record.data_offset, record.data.len()) + .map_err(|err| DnsError(err.to_string()))?; + let value = match decoded { + wire::DecodedRdata::Address(address) => address.to_string(), + wire::DecodedRdata::Name(name) => name, + wire::DecodedRdata::Text(text) => text, + wire::DecodedRdata::Mx { + preference, + exchange, + } => format!("{preference} {exchange}"), + wire::DecodedRdata::Soa { + ns, + mailbox, + serial, + refresh, + retry, + expire, + minimum, + } => format!( + "{ns} {mailbox} serial={serial} refresh={refresh} retry={retry} expire={expire} minttl={minimum}" + ), + wire::DecodedRdata::Srv { + priority, + weight, + port, + target, + } => format!("{priority} {weight} {port} {target}"), + wire::DecodedRdata::Raw(raw) => generic_rdata(raw), }; Ok(value) } -fn parse_txt_rdata(raw: &[u8]) -> String { - let mut parts = Vec::new(); - let mut offset = 0; - while offset < raw.len() { - let len = usize::from(raw[offset]); - offset += 1; - if offset + len > raw.len() { - parts.push(String::from_utf8_lossy(&raw[offset - 1..]).into_owned()); - break; - } - parts.push(String::from_utf8_lossy(&raw[offset..offset + len]).into_owned()); - offset += len; - } - parts.join(" ") -} - -fn parse_soa_rdata(packet: &[u8], offset: usize) -> Result { - let (ns, mut next) = - wire::read_name(packet, offset).map_err(|err| DnsError(err.to_string()))?; - let (mbox, next_after_mbox) = - wire::read_name(packet, next).map_err(|err| DnsError(err.to_string()))?; - next = next_after_mbox; - let serial = wire::read_u32(packet, next).map_err(|err| DnsError(err.to_string()))?; - let refresh = wire::read_u32(packet, next + 4).map_err(|err| DnsError(err.to_string()))?; - let retry = wire::read_u32(packet, next + 8).map_err(|err| DnsError(err.to_string()))?; - let expire = wire::read_u32(packet, next + 12).map_err(|err| DnsError(err.to_string()))?; - let min_ttl = wire::read_u32(packet, next + 16).map_err(|err| DnsError(err.to_string()))?; - Ok(format!( - "{ns} {mbox} serial={serial} refresh={refresh} retry={retry} expire={expire} minttl={min_ttl}" - )) -} - fn generic_rdata(raw: &[u8]) -> String { format!(r"\# {} {}", raw.len(), hex_encode(raw)) } @@ -918,6 +877,23 @@ mod tests { response } + #[test] + fn wire_response_rejects_malformed_rdata_before_adjacent_record() { + let query = wire::build_query(0x1234, "example.com", wire::TYPE_MX).unwrap(); + let response = wire_response( + &query, + vec![ + (wire::TYPE_MX, 30, vec![0, 10, 3, b'm']), + (wire::TYPE_MX, 30, vec![0, 10, 0]), + ], + ); + + let err = doh_records_from_wire_response(&response, 0x1234, "example.com", wire::TYPE_MX) + .unwrap_err(); + + assert!(err.to_string().contains("short DNS name label")); + } + #[test] fn wire_response_rejects_unrelated_answer_owner() { let query = wire::build_query(0x1234, "example.com", wire::TYPE_A).unwrap(); diff --git a/src/dns/inspect/rdata.rs b/src/dns/inspect/rdata.rs index 184bf605..1a02f5f4 100644 --- a/src/dns/inspect/rdata.rs +++ b/src/dns/inspect/rdata.rs @@ -38,88 +38,43 @@ pub(super) fn resource_value( offset: usize, len: usize, ) -> Result, FetchError> { - let rdata = &packet[offset..offset + len]; - let value = match typ { - DNS_TYPE_A if len == 4 => { - IpAddr::from([rdata[0], rdata[1], rdata[2], rdata[3]]).to_string() - } - DNS_TYPE_AAAA if len == 16 => { - let mut octets = [0u8; 16]; - octets.copy_from_slice(rdata); - IpAddr::from(octets).to_string() - } - DNS_TYPE_CNAME | DNS_TYPE_NS => { - wire::read_name(packet, offset) - .map_err(|err| FetchError::Message(err.to_string()))? - .0 - } - DNS_TYPE_TXT => parse_txt_rdata(rdata), - DNS_TYPE_MX if len >= 3 => { - let pref = wire::read_u16(packet, offset) - .map_err(|err| FetchError::Message(err.to_string()))?; - let name = wire::read_name(packet, offset + 2) - .map_err(|err| FetchError::Message(err.to_string()))? - .0; - format!("{pref} {name}") - } - DNS_TYPE_SOA => parse_soa_rdata(packet, offset)?, - DNS_TYPE_SRV if len >= 7 => { - let priority = wire::read_u16(packet, offset) - .map_err(|err| FetchError::Message(err.to_string()))?; - let weight = wire::read_u16(packet, offset + 2) - .map_err(|err| FetchError::Message(err.to_string()))?; - let port = wire::read_u16(packet, offset + 4) - .map_err(|err| FetchError::Message(err.to_string()))?; - let target = wire::read_name(packet, offset + 6) - .map_err(|err| FetchError::Message(err.to_string()))? - .0; - format!("{priority} {weight} {port} {target}") - } - DNS_TYPE_SVCB | DNS_TYPE_HTTPS => crate::dns::svcb::format_rdata(rdata) - .unwrap_or_else(|| format!("0x{}", hex_encode(rdata))), - DNS_TYPE_CAA => format_caa(rdata), - _ => return Ok(None), + let decoded = wire::decode_rdata(packet, typ, offset, len) + .map_err(|err| FetchError::Message(err.to_string()))?; + let value = match decoded { + wire::DecodedRdata::Address(address) => address.to_string(), + wire::DecodedRdata::Name(name) => name, + wire::DecodedRdata::Text(text) => text, + wire::DecodedRdata::Mx { + preference, + exchange, + } => format!("{preference} {exchange}"), + wire::DecodedRdata::Soa { + ns, + mailbox, + serial, + refresh, + retry, + expire, + minimum, + } => format!( + "{ns} {mailbox} serial={serial} refresh={refresh} retry={retry} expire={expire} minttl={minimum}" + ), + wire::DecodedRdata::Srv { + priority, + weight, + port, + target, + } => format!("{priority} {weight} {port} {target}"), + wire::DecodedRdata::Raw(raw) => match typ { + DNS_TYPE_SVCB | DNS_TYPE_HTTPS => crate::dns::svcb::format_rdata(raw) + .ok_or_else(|| FetchError::Message("malformed DNS RDATA".to_string()))?, + DNS_TYPE_CAA => format_caa(raw), + _ => return Ok(None), + }, }; Ok(Some(value)) } -fn parse_txt_rdata(raw: &[u8]) -> String { - let mut parts = Vec::new(); - let mut offset = 0; - while offset < raw.len() { - let len = usize::from(raw[offset]); - offset += 1; - if offset + len > raw.len() { - parts.push(String::from_utf8_lossy(&raw[offset - 1..]).into_owned()); - break; - } - parts.push(String::from_utf8_lossy(&raw[offset..offset + len]).into_owned()); - offset += len; - } - parts.join(" ") -} - -fn parse_soa_rdata(packet: &[u8], offset: usize) -> Result { - let (ns, mut next) = - wire::read_name(packet, offset).map_err(|err| FetchError::Message(err.to_string()))?; - let (mbox, next_after_mbox) = - wire::read_name(packet, next).map_err(|err| FetchError::Message(err.to_string()))?; - next = next_after_mbox; - let serial = - wire::read_u32(packet, next).map_err(|err| FetchError::Message(err.to_string()))?; - let refresh = - wire::read_u32(packet, next + 4).map_err(|err| FetchError::Message(err.to_string()))?; - let retry = - wire::read_u32(packet, next + 8).map_err(|err| FetchError::Message(err.to_string()))?; - let expire = - wire::read_u32(packet, next + 12).map_err(|err| FetchError::Message(err.to_string()))?; - let min_ttl = - wire::read_u32(packet, next + 16).map_err(|err| FetchError::Message(err.to_string()))?; - Ok(format!( - "{ns} {mbox} serial={serial} refresh={refresh} retry={retry} expire={expire} minttl={min_ttl}" - )) -} - pub(super) fn type_label(typ: u16) -> String { match typ { DNS_TYPE_A => "A".to_string(), diff --git a/src/dns/wire.rs b/src/dns/wire.rs index a594b58d..87b10860 100644 --- a/src/dns/wire.rs +++ b/src/dns/wire.rs @@ -1,4 +1,5 @@ use std::fmt; +use std::net::IpAddr; pub(crate) const TYPE_A: u16 = 1; pub(crate) const TYPE_NS: u16 = 2; @@ -47,6 +48,183 @@ pub(crate) struct ResourceRecord<'a> { pub(crate) data: &'a [u8], } +pub(crate) enum DecodedRdata<'a> { + Address(IpAddr), + Name(String), + Text(String), + Mx { + preference: u16, + exchange: String, + }, + Soa { + ns: String, + mailbox: String, + serial: u32, + refresh: u32, + retry: u32, + expire: u32, + minimum: u32, + }, + Srv { + priority: u16, + weight: u16, + port: u16, + target: String, + }, + Raw(&'a [u8]), +} + +pub(crate) fn decode_rdata<'a>( + packet: &'a [u8], + typ: u16, + offset: usize, + len: usize, +) -> Result, WireError> { + let end = offset + .checked_add(len) + .filter(|&end| end <= packet.len()) + .ok_or_else(|| WireError("short DNS resource".to_string()))?; + let raw = &packet[offset..end]; + let mut reader = RdataReader { + packet, + pos: offset, + end, + }; + + match typ { + TYPE_A if len == 4 => Ok(DecodedRdata::Address(IpAddr::from([ + raw[0], raw[1], raw[2], raw[3], + ]))), + TYPE_A => Err(malformed_rdata(typ)), + TYPE_AAAA if len == 16 => { + let mut octets = [0u8; 16]; + octets.copy_from_slice(raw); + Ok(DecodedRdata::Address(IpAddr::from(octets))) + } + TYPE_AAAA => Err(malformed_rdata(typ)), + TYPE_CNAME | TYPE_NS => { + let name = reader.read_name()?; + reader.finish(typ)?; + Ok(DecodedRdata::Name(name)) + } + TYPE_TXT => Ok(DecodedRdata::Text(parse_txt_rdata(raw, typ)?)), + TYPE_MX => { + let preference = reader.read_u16()?; + let exchange = reader.read_name()?; + reader.finish(typ)?; + Ok(DecodedRdata::Mx { + preference, + exchange, + }) + } + TYPE_SOA => { + let ns = reader.read_name()?; + let mailbox = reader.read_name()?; + let serial = reader.read_u32()?; + let refresh = reader.read_u32()?; + let retry = reader.read_u32()?; + let expire = reader.read_u32()?; + let minimum = reader.read_u32()?; + reader.finish(typ)?; + Ok(DecodedRdata::Soa { + ns, + mailbox, + serial, + refresh, + retry, + expire, + minimum, + }) + } + TYPE_SRV => { + let priority = reader.read_u16()?; + let weight = reader.read_u16()?; + let port = reader.read_u16()?; + let target = reader.read_name()?; + reader.finish(typ)?; + Ok(DecodedRdata::Srv { + priority, + weight, + port, + target, + }) + } + TYPE_CAA if len >= 2 && usize::from(raw[1]) <= len - 2 => Ok(DecodedRdata::Raw(raw)), + TYPE_CAA => Err(malformed_rdata(typ)), + TYPE_SVCB | TYPE_HTTPS if crate::dns::svcb::parse_rdata(raw).is_some() => { + Ok(DecodedRdata::Raw(raw)) + } + TYPE_SVCB | TYPE_HTTPS => Err(malformed_rdata(typ)), + _ => Ok(DecodedRdata::Raw(raw)), + } +} + +struct RdataReader<'a> { + packet: &'a [u8], + pos: usize, + end: usize, +} + +impl RdataReader<'_> { + fn read_u16(&mut self) -> Result { + let bytes = self.read_bytes(2)?; + Ok(u16::from_be_bytes([bytes[0], bytes[1]])) + } + + fn read_u32(&mut self) -> Result { + let bytes = self.read_bytes(4)?; + Ok(u32::from_be_bytes([bytes[0], bytes[1], bytes[2], bytes[3]])) + } + + fn read_name(&mut self) -> Result { + let (name, next) = read_name_bounded(self.packet, self.pos, self.end)?; + self.pos = next; + Ok(name) + } + + fn read_bytes(&mut self, len: usize) -> Result<&[u8], WireError> { + let end = self + .pos + .checked_add(len) + .filter(|&end| end <= self.end) + .ok_or_else(|| WireError("short DNS RDATA".to_string()))?; + let bytes = &self.packet[self.pos..end]; + self.pos = end; + Ok(bytes) + } + + fn finish(&self, typ: u16) -> Result<(), WireError> { + if self.pos == self.end { + Ok(()) + } else { + Err(malformed_rdata(typ)) + } + } +} + +fn malformed_rdata(typ: u16) -> WireError { + WireError(format!("malformed DNS RDATA for type {typ}")) +} + +fn parse_txt_rdata(raw: &[u8], typ: u16) -> Result { + if raw.is_empty() { + return Err(malformed_rdata(typ)); + } + let mut parts = Vec::new(); + let mut offset = 0; + while offset < raw.len() { + let len = usize::from(raw[offset]); + offset += 1; + let end = offset + .checked_add(len) + .filter(|&end| end <= raw.len()) + .ok_or_else(|| malformed_rdata(typ))?; + parts.push(String::from_utf8_lossy(&raw[offset..end]).into_owned()); + offset = end; + } + Ok(parts.join(" ")) +} + pub(crate) fn build_query(id: u16, host: &str, dns_type: u16) -> Result, WireError> { let mut raw = Vec::with_capacity(512); raw.extend_from_slice(&id.to_be_bytes()); @@ -175,10 +353,11 @@ fn parse_response_inner<'a>( { continue; } - let (target, next) = read_name(raw, record.data_offset)?; - if next != record.data_offset + record.data.len() { - return Err(WireError("invalid DNS CNAME resource".to_string())); - } + let (target, _) = read_name_bounded( + raw, + record.data_offset, + record.data_offset + record.data.len(), + )?; if !reachable.iter().any(|name| names_equal(name, &target)) { reachable.push(target); changed = true; @@ -198,6 +377,17 @@ pub(crate) fn names_equal(left: &str, right: &str) -> bool { } pub(crate) fn read_name(packet: &[u8], offset: usize) -> Result<(String, usize), WireError> { + read_name_bounded(packet, offset, packet.len()) +} + +fn read_name_bounded( + packet: &[u8], + offset: usize, + end: usize, +) -> Result<(String, usize), WireError> { + if offset > end || end > packet.len() { + return Err(WireError("short DNS name".to_string())); + } let mut labels = Vec::new(); let mut pos = offset; let mut next = offset; @@ -205,12 +395,12 @@ pub(crate) fn read_name(packet: &[u8], offset: usize) -> Result<(String, usize), let mut jumps = 0usize; loop { - if pos >= packet.len() { + if pos >= packet.len() || (!jumped && pos >= end) { return Err(WireError("short DNS name".to_string())); } let len = packet[pos]; if len & 0xc0 == 0xc0 { - if pos + 1 >= packet.len() { + if pos + 1 >= end { return Err(WireError("short DNS name pointer".to_string())); } let pointer = usize::from(u16::from_be_bytes([len & 0x3f, packet[pos + 1]])); @@ -236,7 +426,7 @@ pub(crate) fn read_name(packet: &[u8], offset: usize) -> Result<(String, usize), break; } let len = usize::from(len); - if pos + len > packet.len() { + if pos + len > end { return Err(WireError("short DNS name label".to_string())); } labels.push(String::from_utf8_lossy(&packet[pos..pos + len]).into_owned()); @@ -323,6 +513,85 @@ mod tests { assert_eq!(err.to_string(), "mismatched DNS response question"); } + #[test] + fn malformed_rdata_cannot_consume_adjacent_record() { + let cases = [ + (TYPE_CNAME, vec![3, b'x']), + (TYPE_NS, vec![3, b'x']), + (TYPE_MX, vec![0, 10, 3, b'm']), + (TYPE_SRV, vec![0, 1, 0, 2, 0, 3, 3, b's']), + (TYPE_SOA, vec![1, b'n']), + (TYPE_TXT, vec![3, b't']), + ]; + + for (typ, malformed) in cases { + let query = build_query(0x1234, "example.com", typ).unwrap(); + let (_, question_end) = read_name(&query, 12).unwrap(); + let mut response = Vec::new(); + response.extend_from_slice(&0x1234u16.to_be_bytes()); + response.extend_from_slice(&0x8180u16.to_be_bytes()); + response.extend_from_slice(&1u16.to_be_bytes()); + response.extend_from_slice(&2u16.to_be_bytes()); + response.extend_from_slice(&[0, 0, 0, 0]); + response.extend_from_slice(&query[12..question_end + 4]); + let adjacent = match typ { + TYPE_CNAME | TYPE_NS => vec![0], + TYPE_MX => vec![0, 10, 0], + TYPE_SRV => vec![0, 1, 0, 2, 0, 3, 0], + TYPE_SOA => { + let mut data = vec![0, 0]; + data.extend_from_slice(&[0; 20]); + data + } + TYPE_TXT => vec![0], + _ => unreachable!(), + }; + for data in [malformed, adjacent] { + response.extend_from_slice(&[0xc0, 0x0c]); + response.extend_from_slice(&typ.to_be_bytes()); + response.extend_from_slice(&CLASS_IN.to_be_bytes()); + response.extend_from_slice(&30u32.to_be_bytes()); + response.extend_from_slice(&(data.len() as u16).to_be_bytes()); + response.extend_from_slice(&data); + } + + let records = parse_response_without_id(&response, "example.com", typ, CLASS_IN); + if matches!(typ, TYPE_CNAME) { + assert!(records.is_err(), "malformed {typ} RDATA was accepted"); + continue; + } + let records = records.unwrap(); + assert_eq!(records.len(), 2); + assert!( + decode_rdata( + &response, + records[0].typ, + records[0].data_offset, + records[0].data.len() + ) + .is_err(), + "malformed {typ} RDATA was accepted" + ); + assert!( + decode_rdata( + &response, + records[1].typ, + records[1].data_offset, + records[1].data.len() + ) + .is_ok(), + "adjacent {typ} RDATA was rejected" + ); + } + } + + #[test] + fn bounded_rdata_decoder_rejects_extra_bytes_after_name() { + let raw = [0, 0, 0xc0, 0x0c]; + + assert!(decode_rdata(&raw, TYPE_CNAME, 0, raw.len()).is_err()); + } + #[test] fn build_query_advertises_edns0_udp_payload_size() { let query = build_query(0x1234, "example.com", TYPE_A).unwrap();