// DNS Resource Record data (RFC 1035 and EDNS(0)).

///|
pub enum RData {
  A(Int)
  AAAA(Int, Int, Int, Int)
  CNAME(String)
  NS(String)
  PTR(String)
  MX(Int, String)
  TXT(Array[String])
  SOA(String, String, UInt, UInt, UInt, UInt, UInt)
  SRV(UInt, UInt, UInt, String)
  // OPT is a pseudo-RR, but modelling its RDATA here makes it possible for a
  // Message to preserve and emit EDNS records in the additional section.
  OPT(Array[OptOption])
  Unknown(Array[Byte])
}

///|
pub fn RData::rtype(self : RData) -> UInt16 {
  match self {
    A(_) => qtype_a
    AAAA(_, _, _, _) => qtype_aaaa
    CNAME(_) => qtype_cname
    NS(_) => qtype_ns
    PTR(_) => qtype_ptr
    MX(_, _) => qtype_mx
    TXT(_) => qtype_txt
    SOA(_, _, _, _, _, _, _) => qtype_soa
    SRV(_, _, _, _) => qtype_srv
    OPT(_) => qtype_opt
    Unknown(_) => 0
  }
}

///|
fn validate_rdata_fields(data : RData) -> Result[Unit, String] {
  match data {
    MX(preference, _) =>
      if preference < 0 || preference > 65535 {
        Err("MX preference must fit in the unsigned 16-bit wire range")
      } else {
        Ok(())
      }
    TXT(strings) =>
      if strings.length() == 0 {
        Err("TXT RDATA must contain at least one character-string")
      } else {
        Ok(())
      }
    SRV(priority, weight, port, _) =>
      if priority > 65535 || weight > 65535 || port > 65535 {
        Err("SRV field exceeds 16-bit wire range")
      } else {
        Ok(())
      }
    _ => Ok(())
  }
}

///|
fn rdata_matches_rtype(rtype : UInt16, data : RData) -> Bool {
  match data {
    A(_) => rtype == qtype_a
    AAAA(_, _, _, _) => rtype == qtype_aaaa
    CNAME(_) => rtype == qtype_cname
    NS(_) => rtype == qtype_ns
    PTR(_) => rtype == qtype_ptr
    MX(_, _) => rtype == qtype_mx
    TXT(_) => rtype == qtype_txt
    SOA(_, _, _, _, _, _, _) => rtype == qtype_soa
    SRV(_, _, _, _) => rtype == qtype_srv
    OPT(_) => rtype == qtype_opt
    // Unknown preserves opaque RDATA only for record types this library does
    // not model. Allowing Unknown for A/MX/etc. would bypass their wire shape.
    Unknown(_) =>
      rtype != qtype_a &&
      rtype != qtype_aaaa &&
      rtype != qtype_cname &&
      rtype != qtype_ns &&
      rtype != qtype_ptr &&
      rtype != qtype_mx &&
      rtype != qtype_txt &&
      rtype != qtype_soa &&
      rtype != qtype_srv &&
      rtype != qtype_opt
  }
}

///|
fn validate_rr_rdata(rtype : UInt16, data : RData) -> Result[Unit, String] {
  if !rdata_matches_rtype(rtype, data) {
    return Err("resource-record TYPE does not match its RDATA variant")
  }
  validate_rdata_fields(data)
}

///|
fn append_u16(out : Array[Byte], value : Int) -> Unit {
  out.push(((value >> 8) & 0xFF).to_byte())
  out.push((value & 0xFF).to_byte())
}

///|
fn append_u32(out : Array[Byte], value : UInt) -> Unit {
  let int_value = value.reinterpret_as_int()
  out.push(((int_value >> 24) & 0xFF).to_byte())
  out.push(((int_value >> 16) & 0xFF).to_byte())
  out.push(((int_value >> 8) & 0xFF).to_byte())
  out.push((int_value & 0xFF).to_byte())
}

///|
fn append_bytes(out : Array[Byte], bytes : Array[Byte]) -> Unit {
  for byte in bytes {
    out.push(byte)
  }
}

///|
fn rdata_name(
  bytes : Array[Byte],
  offset : Int,
  msg_start : Int,
  end : Int,
  description : String,
) -> Result[(String, Int), String] {
  match decode_name_in_rdata(bytes, offset, msg_start, end) {
    Ok(value) => Ok(value)
    Err(err) => Err(description + ": " + err)
  }
}

