///|
fn read_byte(bytes : Bytes, offset : Array[Int]) -> Byte raise CborError {
  if offset[0] >= bytes.length() {
    raise UnexpectedEOF
  }
  let b = bytes[offset[0]]
  offset[0] = offset[0] + 1
  b
}

///|
fn peek_byte(bytes : Bytes, offset : Array[Int]) -> Byte raise CborError {
  if offset[0] >= bytes.length() {
    raise UnexpectedEOF
  }
  bytes[offset[0]]
}

///|
fn read_uint16(bytes : Bytes, offset : Array[Int]) -> UInt64 raise CborError {
  let b1 = read_byte(bytes, offset).to_int64()
  let b2 = read_byte(bytes, offset).to_int64()
  ((b1 << 8) | b2).reinterpret_as_uint64()
}

///|
fn read_uint32(bytes : Bytes, offset : Array[Int]) -> UInt64 raise CborError {
  let b1 = read_byte(bytes, offset).to_int64()
  let b2 = read_byte(bytes, offset).to_int64()
  let b3 = read_byte(bytes, offset).to_int64()
  let b4 = read_byte(bytes, offset).to_int64()
  ((b1 << 24) | (b2 << 16) | (b3 << 8) | b4).reinterpret_as_uint64()
}

///|
fn read_uint64(bytes : Bytes, offset : Array[Int]) -> UInt64 raise CborError {
  let b1 = read_byte(bytes, offset).to_int64()
  let b2 = read_byte(bytes, offset).to_int64()
  let b3 = read_byte(bytes, offset).to_int64()
  let b4 = read_byte(bytes, offset).to_int64()
  let b5 = read_byte(bytes, offset).to_int64()
  let b6 = read_byte(bytes, offset).to_int64()
  let b7 = read_byte(bytes, offset).to_int64()
  let b8 = read_byte(bytes, offset).to_int64()
  ((b1 << 56) |
  (b2 << 48) |
  (b3 << 40) |
  (b4 << 32) |
  (b5 << 24) |
  (b6 << 16) |
  (b7 << 8) |
  b8).reinterpret_as_uint64()
}

///|
fn read_argument(
  bytes : Bytes,
  offset : Array[Int],
  additional_info : Int,
) -> UInt64 raise CborError {
  if additional_info < 24 {
    additional_info.to_uint64()
  } else if additional_info == 24 {
    read_byte(bytes, offset).to_uint64()
  } else if additional_info == 25 {
    read_uint16(bytes, offset)
  } else if additional_info == 26 {
    read_uint32(bytes, offset)
  } else if additional_info == 27 {
    read_uint64(bytes, offset)
  } else {
    raise InvalidAdditionalInfo(additional_info.to_byte())
  }
}

///|
fn ensure_preferred_argument(
  additional_info : Int,
  value : UInt64,
) -> Unit raise CborError {
  if additional_info == 24 && value < 24 {
    raise NonCanonicalEncoding(
      "argument must use the direct additional-info form",
    )
  } else if additional_info == 25 && value <= 0xFF {
    raise NonCanonicalEncoding(
      "argument must use the 8-bit additional-info form",
    )
  } else if additional_info == 26 && value <= 0xFFFF {
    raise NonCanonicalEncoding(
      "argument must use the 16-bit additional-info form",
    )
  } else if additional_info == 27 && value <= 0xFFFFFFFF {
    raise NonCanonicalEncoding(
      "argument must use the 32-bit additional-info form",
    )
  }
}

///|
fn read_canonical_argument(
  bytes : Bytes,
  offset : Array[Int],
  additional_info : Int,
) -> UInt64 raise CborError {
  let value = read_argument(bytes, offset, additional_info)
  ensure_preferred_argument(additional_info, value)
  value
}

///|
fn decode_positive_integer(value : UInt64) -> CborValue {
  if value > 0x7FFFFFFFFFFFFFFFUL {
    Unsigned(value)
  } else {
    Integer(value.reinterpret_as_int64())
  }
}

///|
fn read_length(
  bytes : Bytes,
  offset : Array[Int],
  additional_info : Int,
) -> Int raise CborError {
  let value = read_canonical_argument(bytes, offset, additional_info)
  if value > 0x7FFFFFFFFFFFFFFFUL {
    raise SemanticError("length exceeds Int range")
  }
  value.to_int()
}

///|
fn decode_negative_integer(value : UInt64) -> CborValue raise CborError {
  if value > 0x7FFFFFFFFFFFFFFFUL {
    raise SemanticError("negative integer exceeds Int64 range")
  }
  Integer(-1L - value.reinterpret_as_int64())
}

