// 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 "components/unexportable_keys/mojom/unexportable_key_service_proxied.h"

#include <memory>
#include <optional>
#include <type_traits>
#include <utility>
#include <vector>

#include "base/containers/to_vector.h"
#include "base/functional/bind.h"
#include "base/no_destructor.h"
#include "base/test/bind.h"
#include "base/test/gmock_expected_support.h"
#include "base/test/task_environment.h"
#include "base/test/test_future.h"
#include "base/types/expected.h"
#include "base/types/optional_util.h"
#include "base/unguessable_token.h"
#include "components/unexportable_keys/background_task_priority.h"
#include "components/unexportable_keys/mojom/unexportable_key_service.mojom.h"
#include "components/unexportable_keys/service_error.h"
#include "components/unexportable_keys/unexportable_key_id.h"
#include "crypto/signature_verifier.h"
#include "mojo/public/cpp/bindings/receiver.h"
#include "mojo/public/cpp/bindings/remote.h"
#include "testing/gmock/include/gmock/gmock.h"
#include "testing/gtest/include/gtest/gtest.h"

namespace unexportable_keys {

using ::base::test::ErrorIs;
using ::base::test::ValueIs;
using ::testing::ElementsAre;
using ::testing::ElementsAreArray;
using ::testing::IsEmpty;
using ::testing::UnorderedElementsAre;
using ::testing::UnorderedElementsAreArray;

namespace {

constexpr auto kTestSubjectPublicKeyInfo = std::to_array<uint8_t>({1, 2, 3, 4});
constexpr auto kTestWrappedKey = std::to_array<uint8_t>({5, 6, 7, 8});
constexpr std::string_view kTestKeyTag = "test_key_tag";

constexpr auto kTestWrappedAttestationKey =
    std::to_array<uint8_t>({0x11, 0x22, 0x33});
constexpr auto kTestAttestationAlgorithms =
    std::to_array<crypto::SignatureVerifier::SignatureAlgorithm>({
        crypto::SignatureVerifier::SignatureAlgorithm::RSA_PKCS1_SHA256,
        crypto::SignatureVerifier::SignatureAlgorithm::ECDSA_SHA256,
    });
constexpr auto kTestChallenge = std::to_array<uint8_t>({1, 2, 3});

const crypto::AttestationStatement& GetTestAttestationStatement() {
  static const base::NoDestructor<crypto::AttestationStatement> statement({
      .format = crypto::AttestationStatement::Format::kTpm,
      .statement = {0x01, 0x02},
      .signature = {0x03, 0x04},
  });
  return *statement;
}

mojom::NewKeyMetadataPtr ToNewKeyMetadata(CachedKeyData cache_data) {
  auto metadata = mojom::NewKeyMetadata::New();
  metadata->subject_public_key_info =
      std::move(cache_data.subject_public_key_info);
  metadata->wrapped_key = std::move(cache_data.wrapped_key);
  metadata->algorithm = cache_data.algorithm;
  metadata->key_tag = base::OptionalFromExpected(std::move(cache_data.key_tag));
  metadata->creation_time =
      base::OptionalFromExpected(cache_data.creation_time);
  return metadata;
}

template <typename MojomType, typename KeyIdType>
mojo::StructPtr<MojomType> ToMojomKeyDataImpl(KeyIdType key_id,
                                              CachedKeyData cache_data) {
  auto data = MojomType::New();
  data->key_id = key_id;
  data->metadata = ToNewKeyMetadata(std::move(cache_data));
  return data;
}

mojom::NewSigningKeyDataPtr ToMojomKeyData(UnexportableSigningKeyId key_id,
                                           CachedKeyData cache_data) {
  return ToMojomKeyDataImpl<mojom::NewSigningKeyData>(key_id,
                                                      std::move(cache_data));
}

mojom::NewAttestationKeyDataPtr ToMojomKeyData(
    UnexportableAttestationKeyId key_id,
    CachedKeyData cache_data) {
  return ToMojomKeyDataImpl<mojom::NewAttestationKeyData>(
      key_id, std::move(cache_data));
}

template <
    typename NewKeyDataPtrType,
    typename KeyIdType = decltype(std::declval<NewKeyDataPtrType>()->key_id)>
ServiceErrorOr<NewKeyDataPtrType> GenerateKeyImpl(
    std::optional<ServiceErrorOr<NewKeyDataPtrType>> response,
    base::span<const crypto::SignatureVerifier::SignatureAlgorithm>
        acceptable_algorithms) {
  if (response) {
    return std::move(*response);
  }
  if (acceptable_algorithms.empty()) {
    return base::unexpected(ServiceError::kAlgorithmNotSupported);
  }
  return ToMojomKeyData(
      KeyIdType(),
      {
          .subject_public_key_info = base::ToVector(kTestSubjectPublicKeyInfo),
          .wrapped_key = base::ToVector(kTestWrappedKey),
          .algorithm = acceptable_algorithms[0],
          .key_tag = std::string(kTestKeyTag),
      });
}

template <
    typename NewKeyDataPtrType,
    typename KeyIdType = decltype(std::declval<NewKeyDataPtrType>()->key_id)>
ServiceErrorOr<NewKeyDataPtrType> FromWrappedKeyImpl(
    std::optional<ServiceErrorOr<NewKeyDataPtrType>> response,
    base::span<const uint8_t> wrapped_key) {
  if (response) {
    return std::move(*response);
  }
  if (wrapped_key.empty()) {
    return base::unexpected(ServiceError::kKeyNotFound);
  }
  return ToMojomKeyData(
      KeyIdType(),
      {
          .subject_public_key_info = base::ToVector(kTestSubjectPublicKeyInfo),
          .wrapped_key = base::ToVector(wrapped_key),
          .algorithm =
              crypto::SignatureVerifier::SignatureAlgorithm::ECDSA_SHA256,
          .key_tag = std::string(kTestKeyTag),
      });
}

class FakeUnexportableKeyServiceProxy : public mojom::UnexportableKeyService {
 public:
  FakeUnexportableKeyServiceProxy() = default;
  ~FakeUnexportableKeyServiceProxy() override = default;