///|
pub fn decode_rdata(
  rtype : UInt16,
  bytes : Array[Byte],
  offset : Int,
  rdlength : UInt16,
  msg_start : Int,
) -> Result[(RData, Int), String] {
  let length = rdlength.to_int()
  match wire_check_range(bytes, offset, length) {
    Err(err) => return Err("DNS RDATA: " + err)
    Ok(_) => ()
  }
  let end = offset + length
  match rtype {
    1 => {
      if length != 4 {
        return Err("A RDATA must be exactly 4 octets")
      }
      let ip = (bytes[offset].to_int() << 24) |
        (bytes[offset + 1].to_int() << 16) |
        (bytes[offset + 2].to_int() << 8) |
        bytes[offset + 3].to_int()
      Ok((A(ip), end))
    }
    28 => {
      if length != 16 {
        return Err("AAAA RDATA must be exactly 16 octets")
      }
      let read_word = fn(start : Int) -> Int {
        (bytes[start].to_int() << 24) |
        (bytes[start + 1].to_int() << 16) |
        (bytes[start + 2].to_int() << 8) |
        bytes[start + 3].to_int()
      }
      Ok(
        (
          AAAA(
            read_word(offset),
            read_word(offset + 4),
            read_word(offset + 8),
            read_word(offset + 12),
          ),
          end,
        ),
      )
    }
    2 | 5 | 12 => {
      let (name, next) = match
        rdata_name(bytes, offset, msg_start, end, "domain-name RDATA") {
        Ok(value) => value
        Err(err) => return Err(err)
      }
      if next != end {
        return Err("domain-name RDATA has trailing octets")
      }
      let data = if rtype == qtype_ns {
        NS(name)
      } else if rtype == qtype_cname {
        CNAME(name)
      } else {
        PTR(name)
      }
      Ok((data, end))
    }
    15 => {
      if length < 3 {
        return Err("MX RDATA is too short")
      }
      let preference = (bytes[offset].to_int() << 8) |
        bytes[offset + 1].to_int()
      let (exchange, next) = match
        rdata_name(bytes, offset + 2, msg_start, end, "MX exchange") {
        Ok(value) => value
        Err(err) => return Err(err)
      }
      if next != end {
        return Err("MX RDATA has trailing octets")
      }
      Ok((MX(preference, exchange), end))
    }
    16 => {
      if length == 0 {
        return Err("TXT RDATA must contain at least one character-string")
      }
      let strings : Array[String] = Array::new(capacity=4)
      let pos = Ref(offset)
      while pos.val < end {
        let string_len = bytes[pos.val].to_int()
        pos.val = pos.val + 1
        if string_len > end - pos.val {
          return Err("truncated TXT character-string")
        }
        let chars = Array::make(string_len, ' ')
        for i in 0.. {
      let (mname, after_mname) = match
        rdata_name(bytes, offset, msg_start, end, "SOA mname") {
        Ok(value) => value
        Err(err) => return Err(err)
      }
      let (rname, after_rname) = match
        rdata_name(bytes, after_mname, msg_start, end, "SOA rname") {
        Ok(value) => value
        Err(err) => return Err(err)
      }
      if end - after_rname != 20 {
        return Err("SOA RDATA must contain exactly five 32-bit fields")
      }
      let serial = match wire_get_u32(bytes, after_rname) {
        Ok((value, _)) => value
        Err(err) => return Err(err)
      }
      let refresh = match wire_get_u32(bytes, after_rname + 4) {
        Ok((value, _)) => value
        Err(err) => return Err(err)
      }
      let retry = match wire_get_u32(bytes, after_rname + 8) {
        Ok((value, _)) => value
        Err(err) => return Err(err)
      }
      let expire = match wire_get_u32(bytes, after_rname + 12) {
        Ok((value, _)) => value
        Err(err) => return Err(err)
      }
      let minimum = match wire_get_u32(bytes, after_rname + 16) {
        Ok((value, _)) => value
        Err(err) => return Err(err)
      }
      Ok((SOA(mname, rname, serial, refresh, retry, expire, minimum), end))
    }
    33 => {
      if length < 7 {
        return Err("SRV RDATA is too short")
      }
      let priority = ((bytes[offset].to_int() << 8) | bytes[offset + 1].to_int()).reinterpret_as_uint()
      let weight = ((bytes[offset + 2].to_int() << 8) |
      bytes[offset + 3].to_int()).reinterpret_as_uint()
      let port = ((bytes[offset + 4].to_int() << 8) | bytes[offset + 5].to_int()).reinterpret_as_uint()
      let (target, next) = match
        rdata_name(bytes, offset + 6, msg_start, end, "SRV target") {
        Ok(value) => value
        Err(err) => return Err(err)
      }
      if next != end {
        return Err("SRV RDATA has trailing octets")
      }
      Ok((SRV(priority, weight, port, target), end))
    }
    41 => {
      let options = match decode_opt_options(bytes, offset, end) {
        Ok(value) => value
        Err(err) => return Err("OPT RDATA: " + err)
      }
      Ok((OPT(options), end))
    }
    _ =>
      match wire_copy_range(bytes, offset, length) {
        Ok(raw) => Ok((Unknown(raw), end))
        Err(err) => Err(err)
      }
  }
}

///|
pub fn RData::encode_checked(self : RData) -> Result[Array[Byte], String] {
  match validate_rdata_fields(self) {
    Ok(_) => ()
    Err(error) => return Err(error)
  }
  let out : Array[Byte] = Array::new(capacity=32)
  match self {
    A(ip) => {
      out.push(((ip >> 24) & 0xFF).to_byte())
      out.push(((ip >> 16) & 0xFF).to_byte())
      out.push(((ip >> 8) & 0xFF).to_byte())
      out.push((ip & 0xFF).to_byte())
    }
    AAAA(w1, w2, w3, w4) => {
      append_u32(out, w1.reinterpret_as_uint())
      append_u32(out, w2.reinterpret_as_uint())
      append_u32(out, w3.reinterpret_as_uint())
      append_u32(out, w4.reinterpret_as_uint())
    }
    CNAME(name) | NS(name) | PTR(name) => {
      let name_bytes = match encode_name_checked(name) {
        Ok(bytes) => bytes
        Err(err) => return Err(err)
      }
      append_bytes(out, name_bytes)
    }
    MX(preference, exchange) => {
      append_u16(out, preference)
      let name_bytes = match encode_name_checked(exchange) {
        Ok(bytes) => bytes
        Err(err) => return Err(err)
      }
      append_bytes(out, name_bytes)
    }
    TXT(strings) =>
      for string in strings {
        if string.length() > 255 {
          return Err("TXT character-string exceeds 255 octets")
        }
        out.push(string.length().to_byte())
        for i in 0.. 0xFF {
            return Err("TXT character-string contains a non-octet character")
          }
          out.push(string[i].to_byte())
        }
      }
    SOA(mname, rname, serial, refresh, retry, expire, minimum) => {
      let mname_bytes = match encode_name_checked(mname) {
        Ok(bytes) => bytes
        Err(err) => return Err(err)
      }
      let rname_bytes = match encode_name_checked(rname) {
        Ok(bytes) => bytes
        Err(err) => return Err(err)
      }
      append_bytes(out, mname_bytes)
      append_bytes(out, rname_bytes)
      append_u32(out, serial)
      append_u32(out, refresh)
      append_u32(out, retry)
      append_u32(out, expire)
      append_u32(out, minimum)
    }
    SRV(priority, weight, port, target) => {
      if priority > 65535 || weight > 65535 || port > 65535 {
        return Err("SRV field exceeds 16-bit wire range")
      }
      append_u16(out, priority.reinterpret_as_int())
      append_u16(out, weight.reinterpret_as_int())
      append_u16(out, port.reinterpret_as_int())
      let target_bytes = match encode_name_checked(target) {
        Ok(bytes) => bytes
        Err(err) => return Err(err)
      }
      append_bytes(out, target_bytes)
    }
    OPT(options) => {
      let option_bytes = match encode_opt_options(options) {
        Ok(bytes) => bytes
        Err(err) => return Err(err)
      }
      append_bytes(out, option_bytes)
    }
    Unknown(raw) => append_bytes(out, raw)
  }
  if out.length() > 65535 {
    Err("RDATA exceeds 65535 octets")
  } else {
    Ok(out)
  }
}

///|
pub fn RData::encode(self : RData) -> Array[Byte] {
  match self.encode_checked() {
    Ok(bytes) => bytes
    Err(error) => abort(error)
  }
}

// Writes RDATA into a message builder, using the same compression table as
// owners and questions.  RDLENGTH is filled by the RR writer after this call.

///|
fn write_rdata_compressed(
  out : Array[Byte],
  offsets : Map[String, Int],
  data : RData,
) -> Result[Unit, String] {
  match validate_rdata_fields(data) {
    Ok(_) => ()
    Err(error) => return Err(error)
  }
  match data {
    CNAME(name) | NS(name) | PTR(name) =>
      write_name_compressed(out, offsets, name)
    MX(preference, exchange) => {
      append_u16(out, preference)
      write_name_compressed(out, offsets, exchange)
    }
    SOA(mname, rname, serial, refresh, retry, expire, minimum) => {
      match write_name_compressed(out, offsets, mname) {
        Err(err) => return Err(err)
        Ok(_) => ()
      }
      match write_name_compressed(out, offsets, rname) {
        Err(err) => return Err(err)
        Ok(_) => ()
      }
      append_u32(out, serial)
      append_u32(out, refresh)
      append_u32(out, retry)
      append_u32(out, expire)
      append_u32(out, minimum)
      Ok(())
    }
    SRV(priority, weight, port, target) => {
      append_u16(out, priority.reinterpret_as_int())
      append_u16(out, weight.reinterpret_as_int())
      append_u16(out, port.reinterpret_as_int())
      // RFC 2782 requires the SRV Target field to be emitted without DNS name
      // compression even when an identical suffix already has an offset.
      let target_bytes = match encode_name_checked(target) {
        Ok(value) => value
        Err(error) => return Err(error)
      }
      append_bytes(out, target_bytes)
      Ok(())
    }
    _ => {
      let bytes = match data.encode_checked() {
        Ok(value) => value
        Err(err) => return Err(err)
      }
      append_bytes(out, bytes)
      Ok(())
    }
  }
}