///|
fn decode_half_float(bits : UInt64) -> Double {
  let sign = if (bits & 0x8000UL) == 0UL { 1.0 } else { -1.0 }
  let exponent = ((bits >> 10) & 0x1FUL).to_int()
  let fraction = (bits & 0x03FFUL).to_int()
  if exponent == 0 {
    if fraction == 0 {
      sign * 0.0
    } else {
      sign * (fraction.to_double() / 1024.0) * @math.pow(2.0, -14.0)
    }
  } else if exponent == 31 {
    if fraction == 0 {
      sign * (1.0 / 0.0)
    } else {
      0.0 / 0.0
    }
  } else {
    sign *
    (1.0 + fraction.to_double() / 1024.0) *
    @math.pow(2.0, (exponent - 15).to_double())
  }
}

///|
fn is_continuation_byte(byte : Int) -> Bool {
  (byte & 0xC0) == 0x80
}

///|
fn decode_utf8(bytes : Bytes, start : Int, len : Int) -> String raise CborError {
  let builder = StringBuilder::new()
  let mut i = start
  let end = start + len
  if end > bytes.length() {
    raise UnexpectedEOF
  }
  while i < end {
    let b1 = bytes[i].to_int()
    if b1 < 0x80 {
      match Int::to_char(b1) {
        Some(char) => builder.write_char(char)
        None => raise InvalidUtf8
      }
      i = i + 1
    } else if (b1 & 0xE0) == 0xC0 {
      if i + 1 >= end {
        raise InvalidUtf8
      }
      let b2 = bytes[i + 1].to_int()
      if b1 < 0xC2 || !is_continuation_byte(b2) {
        raise InvalidUtf8
      }
      let code = ((b1 & 0x1F) << 6) | (b2 & 0x3F)
      match Int::to_char(code) {
        Some(char) => builder.write_char(char)
        None => raise InvalidUtf8
      }
      i = i + 2
    } else if (b1 & 0xF0) == 0xE0 {
      if i + 2 >= end {
        raise InvalidUtf8
      }
      let b2 = bytes[i + 1].to_int()
      let b3 = bytes[i + 2].to_int()
      if !is_continuation_byte(b2) || !is_continuation_byte(b3) {
        raise InvalidUtf8
      }
      if (b1 == 0xE0 && b2 < 0xA0) || (b1 == 0xED && b2 >= 0xA0) {
        raise InvalidUtf8
      }
      let code = ((b1 & 0x0F) << 12) | ((b2 & 0x3F) << 6) | (b3 & 0x3F)
      if code < 0x800 || (code >= 0xD800 && code <= 0xDFFF) {
        raise InvalidUtf8
      }
      match Int::to_char(code) {
        Some(char) => builder.write_char(char)
        None => raise InvalidUtf8
      }
      i = i + 3
    } else if (b1 & 0xF8) == 0xF0 {
      if i + 3 >= end {
        raise InvalidUtf8
      }
      let b2 = bytes[i + 1].to_int()
      let b3 = bytes[i + 2].to_int()
      let b4 = bytes[i + 3].to_int()
      if !is_continuation_byte(b2) ||
        !is_continuation_byte(b3) ||
        !is_continuation_byte(b4) {
        raise InvalidUtf8
      }
      if b1 > 0xF4 || (b1 == 0xF0 && b2 < 0x90) || (b1 == 0xF4 && b2 >= 0x90) {
        raise InvalidUtf8
      }
      let code = ((b1 & 0x07) << 18) |
        ((b2 & 0x3F) << 12) |
        ((b3 & 0x3F) << 6) |
        (b4 & 0x3F)
      if code < 0x10000 || code > 0x10FFFF {
        raise InvalidUtf8
      }
      match Int::to_char(code) {
        Some(char) => builder.write_char(char)
        None => raise InvalidUtf8
      }
      i = i + 4
    } else {
      raise InvalidUtf8
    }
  }
  builder.to_string()
}

///|
fn decode_indefinite_bytes(
  bytes : Bytes,
  offset : Array[Int],
) -> Bytes raise CborError {
  let buf = Buffer()
  while true {
    let initial = peek_byte(bytes, offset)
    if initial == 0xFF {
      ignore(read_byte(bytes, offset))
      return buf.to_bytes()
    }
    let major = initial.to_int() >> 5
    let additional_info = initial.to_int() & 0x1F
    if major != 2 || additional_info == 31 {
      raise InvalidIndefiniteChunk(initial)
    }
    match decode_internal(bytes, offset, false) {
      Some(Bytes(chunk)) => buf.write_bytes(chunk)
      Some(_) => raise InvalidIndefiniteChunk(initial)
      None => raise UnexpectedBreak
    }
  }
  raise SemanticError("unreachable indefinite bytes loop exit")
}

