///|
const EXT_SUPPORTED_GROUPS : UInt16 = 10

///|
const EXT_EC_POINT_FORMATS : UInt16 = 11

///|
const EXT_SIGNATURE_ALGORITHMS : UInt16 = 13

///|
const EXT_USE_SRTP : UInt16 = 14

///|
const EXT_EXTENDED_MASTER_SECRET : UInt16 = 23

///|
const EXT_RENEGOTIATION_INFO : UInt16 = 0xff01

///|
pub struct HelloExtensions {
  supported_groups : Array[UInt16]
  point_formats : Array[Byte]
  signature_algorithms : Array[UInt16]
  srtp_profiles : Array[UInt16]
  extended_master_secret : Bool
  renegotiation_info : Bool
  unknown : Array[(UInt16, Bytes)]
} derive(Debug, Eq)

///|
pub fn HelloExtensions::webrtc(
  srtp_profiles? : Array[SrtpProtectionProfile] = [],
) -> HelloExtensions {
  {
    supported_groups: [23, 29, 24],
    point_formats: [0],
    signature_algorithms: [0x0403, 0x0401],
    srtp_profiles: srtp_profiles.map(profile => profile.code()),
    extended_master_secret: true,
    renegotiation_info: true,
    unknown: [],
  }
}

///|
fn HelloExtensions::server(selected_srtp_profile : UInt16?) -> HelloExtensions {
  {
    supported_groups: [],
    point_formats: [0],
    signature_algorithms: [],
    srtp_profiles: match selected_srtp_profile {
      Some(profile) => [profile]
      None => []
    },
    extended_master_secret: true,
    renegotiation_info: true,
    unknown: [],
  }
}

///|
pub fn HelloExtensions::supports_group(
  self : HelloExtensions,
  group : UInt16,
) -> Bool {
  self.supported_groups.contains(group)
}

///|
pub fn HelloExtensions::signature_algorithms(
  self : HelloExtensions,
) -> Array[UInt16] {
  self.signature_algorithms.copy()
}

///|
pub fn HelloExtensions::srtp_profile_codes(
  self : HelloExtensions,
) -> Array[UInt16] {
  self.srtp_profiles.copy()
}

///|
pub fn HelloExtensions::uses_extended_master_secret(
  self : HelloExtensions,
) -> Bool {
  self.extended_master_secret
}

///|
fn hs_read_u8(reader : @codec.Reader, context : String) -> Byte raise DtlsError {
  reader.read_u8() catch {
    InvalidLength(length) =>
      raise InvalidHandshake("\{context}: invalid length \{length}")
    Truncated(needed~, remaining~) =>
      raise InvalidHandshake(
        "\{context}: need \{needed} bytes, have \{remaining}",
      )
  }
}

///|
fn hs_read_u16(
  reader : @codec.Reader,
  context : String,
) -> UInt16 raise DtlsError {
  reader.read_u16_be() catch {
    InvalidLength(length) =>
      raise InvalidHandshake("\{context}: invalid length \{length}")
    Truncated(needed~, remaining~) =>
      raise InvalidHandshake(
        "\{context}: need \{needed} bytes, have \{remaining}",
      )
  }
}

///|
fn hs_read_u24(
  reader : @codec.Reader,
  context : String,
) -> UInt raise DtlsError {
  reader.read_u24_be() catch {
    InvalidLength(length) =>
      raise InvalidHandshake("\{context}: invalid length \{length}")
    Truncated(needed~, remaining~) =>
      raise InvalidHandshake(
        "\{context}: need \{needed} bytes, have \{remaining}",
      )
  }
}

///|
fn hs_read_bytes(
  reader : @codec.Reader,
  length : Int,
  context : String,
) -> Bytes raise DtlsError {
  reader.read_bytes(length) catch {
    InvalidLength(invalid_length) =>
      raise InvalidHandshake("\{context}: invalid length \{invalid_length}")
    Truncated(needed~, remaining~) =>
      raise InvalidHandshake(
        "\{context}: need \{needed} bytes, have \{remaining}",
      )
  }
}

///|
fn write_extension(
  writer : @codec.Writer,
  extension_type : UInt16,
  value : Bytes,
) -> Unit raise DtlsError {
  if value.length() > 0xffff {
    raise InvalidHandshake("DTLS extension exceeds 65535 bytes")
  }
  writer.write_u16_be(extension_type)
  writer.write_u16_be(value.length().to_uint16())
  writer.write_bytes(value)
}

///|
fn u16_vector(values : Array[UInt16]) -> Bytes raise DtlsError {
  if values.length() > 0x7fff {
    raise InvalidHandshake("16-bit vector contains too many values")
  }
  let writer = dtls_writer(2 + values.length() * 2, "DTLS extension")
  writer.write_u16_be((values.length() * 2).to_uint16())
  for value in values {
    writer.write_u16_be(value)
  }
  writer.finish()
}

