///|
#cfg(target="native")
priv suberror RsaPssError {
  RsaPssBadPublicKey
  RsaPssSignatureOutOfRange
  RsaPssVerificationFailed
} derive(Debug, ToJson)

///|
#cfg(target="native")
#warnings("-unused_field")
priv struct RsaPublicKey {
  modulus : Bytes
  exponent : Int
}

///|
#cfg(target="native")
fn bigint_normalize(value : Bytes) -> Bytes {
  let mut start = 0
  while start + 1 < value.length() && value[start] == b'\x00' {
    start = start + 1
  }
  tls13_copy_slice(value, start, value.length())
}

///|
#cfg(target="native")
fn bigint_compare(a : Bytes, b : Bytes) -> Int {
  let a = bigint_normalize(a)
  let b = bigint_normalize(b)
  if a.length() < b.length() {
    -1
  } else if a.length() > b.length() {
    1
  } else {
    for i in 0.. b[i] {
        return 1
      }
    }
    0
  }
}

///|
#cfg(target="native")
fn bigint_add(a : Bytes, b : Bytes) -> Bytes {
  let max_len = if a.length() > b.length() { a.length() } else { b.length() }
  let out = FixedArray::make(max_len + 1, b'\x00')
  let mut carry = 0
  for i in 0.. i { a[a.length() - 1 - i].to_int() } else { 0 }
    let bi = if b.length() > i { b[b.length() - 1 - i].to_int() } else { 0 }
    let sum = ai + bi + carry
    out[max_len - i] = (sum & 0xff).to_byte()
    carry = sum >> 8
  }
  out[0] = carry.to_byte()
  bigint_normalize(out.unsafe_reinterpret_as_bytes())
}

///|
#cfg(target="native")
fn bigint_sub(a : Bytes, b : Bytes) -> Bytes raise {
  guard bigint_compare(a, b) >= 0 else { raise RsaPssBadPublicKey }
  let max_len = a.length()
  let out = FixedArray::make(max_len, b'\x00')
  let mut borrow = 0
  for i in 0.. i { b[b.length() - 1 - i].to_int() } else { 0 }
    let mut diff = ai - bi - borrow
    if diff < 0 {
      diff = diff + 256
      borrow = 1
    } else {
      borrow = 0
    }
    out[max_len - 1 - i] = diff.to_byte()
  }
  bigint_normalize(out.unsafe_reinterpret_as_bytes())
}

///|
#cfg(target="native")
fn bigint_add_mod(a : Bytes, b : Bytes, modulus : Bytes) -> Bytes raise {
  let mut sum = bigint_add(a, b)
  if bigint_compare(sum, modulus) >= 0 {
    sum = bigint_sub(sum, modulus)
  }
  sum
}

///|
#cfg(target="native")
fn bigint_mul_mod(a : Bytes, b : Bytes, modulus : Bytes) -> Bytes raise {
  let mut result = b"\x00"
  let mut addend = bigint_normalize(a)
  for byte_index = b.length() - 1; byte_index >= 0; byte_index = byte_index - 1 {
    let byte = b[byte_index].to_int()
    for bit in 0..<8 {
      if ((byte >> bit) & 1) == 1 {
        result = bigint_add_mod(result, addend, modulus)
      }
      addend = bigint_add_mod(addend, addend, modulus)
    }
  }
  result
}

///|
#cfg(target="native")
fn bigint_pow_mod(base : Bytes, exponent : Int, modulus : Bytes) -> Bytes raise {
  let mut result = b"\x01"
  let mut base = bigint_normalize(base)
  let mut exponent = exponent
  while exponent > 0 {
    if (exponent & 1) == 1 {
      result = bigint_mul_mod(result, base, modulus)
    }
    exponent = exponent >> 1
    if exponent > 0 {
      base = bigint_mul_mod(base, base, modulus)
    }
  }
  result
}