///|
fn decode_indefinite_text(
  bytes : Bytes,
  offset : Array[Int],
) -> String raise CborError {
  let builder = StringBuilder::new()
  while true {
    let initial = peek_byte(bytes, offset)
    if initial == 0xFF {
      ignore(read_byte(bytes, offset))
      return builder.to_string()
    }
    let major = initial.to_int() >> 5
    let additional_info = initial.to_int() & 0x1F
    if major != 3 || additional_info == 31 {
      raise InvalidIndefiniteChunk(initial)
    }
    match decode_internal(bytes, offset, false) {
      Some(Text(chunk)) => builder.write_string(chunk)
      Some(_) => raise InvalidIndefiniteChunk(initial)
      None => raise UnexpectedBreak
    }
  }
  raise SemanticError("unreachable indefinite text loop exit")
}

///|
fn decode_internal(
  bytes : Bytes,
  offset : Array[Int],
  allow_break : Bool,
) -> CborValue? raise CborError {
  let initial = read_byte(bytes, offset).to_int()
  let major = initial >> 5
  let additional_info = initial & 0x1F

  if major == 0 {
    let val = read_canonical_argument(bytes, offset, additional_info)
    Some(decode_positive_integer(val))
  } else if major == 1 {
    let val = read_canonical_argument(bytes, offset, additional_info)
    Some(decode_negative_integer(val))
  } else if major == 2 {
    if additional_info == 31 {
      Some(Bytes(decode_indefinite_bytes(bytes, offset)))
    } else {
      let len = read_length(bytes, offset, additional_info)
      let end = offset[0] + len
      if end > bytes.length() {
        raise UnexpectedEOF
      }
      let buf = Buffer()
      for i = 0; i < len; i = i + 1 {
        buf.write_byte(bytes[offset[0] + i])
      }
      offset[0] = end
      Some(Bytes(buf.to_bytes()))
    }
  } else if major == 3 {
    if additional_info == 31 {
      Some(Text(decode_indefinite_text(bytes, offset)))
    } else {
      let len = read_length(bytes, offset, additional_info)
      let s = decode_utf8(bytes, offset[0], len)
      offset[0] = offset[0] + len
      Some(Text(s))
    }
  } else if major == 4 {
    let arr = []
    if additional_info == 31 {
      while true {
        match decode_internal(bytes, offset, true) {
          Some(item) => arr.push(item)
          None => break
        }
      }
    } else {
      let len = read_length(bytes, offset, additional_info)
      for i = 0; i < len; i = i + 1 {
        arr.push(decode_internal(bytes, offset, false).unwrap())
      }
    }
    Some(Array(arr))
  } else if major == 5 {
    let map = []
    if additional_info == 31 {
      while true {
        match decode_internal(bytes, offset, true) {
          Some(key) =>
            match decode_internal(bytes, offset, true) {
              Some(val) => map.push((key, val))
              None => raise UnexpectedBreak
            }
          None => break
        }
      }
    } else {
      let len = read_length(bytes, offset, additional_info)
      for i = 0; i < len; i = i + 1 {
        let key = decode_internal(bytes, offset, false).unwrap()
        let val = decode_internal(bytes, offset, false).unwrap()
        map.push((key, val))
      }
    }
    Some(Map(map))
  } else if major == 6 {
    let tag_val = read_canonical_argument(bytes, offset, additional_info)
    let item = decode_internal(bytes, offset, false).unwrap()
    Some(Tag(tag_val, item))
  } else if additional_info < 24 {
    Some(Simple(additional_info.to_byte()))
  } else if additional_info == 24 {
    let b = read_byte(bytes, offset)
    if b < 24 {
      raise NonCanonicalEncoding(
        "simple value must use the direct additional-info form",
      )
    } else if b < 32 {
      raise InvalidAdditionalInfo(b)
    } else {
      Some(Simple(b))
    }
  } else if additional_info == 25 {
    let val = read_uint16(bytes, offset)
    Some(Float64(decode_half_float(val)))
  } else if additional_info == 26 {
    let val = read_uint32(bytes, offset).to_int()
    let float32 = Float::reinterpret_from_int(val)
    Some(Float64(float32.to_double()))
  } else if additional_info == 27 {
    let val = read_uint64(bytes, offset).reinterpret_as_int64()
    Some(Float64(val.reinterpret_as_double()))
  } else if additional_info == 31 {
    if allow_break {
      None
    } else {
      raise UnexpectedBreak
    }
  } else {
    raise InvalidAdditionalInfo(additional_info.to_byte())
  }
}

///|
/// Decode a CBOR value from a byte array starting at the given offset.
pub fn decode_from_offset(
  bytes : Bytes,
  offset : Array[Int],
) -> CborValue raise CborError {
  match decode_internal(bytes, offset, false) {
    Some(value) => value
    None => raise UnexpectedBreak
  }
}

///|
/// Decode a CBOR value from a byte array.
pub fn decode(bytes : Bytes) -> CborValue raise CborError {
  let offset = [0]
  let value = decode_from_offset(bytes, offset)
  if offset[0] != bytes.length() {
    raise TrailingBytes(bytes.length() - offset[0])
  }
  value
}