  void GenerateSigningKey(
      const std::vector<crypto::SignatureVerifier::SignatureAlgorithm>&
          acceptable_algorithms,
      BackgroundTaskPriority priority,
      GenerateSigningKeyCallback callback) override {
    std::move(callback).Run(GenerateKeyImpl(
        std::exchange(generate_response_, {}), acceptable_algorithms));
  }

  void FromWrappedSigningKey(const std::vector<uint8_t>& wrapped_key,
                             BackgroundTaskPriority priority,
                             FromWrappedSigningKeyCallback callback) override {
    std::move(callback).Run(FromWrappedKeyImpl(
        std::exchange(from_wrapped_response_, {}), wrapped_key));
  }

  void GenerateAttestationKey(
      const std::vector<crypto::SignatureVerifier::SignatureAlgorithm>&
          acceptable_algorithms,
      BackgroundTaskPriority priority,
      GenerateAttestationKeyCallback callback) override {
    std::move(callback).Run(
        GenerateKeyImpl(std::exchange(generate_attestation_response_, {}),
                        acceptable_algorithms));
  }

  void FromWrappedAttestationKey(
      const std::vector<uint8_t>& wrapped_key,
      BackgroundTaskPriority priority,
      FromWrappedAttestationKeyCallback callback) override {
    std::move(callback).Run(FromWrappedKeyImpl(
        std::exchange(from_wrapped_attestation_response_, {}), wrapped_key));
  }

  void Sign(const UnexportableSigningKeyId& key_id,
            const std::vector<uint8_t>& data,
            BackgroundTaskPriority priority,
            SignCallback callback) override {
    if (sign_response_) {
      std::move(callback).Run(std::move(sign_response_.value()));
      sign_response_.reset();
    } else if (data.empty()) {
      std::move(callback).Run(base::unexpected(ServiceError::kKeyNotFound));
    } else {
      std::vector<uint8_t> signature = {0x11, 0x22, 0x33, 0x44};
      std::move(callback).Run(std::move(signature));
    }
  }

  void Certify(const UnexportableAttestationKeyId& attestation_key_id,
               const UnexportableSigningKeyId& signing_key_id,
               const std::vector<uint8_t>& challenge,
               BackgroundTaskPriority priority,
               CertifyCallback callback) override {
    if (certify_response_) {
      std::move(callback).Run(*std::exchange(certify_response_, {}));
    } else {
      std::move(callback).Run(GetTestAttestationStatement());
    }
  }

  void GetAllKeysForGarbageCollection(
      BackgroundTaskPriority priority,
      GetAllKeysForGarbageCollectionCallback callback) override {
    if (get_all_keys_response_) {
      std::move(callback).Run(std::move(get_all_keys_response_.value()));
      get_all_keys_response_.reset();
    } else {
      std::move(callback).Run(
          base::ok(std::vector<mojom::NewSigningKeyDataPtr>()));
    }
  }

  void DeleteKeys(const std::vector<UnexportableSigningKeyId>& key_ids,
                  BackgroundTaskPriority priority,
                  DeleteKeysCallback callback) override {
    std::move(callback).Run(
        *std::exchange(delete_keys_response_, std::nullopt));
  }

  void DeleteAllKeys(DeleteAllKeysCallback callback) override {
    std::move(callback).Run(std::move(delete_all_keys_response_.value()));
    delete_all_keys_response_.reset();
  }

  void SetGenerateResponse(
      ServiceErrorOr<mojom::NewSigningKeyDataPtr> response) {
    generate_response_ = std::move(response);
  }

  void SetFromWrappedResponse(
      ServiceErrorOr<mojom::NewSigningKeyDataPtr> response) {
    from_wrapped_response_ = std::move(response);
  }

  void SetSignResponse(ServiceErrorOr<std::vector<uint8_t>> response) {
    sign_response_ = std::move(response);
  }

  void SetGetAllKeysForGarbageCollectionResponse(
      ServiceErrorOr<std::vector<mojom::NewSigningKeyDataPtr>> response) {
    get_all_keys_response_ = std::move(response);
  }

  void SetDeleteKeysResponse(ServiceErrorOr<uint64_t> response) {
    delete_keys_response_ = std::move(response);
  }

  void SetDeleteAllKeysResponse(ServiceErrorOr<uint64_t> response) {
    delete_all_keys_response_ = std::move(response);
  }

  void SetGenerateAttestationResponse(
      ServiceErrorOr<mojom::NewAttestationKeyDataPtr> response) {
    generate_attestation_response_ = std::move(response);
  }

