///|
priv enum RecordCipherMode {
  AeadCipher(
    algorithm~ : @crypto.AeadAlgorithm,
    tag_length~ : Int,
    explicit_nonce~ : Bool
  )
  CbcCipher
}

///|
struct RecordCipher {
  provider : @crypto.Provider
  mode : RecordCipherMode
  write_key : @crypto.Secret
  read_key : @crypto.Secret
  write_mac_key : @crypto.Secret
  read_mac_key : @crypto.Secret
  write_iv : Bytes
  read_iv : Bytes
  replay_windows : Map[UInt16, @replay.ReplayWindow]
  replay_window : Int
}

///|
fn cipher_mode(suite : CipherSuite) -> RecordCipherMode {
  match suite {
    EcdheEcdsaAes128Ccm | PskAes128Ccm =>
      AeadCipher(algorithm=Aes128Ccm, tag_length=16, explicit_nonce=true)
    EcdheEcdsaAes128Ccm8 | PskAes128Ccm8 =>
      AeadCipher(algorithm=Aes128Ccm, tag_length=8, explicit_nonce=true)
    EcdheEcdsaAes128GcmSha256 | EcdheRsaAes128GcmSha256 | PskAes128GcmSha256 =>
      AeadCipher(algorithm=Aes128Gcm, tag_length=16, explicit_nonce=true)
    EcdheRsaChacha20Poly1305Sha256 | EcdheEcdsaChacha20Poly1305Sha256 =>
      AeadCipher(
        algorithm=Chacha20Poly1305,
        tag_length=16,
        explicit_nonce=false,
      )
    EcdheEcdsaAes256CbcSha | EcdheRsaAes256CbcSha => CbcCipher
  }
}

///|
fn expected_key_lengths(suite : CipherSuite) -> (Int, Int, Int) {
  match suite {
    EcdheEcdsaAes256CbcSha | EcdheRsaAes256CbcSha => (32, 0, 20)
    EcdheRsaChacha20Poly1305Sha256 | EcdheEcdsaChacha20Poly1305Sha256 =>
      (32, 12, 0)
    _ => (16, 4, 0)
  }
}

///|
fn RecordCipher::new(
  keys : EncryptionKeys,
  role : Role,
  suite? : CipherSuite = EcdheEcdsaAes128GcmSha256,
  replay_window? : Int = 64,
) -> RecordCipher raise DtlsError {
  if role == Auto {
    raise HandshakeFailed("DTLS record cipher requires a fixed role")
  }
  if replay_window < 1 || replay_window > 64 {
    raise HandshakeFailed("DTLS replay window must be between 1 and 64")
  }
  let (write_key, read_key, write_mac_key, read_mac_key, write_iv, read_iv) = if role ==
    Client {
    (
      keys.client_write_key,
      keys.server_write_key,
      keys.client_mac_key,
      keys.server_mac_key,
      keys.client_write_iv,
      keys.server_write_iv,
    )
  } else {
    (
      keys.server_write_key,
      keys.client_write_key,
      keys.server_mac_key,
      keys.client_mac_key,
      keys.server_write_iv,
      keys.client_write_iv,
    )
  }
  let (key_length, iv_length, mac_length) = expected_key_lengths(suite)
  if write_key.length() != key_length ||
    read_key.length() != key_length ||
    write_iv.length() != iv_length ||
    read_iv.length() != iv_length ||
    write_mac_key.length() != mac_length ||
    read_mac_key.length() != mac_length {
    raise HandshakeFailed(
      "invalid key material for DTLS cipher suite \{suite.code()}",
    )
  }
  {
    provider: crypto_provider(),
    mode: cipher_mode(suite),
    write_key: @crypto.Secret::from_bytes(write_key),
    read_key: @crypto.Secret::from_bytes(read_key),
    write_mac_key: @crypto.Secret::from_bytes(write_mac_key),
    read_mac_key: @crypto.Secret::from_bytes(read_mac_key),
    write_iv,
    read_iv,
    replay_windows: Map([]),
    replay_window,
  }
}

///|
fn epoch_sequence(epoch : UInt16, sequence_number : UInt64) -> Bytes {
  Bytes::from_array([
    (epoch >> 8).to_byte(),
    epoch.to_byte(),
    (sequence_number >> 40).to_byte(),
    (sequence_number >> 32).to_byte(),
    (sequence_number >> 24).to_byte(),
    (sequence_number >> 16).to_byte(),
    (sequence_number >> 8).to_byte(),
    sequence_number.to_byte(),
  ])
}

