///|
const DNS_HEADER_LENGTH : Int = 12

///|
const DNS_TYPE_A : UInt16 = 1

///|
const DNS_TYPE_AAAA : UInt16 = 28

///|
const DNS_CLASS_IN : UInt16 = 1

///|
const DNS_TYPE_ANY : UInt16 = 255

///|
fn normalize_name(value : String) -> String raise MdnsError {
  let value = value.to_lower()
  let value = if value.has_suffix(".") {
    value[0:value.length() - 1].to_owned()
  } else {
    value
  }
  if value.is_empty() || value.length() > 253 || !value.has_suffix(".local") {
    raise InvalidName("mDNS name must end in .local")
  }
  let labels = value.split(".").to_array()
  for label in labels {
    if label.is_empty() || label.length() > 63 {
      raise InvalidName("mDNS label must contain 1 through 63 characters")
    }
    if label[0] == '-' || label[label.length() - 1] == '-' {
      raise InvalidName("mDNS label cannot begin or end with a hyphen")
    }
    for character in label.code_units() {
      if !((character >= 'a' && character <= 'z') ||
        (character >= '0' && character <= '9') ||
        character == '-') {
        raise InvalidName("mDNS name contains an invalid character")
      }
    }
  }
  value
}

///|
fn encode_name(name : String) -> Bytes raise MdnsError {
  let result : Array[Byte] = []
  for label in name.split(".") {
    let encoded = @utf8.encode(label)
    if encoded.is_empty() || encoded.length() > 63 {
      raise InvalidName("invalid mDNS label length")
    }
    result.push(encoded.length().to_byte())
    for byte in encoded {
      result.push(byte)
    }
  }
  result.push(0)
  Bytes::from_array(result)
}

///|
fn dns_writer(capacity : Int) -> @codec.Writer raise MdnsError {
  @codec.Writer::new(capacity~) catch {
    InvalidLength(length) =>
      raise InvalidPacket("invalid DNS writer capacity \{length}")
    Truncated(needed~, remaining~) =>
      raise InvalidPacket(
        "unexpected DNS writer truncation \{needed}/\{remaining}",
      )
  }
}

///|
fn encode_query(name : String) -> Bytes raise MdnsError {
  let encoded_name = encode_name(name)
  let writer = dns_writer(DNS_HEADER_LENGTH + (encoded_name.length() + 4) * 2)
  writer.write_u16_be(0)
  writer.write_u16_be(0)
  writer.write_u16_be(2)
  writer.write_u16_be(0)
  writer.write_u16_be(0)
  writer.write_u16_be(0)
  writer.write_bytes(encoded_name)
  writer.write_u16_be(DNS_TYPE_A)
  writer.write_u16_be(DNS_CLASS_IN | 0x8000)
  writer.write_bytes(encoded_name)
  writer.write_u16_be(DNS_TYPE_AAAA)
  writer.write_u16_be(DNS_CLASS_IN | 0x8000)
  writer.finish()
}

///|
fn read_u16_at(data : Bytes, offset : Int) -> UInt16 raise MdnsError {
  if offset < 0 || offset + 2 > data.length() {
    raise InvalidPacket("truncated DNS 16-bit field")
  }
  (data[offset].to_uint16() << 8) | data[offset + 1].to_uint16()
}

///|
fn read_u32_at(data : Bytes, offset : Int) -> UInt raise MdnsError {
  if offset < 0 || offset + 4 > data.length() {
    raise InvalidPacket("truncated DNS 32-bit field")
  }
  (data[offset].to_uint() << 24) |
  (data[offset + 1].to_uint() << 16) |
  (data[offset + 2].to_uint() << 8) |
  data[offset + 3].to_uint()
}

