// DNS wire-format primitives.  Every reader is bounds checked; callers must
// propagate the Result instead of turning malformed input into an empty value.

///|
pub fn wire_check_range(
  bytes : Array[Byte],
  offset : Int,
  length : Int,
) -> Result[Unit, String] {
  if offset < 0 ||
    length < 0 ||
    offset > bytes.length() ||
    length > bytes.length() - offset {
    Err("truncated DNS wire data")
  } else {
    Ok(())
  }
}

///|
pub fn wire_get_u16(
  bytes : Array[Byte],
  offset : Int,
) -> Result[(UInt16, Int), String] {
  match wire_check_range(bytes, offset, 2) {
    Err(err) => Err(err)
    Ok(_) => {
      let hi = bytes[offset].to_int()
      let lo = bytes[offset + 1].to_int()
      Ok((((hi << 8) | lo).to_uint16(), offset + 2))
    }
  }
}

///|
pub fn wire_get_u32(
  bytes : Array[Byte],
  offset : Int,
) -> Result[(UInt, Int), String] {
  match wire_check_range(bytes, offset, 4) {
    Err(err) => Err(err)
    Ok(_) => {
      let b0 = bytes[offset].to_int()
      let b1 = bytes[offset + 1].to_int()
      let b2 = bytes[offset + 2].to_int()
      let b3 = bytes[offset + 3].to_int()
      Ok(
        (
          ((b0 << 24) | (b1 << 16) | (b2 << 8) | b3).reinterpret_as_uint(),
          offset + 4,
        ),
      )
    }
  }
}

///|
pub fn wire_put_u16(
  buf : Array[Byte],
  offset : Int,
  val : UInt16,
) -> Result[Int, String] {
  match wire_check_range(buf, offset, 2) {
    Err(err) => Err(err)
    Ok(_) => {
      buf[offset] = ((val >> 8) & 0xFF).to_byte()
      buf[offset + 1] = (val & 0xFF).to_byte()
      Ok(offset + 2)
    }
  }
}

///|
pub fn wire_put_u32(
  buf : Array[Byte],
  offset : Int,
  val : UInt,
) -> Result[Int, String] {
  match wire_check_range(buf, offset, 4) {
    Err(err) => Err(err)
    Ok(_) => {
      let v = val.reinterpret_as_int()
      buf[offset] = ((v >> 24) & 0xFF).to_byte()
      buf[offset + 1] = ((v >> 16) & 0xFF).to_byte()
      buf[offset + 2] = ((v >> 8) & 0xFF).to_byte()
      buf[offset + 3] = (v & 0xFF).to_byte()
      Ok(offset + 4)
    }
  }
}

///|
pub fn wire_is_pointer(b : Byte) -> Bool {
  (b.to_int() & 0xC0) == 0xC0
}

///|
pub fn wire_get_pointer(
  bytes : Array[Byte],
  offset : Int,
) -> Result[Int, String] {
  match wire_check_range(bytes, offset, 2) {
    Err(err) => Err(err)
    Ok(_) => {
      let first = bytes[offset].to_int()
      if !wire_is_pointer(bytes[offset]) {
        Err("DNS compression pointer has invalid tag")
      } else {
        Ok(((first & 0x3F) << 8) | bytes[offset + 1].to_int())
      }
    }
  }
}

///|
pub fn wire_put_pointer(
  buf : Array[Byte],
  offset : Int,
  target : Int,
) -> Result[Int, String] {
  if target < 0 || target >= 0x4000 {
    return Err("DNS compression pointer target is outside 14-bit range")
  }
  match wire_check_range(buf, offset, 2) {
    Err(err) => Err(err)
    Ok(_) => {
      buf[offset] = (0xC0 | ((target >> 8) & 0x3F)).to_byte()
      buf[offset + 1] = (target & 0xFF).to_byte()
      Ok(offset + 2)
    }
  }
}

///|
pub fn wire_copy_range(
  bytes : Array[Byte],
  offset : Int,
  length : Int,
) -> Result[Array[Byte], String] {
  match wire_check_range(bytes, offset, length) {
    Err(err) => Err(err)
    Ok(_) => {
      let out = Array::make(length, (0).to_byte())
      for i in 0.. Int {
  let total = Ref(header_bytes.length())
  for section in sections {
    total.val = total.val + section.length()
  }
  total.val
}

///|
pub fn wire_concat(segments : Array[Array[Byte]]) -> Array[Byte] {
  let total = Ref(0)
  for seg in segments {
    total.val = total.val + seg.length()
  }
  let buf = Array::make(total.val, (0).to_byte())
  let pos = Ref(0)
  for seg in segments {
    for byte in seg {
      buf[pos.val] = byte
      pos.val = pos.val + 1
    }
  }
  buf
}

///|
fn hex_char(v : Int) -> String {
  if v < 10 {
    "0123456789"[v].unsafe_to_char().to_string()
  } else {
    "abcdef"[v - 10].unsafe_to_char().to_string()
  }
}

///|
pub fn wire_hex_dump(bytes : Array[Byte], max_len? : Int = 64) -> String {
  let result = Ref("")
  let limit = if bytes.length() < max_len { bytes.length() } else { max_len }
  for i in 0.. 0 && i % 16 == 0 {
      result.val = result.val + "\n"
    } else if i > 0 {
      result.val = result.val + " "
    }
    let b = bytes[i].to_int()
    result.val = result.val + hex_char((b >> 4) & 0xF) + hex_char(b & 0xF)
  }
  if bytes.length() > limit {
    result.val = result.val +
      "... (" +
      bytes.length().to_string() +
      " bytes total)"
  }
  result.val
}