///|
fn record_aad(
  content_type : ContentType,
  version : ProtocolVersion,
  epoch : UInt16,
  sequence_number : UInt64,
  plaintext_length : Int,
) -> Bytes raise DtlsError {
  if plaintext_length < 0 || plaintext_length > 0xffff {
    raise InvalidRecord("invalid DTLS plaintext length")
  }
  let result = epoch_sequence(epoch, sequence_number).to_array()
  result.push(content_type.code())
  result.push(version.major())
  result.push(version.minor())
  result.push((plaintext_length >> 8).to_byte())
  result.push(plaintext_length.to_byte())
  Bytes::from_array(result)
}

///|
fn explicit_nonce(fixed_iv : Bytes, explicit : Bytes) -> Bytes {
  append_bytes(fixed_iv, explicit)
}

///|
fn chacha_nonce(
  fixed_iv : Bytes,
  epoch : UInt16,
  sequence_number : UInt64,
) -> Bytes raise DtlsError {
  if fixed_iv.length() != 12 {
    raise InvalidRecord("ChaCha20-Poly1305 fixed IV must contain 12 bytes")
  }
  let result = fixed_iv.to_array()
  let sequence = epoch_sequence(epoch, sequence_number)
  for index = 0; index < 8; index = index + 1 {
    result[index + 4] = result[index + 4] ^ sequence[index]
  }
  Bytes::from_array(result)
}

///|
fn endpoint_aead_nonce(
  fixed_iv : Bytes,
  explicit : Bytes,
  epoch : UInt16,
  sequence_number : UInt64,
  has_explicit : Bool,
) -> Bytes raise DtlsError {
  if has_explicit {
    explicit_nonce(fixed_iv, explicit)
  } else {
    chacha_nonce(fixed_iv, epoch, sequence_number)
  }
}

///|
fn cbc_plaintext(
  provider : @crypto.Provider,
  mac_key : @crypto.Secret,
  content_type : ContentType,
  version : ProtocolVersion,
  epoch : UInt16,
  sequence_number : UInt64,
  plaintext : Bytes,
) -> Bytes raise DtlsError {
  let aad = record_aad(
    content_type,
    version,
    epoch,
    sequence_number,
    plaintext.length(),
  )
  let mac = crypto_operation(() => {
    provider.hmac(Sha1, mac_key, append_bytes(aad, plaintext))
  })
  let result = append_bytes(plaintext, mac).to_array()
  let padding_count = 16 - result.length() % 16
  let padding = (padding_count - 1).to_byte()
  for index = 0; index < padding_count; index = index + 1 {
    result.push(padding)
  }
  Bytes::from_array(result)
}

///|
fn RecordCipher::seal(
  self : RecordCipher,
  content_type~ : ContentType,
  version? : ProtocolVersion = Dtls12,
  epoch~ : UInt16,
  sequence_number~ : UInt64,
  plaintext~ : Bytes,
) -> Record raise DtlsError {
  if sequence_number > 0x0000ffffffffffffUL {
    raise InvalidRecord("DTLS record sequence number exceeds 48 bits")
  }
  let payload = match self.mode {
    AeadCipher(algorithm~, tag_length~, explicit_nonce=has_explicit) => {
      let explicit = if has_explicit {
        epoch_sequence(epoch, sequence_number)
      } else {
        b""
      }
      let aad = record_aad(
        content_type,
        version,
        epoch,
        sequence_number,
        plaintext.length(),
      )
      let nonce = endpoint_aead_nonce(
        self.write_iv,
        explicit,
        epoch,
        sequence_number,
        has_explicit,
      )
      let encrypted = crypto_operation(() => {
        self.provider.aead_seal(
          algorithm,
          self.write_key,
          nonce,
          aad,
          plaintext,
          tag_length~,
        )
      })
      append_three(explicit, encrypted.ciphertext(), encrypted.tag())
    }
    CbcCipher => {
      let iv = crypto_operation(() => self.provider.random_bytes(16))
      let padded = cbc_plaintext(
        self.provider,
        self.write_mac_key,
        content_type,
        version,
        epoch,
        sequence_number,
        plaintext,
      )
      let encrypted = crypto_operation(() => {
        self.provider.cipher_encrypt(Aes256Cbc, self.write_key, iv, padded)
      })
      append_bytes(iv, encrypted)
    }
  }
  Record::new(content_type~, version~, epoch~, sequence_number~, payload~)
}

///|
fn RecordCipher::mark_replay(
  self : RecordCipher,
  epoch : UInt16,
  sequence_number : UInt64,
) -> Unit raise DtlsError {
  let window = match self.replay_windows.get(epoch) {
    Some(window) => window
    None => {
      let window = @replay.ReplayWindow::new(width=self.replay_window) catch {
        InvalidWidth(width) =>
          raise InvalidRecord("invalid replay window width \{width}")
      }
      self.replay_windows[epoch] = window
      window
    }
  }
  match window.check_and_mark(sequence_number) {
    AcceptedNew | AcceptedOutOfOrder => ()
    Duplicate | TooOld => raise ReplayRejected
  }
}