///|
fn decode_name(data : Bytes, start : Int) -> (String, Int) raise MdnsError {
  if start < 0 || start >= data.length() {
    raise InvalidPacket("DNS name starts outside the packet")
  }
  let labels : Array[String] = []
  let mut cursor = start
  let mut next_offset = -1
  let mut jumps = 0
  let mut total_length = 0
  while true {
    if cursor >= data.length() {
      raise InvalidPacket("truncated DNS name")
    }
    let length = data[cursor].to_uint()
    if (length & 0xc0U) == 0xc0U {
      if cursor + 2 > data.length() {
        raise InvalidPacket("truncated DNS compression pointer")
      }
      let pointer = (((length & 0x3fU) << 8) | data[cursor + 1].to_uint()).reinterpret_as_int()
      if pointer >= cursor || pointer >= data.length() {
        raise InvalidPacket("invalid forward DNS compression pointer")
      }
      if next_offset < 0 {
        next_offset = cursor + 2
      }
      cursor = pointer
      jumps += 1
      if jumps > 32 {
        raise InvalidPacket("too many DNS compression pointers")
      }
      continue
    }
    if (length & 0xc0U) != 0U {
      raise InvalidPacket("unsupported DNS label encoding")
    }
    cursor += 1
    if length == 0U {
      if next_offset < 0 {
        next_offset = cursor
      }
      return (labels.join(".").to_lower(), next_offset)
    }
    let label_length = length.reinterpret_as_int()
    if label_length > 63 || cursor + label_length > data.length() {
      raise InvalidPacket("invalid DNS label length")
    }
    let label = @utf8.decode(data[cursor:cursor + label_length]) catch {
      Malformed(_) => raise InvalidPacket("DNS label is not valid UTF-8")
    }
    if label.contains(".") {
      raise InvalidPacket("DNS label contains a dot")
    }
    labels.push(label)
    total_length += label_length + 1
    if total_length > 254 {
      raise InvalidPacket("decoded DNS name exceeds 254 bytes")
    }
    cursor += label_length
  }
  raise InvalidPacket("unreachable DNS name parser state")
}

///|
priv struct DnsQuestion {
  name : String
  resource_type : UInt16
}

///|
fn skip_resource_records(
  data : Bytes,
  start : Int,
  count : Int,
) -> Int raise MdnsError {
  let mut offset = start
  for _record = 0; _record < count; _record = _record + 1 {
    let (_, next) = decode_name(data, offset)
    offset = next
    if offset + 10 > data.length() {
      raise InvalidPacket("truncated DNS resource header")
    }
    let data_length = read_u16_at(data, offset + 8).to_int()
    offset += 10
    if offset + data_length > data.length() {
      raise InvalidPacket("truncated DNS resource data")
    }
    offset += data_length
  }
  offset
}

///|
fn parse_query(data : Bytes) -> Array[DnsQuestion] raise MdnsError {
  if data.length() < DNS_HEADER_LENGTH {
    raise InvalidPacket("DNS packet is shorter than its header")
  }
  let flags = read_u16_at(data, 2)
  if (flags & 0x8000) != 0 || (flags & 0x7800) != 0 {
    raise InvalidPacket("packet is not a standard DNS query")
  }
  let question_count = read_u16_at(data, 4).to_int()
  let resource_count = read_u16_at(data, 6).to_int() +
    read_u16_at(data, 8).to_int() +
    read_u16_at(data, 10).to_int()
  if question_count < 1 || question_count > 32 || resource_count > 128 {
    raise InvalidPacket("DNS query contains too many records")
  }
  let questions : Array[DnsQuestion] = []
  let mut offset = DNS_HEADER_LENGTH
  for _question = 0; _question < question_count; _question = _question + 1 {
    let (name, next) = decode_name(data, offset)
    offset = next
    if offset + 4 > data.length() {
      raise InvalidPacket("truncated DNS question")
    }
    let resource_type = read_u16_at(data, offset)
    let resource_class = read_u16_at(data, offset + 2)
    offset += 4
    if (resource_class & 0x7fff) == DNS_CLASS_IN {
      questions.push({ name, resource_type, })
    }
  }
  offset = skip_resource_records(data, offset, resource_count)
  if offset != data.length() {
    raise InvalidPacket("DNS query contains trailing bytes")
  }
  questions
}

///|
fn encode_answer_record(
  name : String,
  address : @transport.IpAddress,
  ttl_seconds : UInt,
) -> Bytes raise MdnsError {
  let encoded_name = encode_name(name)
  let address_bytes = address.to_bytes()
  let resource_type : UInt16 = match address {
    V4(_) => DNS_TYPE_A
    V6(_) => DNS_TYPE_AAAA
  }
  let writer = dns_writer(encoded_name.length() + 10 + address_bytes.length())
  writer.write_bytes(encoded_name)
  writer.write_u16_be(resource_type)
  writer.write_u16_be(DNS_CLASS_IN | 0x8000)
  writer.write_u32_be(ttl_seconds)
  writer.write_u16_be(address_bytes.length().to_uint16())
  writer.write_bytes(address_bytes)
  writer.finish()
}

