use crate::bail;
use crate::io::{CursorExt, DNSReadExt, SeekExt};
use crate::limits;
use crate::types::Record;
use crate::types::*;
use byteorder::{ReadBytesExt, BE};
use num_traits::FromPrimitive;
use rand::RngExt;
use std::convert::TryFrom;
use std::io;
use std::io::BufRead;
use std::io::Cursor;
#[derive(Copy, Clone, PartialEq)]
enum RecordSection {
Answers,
Authorities,
Additionals,
}
pub(crate) struct MessageParser<'a> {
cur: Cursor<&'a [u8]>,
m: Message,
}
impl<'a> MessageParser<'a> {
fn new(buf: &[u8]) -> MessageParser<'_> {
MessageParser {
cur: Cursor::new(buf),
m: Message::default(),
}
}
fn parse(mut self) -> io::Result<Message> {
self.m.id = self.cur.read_u16::<BE>()?;
let b = self.cur.read_u8()?;
self.m.qr = QR::from(0b1000_0000 & b != 0);
let opcode = (0b0111_1000 & b) >> 3;
self.m.aa = (0b0000_0100 & b) != 0;
self.m.tc = (0b0000_0010 & b) != 0;
self.m.rd = (0b0000_0001 & b) != 0;
self.m.opcode = match FromPrimitive::from_u8(opcode) {
Some(t) => t,
None => bail!(InvalidData, "invalid Opcode({})", opcode),
};
let b = self.cur.read_u8()?;
self.m.ra = (0b1000_0000 & b) != 0;
self.m.z = (0b0100_0000 & b) != 0; self.m.ad = (0b0010_0000 & b) != 0;
self.m.cd = (0b0001_0000 & b) != 0;
let rcode = 0b0000_1111 & b;
self.m.rcode = match FromPrimitive::from_u8(rcode) {
Some(t) => t,
None => bail!(InvalidData, "invalid RCode({})", opcode),
};
let qd_count = self.cur.read_u16::<BE>()?;
let an_count = self.cur.read_u16::<BE>()?;
let ns_count = self.cur.read_u16::<BE>()?;
let ar_count = self.cur.read_u16::<BE>()?;
self.read_questions(qd_count)?;
self.read_records(an_count, RecordSection::Answers)?;
self.read_records(ns_count, RecordSection::Authorities)?;
self.read_records(ar_count, RecordSection::Additionals)?;
if self.cur.remaining()? > 0 {
bail!(
Other,
"finished parsing with {} bytes left over",
self.cur.remaining()?
);
}
Ok(self.m)
}
fn read_questions(&mut self, count: u16) -> io::Result<()> {
self.m.questions.reserve_exact(count.into());
for _ in 0..count {
let name = self.cur.read_qname()?;
let r#type = self.cur.read_type()?;
let class = self.cur.read_class()?;
self.m.questions.push(Question {
name,
r#type,
class,
});
}
Ok(())
}
fn read_records(&mut self, count: u16, section: RecordSection) -> io::Result<()> {
let records = match section {
RecordSection::Answers => &mut self.m.answers,
RecordSection::Authorities => &mut self.m.authoritys,
RecordSection::Additionals => &mut self.m.additionals,
};
records.reserve_exact(count.into());
for _ in 0..count {
let name = self.cur.read_qname()?;
let r#type = self.cur.read_type()?;
if section == RecordSection::Additionals && r#type == Type::OPT {
if self.m.extension.is_some() {
bail!(
InvalidData,
"multiple EDNS(0) extensions. Expected only one."
);
}
let ext = Extension::parse_internal(&mut self.cur, name, r#type)?;
self.m.extension = Some(ext);
} else {
let class = self.cur.read_class()?;
let record = Record::parse(&mut self.cur, name, r#type, class)?;
records.push(record);
}
}
Ok(())
}
}
impl Default for Message {
fn default() -> Self {
Message {
id: Message::random_id(),
rd: true,
tc: false,
aa: false,
opcode: Opcode::Query,
qr: QR::Query,
rcode: Rcode::NoError,
cd: false,
ad: true,
z: false,
ra: false,
questions: Vec::default(),
answers: Vec::default(),
authoritys: Vec::default(),
additionals: Vec::default(),
extension: None,
stats: None,
}
}
}
impl Message {
pub fn random_id() -> u16 {
rand::rng().random()
}
pub fn from_slice(buf: &[u8]) -> io::Result<Message> {
MessageParser::new(buf).parse()
}
fn normalise_domain(&mut self, domain: &str) -> Result<String, idna::Errors> {
let ascii = idna::domain_to_ascii(domain)?;
let (mut unicode, result) = idna::domain_to_unicode(&ascii);
match result {
Ok(_) => {
if !unicode.ends_with('.') {
unicode.push('.')
}
Ok(unicode)
}
Err(errors) => Err(errors),
}
}
#[deprecated(note = "use try_add_question to handle invalid domains")]
pub fn add_question(&mut self, domain: &str, r#type: Type, class: Class) {
self.try_add_question(domain, r#type, class)
.expect("invalid domain");
}
pub fn try_add_question(
&mut self,
domain: &str,
r#type: Type,
class: Class,
) -> Result<(), crate::Error> {
let domain = self
.normalise_domain(domain)
.map_err(|error| crate::Error::InvalidArgument(error.to_string()))?;
let ascii_domain = idna::domain_to_ascii(&domain)
.map_err(|error| crate::Error::InvalidArgument(error.to_string()))?;
limits::validate_ascii_name(&ascii_domain)
.map_err(|error| crate::Error::InvalidArgument(error.to_string()))?;
let q = Question {
name: domain,
r#type,
class,
};
self.questions.push(q);
Ok(())
}
pub fn set_extension(&mut self, ext: Extension) {
self.extension = Some(ext);
}
#[deprecated(
note = "use set_extension because a message can contain at most one EDNS(0) extension"
)]
pub fn add_extension(&mut self, ext: Extension) {
self.set_extension(ext);
}
pub fn to_vec(&self) -> io::Result<Vec<u8>> {
let mut req = Vec::<u8>::with_capacity(512);
self.append_to_vec(&mut req)?;
Ok(req)
}
pub fn append_to_vec(&self, buf: &mut Vec<u8>) -> io::Result<()> {
buf.extend_from_slice(&self.id.to_be_bytes());
let mut b = 0_u8;
b |= if bool::from(self.qr) { 0b1000_0000 } else { 0 };
b |= ((self.opcode as u8) << 3) & 0b0111_1000;
b |= if self.aa { 0b0000_0100 } else { 0 };
b |= if self.tc { 0b0000_0010 } else { 0 };
b |= if self.rd { 0b0000_0001 } else { 0 };
buf.push(b);
let mut b = 0_u8;
b |= if self.ra { 0b1000_0000 } else { 0 };
b |= if self.z { 0b0100_0000 } else { 0 };
b |= if self.ad { 0b0010_0000 } else { 0 };
b |= if self.cd { 0b0001_0000 } else { 0 };
b |= (self.rcode as u8) & 0b0000_1111;
buf.push(b);
limits::validate_section_count(self.questions.len())?;
limits::validate_section_count(self.answers.len())?;
limits::validate_section_count(self.authoritys.len())?;
limits::validate_section_count(self.additionals.len() + self.extension.is_some() as usize)?;
let ar_count = self.additionals.len() as u16 + self.extension.is_some() as u16;
buf.extend_from_slice(&(self.questions.len() as u16).to_be_bytes());
buf.extend_from_slice(&(self.answers.len() as u16).to_be_bytes());
buf.extend_from_slice(&(self.authoritys.len() as u16).to_be_bytes());
buf.extend_from_slice(&ar_count.to_be_bytes());
for question in &self.questions {
question.append_to_vec(buf)?;
}
for record in self
.answers
.iter()
.chain(self.authoritys.iter())
.chain(self.additionals.iter())
{
record.append_to_vec(buf)?;
}
if let Some(e) = &self.extension {
e.append_to_vec(buf)?
}
Ok(())
}
pub(crate) fn append_qname_to_vec(buf: &mut Vec<u8>, domain: &str) -> io::Result<()> {
let domain = match idna::domain_to_ascii(domain) {
Err(e) => {
bail!(InvalidData, "invalid dns name '{0}': {1}", domain, e);
}
Ok(domain) => domain,
};
if !domain.is_empty() && domain != "." {
limits::validate_ascii_name(&domain)?;
for label in domain.split_terminator('.') {
buf.push(label.len() as u8);
buf.extend_from_slice(label.as_bytes());
}
}
buf.push(0);
Ok(())
}
}
impl Question {
pub fn append_to_vec(&self, buf: &mut Vec<u8>) -> io::Result<()> {
Message::append_qname_to_vec(buf, &self.name)?;
buf.extend_from_slice(&(self.r#type as u16).to_be_bytes());
buf.extend_from_slice(&(self.class as u16).to_be_bytes());
Ok(())
}
}
impl TryFrom<&[u8]> for Message {
type Error = io::Error;
fn try_from(buf: &[u8]) -> Result<Self, Self::Error> {
Self::from_slice(buf)
}
}
impl Extension {
#[deprecated(note = "this low-level cursor parser is retained for compatibility")]
pub fn parse(cur: &mut Cursor<&[u8]>, domain: String, r#type: Type) -> io::Result<Extension> {
Self::parse_internal(cur, domain, r#type)
}
pub(crate) fn parse_internal(
cur: &mut Cursor<&[u8]>,
domain: String,
r#type: Type,
) -> io::Result<Extension> {
if r#type != Type::OPT {
bail!(InvalidInput, "expected EDNS(0) OPT record");
}
if domain != "." {
bail!(
InvalidData,
"expected root domain for EDNS(0) extension, got '{}'",
domain
);
}
let payload_size = cur.read_u16::<BE>()?;
let extend_rcode = cur.read_u8()?;
let version = cur.read_u8()?;
let b = cur.read_u8()?;
let dnssec_ok = b & 0b1000_0000 == 0b1000_0000;
let _z = cur.read_u8()?;
let rd_len = cur.read_u16::<BE>()?;
if cur.remaining()? < u64::from(rd_len) {
bail!(InvalidData, "EDNS(0) data exceeds the remaining message");
}
let rd_len = usize::from(rd_len);
let pos = usize::try_from(cur.position()).map_err(|_| {
io::Error::new(
io::ErrorKind::InvalidData,
"EDNS(0) cursor position is invalid",
)
})?;
let end = pos
.checked_add(rd_len)
.ok_or_else(|| io::Error::new(io::ErrorKind::InvalidData, "EDNS(0) length overflow"))?;
let mut option_cur = cur.sub_cursor(pos, end)?;
let mut options = Vec::new();
while option_cur.position() < rd_len as u64 {
options.push(EdnsOption::parse(&mut option_cur)?);
}
cur.consume(rd_len);
Ok(Extension {
payload_size,
extend_rcode,
version,
dnssec_ok,
options,
})
}
pub fn append_to_vec(&self, buf: &mut Vec<u8>) -> io::Result<()> {
buf.push(0); buf.extend_from_slice(&(Type::OPT as u16).to_be_bytes()); buf.extend_from_slice(&self.payload_size.to_be_bytes());
buf.push(self.extend_rcode); buf.push(self.version);
let mut b = 0_u8;
b |= if self.dnssec_ok { 0b1000_0000 } else { 0 };
buf.push(b);
buf.push(0);
let mut option_data = Vec::new();
for option in &self.options {
option.append_to_vec(&mut option_data)?;
}
let option_data_len = u16::try_from(option_data.len()).map_err(|_| {
io::Error::new(
io::ErrorKind::InvalidData,
"EDNS(0) option data is too long",
)
})?;
buf.extend_from_slice(&option_data_len.to_be_bytes());
buf.extend_from_slice(&option_data);
Ok(())
}
#[deprecated(note = "use append_to_vec")]
pub fn write(&self, buf: &mut Vec<u8>) -> io::Result<()> {
self.append_to_vec(buf)
}
}
#[cfg(test)]
mod tests {
use super::Message;
use crate::{
Class, EdnsOption, Extension, Question, Record, Resource, Type, MX, SOA, SRV, TXT,
};
use std::convert::TryFrom;
#[test]
fn truncated_dns_messages_return_errors() {
let cases = [
vec![],
vec![0, 0, 0, 0, 0, 1, 0, 0, 0, 0, 0, 0],
vec![0, 0, 0, 0, 0, 1, 0, 0, 0, 0, 0, 0, 3, b'a'],
];
for input in cases {
assert!(
Message::from_slice(&input).is_err(),
"accepted truncated input"
);
}
}
#[test]
fn truncated_edns_options_return_error() {
let input = [
0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 0, 1, 0, 0, 41, 0x10, 0, 0, 0, 0, 0, 1, ];
assert!(Message::from_slice(&input).is_err());
}
#[test]
fn to_vec_encodes_edns_options() {
let mut message = Message {
id: 0x1234,
..Default::default()
};
message.set_extension(
Extension::default()
.with_option(EdnsOption::nsid(b"abc".to_vec()))
.with_option(EdnsOption::client_subnet(
"192.0.2.129".parse().unwrap(),
24,
0,
)),
);
let encoded = message.to_vec().expect("EDNS options should encode");
assert_eq!(
encoded,
vec![
0x12, 0x34, 0x01, 0x20, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x01, 0x00, 0x00, 0x29, 0x10, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x12, 0x00, 0x03, 0x00, 0x03, b'a', b'b', b'c', 0x00, 0x08, 0x00, 0x07, 0x00, 0x01, 24, 0, 192, 0, 2, ]
);
let decoded = Message::from_slice(&encoded).expect("encoded message should parse");
assert_eq!(
decoded.extension.expect("extension should parse").options,
message.extension.unwrap().options
);
}
#[test]
fn append_to_vec_appends_encoded_message() {
let mut message = Message {
id: 0x1234,
..Default::default()
};
message
.try_add_question("example.com", Type::A, Class::Internet)
.expect("question should be valid");
let encoded = message.to_vec().expect("message should encode");
let mut buf = vec![0xaa, 0xbb];
message
.append_to_vec(&mut buf)
.expect("message should append to buffer");
assert_eq!(&buf[..2], &[0xaa, 0xbb]);
assert_eq!(&buf[2..], encoded);
let decoded = Message::try_from(&buf[2..]).expect("appended message should decode");
assert_eq!(decoded.questions, message.questions);
}
#[test]
fn question_append_to_vec_encodes_question() {
let question = Question {
name: "example.com.".to_string(),
r#type: Type::A,
class: Class::Internet,
};
let mut buf = Vec::new();
question
.append_to_vec(&mut buf)
.expect("question should encode");
assert_eq!(
buf,
vec![7, b'e', b'x', b'a', b'm', b'p', b'l', b'e', 3, b'c', b'o', b'm', 0, 0, 1, 0, 1,]
);
}
#[test]
fn to_vec_encodes_extension_options() {
let mut message = Message {
id: 0x1234,
..Default::default()
};
message.set_extension(Extension::default().with_option(EdnsOption::nsid(Vec::new())));
let encoded = message
.to_vec()
.expect("EDNS extension options should encode");
assert_eq!(&encoded[10..12], &[0x00, 0x01]);
assert_eq!(
&encoded[12..],
&[0, 0, 41, 0x10, 0, 0, 0, 0, 0, 0, 4, 0, 3, 0, 0]
);
}
#[test]
fn to_vec_rejects_malformed_edns_options() {
let mut message = Message::default();
message.set_extension(Extension::default().with_option(EdnsOption::client_subnet(
"192.0.2.1".parse().unwrap(),
33,
0,
)));
assert!(message.to_vec().is_err());
}
#[test]
fn truncated_record_data_returns_error() {
let input = [
0, 0, 0, 0, 0, 0, 0, 1, 0, 0, 1, 0, 0, 0, 0, 0, 0, 4, 192, 0, 2, ];
assert!(Message::from_slice(&input).is_err());
}
#[test]
fn try_add_question_rejects_invalid_domain() {
let mut message = Message::default();
let invalid_domain = format!("{}.example.com", "a".repeat(64));
assert!(message
.try_add_question(&invalid_domain, Type::A, Class::Internet)
.is_err());
assert!(message.questions.is_empty());
}
#[test]
fn to_vec_rejects_oversized_domain_name() {
let domain = vec!["a".repeat(63); 4].join(".");
let mut message = Message::default();
message.questions.push(Question {
name: domain,
r#type: Type::A,
class: Class::Internet,
});
assert!(message.to_vec().is_err());
}
#[test]
fn to_vec_round_trips_records() {
let mut message = Message::default();
message.answers.push(Record::new(
"example.com.",
Class::Internet,
std::time::Duration::from_secs(60),
Resource::A("192.0.2.1".parse().unwrap()),
));
let encoded = message.to_vec().expect("records should be encoded");
let decoded = Message::from_slice(&encoded).expect("encoded records should parse");
assert_eq!(decoded.answers, message.answers);
}
#[test]
fn to_vec_round_trips_supported_resources() {
let resources = vec![
Resource::A("192.0.2.1".parse().unwrap()),
Resource::AAAA("2001:db8::1".parse().unwrap()),
Resource::CNAME("target.example.com.".to_string()),
Resource::NS("ns.example.com.".to_string()),
Resource::PTR("ptr.example.com.".to_string()),
Resource::TXT(TXT::from("text")),
Resource::SPF(TXT::from("v=spf1 -all")),
Resource::MX(MX {
preference: 10,
exchange: "mail.example.com.".to_string(),
}),
Resource::SOA(SOA {
mname: "ns.example.com.".to_string(),
rname: "admin@example.com".to_string(),
serial: 1,
refresh: std::time::Duration::from_secs(2),
retry: std::time::Duration::from_secs(3),
expire: std::time::Duration::from_secs(4),
minimum: std::time::Duration::from_secs(5),
}),
Resource::SRV(SRV {
priority: 1,
weight: 2,
port: 443,
name: "service.example.com.".to_string(),
}),
];
for resource in resources {
let mut message = Message::default();
message.answers.push(Record::new(
"example.com.",
Class::Internet,
std::time::Duration::from_secs(60),
resource.clone(),
));
let encoded = message.to_vec().expect("resource should be encoded");
let decoded = Message::from_slice(&encoded).expect("encoded resource should parse");
let decoded_resource = match &decoded.answers[0].resource {
Resource::SOA(soa) => Resource::SOA(SOA {
rname: soa.rname.trim_end_matches('.').to_string(),
..soa.clone()
}),
resource => resource.clone(),
};
assert_eq!(decoded_resource, resource);
}
}
}