// Copyright 2025 The Chromium Authors
// Use of this source code is governed by a BSD-style license that can be
// found in the LICENSE file.

#include "crypto/keypair.h"

#include "base/containers/to_vector.h"
#include "base/logging.h"
#include "crypto/openssl_util.h"
#include "third_party/boringssl/src/include/openssl/base.h"
#include "third_party/boringssl/src/include/openssl/bn.h"
#include "third_party/boringssl/src/include/openssl/bytestring.h"
#include "third_party/boringssl/src/include/openssl/curve25519.h"
#include "third_party/boringssl/src/include/openssl/ec.h"
#include "third_party/boringssl/src/include/openssl/evp.h"
#include "third_party/boringssl/src/include/openssl/mem.h"
#include "third_party/boringssl/src/include/openssl/mlkem.h"
#include "third_party/boringssl/src/include/openssl/rsa.h"

namespace crypto::keypair {

namespace {

bssl::UniquePtr<EVP_PKEY> GenerateRsa(size_t bits) {
  OpenSSLErrStackTracer err_tracer(FROM_HERE);
  bssl::UniquePtr<EVP_PKEY> key(EVP_RSA_gen(bits));
  CHECK(key);
  return key;
}

bssl::UniquePtr<EVP_PKEY> GeneratePkey(const EVP_PKEY_ALG* alg) {
  OpenSSLErrStackTracer err_tracer(FROM_HERE);
  bssl::UniquePtr<EVP_PKEY> key(EVP_PKEY_generate_from_alg(alg));
  CHECK(key);
  return key;
}

base::span<const EVP_PKEY_ALG* const> SupportedEvpAlgorithms() {
  static const EVP_PKEY_ALG* kAlgs[] = {
      EVP_pkey_rsa(),        EVP_pkey_ec_p256(),   EVP_pkey_ec_p384(),
      EVP_pkey_ec_p521(),    EVP_pkey_ed25519(),   EVP_pkey_x25519(),
      EVP_pkey_ml_dsa_44(),  EVP_pkey_ml_dsa_65(), EVP_pkey_ml_dsa_87(),
      EVP_pkey_ml_kem_768(),
  };
  return kAlgs;
}

std::vector<uint8_t> CBBToVector(CBB* cbb) {
  return base::ToVector(CbbAsSpan(cbb));
}

std::vector<uint8_t> ExportEVPPublicKey(EVP_PKEY* pkey) {
  OpenSSLErrStackTracer err_tracer(FROM_HERE);
  bssl::ScopedCBB cbb;

  CHECK(CBB_init(cbb.get(), 0));
  CHECK(EVP_marshal_public_key(cbb.get(), pkey));
  return CBBToVector(cbb.get());
}

bssl::UniquePtr<EVP_PKEY> EVP_PKEYFromEcPoint(const EVP_PKEY_ALG* alg,
                                              base::span<const uint8_t> p) {
  if (p.empty()) {
    return nullptr;
  }
  // This function supports both uncompressed and compressed.
  return bssl::UniquePtr<EVP_PKEY>(
      p[0] == 0x04
          ? EVP_PKEY_from_ec_uncompressed_point(alg, p.data(), p.size())
          : EVP_PKEY_from_ec_compressed_point(alg, p.data(), p.size()));
}

std::vector<uint8_t> EvpToUncompressedX962Point(const EVP_PKEY* key) {
  OpenSSLErrStackTracer err_tracer(FROM_HERE);
  bssl::ScopedCBB cbb;
  CHECK(CBB_init(cbb.get(), 255));
  CHECK(EVP_PKEY_marshal_ec_uncompressed_point(cbb.get(), key));
  return CBBToVector(cbb.get());
}

std::vector<uint8_t> EvpToRawPublicKey(EVP_PKEY* key) {
  size_t len = 0;
  CHECK(EVP_PKEY_get_raw_public_key(key, nullptr, &len));
  std::vector<uint8_t> result(len);
  CHECK(EVP_PKEY_get_raw_public_key(key, result.data(), &len));
  CHECK(len == result.size());
  return result;
}

}  // namespace

PrivateKey::PrivateKey(bssl::UniquePtr<EVP_PKEY> key, crypto::SubtlePassKey)
    : PrivateKey(std::move(key)) {}
PrivateKey::~PrivateKey() = default;
PrivateKey::PrivateKey(PrivateKey&& other) = default;
PrivateKey::PrivateKey(const PrivateKey& other)
    : key_(bssl::UpRef(const_cast<PrivateKey&>(other).key())) {}
PrivateKey& PrivateKey::operator=(PrivateKey&& other) = default;
PrivateKey& PrivateKey::operator=(const PrivateKey& other) {
  key_ = bssl::UpRef(const_cast<PrivateKey&>(other).key());
  return *this;
}

// static
PrivateKey PrivateKey::GenerateRsa2048() {
  return PrivateKey(GenerateRsa(2048));
}

// static
PrivateKey PrivateKey::GenerateEcP256() {
  return PrivateKey(GeneratePkey(EVP_pkey_ec_p256()));
}

// static
PrivateKey PrivateKey::GenerateEcP384() {
  return PrivateKey(GeneratePkey(EVP_pkey_ec_p384()));
}

// static
PrivateKey PrivateKey::GenerateEd25519() {
  return PrivateKey(GeneratePkey(EVP_pkey_ed25519()));
}

// static
PrivateKey PrivateKey::GenerateX25519() {
  return PrivateKey(GeneratePkey(EVP_pkey_x25519()));
}

// static
PrivateKey PrivateKey::GenerateMldsa44() {
  return PrivateKey(GeneratePkey(EVP_pkey_ml_dsa_44()));
}

// static
PrivateKey PrivateKey::GenerateMldsa65() {
  return PrivateKey(GeneratePkey(EVP_pkey_ml_dsa_65()));
}

// static
PrivateKey PrivateKey::GenerateMldsa87() {
  return PrivateKey(GeneratePkey(EVP_pkey_ml_dsa_87()));
}

// static
PrivateKey PrivateKey::GenerateMlkem768() {
  return PrivateKey(GeneratePkey(EVP_pkey_ml_kem_768()));
}

// static
std::optional<PrivateKey> PrivateKey::FromPrivateKeyInfo(
    base::span<const uint8_t> pki) {
  OpenSSLErrStackTracer err_tracer(FROM_HERE);

  auto algs = SupportedEvpAlgorithms();
  bssl::UniquePtr<EVP_PKEY> pkey(EVP_PKEY_from_private_key_info(
      pki.data(), pki.size(), algs.data(), algs.size()));
  if (!pkey) {
    LOG(WARNING) << "Malformed PrivateKeyInfo or trailing data";
    return std::nullopt;
  }

  return std::optional<PrivateKey>(PrivateKey(std::move(pkey)));
}

// static
std::optional<PrivateKey> PrivateKey::FromRSAPrivateKey(
    base::span<const uint8_t> key) {
  OpenSSLErrStackTracer err_tracer(FROM_HERE);
  bssl::UniquePtr<EVP_PKEY> pkey(
      EVP_PKEY_from_rsa_private_key(EVP_pkey_rsa(), key.data(), key.size()));
  if (!pkey) {
    return std::nullopt;
  }

  return PrivateKey(std::move(pkey));
}

// static
std::optional<PrivateKey> PrivateKey::FromEcP256PrivateKey(
    base::span<const uint8_t> key) {
  OpenSSLErrStackTracer err_tracer(FROM_HERE);

  CBS cbs(key);
  bssl::UniquePtr<EC_KEY> ec(EC_KEY_parse_private_key(&cbs, EC_group_p256()));
  if (!ec || CBS_len(&cbs) != 0) {
    return std::nullopt;
  }

  bssl::UniquePtr<EVP_PKEY> pkey(EVP_PKEY_new());
  CHECK(pkey);
  CHECK(EVP_PKEY_set1_EC_KEY(pkey.get(), ec.get()));

  return PrivateKey(std::move(pkey));
}

// static
std::optional<PrivateKey> PrivateKey::FromEcP256PrivateScalar(
    base::span<const uint8_t> key) {
  OpenSSLErrStackTracer err_tracer(FROM_HERE);
  bssl::UniquePtr<EVP_PKEY> pkey(EVP_PKEY_from_ec_private_scalar(
      EVP_pkey_ec_p256(), key.data(), key.size()));
  if (!pkey) {
    return std::nullopt;
  }

  return PrivateKey(std::move(pkey));
}

// static
PrivateKey PrivateKey::FromEd25519PrivateKey(
    base::span<const uint8_t, 32> key) {
  OpenSSLErrStackTracer err_tracer(FROM_HERE);
  bssl::UniquePtr<EVP_PKEY> pkey(EVP_PKEY_from_raw_private_key(
      EVP_pkey_ed25519(), key.data(), key.size()));
  CHECK(pkey);
  return PrivateKey(std::move(pkey));
}

// static
PrivateKey PrivateKey::FromX25519PrivateKey(base::span<const uint8_t, 32> key) {
  OpenSSLErrStackTracer err_tracer(FROM_HERE);
  bssl::UniquePtr<EVP_PKEY> pkey(
      EVP_PKEY_from_raw_private_key(EVP_pkey_x25519(), key.data(), key.size()));
  CHECK(pkey);
  return PrivateKey(std::move(pkey));
}

// static
PrivateKey PrivateKey::FromMldsa44PrivateKey(
    base::span<const uint8_t, 32> seed) {
  bssl::UniquePtr<EVP_PKEY> pkey(EVP_PKEY_from_private_seed(
      EVP_pkey_ml_dsa_44(), seed.data(), seed.size()));
  CHECK(pkey);
  return PrivateKey(std::move(pkey));
}

// static
PrivateKey PrivateKey::FromMldsa65PrivateKey(
    base::span<const uint8_t, 32> seed) {
  bssl::UniquePtr<EVP_PKEY> pkey(EVP_PKEY_from_private_seed(
      EVP_pkey_ml_dsa_65(), seed.data(), seed.size()));
  CHECK(pkey);
  return PrivateKey(std::move(pkey));
}

// static
PrivateKey PrivateKey::FromMldsa87PrivateKey(
    base::span<const uint8_t, 32> seed) {
  bssl::UniquePtr<EVP_PKEY> pkey(EVP_PKEY_from_private_seed(
      EVP_pkey_ml_dsa_87(), seed.data(), seed.size()));
  CHECK(pkey);
  return PrivateKey(std::move(pkey));
}

// static
PrivateKey PrivateKey::FromMlkem768PrivateKey(
    base::span<const uint8_t, 64> key) {
  static_assert(std::size(key) == MLKEM_SEED_BYTES);

  bssl::UniquePtr<EVP_PKEY> pkey(EVP_PKEY_from_private_seed(
      EVP_pkey_ml_kem_768(), key.data(), key.size()));
  CHECK(pkey);
  return PrivateKey(std::move(pkey));
}

std::vector<uint8_t> PrivateKey::ToPrivateKeyInfo() const {
  OpenSSLErrStackTracer err_tracer(FROM_HERE);
  bssl::ScopedCBB cbb;

  CHECK(CBB_init(cbb.get(), 0));
  CHECK(EVP_marshal_private_key(cbb.get(), key_.get()));
  return CBBToVector(cbb.get());
}

std::vector<uint8_t> PrivateKey::ToRSAPrivateKey() const {
  CHECK(IsRsa());
  OpenSSLErrStackTracer err_tracer(FROM_HERE);
  bssl::ScopedCBB cbb;

  CHECK(CBB_init(cbb.get(), 0));
  CHECK(EVP_PKEY_marshal_rsa_private_key(cbb.get(), key_.get()));
  return CBBToVector(cbb.get());
}

std::vector<uint8_t> PrivateKey::ToEcP256PrivateKey() const {
  CHECK(IsEcP256());
  OpenSSLErrStackTracer err_tracer(FROM_HERE);
  bssl::ScopedCBB cbb;

  CHECK(CBB_init(cbb.get(), 0));
  EC_KEY* ec = EVP_PKEY_get0_EC_KEY(key_.get());
  CHECK(ec);
  CHECK(EC_KEY_marshal_private_key(cbb.get(), ec,
                                   EC_PKEY_NO_PARAMETERS | EC_PKEY_NO_PUBKEY));
  return CBBToVector(cbb.get());
}

std::array<uint8_t, 32> PrivateKey::ToEcP256PrivateScalar() const {
  CHECK(IsEcP256());
  std::array<uint8_t, 32> result;
  CBB cbb;
  CBB_init_fixed(&cbb, result.data(), result.size());
  CHECK(EVP_PKEY_marshal_ec_private_scalar(&cbb, key_.get()));
  CHECK(CBB_len(&cbb) == result.size());
  return result;
}

std::array<uint8_t, 32> PrivateKey::ToEd25519PrivateKey() const {
  CHECK(IsEd25519());
  std::array<uint8_t, 32> result;
  size_t len = std::size(result);
  CHECK(EVP_PKEY_get_raw_private_key(key_.get(), result.data(), &len));
  CHECK(len == std::size(result));
  return result;
}

std::array<uint8_t, 32> PrivateKey::ToX25519PrivateKey() const {
  CHECK(IsX25519());
  std::array<uint8_t, 32> result;
  size_t len = std::size(result);
  CHECK(EVP_PKEY_get_raw_private_key(key_.get(), result.data(), &len));
  CHECK(len == std::size(result));
  return result;
}

std::array<uint8_t, 32> PrivateKey::ToMldsa44PrivateKey() const {
  CHECK(IsMldsa44());
  std::array<uint8_t, 32> result;
  size_t len = std::size(result);
  CHECK(EVP_PKEY_get_private_seed(key_.get(), result.data(), &len));
  CHECK(len == std::size(result));
  return result;
}

std::array<uint8_t, 32> PrivateKey::ToMldsa65PrivateKey() const {
  CHECK(IsMldsa65());
  std::array<uint8_t, 32> result;
  size_t len = std::size(result);
  CHECK(EVP_PKEY_get_private_seed(key_.get(), result.data(), &len));
  CHECK(len == std::size(result));
  return result;
}

std::array<uint8_t, 32> PrivateKey::ToMldsa87PrivateKey() const {
  CHECK(IsMldsa87());
  std::array<uint8_t, 32> result;
  size_t len = std::size(result);
  CHECK(EVP_PKEY_get_private_seed(key_.get(), result.data(), &len));
  CHECK(len == std::size(result));
  return result;
}

std::array<uint8_t, 64> PrivateKey::ToMlkem768PrivateKey() const {
  CHECK(IsMlkem768());
  std::array<uint8_t, 64> result;
  static_assert(std::size(result) == MLKEM_SEED_BYTES);
  size_t len = std::size(result);
  CHECK(EVP_PKEY_get_private_seed(key_.get(), result.data(), &len));
  CHECK(len == std::size(result));
  return result;
}

std::vector<uint8_t> PrivateKey::ToSubjectPublicKeyInfo() const {
  return ExportEVPPublicKey(key_.get());
}

std::vector<uint8_t> PrivateKey::ToUncompressedX962Point() const {
  return EvpToUncompressedX962Point(key_.get());
}

std::array<uint8_t, 32> PrivateKey::ToEd25519PublicKey() const {
  CHECK(IsEd25519());
  std::array<uint8_t, 32> result;
  size_t len = std::size(result);
  CHECK(EVP_PKEY_get_raw_public_key(key_.get(), result.data(), &len));
  CHECK(len == std::size(result));
  return result;
}

std::array<uint8_t, 32> PrivateKey::ToX25519PublicKey() const {
  CHECK(IsX25519());
  std::array<uint8_t, 32> result;
  size_t len = std::size(result);
  CHECK(EVP_PKEY_get_raw_public_key(key_.get(), result.data(), &len));
  CHECK(len == std::size(result));
  return result;
}

std::vector<uint8_t> PrivateKey::ToMldsa44PublicKey() const {
  CHECK(IsMldsa44());
  return EvpToRawPublicKey(key_.get());
}

std::vector<uint8_t> PrivateKey::ToMldsa65PublicKey() const {
  CHECK(IsMldsa65());
  return EvpToRawPublicKey(key_.get());
}

std::vector<uint8_t> PrivateKey::ToMldsa87PublicKey() const {
  CHECK(IsMldsa87());
  return EvpToRawPublicKey(key_.get());
}

std::array<uint8_t, 1184> PrivateKey::ToMlkem768PublicKey() const {
  CHECK(IsMlkem768());
  std::array<uint8_t, 1184> result;
  static_assert(std::size(result) == MLKEM768_PUBLIC_KEY_BYTES);
  size_t len = std::size(result);
  CHECK(EVP_PKEY_get_raw_public_key(key_.get(), result.data(), &len));
  CHECK(len == std::size(result));
  return result;
}

bool PrivateKey::IsRsa() const {
  return EVP_PKEY_id(key_.get()) == EVP_PKEY_RSA;
}

bool PrivateKey::IsEc() const {
  return EVP_PKEY_id(key_.get()) == EVP_PKEY_EC;
}

bool PrivateKey::IsEd25519() const {
  return EVP_PKEY_id(key_.get()) == EVP_PKEY_ED25519;
}

bool PrivateKey::IsX25519() const {
  return EVP_PKEY_id(key_.get()) == EVP_PKEY_X25519;
}

bool PrivateKey::IsMldsa44() const {
  return EVP_PKEY_id(key_.get()) == EVP_PKEY_ML_DSA_44;
}

bool PrivateKey::IsMldsa65() const {
  return EVP_PKEY_id(key_.get()) == EVP_PKEY_ML_DSA_65;
}

bool PrivateKey::IsMldsa87() const {
  return EVP_PKEY_id(key_.get()) == EVP_PKEY_ML_DSA_87;
}

bool PrivateKey::IsMlkem768() const {
  return EVP_PKEY_id(key_.get()) == EVP_PKEY_ML_KEM_768;
}

bool PrivateKey::IsEcP256() const {
  return EVP_PKEY_get_ec_curve_nid(key_.get()) == NID_X9_62_prime256v1;
}

bool PrivateKey::IsEcP384() const {
  return EVP_PKEY_get_ec_curve_nid(key_.get()) == NID_secp384r1;
}

PrivateKey::PrivateKey(bssl::UniquePtr<EVP_PKEY> key) : key_(std::move(key)) {}

PublicKey::PublicKey(bssl::UniquePtr<EVP_PKEY> key, crypto::SubtlePassKey)
    : PublicKey(std::move(key)) {}
PublicKey::~PublicKey() = default;
PublicKey::PublicKey(PublicKey&& other) = default;
PublicKey::PublicKey(const PublicKey& other)
    : key_(bssl::UpRef(const_cast<PublicKey&>(other).key())) {}
PublicKey& PublicKey::operator=(PublicKey&& other) = default;
PublicKey& PublicKey::operator=(const PublicKey& other) {
  key_ = bssl::UpRef(const_cast<PublicKey&>(other).key());
  return *this;
}

// static
PublicKey PublicKey::FromPrivateKey(const PrivateKey& key) {
  bssl::UniquePtr<EVP_PKEY> pub(EVP_PKEY_copy_public(key.key()));
  CHECK(pub);
  return PublicKey(std::move(pub));
}

// static
std::optional<PublicKey> PublicKey::FromSubjectPublicKeyInfo(
    base::span<const uint8_t> spki) {
  OpenSSLErrStackTracer err_tracer(FROM_HERE);

  auto algs = SupportedEvpAlgorithms();
  bssl::UniquePtr<EVP_PKEY> pkey(EVP_PKEY_from_subject_public_key_info(
      spki.data(), spki.size(), algs.data(), algs.size()));
  if (!pkey) {
    LOG(WARNING) << "Malformed SubjectPublicKeyInfo or trailing data";
    return std::nullopt;
  }

  return std::optional<PublicKey>(PublicKey(std::move(pkey)));
}

std::optional<PublicKey> PublicKey::FromRsaPublicKeyComponents(
    base::span<const uint8_t> n,
    base::span<const uint8_t> e) {
  bssl::UniquePtr<BIGNUM> bn_n(BN_bin2bn(n.data(), n.size(), nullptr));
  bssl::UniquePtr<BIGNUM> bn_e(BN_bin2bn(e.data(), e.size(), nullptr));
  if (!bn_n || !bn_e) {
    return std::nullopt;
  }

  bssl::UniquePtr<RSA> rsa(RSA_new_public_key(bn_n.get(), bn_e.get()));
  if (!rsa) {
    return std::nullopt;
  }

  // The only failure mode for EVP_PKEY_new() is memory allocation failures,
  // and the only failure mode for EVP_PKEY_set1_RSA() is being passed a null
  // key or RSA object.
  bssl::UniquePtr<EVP_PKEY> pkey(EVP_PKEY_new());
  CHECK(pkey);
  CHECK(EVP_PKEY_set1_RSA(pkey.get(), rsa.get()));
  return PublicKey(std::move(pkey));
}

// static
std::optional<PublicKey> PublicKey::FromEcP256Point(
    base::span<const uint8_t> p) {
  auto key = EVP_PKEYFromEcPoint(EVP_pkey_ec_p256(), p);
  if (!key) {
    return std::nullopt;
  }
  return PublicKey(std::move(key));
}

// static
std::optional<PublicKey> PublicKey::FromEcP384Point(
    base::span<const uint8_t> p) {
  auto key = EVP_PKEYFromEcPoint(EVP_pkey_ec_p384(), p);
  if (!key) {
    return std::nullopt;
  }
  return PublicKey(std::move(key));
}

// static
PublicKey PublicKey::FromEd25519PublicKey(base::span<const uint8_t, 32> key) {
  static_assert(std::size(key) == ED25519_PUBLIC_KEY_LEN);

  bssl::UniquePtr<EVP_PKEY> pkey(
      EVP_PKEY_from_raw_public_key(EVP_pkey_ed25519(), key.data(), key.size()));
  CHECK(pkey);
  return PublicKey(std::move(pkey));
}

// static
PublicKey PublicKey::FromX25519PublicKey(base::span<const uint8_t, 32> key) {
  static_assert(std::size(key) == X25519_PUBLIC_VALUE_LEN);

  bssl::UniquePtr<EVP_PKEY> pkey(
      EVP_PKEY_from_raw_public_key(EVP_pkey_x25519(), key.data(), key.size()));
  CHECK(pkey);
  return PublicKey(std::move(pkey));
}

// static
std::optional<PublicKey> PublicKey::FromMldsa44PublicKey(
    base::span<const uint8_t> key) {
  bssl::UniquePtr<EVP_PKEY> pkey(EVP_PKEY_from_raw_public_key(
      EVP_pkey_ml_dsa_44(), key.data(), key.size()));
  if (!pkey) {
    return std::nullopt;
  }
  return PublicKey(std::move(pkey));
}

// static
std::optional<PublicKey> PublicKey::FromMldsa65PublicKey(
    base::span<const uint8_t> key) {
  bssl::UniquePtr<EVP_PKEY> pkey(EVP_PKEY_from_raw_public_key(
      EVP_pkey_ml_dsa_65(), key.data(), key.size()));
  if (!pkey) {
    return std::nullopt;
  }
  return PublicKey(std::move(pkey));
}

// static
std::optional<PublicKey> PublicKey::FromMldsa87PublicKey(
    base::span<const uint8_t> key) {
  bssl::UniquePtr<EVP_PKEY> pkey(EVP_PKEY_from_raw_public_key(
      EVP_pkey_ml_dsa_87(), key.data(), key.size()));
  if (!pkey) {
    return std::nullopt;
  }
  return PublicKey(std::move(pkey));
}

// static
std::optional<PublicKey> PublicKey::FromMlkem768PublicKey(
    base::span<const uint8_t, 1184> key) {
  static_assert(std::size(key) == MLKEM768_PUBLIC_KEY_BYTES);

  bssl::UniquePtr<EVP_PKEY> pkey(EVP_PKEY_from_raw_public_key(
      EVP_pkey_ml_kem_768(), key.data(), key.size()));
  if (!pkey) {
    return std::nullopt;
  }
  return PublicKey(std::move(pkey));
}

std::vector<uint8_t> PublicKey::ToSubjectPublicKeyInfo() const {
  return ExportEVPPublicKey(key_.get());
}

std::vector<uint8_t> PublicKey::ToUncompressedX962Point() const {
  return EvpToUncompressedX962Point(key_.get());
}

std::array<uint8_t, 32> PublicKey::ToEd25519PublicKey() const {
  CHECK(IsEd25519());
  std::array<uint8_t, 32> result;
  size_t len = std::size(result);
  CHECK(EVP_PKEY_get_raw_public_key(key_.get(), result.data(), &len));
  CHECK(len == std::size(result));
  return result;
}

std::array<uint8_t, 32> PublicKey::ToX25519PublicKey() const {
  CHECK(IsX25519());
  std::array<uint8_t, 32> result;
  size_t len = std::size(result);
  CHECK(EVP_PKEY_get_raw_public_key(key_.get(), result.data(), &len));
  CHECK(len == std::size(result));
  return result;
}

std::vector<uint8_t> PublicKey::ToMldsa44PublicKey() const {
  CHECK(IsMldsa44());
  return EvpToRawPublicKey(key_.get());
}

std::vector<uint8_t> PublicKey::ToMldsa65PublicKey() const {
  CHECK(IsMldsa65());
  return EvpToRawPublicKey(key_.get());
}

std::vector<uint8_t> PublicKey::ToMldsa87PublicKey() const {
  CHECK(IsMldsa87());
  return EvpToRawPublicKey(key_.get());
}

std::array<uint8_t, 1184> PublicKey::ToMlkem768PublicKey() const {
  CHECK(IsMlkem768());
  std::array<uint8_t, 1184> result;
  static_assert(std::size(result) == MLKEM768_PUBLIC_KEY_BYTES);
  size_t len = std::size(result);
  CHECK(EVP_PKEY_get_raw_public_key(key_.get(), result.data(), &len));
  CHECK(len == std::size(result));
  return result;
}

std::vector<uint8_t> PublicKey::GetRsaExponent() const {
  CHECK(IsRsa());
  RSA* rsa = EVP_PKEY_get0_RSA(key_.get());
  const BIGNUM* e = RSA_get0_e(rsa);
  std::vector<uint8_t> result(BN_num_bytes(e));
  BN_bn2bin(e, result.data());
  return result;
}

std::vector<uint8_t> PublicKey::GetRsaModulus() const {
  CHECK(IsRsa());
  RSA* rsa = EVP_PKEY_get0_RSA(key_.get());
  const BIGNUM* n = RSA_get0_n(rsa);
  std::vector<uint8_t> result(BN_num_bytes(n));
  BN_bn2bin(n, result.data());
  return result;
}

bool PublicKey::IsRsa() const {
  return EVP_PKEY_id(key_.get()) == EVP_PKEY_RSA;
}

bool PublicKey::IsEc() const {
  return EVP_PKEY_id(key_.get()) == EVP_PKEY_EC;
}

bool PublicKey::IsEd25519() const {
  return EVP_PKEY_id(key_.get()) == EVP_PKEY_ED25519;
}

bool PublicKey::IsX25519() const {
  return EVP_PKEY_id(key_.get()) == EVP_PKEY_X25519;
}

bool PublicKey::IsMldsa44() const {
  return EVP_PKEY_id(key_.get()) == EVP_PKEY_ML_DSA_44;
}

bool PublicKey::IsMldsa65() const {
  return EVP_PKEY_id(key_.get()) == EVP_PKEY_ML_DSA_65;
}

bool PublicKey::IsMldsa87() const {
  return EVP_PKEY_id(key_.get()) == EVP_PKEY_ML_DSA_87;
}

bool PublicKey::IsMlkem768() const {
  return EVP_PKEY_id(key_.get()) == EVP_PKEY_ML_KEM_768;
}

bool PublicKey::IsEcP256() const {
  return EVP_PKEY_get_ec_curve_nid(key_.get()) == NID_X9_62_prime256v1;
}

bool PublicKey::IsEcP384() const {
  return EVP_PKEY_get_ec_curve_nid(key_.get()) == NID_secp384r1;
}

PublicKey::PublicKey(bssl::UniquePtr<EVP_PKEY> key) : key_(std::move(key)) {}

}  // namespace crypto::keypair