  void SetFromWrappedAttestationResponse(
      ServiceErrorOr<mojom::NewAttestationKeyDataPtr> response) {
    from_wrapped_attestation_response_ = std::move(response);
  }

  void SetCertifyResponse(
      ServiceErrorOr<crypto::AttestationStatement> response) {
    certify_response_ = std::move(response);
  }

 private:
  std::optional<ServiceErrorOr<mojom::NewSigningKeyDataPtr>> generate_response_;
  std::optional<ServiceErrorOr<mojom::NewSigningKeyDataPtr>>
      from_wrapped_response_;
  std::optional<ServiceErrorOr<std::vector<uint8_t>>> sign_response_;
  std::optional<ServiceErrorOr<std::vector<mojom::NewSigningKeyDataPtr>>>
      get_all_keys_response_;
  std::optional<ServiceErrorOr<uint64_t>> delete_keys_response_;
  std::optional<ServiceErrorOr<uint64_t>> delete_all_keys_response_;
  std::optional<ServiceErrorOr<mojom::NewAttestationKeyDataPtr>>
      generate_attestation_response_;
  std::optional<ServiceErrorOr<mojom::NewAttestationKeyDataPtr>>
      from_wrapped_attestation_response_;
  std::optional<ServiceErrorOr<crypto::AttestationStatement>> certify_response_;
};

class UnexportableKeyServiceProxiedTest : public ::testing::Test {
 protected:
  UnexportableSigningKeyId GenerateSigningKeyOrDie() {
    base::test::TestFuture<ServiceErrorOr<UnexportableSigningKeyId>> future;
    proxied_service_.GenerateSigningKeySlowlyAsync(
        {crypto::SignatureVerifier::SignatureAlgorithm::RSA_PKCS1_SHA256},
        BackgroundTaskPriority::kUserVisible, future.GetCallback());
    return future.Get().value();
  }

  UnexportableAttestationKeyId GenerateAttestationKeyOrDie() {
    base::test::TestFuture<ServiceErrorOr<UnexportableAttestationKeyId>> future;
    proxied_service_.GenerateAttestationKeySlowlyAsync(
        {crypto::SignatureVerifier::SignatureAlgorithm::RSA_PKCS1_SHA256},
        BackgroundTaskPriority::kUserVisible, future.GetCallback());
    return future.Get().value();
  }

