// 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),
  }
}