///|
#cfg(target="native")
fn bigint_left_pad(value : Bytes, len : Int) -> Bytes raise {
  let value = bigint_normalize(value)
  guard value.length() <= len else { raise RsaPssBadPublicKey }
  let out = FixedArray::make(len, b'\x00')
  let start = len - value.length()
  for i in 0.. Bytes raise {
  x509_der_expect_tag(element, 0x02)
  guard element.content_start < element.content_end else {
    raise RsaPssBadPublicKey
  }
  let value = x509_der_content(data, element)
  bigint_normalize(value)
}

///|
#cfg(target="native")
fn rsa_der_integer_int(data : Bytes, element : X509DerElement) -> Int raise {
  let value = rsa_der_integer_bytes(data, element)
  guard value.length() <= 4 else { raise RsaPssBadPublicKey }
  let mut out = 0
  for b in value {
    out = (out << 8) | b.to_int()
  }
  out
}

///|
#cfg(target="native")
fn rsa_parse_public_key(public_key : Bytes) -> RsaPublicKey raise {
  let root = x509_der_read_element(public_key, 0)
  x509_der_expect_tag(root, x509_tag_sequence)
  guard root.end == public_key.length() else { raise RsaPssBadPublicKey }
  let fields = x509_der_children(public_key, root)
  guard fields.length() >= 2 else { raise RsaPssBadPublicKey }
  {
    modulus: rsa_der_integer_bytes(public_key, fields[0]),
    exponent: rsa_der_integer_int(public_key, fields[1]),
  }
}

///|
#cfg(target="native")
fn rsa_modulus_bit_length(modulus : Bytes) -> Int raise {
  let modulus = bigint_normalize(modulus)
  guard modulus.length() > 0 else { raise RsaPssBadPublicKey }
  let first = modulus[0].to_int()
  let mut bits = 8
  while bits > 0 && ((first >> (bits - 1)) & 1) == 0 {
    bits = bits - 1
  }
  (modulus.length() - 1) * 8 + bits
}

///|
#cfg(target="native")
fn rsa_i2osp(value : Bytes, len : Int) -> Bytes raise {
  bigint_left_pad(value, len)
}

///|
#cfg(target="native")
fn rsa_public_operation(signature : Bytes, key : RsaPublicKey) -> Bytes raise {
  guard bigint_compare(signature, key.modulus) < 0 else {
    raise RsaPssSignatureOutOfRange
  }
  let value = bigint_pow_mod(signature, key.exponent, key.modulus)
  rsa_i2osp(value, key.modulus.length())
}

///|
#cfg(target="native")
fn rsa_mgf1_sha256(seed : Bytes, len : Int) -> Bytes {
  let out = @buffer.new()
  let mut counter = 0
  while out.length() < len {
    let block = @buffer.new()
    block.write_bytes(seed)
    block.write_byte(((counter >> 24) & 0xff).to_byte())
    block.write_byte(((counter >> 16) & 0xff).to_byte())
    block.write_byte(((counter >> 8) & 0xff).to_byte())
    block.write_byte((counter & 0xff).to_byte())
    out.write_bytes(tls13_sha256(block.contents()))
    counter = counter + 1
  }
  tls13_copy_slice(out.contents(), 0, len)
}

///|
#cfg(target="native")
fn rsa_pss_hash_message(message : Bytes, salt : Bytes) -> Bytes {
  let out = @buffer.new()
  out.write_bytes(Bytes::make(8, b'\x00'))
  out.write_bytes(tls13_sha256(message))
  out.write_bytes(salt)
  tls13_sha256(out.contents())
}

///|
#cfg(target="native")
fn rsa_pss_verify_encoded(
  message : Bytes,
  encoded : Bytes,
  em_bits : Int,
) -> Unit raise {
  let hash_len = 32
  let salt_len = 32
  guard encoded.length() >= hash_len + salt_len + 2 else {
    raise RsaPssVerificationFailed
  }
  guard encoded[encoded.length() - 1] == b'\xbc' else {
    raise RsaPssVerificationFailed
  }
  let db_len = encoded.length() - hash_len - 1
  let masked_db = tls13_copy_slice(encoded, 0, db_len)
  let h = tls13_copy_slice(encoded, db_len, db_len + hash_len)
  let unused_bits = encoded.length() * 8 - em_bits
  if unused_bits > 0 {
    let mask = 0xff >> unused_bits
    guard (masked_db[0].to_int() & (0xff ^ mask)) == 0 else {
      raise RsaPssVerificationFailed
    }
  }
  let db_mask = rsa_mgf1_sha256(h, db_len)
  let db = FixedArray::make(db_len, b'\x00')
  for i in 0.. 0 {
    db[0] = (db[0].to_int() & (0xff >> unused_bits)).to_byte()
  }
  let ps_len = db_len - salt_len - 1
  for i in 0.. Unit raise {
  let key = rsa_parse_public_key(public_key)
  let em_bits = rsa_modulus_bit_length(key.modulus) - 1
  let encoded = rsa_public_operation(signature, key)
  rsa_pss_verify_encoded(message, encoded, em_bits)
}