///|
fn encode_registered_response(
  query : Bytes,
  registrations : Map[String, Array[@transport.IpAddress]],
  ttl_seconds : UInt,
) -> Bytes? raise MdnsError {
  let questions = parse_query(query)
  let records : Array[(String, @transport.IpAddress)] = []
  for question in questions {
    match registrations.get(question.name) {
      Some(addresses) =>
        for address in addresses {
          let matches = match address {
            V4(_) =>
              question.resource_type == DNS_TYPE_A ||
              question.resource_type == DNS_TYPE_ANY
            V6(_) =>
              question.resource_type == DNS_TYPE_AAAA ||
              question.resource_type == DNS_TYPE_ANY
          }
          if matches && !records.contains((question.name, address)) {
            records.push((question.name, address))
          }
        }
      None => ()
    }
  }
  if records.is_empty() {
    return None
  }
  let encoded_records : Array[Bytes] = []
  let mut total_length = DNS_HEADER_LENGTH
  for record in records {
    let (name, address) = record
    let encoded = encode_answer_record(name, address, ttl_seconds)
    total_length += encoded.length()
    encoded_records.push(encoded)
  }
  let writer = dns_writer(total_length)
  writer.write_u16_be(read_u16_at(query, 0))
  writer.write_u16_be(0x8400)
  writer.write_u16_be(0)
  writer.write_u16_be(encoded_records.length().to_uint16())
  writer.write_u16_be(0)
  writer.write_u16_be(0)
  for record in encoded_records {
    writer.write_bytes(record)
  }
  Some(writer.finish())
}

///|
priv struct DnsAnswer {
  name : String
  address : @transport.IpAddress
  ttl_seconds : UInt
}

///|
fn parse_response(data : Bytes) -> Array[DnsAnswer] raise MdnsError {
  if data.length() < DNS_HEADER_LENGTH {
    raise InvalidPacket("DNS packet is shorter than its header")
  }
  let flags = read_u16_at(data, 2)
  if (flags & 0x8000) == 0 || (flags & 0x7800) != 0 {
    raise InvalidPacket("packet is not a standard DNS response")
  }
  let question_count = read_u16_at(data, 4).to_int()
  let answer_count = read_u16_at(data, 6).to_int()
  let authority_count = read_u16_at(data, 8).to_int()
  let additional_count = read_u16_at(data, 10).to_int()
  let resource_count = answer_count + authority_count + additional_count
  if question_count > 32 || resource_count > 128 {
    raise InvalidPacket("DNS packet contains too many records")
  }
  let mut offset = DNS_HEADER_LENGTH
  for _question = 0; _question < question_count; _question = _question + 1 {
    let (_, next) = decode_name(data, offset)
    offset = next
    if offset + 4 > data.length() {
      raise InvalidPacket("truncated DNS question")
    }
    offset += 4
  }
  let answers : Array[DnsAnswer] = []
  for _record = 0; _record < resource_count; _record = _record + 1 {
    let (name, next) = decode_name(data, offset)
    offset = next
    if offset + 10 > data.length() {
      raise InvalidPacket("truncated DNS resource header")
    }
    let resource_type = read_u16_at(data, offset)
    let resource_class = read_u16_at(data, offset + 2)
    let ttl_seconds = read_u32_at(data, offset + 4)
    let data_length = read_u16_at(data, offset + 8).to_int()
    offset += 10
    if offset + data_length > data.length() {
      raise InvalidPacket("truncated DNS resource data")
    }
    if (resource_class & 0x7fff) == DNS_CLASS_IN && ttl_seconds > 0U {
      if resource_type == DNS_TYPE_A && data_length == 4 {
        answers.push({
          name,
          address: @transport.IpAddress::v4(
            data[offset],
            data[offset + 1],
            data[offset + 2],
            data[offset + 3],
          ),
          ttl_seconds,
        })
      } else if resource_type == DNS_TYPE_AAAA && data_length == 16 {
        let address = @transport.IpAddress::v6(
          data[offset:offset + 16].to_owned(),
        ) catch {
          InvalidIpv6Length(length) =>
            raise InvalidPacket("invalid AAAA length \{length}")
          InvalidAddress(message) => raise InvalidPacket(message)
        }
        answers.push({ name, address, ttl_seconds, })
      }
    }
    offset += data_length
  }
  if offset != data.length() {
    raise InvalidPacket("DNS packet contains trailing bytes")
  }
  answers
}