///|
fn HelloExtensions::encode(self : HelloExtensions) -> Bytes raise DtlsError {
  let body = dtls_writer(128, "DTLS extensions")
  if !self.supported_groups.is_empty() {
    write_extension(
      body,
      EXT_SUPPORTED_GROUPS,
      u16_vector(self.supported_groups),
    )
  }
  if !self.point_formats.is_empty() {
    if self.point_formats.length() > 0xff {
      raise InvalidHandshake("too many EC point formats")
    }
    let value = [self.point_formats.length().to_byte()]
    for point_format in self.point_formats {
      value.push(point_format)
    }
    write_extension(body, EXT_EC_POINT_FORMATS, Bytes::from_array(value))
  }
  if !self.signature_algorithms.is_empty() {
    write_extension(
      body,
      EXT_SIGNATURE_ALGORITHMS,
      u16_vector(self.signature_algorithms),
    )
  }
  if !self.srtp_profiles.is_empty() {
    let value = u16_vector(self.srtp_profiles).to_array()
    value.push(0)
    write_extension(body, EXT_USE_SRTP, Bytes::from_array(value))
  }
  if self.extended_master_secret {
    write_extension(body, EXT_EXTENDED_MASTER_SECRET, b"")
  }
  if self.renegotiation_info {
    write_extension(body, EXT_RENEGOTIATION_INFO, b"\x00")
  }
  for extension in self.unknown {
    let (extension_type, value) = extension
    write_extension(body, extension_type, value)
  }
  let body = body.finish()
  if body.length() > 0xffff {
    raise InvalidHandshake("DTLS extension block exceeds 65535 bytes")
  }
  let writer = dtls_writer(body.length() + 2, "DTLS extensions")
  writer.write_u16_be(body.length().to_uint16())
  writer.write_bytes(body)
  writer.finish()
}

///|
fn parse_u16_vector(
  value : Bytes,
  context : String,
) -> Array[UInt16] raise DtlsError {
  let reader = @codec.Reader::new(value)
  let byte_length = hs_read_u16(reader, context).to_int()
  if byte_length % 2 != 0 || byte_length != reader.remaining() {
    raise InvalidHandshake("\{context}: malformed 16-bit vector")
  }
  let values : Array[UInt16] = []
  while reader.remaining() > 0 {
    values.push(hs_read_u16(reader, context))
  }
  values
}

///|
fn parse_use_srtp(value : Bytes) -> Array[UInt16] raise DtlsError {
  let reader = @codec.Reader::new(value)
  let profile_bytes = hs_read_u16(reader, "use_srtp profiles").to_int()
  if profile_bytes == 0 ||
    profile_bytes % 2 != 0 ||
    reader.remaining() < profile_bytes + 1 {
    raise InvalidHandshake("malformed use_srtp profile vector")
  }
  let profiles : Array[UInt16] = []
  for offset = 0; offset < profile_bytes; offset = offset + 2 {
    profiles.push(hs_read_u16(reader, "use_srtp profile"))
  }
  let mki_length = hs_read_u8(reader, "use_srtp MKI").to_int()
  ignore(hs_read_bytes(reader, mki_length, "use_srtp MKI"))
  if reader.remaining() != 0 || mki_length != 0 {
    raise InvalidHandshake("WebRTC use_srtp requires an empty MKI")
  }
  profiles
}

///|
fn HelloExtensions::decode(encoded : Bytes) -> HelloExtensions raise DtlsError {
  let reader = @codec.Reader::new(encoded)
  let declared_length = hs_read_u16(reader, "DTLS extensions").to_int()
  if declared_length != reader.remaining() {
    raise InvalidHandshake("DTLS extension block length mismatch")
  }
  let mut supported_groups : Array[UInt16] = []
  let mut point_formats : Array[Byte] = []
  let mut signature_algorithms : Array[UInt16] = []
  let mut srtp_profiles : Array[UInt16] = []
  let mut extended_master_secret = false
  let mut renegotiation_info = false
  let unknown : Array[(UInt16, Bytes)] = []
  let seen : Map[UInt16, Bool] = Map([])
  while reader.remaining() > 0 {
    let extension_type = hs_read_u16(reader, "DTLS extension type")
    let length = hs_read_u16(reader, "DTLS extension length").to_int()
    let value = hs_read_bytes(reader, length, "DTLS extension value")
    if seen.contains(extension_type) {
      raise InvalidHandshake("duplicate DTLS extension \{extension_type}")
    }
    seen[extension_type] = true
    match extension_type {
      EXT_SUPPORTED_GROUPS =>
        supported_groups = parse_u16_vector(value, "supported_groups")
      EXT_EC_POINT_FORMATS => {
        let value_reader = @codec.Reader::new(value)
        let count = hs_read_u8(value_reader, "EC point formats").to_int()
        if count != value_reader.remaining() {
          raise InvalidHandshake("malformed EC point-format vector")
        }
        point_formats = hs_read_bytes(value_reader, count, "EC point formats").to_array()
      }
      EXT_SIGNATURE_ALGORITHMS =>
        signature_algorithms = parse_u16_vector(value, "signature_algorithms")
      EXT_USE_SRTP => srtp_profiles = parse_use_srtp(value)
      EXT_EXTENDED_MASTER_SECRET => {
        if !value.is_empty() {
          raise InvalidHandshake(
            "extended_master_secret extension must be empty",
          )
        }
        extended_master_secret = true
      }
      EXT_RENEGOTIATION_INFO => {
        if value != b"\x00" {
          raise InvalidHandshake("initial renegotiation_info must be empty")
        }
        renegotiation_info = true
      }
      _ => unknown.push((extension_type, value))
    }
  }
  {
    supported_groups,
    point_formats,
    signature_algorithms,
    srtp_profiles,
    extended_master_secret,
    renegotiation_info,
    unknown,
  }
}