///|
/// A cursor over a byte buffer used to decode one or more BER values.
pub struct BerDecoder {
  data : Bytes
  limits : BerLimits
  priv mut pos : Int
} derive(@debug.Debug)

///|
pub fn BerDecoder::new(data : Bytes, limits : BerLimits) -> BerDecoder {
  { data, limits, pos: 0 }
}

///|
pub fn BerDecoder::position(self : BerDecoder) -> Int {
  self.pos
}

///|
pub fn BerDecoder::remaining(self : BerDecoder) -> Int {
  self.data.length() - self.pos
}

///|
pub fn BerDecoder::at_end(self : BerDecoder) -> Bool {
  self.pos >= self.data.length()
}

///|
fn BerDecoder::read_byte_raise(self : BerDecoder) -> Byte raise BerError {
  if self.pos >= self.data.length() {
    raise Truncated
  }
  let b = self.data.get(self.pos).unwrap()
  self.pos = self.pos + 1
  b
}

///|
fn BerDecoder::read_octets_raise(
  self : BerDecoder,
  n : Int,
) -> Bytes raise BerError {
  if n < 0 || self.pos + n > self.data.length() {
    raise Truncated
  }
  let view = self.data.get_view(start=self.pos, end=self.pos + n).unwrap()
  let out = Bytes::from_array(view.to_array())
  self.pos = self.pos + n
  out
}

///|
pub fn BerDecoder::read_tag(self : BerDecoder) -> Result[BerTag, BerError] {
  result_of_ber(fn() raise BerError { self.read_tag_raise() })
}

///|
fn BerDecoder::read_tag_raise(self : BerDecoder) -> BerTag raise BerError {
  let first = self.read_byte_raise()
  let v = first.to_int()
  let class = TagClass::from_bits((v >> 6) & 0x3)
  let constructed = (v & 0x20) != 0
  let low = v & 0x1F
  if low != 0x1F {
    return BerTag::new(class, constructed, low)
  }
  // Long form: base-128 big-endian.
  let mut number = 0
  let mut count = 0
  let mut done = false
  while !done {
    let b = self.read_byte_raise()
    count = count + 1
    if count > 5 {
      raise InvalidTag
    }
    number = (number << 7) | (b.to_int() & 0x7F)
    if (b.to_int() & 0x80) == 0 {
      done = true
    }
  }
  BerTag::new(class, constructed, number)
}

///|
pub fn BerDecoder::read_length(
  self : BerDecoder,
) -> Result[BerLength, BerError] {
  result_of_ber(fn() raise BerError { self.read_length_raise() })
}

///|
fn BerDecoder::read_length_raise(self : BerDecoder) -> BerLength raise BerError {
  let first = self.read_byte_raise()
  let v = first.to_int()
  if v == 0x80 {
    return Indefinite
  }
  if (v & 0x80) == 0 {
    return Definite(v)
  }
  let octets = v & 0x7F
  if octets == 0 || octets > 4 {
    raise InvalidLength
  }
  let mut length = 0
  for _ in 0.. Result[BerValue, BerError] {
  result_of_ber(fn() raise BerError { self.read_value_raise(depth) })
}

///|
fn BerDecoder::read_value_raise(
  self : BerDecoder,
  depth : Int,
) -> BerValue raise BerError {
  if depth > self.limits.max_depth {
    raise TooDeep
  }
  let tag = self.read_tag_raise()
  let length = self.read_length_raise()
  match length {
    Indefinite =>
      if tag.constructed {
        if tag.number == tag_eoc {
          raise UnexpectedEoc
        }
        let children = self.read_children_until_eoc_raise(depth + 1)
        BerValue::constructed(tag, children)
      } else {
        raise IndefinitePrimitive
      }
    Definite(len) =>
      if len > self.limits.max_length {
        raise LengthExceedsLimit(len)
      } else if tag.constructed {
        let end = self.pos + len
        if end > self.data.length() {
          raise Truncated
        }
        let children : Array[BerValue] = []
        while self.pos < end {
          let child = self.read_value_raise(depth + 1)
          children.push(child)
          if children.length() > self.limits.max_elements {
            raise TooManyElements
          }
        }
        if self.pos != end {
          raise Truncated
        }
        BerValue::constructed(tag, children)
      } else {
        let content = self.read_octets_raise(len)
        BerValue::primitive(tag, content)
      }
  }
}

///|
fn BerDecoder::read_children_until_eoc_raise(
  self : BerDecoder,
  depth : Int,
) -> Array[BerValue] raise BerError {
  let children : Array[BerValue] = []
  let mut done = false
  while !done {
    if self.at_end() {
      raise MissingEoc
    }
    // Peek the next tag to detect end-of-content.
    let saved = self.pos
    let tag = self.read_tag_raise()
    if !tag.constructed && tag.class == Universal && tag.number == tag_eoc {
      let length = self.read_length_raise()
      match length {
        Definite(0) => done = true
        _ => raise InvalidLength
      }
    } else {
      self.pos = saved
      let child = self.read_value_raise(depth)
      children.push(child)
      if children.length() > self.limits.max_elements {
        raise TooManyElements
      }
    }
  }
  children
}

///|
/// Decode exactly one top-level BER value. Any trailing octets are reported
/// as `NotSingleValue`.
pub fn decode_ber(
  data : Bytes,
  limits : BerLimits?,
) -> Result[BerValue, BerError] {
  let lim = match limits {
    Some(l) => l
    None => Default::default()
  }
  result_of_ber(fn() raise BerError {
    let decoder = BerDecoder::new(data, lim)
    let value = decoder.read_value_raise(0)
    if !decoder.at_end() {
      raise NotSingleValue
    }
    value
  })
}

///|
/// Decode all BER values in the buffer until it is fully consumed. Used to
/// decode LDAP message streams.
pub fn decode_ber_many(
  data : Bytes,
  limits : BerLimits?,
) -> Result[Array[BerValue], BerError] {
  let lim = match limits {
    Some(l) => l
    None => Default::default()
  }
  result_of_ber(fn() raise BerError {
    let decoder = BerDecoder::new(data, lim)
    let out : Array[BerValue] = []
    while !decoder.at_end() {
      let value = decoder.read_value_raise(0)
      out.push(value)
      if out.length() > lim.max_elements {
        raise TooManyElements
      }
    }
    out
  })
}