// Complete DNS messages: section framing, compression-aware encoding, and
// strict decode error propagation.
///|
pub struct Message {
header : Header
questions : Array[Question]
answers : Array[RR]
authorities : Array[RR]
additionals : Array[RR]
}
///|
fn validate_section_count(
length : Int,
section : String,
) -> Result[UInt16, String] {
if length > max_section_records {
Err(
section +
" exceeds the symmetric codec limit of " +
max_section_records.to_string() +
" entries",
)
} else {
Ok(length.to_uint16())
}
}
///|
fn validate_opt_locations(message : Message) -> Result[Unit, String] {
for record in message.answers {
if record.rtype == qtype_opt {
return Err("OPT record is only valid in the additional section")
}
}
for record in message.authorities {
if record.rtype == qtype_opt {
return Err("OPT record is only valid in the additional section")
}
}
let count = Ref(0)
for record in message.additionals {
if record.rtype == qtype_opt {
count.val = count.val + 1
if count.val > 1 {
return Err("DNS message contains more than one OPT record")
}
match opt_from_rr(record) {
Ok(_) => ()
Err(err) => return Err(err)
}
}
}
Ok(())
}
///|
fn write_question_compressed(
out : Array[Byte],
offsets : Map[String, Int],
question : Question,
) -> Result[Unit, String] {
match write_name_compressed(out, offsets, question.name) {
Err(err) => Err(err)
Ok(_) => {
append_u16(out, question.qtype.to_int())
append_u16(out, question.qclass.to_int())
Ok(())
}
}
}
///|
fn write_rr_compressed(
out : Array[Byte],
offsets : Map[String, Int],
record : RR,
) -> Result[Unit, String] {
match validate_rr_rdata(record.rtype, record.rdata) {
Ok(_) => ()
Err(err) => return Err("DNS RR RDATA: " + err)
}
if record.rtype == qtype_opt && record.name != "" {
return Err("OPT owner name must be the root domain")
}
match write_name_compressed(out, offsets, record.name) {
Err(err) => return Err("DNS RR owner: " + err)
Ok(_) => ()
}
append_u16(out, record.rtype.to_int())
append_u16(out, record.rclass.to_int())
append_u32(out, record.ttl)
let length_offset = out.length()
out.push(0)
out.push(0)
let rdata_offset = out.length()
match write_rdata_compressed(out, offsets, record.rdata) {
Err(err) => return Err("DNS RR RDATA: " + err)
Ok(_) => ()
}
let rdata_length = out.length() - rdata_offset
if rdata_length > 65535 {
return Err("DNS RR RDATA exceeds 65535 octets")
}
out[length_offset] = ((rdata_length >> 8) & 0xFF).to_byte()
out[length_offset + 1] = (rdata_length & 0xFF).to_byte()
Ok(())
}
///|
pub fn Message::encode_checked(self : Message) -> Result[Array[Byte], String] {
let qdcount = match
validate_section_count(self.questions.length(), "question section") {
Ok(value) => value
Err(err) => return Err(err)
}
let ancount = match
validate_section_count(self.answers.length(), "answer section") {
Ok(value) => value
Err(err) => return Err(err)
}
let nscount = match
validate_section_count(self.authorities.length(), "authority section") {
Ok(value) => value
Err(err) => return Err(err)
}
let arcount = match
validate_section_count(self.additionals.length(), "additional section") {
Ok(value) => value
Err(err) => return Err(err)
}
match validate_opt_locations(self) {
Err(err) => return Err(err)
Ok(_) => ()
}
let header = { ..self.header, qdcount, ancount, nscount, arcount }
let out : Array[Byte] = Array::new(capacity=512)
append_bytes(out, header.encode())
let offsets : Map[String, Int] = Map([])
for question in self.questions {
match write_question_compressed(out, offsets, question) {
Err(err) => return Err(err)
Ok(_) => ()
}
}
for record in self.answers {
match write_rr_compressed(out, offsets, record) {
Err(err) => return Err(err)
Ok(_) => ()
}
}
for record in self.authorities {
match write_rr_compressed(out, offsets, record) {
Err(err) => return Err(err)
Ok(_) => ()
}
}
for record in self.additionals {
match write_rr_compressed(out, offsets, record) {
Err(err) => return Err(err)
Ok(_) => ()
}
}
if out.length() > max_dns_message_len {
Err("DNS message exceeds 65535 octets")
} else {
Ok(out)
}
}
///|
pub fn Message::encode(self : Message) -> Array[Byte] {
match self.encode_checked() {
Ok(bytes) => bytes
Err(error) => abort(error)
}
}
///|
pub fn decode_message(bytes : Array[Byte]) -> Result[Message, String] {
if bytes.length() < dns_header_size {
return Err("DNS message is shorter than its 12-octet header")
}
if bytes.length() > max_dns_message_len {
return Err("DNS message exceeds 65535-octet codec limit")
}
let (header, after_header) = match decode_header(bytes) {
Ok(value) => value
Err(err) => return Err(err)
}
let (questions, after_questions) = match
decode_questions(bytes, after_header, header.qdcount, 0) {
Ok(value) => value
Err(err) => return Err(err)
}
let (answers, after_answers) = match
decode_rrs(bytes, after_questions, header.ancount, 0) {
Ok(value) => value
Err(err) => return Err(err)
}
let (authorities, after_authorities) = match
decode_rrs(bytes, after_answers, header.nscount, 0) {
Ok(value) => value
Err(err) => return Err(err)
}
let (additionals, end) = match
decode_rrs(bytes, after_authorities, header.arcount, 0) {
Ok(value) => value
Err(err) => return Err(err)
}
if end != bytes.length() {
return Err("trailing octets after DNS sections")
}
let message = { header, questions, answers, authorities, additionals }
match validate_opt_locations(message) {
Ok(_) => Ok(message)
Err(err) => Err(err)
}
}
///|
pub fn build_query(
id : UInt16,
name : String,
qtype : UInt16,
recurse : Bool,
) -> Message {
let flags = if recurse { flag_rd } else { 0 }
{
header: { id, flags, qdcount: 1, ancount: 0, nscount: 0, arcount: 0 },
questions: [{ name, qtype, qclass: qclass_in }],
answers: Array::new(capacity=0),
authorities: Array::new(capacity=0),
additionals: Array::new(capacity=0),
}
}