///|
fn RecordCipher::open_aead(
  self : RecordCipher,
  record : Record,
  algorithm : @crypto.AeadAlgorithm,
  tag_length : Int,
  has_explicit : Bool,
) -> Bytes raise DtlsError {
  let payload = record.payload
  let explicit_length = if has_explicit { 8 } else { 0 }
  if payload.length() < explicit_length + tag_length {
    raise InvalidRecord("encrypted DTLS record is shorter than nonce and tag")
  }
  let explicit = payload[0:explicit_length].to_owned()
  let ciphertext_length = payload.length() - explicit_length - tag_length
  let ciphertext = payload[explicit_length:explicit_length + ciphertext_length].to_owned()
  let tag = payload[payload.length() - tag_length:].to_owned()
  let aad = record_aad(
    record.header.content_type,
    record.header.version,
    record.header.epoch,
    record.header.sequence_number,
    ciphertext_length,
  )
  let packet = @crypto.AeadSealed::new(ciphertext~, tag~) catch {
    InvalidLength(length) =>
      raise InvalidRecord("invalid DTLS authentication tag length \{length}")
    CryptoUnavailable(message) => raise CryptoUnavailable(message)
    OperationFailed(message) => raise CryptoUnavailable(message)
  }
  let nonce = endpoint_aead_nonce(
    self.read_iv,
    explicit,
    record.header.epoch,
    record.header.sequence_number,
    has_explicit,
  )
  self.provider.aead_open(algorithm, self.read_key, nonce, aad, packet) catch {
    OperationFailed(_) =>
      raise InvalidRecord("DTLS record authentication failed")
    CryptoUnavailable(message) => raise CryptoUnavailable(message)
    InvalidLength(length) =>
      raise InvalidRecord("invalid DTLS cipher input length \{length}")
  }
}

///|
fn RecordCipher::open_cbc(
  self : RecordCipher,
  record : Record,
) -> Bytes raise DtlsError {
  let payload = record.payload
  if payload.length() < 32 || (payload.length() - 16) % 16 != 0 {
    raise InvalidRecord("invalid DTLS AES-CBC record length")
  }
  let iv = payload[0:16].to_owned()
  let encrypted = payload[16:].to_owned()
  let padded = self.provider.cipher_decrypt(
    Aes256Cbc,
    self.read_key,
    iv,
    encrypted,
  ) catch {
    OperationFailed(_) => raise InvalidRecord("DTLS CBC decryption failed")
    CryptoUnavailable(message) => raise CryptoUnavailable(message)
    InvalidLength(length) =>
      raise InvalidRecord("invalid DTLS CBC input length \{length}")
  }
  if padded.is_empty() {
    raise InvalidRecord("DTLS CBC plaintext is empty")
  }
  let padding_count = padded[padded.length() - 1].to_int() + 1
  if padding_count > padded.length() {
    raise InvalidRecord("invalid DTLS CBC padding")
  }
  let mut valid_padding = true
  let padding_value = (padding_count - 1).to_byte()
  for byte in padded[padded.length() - padding_count:] {
    if byte != padding_value {
      valid_padding = false
    }
  }
  if !valid_padding {
    raise InvalidRecord("invalid DTLS CBC padding")
  }
  let authenticated_length = padded.length() - padding_count
  if authenticated_length < 20 {
    raise InvalidRecord("DTLS CBC plaintext omitted its MAC")
  }
  let plaintext_length = authenticated_length - 20
  let plaintext = padded[0:plaintext_length].to_owned()
  let received_mac = padded[plaintext_length:authenticated_length].to_owned()
  let aad = record_aad(
    record.header.content_type,
    record.header.version,
    record.header.epoch,
    record.header.sequence_number,
    plaintext_length,
  )
  let expected_mac = crypto_operation(() => {
    self.provider.hmac(Sha1, self.read_mac_key, append_bytes(aad, plaintext))
  })
  if !self.provider.constant_time_equal(expected_mac, received_mac) {
    raise InvalidRecord("DTLS CBC record authentication failed")
  }
  plaintext
}

///|
fn RecordCipher::open(
  self : RecordCipher,
  record : Record,
) -> Bytes raise DtlsError {
  let plaintext = match self.mode {
    AeadCipher(algorithm~, tag_length~, explicit_nonce=has_explicit) =>
      self.open_aead(record, algorithm, tag_length, has_explicit)
    CbcCipher => self.open_cbc(record)
  }
  self.mark_replay(record.header.epoch, record.header.sequence_number)
  plaintext
}