  base::test::TaskEnvironment task_environment_;
  FakeUnexportableKeyServiceProxy fake_service_;
  mojo::Receiver<mojom::UnexportableKeyService> receiver_{&fake_service_};
  UnexportableKeyServiceProxied proxied_service_{
      receiver_.BindNewPipeAndPassRemote()};
};

TEST_F(UnexportableKeyServiceProxiedTest, GenerateSigningKeySuccess) {
  base::test::TestFuture<ServiceErrorOr<UnexportableSigningKeyId>> future;
  std::vector<crypto::SignatureVerifier::SignatureAlgorithm> algos = {
      crypto::SignatureVerifier::SignatureAlgorithm::RSA_PKCS1_SHA256,
      crypto::SignatureVerifier::SignatureAlgorithm::ECDSA_SHA256};

  proxied_service_.GenerateSigningKeySlowlyAsync(
      algos, BackgroundTaskPriority::kUserVisible, future.GetCallback());

  ASSERT_OK_AND_ASSIGN(UnexportableSigningKeyId key_id, future.Get());

  EXPECT_THAT(proxied_service_.GetSubjectPublicKeyInfo(key_id),
              ValueIs(ElementsAreArray(kTestSubjectPublicKeyInfo)));
  EXPECT_THAT(proxied_service_.GetWrappedKey(key_id),
              ValueIs(ElementsAreArray(kTestWrappedKey)));
  EXPECT_THAT(
      proxied_service_.GetAlgorithm(key_id),
      ValueIs(crypto::SignatureVerifier::SignatureAlgorithm::RSA_PKCS1_SHA256));
  EXPECT_THAT(proxied_service_.GetKeyTag(key_id), ValueIs(kTestKeyTag));
}

TEST_F(UnexportableKeyServiceProxiedTest, GenerateSigningKeyError) {
  fake_service_.SetGenerateResponse(
      base::unexpected(ServiceError::kCryptoApiFailed));

  base::test::TestFuture<ServiceErrorOr<UnexportableSigningKeyId>> future;
  std::vector<crypto::SignatureVerifier::SignatureAlgorithm> algos = {
      crypto::SignatureVerifier::SignatureAlgorithm::RSA_PKCS1_SHA256};

  proxied_service_.GenerateSigningKeySlowlyAsync(
      algos, BackgroundTaskPriority::kUserVisible, future.GetCallback());

  EXPECT_THAT(future.Get(), ErrorIs(ServiceError::kCryptoApiFailed));
}

TEST_F(UnexportableKeyServiceProxiedTest, GenerateSigningKeyEmptyAlgorithms) {
  base::test::TestFuture<ServiceErrorOr<UnexportableSigningKeyId>> future;
  std::vector<crypto::SignatureVerifier::SignatureAlgorithm> algos = {};

  proxied_service_.GenerateSigningKeySlowlyAsync(
      algos, BackgroundTaskPriority::kUserVisible, future.GetCallback());

  EXPECT_THAT(future.Get(), ErrorIs(ServiceError::kAlgorithmNotSupported));
}

TEST_F(UnexportableKeyServiceProxiedTest, GenerateKeyCollision) {
  base::test::TestFuture<ServiceErrorOr<UnexportableSigningKeyId>> future1;
  std::vector<crypto::SignatureVerifier::SignatureAlgorithm> algos = {
      crypto::SignatureVerifier::SignatureAlgorithm::RSA_PKCS1_SHA256};
  proxied_service_.GenerateSigningKeySlowlyAsync(
      algos, BackgroundTaskPriority::kUserVisible, future1.GetCallback());
  ASSERT_TRUE(future1.Wait());
  ASSERT_TRUE(future1.Get().has_value());
  UnexportableSigningKeyId key_id = future1.Get().value();

  mojom::NewSigningKeyDataPtr collision_data = ToMojomKeyData(
      key_id,
      {
          .subject_public_key_info = {9, 9},
          .wrapped_key = {9, 9, 9},
          .algorithm =
              crypto::SignatureVerifier::SignatureAlgorithm::RSA_PKCS1_SHA256,
      });
  fake_service_.SetGenerateResponse(std::move(collision_data));

  base::test::TestFuture<ServiceErrorOr<UnexportableSigningKeyId>> future2;
  proxied_service_.GenerateSigningKeySlowlyAsync(
      algos, BackgroundTaskPriority::kUserVisible, future2.GetCallback());
  ASSERT_TRUE(future2.Wait());
  EXPECT_THAT(future2.Get(), ErrorIs(ServiceError::kKeyCollision));
}

TEST_F(UnexportableKeyServiceProxiedTest, FromWrappedSigningKeySuccess) {
  base::test::TestFuture<ServiceErrorOr<UnexportableSigningKeyId>> future;
  std::vector<uint8_t> wrapped_key = {0x11, 0x22, 0x33};

  proxied_service_.FromWrappedSigningKeySlowlyAsync(
      wrapped_key, BackgroundTaskPriority::kUserVisible, future.GetCallback());

  const ServiceErrorOr<UnexportableSigningKeyId>& result = future.Get();
  ASSERT_TRUE(result.has_value());
  UnexportableSigningKeyId key_id = result.value();

  EXPECT_THAT(proxied_service_.GetSubjectPublicKeyInfo(key_id),
              ValueIs(ElementsAreArray(kTestSubjectPublicKeyInfo)));
  EXPECT_THAT(proxied_service_.GetWrappedKey(key_id),
              ValueIs(ElementsAreArray(wrapped_key)));
  EXPECT_THAT(
      proxied_service_.GetAlgorithm(key_id),
      ValueIs(crypto::SignatureVerifier::SignatureAlgorithm::ECDSA_SHA256));
  EXPECT_THAT(proxied_service_.GetKeyTag(key_id), ValueIs(kTestKeyTag));
}

TEST_F(UnexportableKeyServiceProxiedTest, FromWrappedSigningKeyAlreadyCached) {
  base::test::TestFuture<ServiceErrorOr<UnexportableSigningKeyId>>
      generate_future;
  std::vector<crypto::SignatureVerifier::SignatureAlgorithm> algos = {
      crypto::SignatureVerifier::SignatureAlgorithm::RSA_PKCS1_SHA256};
  proxied_service_.GenerateSigningKeySlowlyAsync(
      algos, BackgroundTaskPriority::kUserVisible,
      generate_future.GetCallback());
  ASSERT_TRUE(generate_future.Get().has_value());
  UnexportableSigningKeyId key_id = generate_future.Get().value();

  ServiceErrorOr<std::vector<uint8_t>> original_spki =
      proxied_service_.GetSubjectPublicKeyInfo(key_id);
  ServiceErrorOr<std::vector<uint8_t>> original_wrapped =
      proxied_service_.GetWrappedKey(key_id);
  ServiceErrorOr<crypto::SignatureVerifier::SignatureAlgorithm> original_algo =
      proxied_service_.GetAlgorithm(key_id);
  ASSERT_TRUE(original_spki.has_value());
  ASSERT_TRUE(original_wrapped.has_value());
  ASSERT_TRUE(original_algo.has_value());

  mojom::NewSigningKeyDataPtr new_key_data = ToMojomKeyData(
      key_id,
      {
          .subject_public_key_info = {99, 99},
          .wrapped_key = {99},
          .algorithm =
              crypto::SignatureVerifier::SignatureAlgorithm::ECDSA_SHA256,
      });

  fake_service_.SetFromWrappedResponse(std::move(new_key_data));

  base::test::TestFuture<ServiceErrorOr<UnexportableSigningKeyId>>
      from_wrapped_future;
  std::vector<uint8_t> wrapped_key = {0xaa, 0xbb};
  proxied_service_.FromWrappedSigningKeySlowlyAsync(
      wrapped_key, BackgroundTaskPriority::kUserVisible,
      from_wrapped_future.GetCallback());

  EXPECT_THAT(from_wrapped_future.Get(), ValueIs(key_id));

  EXPECT_THAT(proxied_service_.GetSubjectPublicKeyInfo(key_id),
              ValueIs(original_spki.value()));
  EXPECT_THAT(proxied_service_.GetWrappedKey(key_id),
              ValueIs(original_wrapped.value()));
  EXPECT_THAT(proxied_service_.GetAlgorithm(key_id),
              ValueIs(original_algo.value()));
}

TEST_F(UnexportableKeyServiceProxiedTest, FromWrappedSigningKeyError) {
  fake_service_.SetFromWrappedResponse(
      base::unexpected(ServiceError::kKeyNotFound));

  base::test::TestFuture<ServiceErrorOr<UnexportableSigningKeyId>> future;
  std::vector<uint8_t> wrapped_key = {0x11, 0x22, 0x33};
  proxied_service_.FromWrappedSigningKeySlowlyAsync(
      wrapped_key, BackgroundTaskPriority::kUserVisible, future.GetCallback());

  EXPECT_THAT(future.Get(), ErrorIs(ServiceError::kKeyNotFound));
}

TEST_F(UnexportableKeyServiceProxiedTest, SignSuccess) {
  base::test::TestFuture<ServiceErrorOr<UnexportableSigningKeyId>>
      generate_future;
  std::vector<crypto::SignatureVerifier::SignatureAlgorithm> algos = {
      crypto::SignatureVerifier::SignatureAlgorithm::RSA_PKCS1_SHA256};
  proxied_service_.GenerateSigningKeySlowlyAsync(
      algos, BackgroundTaskPriority::kUserVisible,
      generate_future.GetCallback());
  ASSERT_TRUE(generate_future.Get().has_value());
  UnexportableSigningKeyId key_id = generate_future.Get().value();

  std::vector<uint8_t> expected_signature = {0xaa, 0xbb, 0xcc, 0xdd};
  fake_service_.SetSignResponse(expected_signature);

  base::test::TestFuture<ServiceErrorOr<std::vector<uint8_t>>> sign_future;
  std::vector<uint8_t> data_to_sign = {1, 2, 3, 4, 5, 6};
  proxied_service_.SignSlowlyAsync(key_id, data_to_sign,
                                   BackgroundTaskPriority::kUserVisible,
                                   sign_future.GetCallback());

  EXPECT_THAT(sign_future.Get(), ValueIs(expected_signature));
}

TEST_F(UnexportableKeyServiceProxiedTest, SignError) {
  base::test::TestFuture<ServiceErrorOr<UnexportableSigningKeyId>>
      generate_future;
  std::vector<crypto::SignatureVerifier::SignatureAlgorithm> algos = {
      crypto::SignatureVerifier::SignatureAlgorithm::RSA_PKCS1_SHA256};
  proxied_service_.GenerateSigningKeySlowlyAsync(
      algos, BackgroundTaskPriority::kUserVisible,
      generate_future.GetCallback());
  ASSERT_TRUE(generate_future.Get().has_value());
  UnexportableSigningKeyId key_id = generate_future.Get().value();

  fake_service_.SetSignResponse(
      base::unexpected(ServiceError::kVerifySignatureFailed));

  base::test::TestFuture<ServiceErrorOr<std::vector<uint8_t>>> sign_future;
  std::vector<uint8_t> data_to_sign = {1, 2, 3, 4, 5, 6};
  proxied_service_.SignSlowlyAsync(key_id, data_to_sign,
                                   BackgroundTaskPriority::kUserVisible,
                                   sign_future.GetCallback());

  EXPECT_THAT(sign_future.Get(), ErrorIs(ServiceError::kVerifySignatureFailed));
}

TEST_F(UnexportableKeyServiceProxiedTest, GettersKeyNotFound) {
  UnexportableSigningKeyId unknown_key_id;

  EXPECT_THAT(proxied_service_.GetSubjectPublicKeyInfo(unknown_key_id),
              ErrorIs(ServiceError::kKeyNotFound));
  EXPECT_THAT(proxied_service_.GetWrappedKey(unknown_key_id),
              ErrorIs(ServiceError::kKeyNotFound));
  EXPECT_THAT(proxied_service_.GetAlgorithm(unknown_key_id),
              ErrorIs(ServiceError::kKeyNotFound));
}

TEST_F(UnexportableKeyServiceProxiedTest, DeleteKeysSuccess) {
  UnexportableSigningKeyId key_id1 = GenerateSigningKeyOrDie();
  UnexportableSigningKeyId key_id2 = GenerateSigningKeyOrDie();
  ASSERT_TRUE(proxied_service_.GetSubjectPublicKeyInfo(key_id1).has_value());
  ASSERT_TRUE(proxied_service_.GetSubjectPublicKeyInfo(key_id2).has_value());

  fake_service_.SetDeleteKeysResponse(base::ok(2));

  base::test::TestFuture<ServiceErrorOr<size_t>> delete_keys_future;
  std::vector<UnexportableSigningKeyId> key_ids = {key_id1, key_id2};
  proxied_service_.DeleteKeysSlowlyAsync(key_ids,
                                         BackgroundTaskPriority::kUserVisible,
                                         delete_keys_future.GetCallback());

  EXPECT_THAT(delete_keys_future.Get(), ValueIs(2));
  EXPECT_THAT(proxied_service_.GetSubjectPublicKeyInfo(key_id1),
              ErrorIs(ServiceError::kKeyNotFound));
  EXPECT_THAT(proxied_service_.GetSubjectPublicKeyInfo(key_id2),
              ErrorIs(ServiceError::kKeyNotFound));
}

TEST_F(UnexportableKeyServiceProxiedTest, DeleteKeysNotFoundInCache) {
  UnexportableSigningKeyId unknown_key_id;

  base::test::TestFuture<ServiceErrorOr<size_t>> delete_keys_future;
  proxied_service_.DeleteKeysSlowlyAsync({unknown_key_id},
                                         BackgroundTaskPriority::kUserVisible,
                                         delete_keys_future.GetCallback());

  EXPECT_THAT(delete_keys_future.Get(), ErrorIs(ServiceError::kKeyNotFound));
}

TEST_F(UnexportableKeyServiceProxiedTest, DeleteKeysErrorFromService) {
  fake_service_.SetDeleteKeysResponse(
      base::unexpected(ServiceError::kCryptoApiFailed));

  const UnexportableSigningKeyId key_id = GenerateSigningKeyOrDie();

  base::test::TestFuture<ServiceErrorOr<size_t>> delete_keys_future;
  std::vector<UnexportableSigningKeyId> key_ids = {key_id};
  proxied_service_.DeleteKeysSlowlyAsync(key_ids,
                                         BackgroundTaskPriority::kUserVisible,
                                         delete_keys_future.GetCallback());

  EXPECT_THAT(delete_keys_future.Get(),
              ErrorIs(ServiceError::kCryptoApiFailed));
  EXPECT_FALSE(proxied_service_.GetSubjectPublicKeyInfo(key_id).has_value());
}

TEST_F(UnexportableKeyServiceProxiedTest, DeleteAllKeysSuccess) {
  UnexportableSigningKeyId key_id1 = GenerateSigningKeyOrDie();
  UnexportableSigningKeyId key_id2 = GenerateSigningKeyOrDie();
  ASSERT_TRUE(proxied_service_.GetSubjectPublicKeyInfo(key_id1).has_value());
  ASSERT_TRUE(proxied_service_.GetSubjectPublicKeyInfo(key_id2).has_value());

  fake_service_.SetDeleteAllKeysResponse(base::ok(2));

  base::test::TestFuture<ServiceErrorOr<size_t>> delete_all_future;
  proxied_service_.DeleteAllKeysSlowlyAsync(delete_all_future.GetCallback());

  EXPECT_THAT(delete_all_future.Get(), ValueIs(2));
  EXPECT_THAT(proxied_service_.GetSubjectPublicKeyInfo(key_id1),
              ErrorIs(ServiceError::kKeyNotFound));
  EXPECT_THAT(proxied_service_.GetSubjectPublicKeyInfo(key_id2),
              ErrorIs(ServiceError::kKeyNotFound));
}

TEST_F(UnexportableKeyServiceProxiedTest, DeleteAllKeysErrorFromService) {
  fake_service_.SetDeleteAllKeysResponse(
      base::unexpected(ServiceError::kCryptoApiFailed));

  const UnexportableSigningKeyId key_id = GenerateSigningKeyOrDie();

  base::test::TestFuture<ServiceErrorOr<size_t>> delete_all_future;
  proxied_service_.DeleteAllKeysSlowlyAsync(delete_all_future.GetCallback());

  EXPECT_THAT(delete_all_future.Get(), ErrorIs(ServiceError::kCryptoApiFailed));
  EXPECT_FALSE(proxied_service_.GetSubjectPublicKeyInfo(key_id).has_value());
}

TEST_F(UnexportableKeyServiceProxiedTest,
       GetAllKeysForGarbageCollectionSuccess) {
  std::vector<mojom::NewSigningKeyDataPtr> key_data_list;
  UnexportableSigningKeyId key_id1;
  UnexportableSigningKeyId key_id2;

  auto create_data = [](UnexportableSigningKeyId id) {
    return ToMojomKeyData(
        id, {
                .subject_public_key_info =
                    base::ToVector(kTestSubjectPublicKeyInfo),
                .wrapped_key = base::ToVector(kTestWrappedKey),
                .algorithm =
                    crypto::SignatureVerifier::SignatureAlgorithm::ECDSA_SHA256,
                .key_tag = std::string(kTestKeyTag),
            });
  };

  key_data_list.push_back(create_data(key_id1));
  key_data_list.push_back(create_data(key_id2));

  fake_service_.SetGetAllKeysForGarbageCollectionResponse(
      base::ok(std::move(key_data_list)));

  base::test::TestFuture<ServiceErrorOr<std::vector<UnexportableSigningKeyId>>>
      future;
  proxied_service_.GetAllKeysForGarbageCollectionSlowlyAsync(
      BackgroundTaskPriority::kUserVisible, future.GetCallback());

  ASSERT_OK_AND_ASSIGN(std::vector<UnexportableSigningKeyId> key_ids,
                       future.Get());
  EXPECT_THAT(key_ids, UnorderedElementsAre(key_id1, key_id2));

  // Verify cache population
  EXPECT_THAT(proxied_service_.GetSubjectPublicKeyInfo(key_id1),
              ValueIs(ElementsAreArray(kTestSubjectPublicKeyInfo)));
  EXPECT_THAT(proxied_service_.GetSubjectPublicKeyInfo(key_id2),
              ValueIs(ElementsAreArray(kTestSubjectPublicKeyInfo)));
}

TEST_F(UnexportableKeyServiceProxiedTest, GetAllKeysForGarbageCollectionEmpty) {
  fake_service_.SetGetAllKeysForGarbageCollectionResponse(
      base::ok(std::vector<mojom::NewSigningKeyDataPtr>()));

  base::test::TestFuture<ServiceErrorOr<std::vector<UnexportableSigningKeyId>>>
      future;
  proxied_service_.GetAllKeysForGarbageCollectionSlowlyAsync(
      BackgroundTaskPriority::kUserVisible, future.GetCallback());

  EXPECT_THAT(future.Get(), ValueIs(IsEmpty()));
}

TEST_F(UnexportableKeyServiceProxiedTest, GetAllKeysForGarbageCollectionError) {
  fake_service_.SetGetAllKeysForGarbageCollectionResponse(
      base::unexpected(ServiceError::kCryptoApiFailed));

  base::test::TestFuture<ServiceErrorOr<std::vector<UnexportableSigningKeyId>>>
      future;
  proxied_service_.GetAllKeysForGarbageCollectionSlowlyAsync(
      BackgroundTaskPriority::kUserVisible, future.GetCallback());

  EXPECT_THAT(future.Get(), ErrorIs(ServiceError::kCryptoApiFailed));
}

TEST_F(UnexportableKeyServiceProxiedTest, GenerateSigningKeyCancelled) {
  base::test::TestFuture<ServiceErrorOr<UnexportableSigningKeyId>> future;
  std::vector<crypto::SignatureVerifier::SignatureAlgorithm> algos = {
      crypto::SignatureVerifier::SignatureAlgorithm::RSA_PKCS1_SHA256};

  proxied_service_.GenerateSigningKeySlowlyAsync(
      algos, BackgroundTaskPriority::kUserVisible, future.GetCallback());

  receiver_.reset();
  EXPECT_THAT(future.Get(), ErrorIs(ServiceError::kOperationCancelled));
}

TEST_F(UnexportableKeyServiceProxiedTest, FromWrappedSigningKeyCancelled) {
  base::test::TestFuture<ServiceErrorOr<UnexportableSigningKeyId>> future;
  std::vector<uint8_t> wrapped_key = {0x11, 0x22, 0x33};

  proxied_service_.FromWrappedSigningKeySlowlyAsync(
      wrapped_key, BackgroundTaskPriority::kUserVisible, future.GetCallback());

  receiver_.reset();
  EXPECT_THAT(future.Get(), ErrorIs(ServiceError::kOperationCancelled));
}

TEST_F(UnexportableKeyServiceProxiedTest, DeleteKeysCancelled) {
  UnexportableSigningKeyId key_id = GenerateSigningKeyOrDie();
  base::test::TestFuture<ServiceErrorOr<size_t>> future;

  proxied_service_.DeleteKeysSlowlyAsync(
      {key_id}, BackgroundTaskPriority::kUserVisible, future.GetCallback());

  receiver_.reset();
  EXPECT_THAT(future.Get(), ErrorIs(ServiceError::kOperationCancelled));
}

TEST_F(UnexportableKeyServiceProxiedTest, DeleteAllKeysCancelled) {
  base::test::TestFuture<ServiceErrorOr<size_t>> future;

  proxied_service_.DeleteAllKeysSlowlyAsync(future.GetCallback());

  receiver_.reset();
  EXPECT_THAT(future.Get(), ErrorIs(ServiceError::kOperationCancelled));
}

TEST_F(UnexportableKeyServiceProxiedTest,
       GetAllKeysForGarbageCollectionCancelled) {
  base::test::TestFuture<ServiceErrorOr<std::vector<UnexportableSigningKeyId>>>
      future;

  proxied_service_.GetAllKeysForGarbageCollectionSlowlyAsync(
      BackgroundTaskPriority::kUserVisible, future.GetCallback());

  receiver_.reset();
  EXPECT_THAT(future.Get(), ErrorIs(ServiceError::kOperationCancelled));
}

TEST_F(UnexportableKeyServiceProxiedTest, SignCancelled) {
  UnexportableSigningKeyId key_id = GenerateSigningKeyOrDie();

  base::test::TestFuture<ServiceErrorOr<std::vector<uint8_t>>> future;
  std::vector<uint8_t> data_to_sign = {1, 2, 3};

  proxied_service_.SignSlowlyAsync(key_id, data_to_sign,
                                   BackgroundTaskPriority::kUserVisible,
                                   future.GetCallback());

  receiver_.reset();
  EXPECT_THAT(future.Get(), ErrorIs(ServiceError::kOperationCancelled));
}

TEST_F(UnexportableKeyServiceProxiedTest, GenerateAttestationKeySuccess) {
  base::test::TestFuture<ServiceErrorOr<UnexportableAttestationKeyId>> future;
  proxied_service_.GenerateAttestationKeySlowlyAsync(
      kTestAttestationAlgorithms, BackgroundTaskPriority::kUserVisible,
      future.GetCallback());

  ASSERT_OK_AND_ASSIGN(UnexportableAttestationKeyId key_id, future.Get());

  EXPECT_THAT(proxied_service_.GetSubjectPublicKeyInfo(key_id),
              ValueIs(ElementsAreArray(kTestSubjectPublicKeyInfo)));
  EXPECT_THAT(proxied_service_.GetWrappedKey(key_id),
              ValueIs(ElementsAreArray(kTestWrappedKey)));
  EXPECT_THAT(
      proxied_service_.GetAlgorithm(key_id),
      ValueIs(crypto::SignatureVerifier::SignatureAlgorithm::RSA_PKCS1_SHA256));
  EXPECT_THAT(proxied_service_.GetKeyTag(key_id), ValueIs(kTestKeyTag));
}

TEST_F(UnexportableKeyServiceProxiedTest, GenerateAttestationKeyError) {
  fake_service_.SetGenerateAttestationResponse(
      base::unexpected(ServiceError::kCryptoApiFailed));

  base::test::TestFuture<ServiceErrorOr<UnexportableAttestationKeyId>> future;
  proxied_service_.GenerateAttestationKeySlowlyAsync(
      kTestAttestationAlgorithms, BackgroundTaskPriority::kUserVisible,
      future.GetCallback());

  EXPECT_THAT(future.Get(), ErrorIs(ServiceError::kCryptoApiFailed));
}

TEST_F(UnexportableKeyServiceProxiedTest, GenerateAttestationKeyCancelled) {
  base::test::TestFuture<ServiceErrorOr<UnexportableAttestationKeyId>> future;
  proxied_service_.GenerateAttestationKeySlowlyAsync(
      kTestAttestationAlgorithms, BackgroundTaskPriority::kUserVisible,
      future.GetCallback());

  receiver_.reset();
  EXPECT_THAT(future.Get(), ErrorIs(ServiceError::kOperationCancelled));
}

TEST_F(UnexportableKeyServiceProxiedTest, FromWrappedAttestationKeySuccess) {
  base::test::TestFuture<ServiceErrorOr<UnexportableAttestationKeyId>> future;
  proxied_service_.FromWrappedAttestationKeySlowlyAsync(
      kTestWrappedAttestationKey, BackgroundTaskPriority::kUserVisible,
      future.GetCallback());

  ASSERT_OK_AND_ASSIGN(UnexportableAttestationKeyId key_id, future.Get());

  EXPECT_THAT(proxied_service_.GetSubjectPublicKeyInfo(key_id),
              ValueIs(ElementsAreArray(kTestSubjectPublicKeyInfo)));
  EXPECT_THAT(proxied_service_.GetWrappedKey(key_id),
              ValueIs(ElementsAreArray(kTestWrappedAttestationKey)));
  EXPECT_THAT(
      proxied_service_.GetAlgorithm(key_id),
      ValueIs(crypto::SignatureVerifier::SignatureAlgorithm::ECDSA_SHA256));
  EXPECT_THAT(proxied_service_.GetKeyTag(key_id), ValueIs(kTestKeyTag));
}

TEST_F(UnexportableKeyServiceProxiedTest, FromWrappedAttestationKeyError) {
  fake_service_.SetFromWrappedAttestationResponse(
      base::unexpected(ServiceError::kKeyNotFound));

  base::test::TestFuture<ServiceErrorOr<UnexportableAttestationKeyId>> future;
  proxied_service_.FromWrappedAttestationKeySlowlyAsync(
      kTestWrappedAttestationKey, BackgroundTaskPriority::kUserVisible,
      future.GetCallback());

  EXPECT_THAT(future.Get(), ErrorIs(ServiceError::kKeyNotFound));
}

TEST_F(UnexportableKeyServiceProxiedTest, FromWrappedAttestationKeyCancelled) {
  base::test::TestFuture<ServiceErrorOr<UnexportableAttestationKeyId>> future;
  proxied_service_.FromWrappedAttestationKeySlowlyAsync(
      kTestWrappedAttestationKey, BackgroundTaskPriority::kUserVisible,
      future.GetCallback());

  receiver_.reset();
  EXPECT_THAT(future.Get(), ErrorIs(ServiceError::kOperationCancelled));
}

TEST_F(UnexportableKeyServiceProxiedTest, CertifySuccess) {
  base::test::TestFuture<ServiceErrorOr<crypto::AttestationStatement>> future;
  proxied_service_.CertifySlowlyAsync(
      GenerateAttestationKeyOrDie(), GenerateSigningKeyOrDie(), kTestChallenge,
      BackgroundTaskPriority::kUserVisible, future.GetCallback());

  ASSERT_OK_AND_ASSIGN(const crypto::AttestationStatement& statement,
                       future.Get());
  const crypto::AttestationStatement& expected_statement =
      GetTestAttestationStatement();
  EXPECT_EQ(statement.format, expected_statement.format);
  EXPECT_EQ(statement.statement, expected_statement.statement);
  EXPECT_EQ(statement.signature, expected_statement.signature);
}

TEST_F(UnexportableKeyServiceProxiedTest, CertifyError) {
  fake_service_.SetCertifyResponse(
      base::unexpected(ServiceError::kCryptoApiFailed));

  base::test::TestFuture<ServiceErrorOr<crypto::AttestationStatement>> future;
  proxied_service_.CertifySlowlyAsync(
      GenerateAttestationKeyOrDie(), GenerateSigningKeyOrDie(), kTestChallenge,
      BackgroundTaskPriority::kUserVisible, future.GetCallback());

  EXPECT_THAT(future.Get(), ErrorIs(ServiceError::kCryptoApiFailed));
}

TEST_F(UnexportableKeyServiceProxiedTest, CertifyCancelled) {
  base::test::TestFuture<ServiceErrorOr<crypto::AttestationStatement>> future;
  proxied_service_.CertifySlowlyAsync(
      GenerateAttestationKeyOrDie(), GenerateSigningKeyOrDie(), kTestChallenge,
      BackgroundTaskPriority::kUserVisible, future.GetCallback());

  receiver_.reset();
  EXPECT_THAT(future.Get(), ErrorIs(ServiceError::kOperationCancelled));
}
}  // namespace
}  // namespace unexportable_keys
