// 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_proxy_impl.h"

#include <cstdint>
#include <optional>
#include <utility>

#include "base/check_deref.h"
#include "base/numerics/safe_conversions.h"
#include "base/time/time.h"
#include "base/types/expected.h"
#include "base/types/expected_macros.h"
#include "base/unguessable_token.h"
#include "build/build_config.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 "crypto/unexportable_key.h"
#include "mojo/public/cpp/bindings/enum_traits.h"
#include "mojo/public/cpp/bindings/pending_receiver.h"

namespace unexportable_keys {

namespace {

// For operations requiring stateful keys `ServiceError::kOperationNotSupported`
// might be returned on platforms where keys are stateless. Since this is not an
// actual error when retrieving the key, treat it simply as a missing key.
template <typename T>
ServiceErrorOr<std::optional<T>> AdaptOperationNotSupported(
    ServiceErrorOr<T> result) {
  using ServiceErrorOrOpt = ServiceErrorOr<std::optional<T>>;
  return ServiceErrorOrOpt(std::move(result)).or_else([](ServiceError error) {
    return error == ServiceError::kOperationNotSupported
               ? ServiceErrorOrOpt(std::nullopt)
               : base::unexpected(error);
  });
}

ServiceErrorOr<uint64_t> AdaptSizeType(ServiceErrorOr<size_t> result) {
  return result.transform(
      [](size_t r) { return base::strict_cast<uint64_t>(r); });
}

ServiceErrorOr<mojom::NewKeyMetadataPtr> PopulateKeyMetadata(
    unexportable_keys::UnexportableKeyService& unexportable_key_service,
    UnexportableSigningKeyId key_id) {
  auto metadata = mojom::NewKeyMetadata::New();
  ASSIGN_OR_RETURN(metadata->algorithm,
                   unexportable_key_service.GetAlgorithm(key_id));
  ASSIGN_OR_RETURN(metadata->wrapped_key,
                   unexportable_key_service.GetWrappedKey(key_id));
  ASSIGN_OR_RETURN(metadata->subject_public_key_info,
                   unexportable_key_service.GetSubjectPublicKeyInfo(key_id));
  ASSIGN_OR_RETURN(
      metadata->key_tag,
      AdaptOperationNotSupported(unexportable_key_service.GetKeyTag(key_id)));
  ASSIGN_OR_RETURN(metadata->creation_time,
                   AdaptOperationNotSupported(
                       unexportable_key_service.GetCreationTime(key_id)));
  return metadata;
}

ServiceErrorOr<mojom::NewSigningKeyDataPtr> PopulateNewSigningKeyData(
    unexportable_keys::UnexportableKeyService& unexportable_key_service,
    const ServiceErrorOr<UnexportableSigningKeyId> error_or_key_id) {
  auto data = mojom::NewSigningKeyData::New();
  ASSIGN_OR_RETURN(data->key_id, error_or_key_id);
  ASSIGN_OR_RETURN(data->metadata,
                   PopulateKeyMetadata(unexportable_key_service, data->key_id));
  return data;
}

ServiceErrorOr<mojom::NewAttestationKeyDataPtr> PopulateNewAttestationKeyData(
    unexportable_keys::UnexportableKeyService& unexportable_key_service,
    const ServiceErrorOr<UnexportableAttestationKeyId> error_or_key_id) {
  auto data = mojom::NewAttestationKeyData::New();
  ASSIGN_OR_RETURN(data->key_id, error_or_key_id);
  ASSIGN_OR_RETURN(data->metadata,
                   PopulateKeyMetadata(unexportable_key_service, data->key_id));
  return data;
}

ServiceErrorOr<std::vector<mojom::NewSigningKeyDataPtr>> PopulateAllNewKeyData(
    unexportable_keys::UnexportableKeyService& unexportable_key_service,
    ServiceErrorOr<std::vector<UnexportableSigningKeyId>> error_or_key_ids) {
  ASSIGN_OR_RETURN(std::vector<UnexportableSigningKeyId> key_ids,
                   error_or_key_ids);
  std::vector<mojom::NewSigningKeyDataPtr> new_key_data;
  new_key_data.reserve(key_ids.size());
  for (UnexportableSigningKeyId key_id : key_ids) {
    ASSIGN_OR_RETURN(
        new_key_data.emplace_back(),
        PopulateNewSigningKeyData(unexportable_key_service, key_id));
  }
  return new_key_data;
}
}  // namespace

UnexportableKeyServiceProxyImpl::UnexportableKeyServiceProxyImpl(
    unexportable_keys::UnexportableKeyService* unexportable_key_service,
    mojo::PendingReceiver<mojom::UnexportableKeyService> receiver)
    : receiver_(this, std::move(receiver)),
      unexportable_key_service_(CHECK_DEREF(unexportable_key_service)) {}

UnexportableKeyServiceProxyImpl::~UnexportableKeyServiceProxyImpl() = default;

void unexportable_keys::UnexportableKeyServiceProxyImpl::GenerateSigningKey(
    const std::vector<crypto::SignatureVerifier::SignatureAlgorithm>&
        acceptable_algorithms,
    BackgroundTaskPriority priority,
    GenerateSigningKeyCallback callback) {
  unexportable_key_service_->GenerateSigningKeySlowlyAsync(
      acceptable_algorithms, priority,
      base::BindOnce(&UnexportableKeyServiceProxyImpl::OnSigningKeyCreated,
                     weak_ptr_factory_.GetWeakPtr(), std::move(callback)));
}

void unexportable_keys::UnexportableKeyServiceProxyImpl::FromWrappedSigningKey(
    const std::vector<uint8_t>& wrapped_key,
    BackgroundTaskPriority priority,
    FromWrappedSigningKeyCallback callback) {
  if (wrapped_key.size() > kMaxWrappedKeySize) {
    receiver_.ReportBadMessage("wrapped key size too long");
    return;
  }
  unexportable_key_service_->FromWrappedSigningKeySlowlyAsync(
      wrapped_key, priority,
      base::BindOnce(&UnexportableKeyServiceProxyImpl::OnSigningKeyCreated,
                     weak_ptr_factory_.GetWeakPtr(), std::move(callback)));
}

void unexportable_keys::UnexportableKeyServiceProxyImpl::GenerateAttestationKey(
    const std::vector<crypto::SignatureVerifier::SignatureAlgorithm>&
        acceptable_algorithms,
    BackgroundTaskPriority priority,
    GenerateAttestationKeyCallback callback) {
  unexportable_key_service_->GenerateAttestationKeySlowlyAsync(
      acceptable_algorithms, priority,
      base::BindOnce(&UnexportableKeyServiceProxyImpl::OnAttestationKeyCreated,
                     weak_ptr_factory_.GetWeakPtr(), std::move(callback)));
}

void unexportable_keys::UnexportableKeyServiceProxyImpl::
    FromWrappedAttestationKey(const std::vector<uint8_t>& wrapped_key,
                              BackgroundTaskPriority priority,
                              FromWrappedAttestationKeyCallback callback) {
  if (wrapped_key.size() > kMaxWrappedKeySize) {
    receiver_.ReportBadMessage("wrapped key size too long");
    return;
  }
  unexportable_key_service_->FromWrappedAttestationKeySlowlyAsync(
      wrapped_key, priority,
      base::BindOnce(&UnexportableKeyServiceProxyImpl::OnAttestationKeyCreated,
                     weak_ptr_factory_.GetWeakPtr(), std::move(callback)));
}

void unexportable_keys::UnexportableKeyServiceProxyImpl::Sign(
    const UnexportableSigningKeyId& key_id,
    const std::vector<uint8_t>& data,
    BackgroundTaskPriority priority,
    SignCallback callback) {
  unexportable_key_service_->SignSlowlyAsync(key_id, data, priority,
                                             std::move(callback));
}

void unexportable_keys::UnexportableKeyServiceProxyImpl::Certify(
    const UnexportableAttestationKeyId& attestation_key_id,
    const UnexportableSigningKeyId& signing_key_id,
    const std::vector<uint8_t>& challenge,
    BackgroundTaskPriority priority,
    CertifyCallback callback) {
  unexportable_key_service_->CertifySlowlyAsync(attestation_key_id,
                                                signing_key_id, challenge,
                                                priority, std::move(callback));
}

void unexportable_keys::UnexportableKeyServiceProxyImpl::
    GetAllKeysForGarbageCollection(
        BackgroundTaskPriority priority,
        GetAllKeysForGarbageCollectionCallback callback) {
  unexportable_key_service_->GetAllKeysForGarbageCollectionSlowlyAsync(
      priority,
      base::BindOnce(
          &UnexportableKeyServiceProxyImpl::OnGetAllKeysForGarbageCollection,
          weak_ptr_factory_.GetWeakPtr(), std::move(callback)));
}

void unexportable_keys::UnexportableKeyServiceProxyImpl::DeleteKeys(
    const std::vector<UnexportableSigningKeyId>& key_ids,
    BackgroundTaskPriority priority,
    DeleteKeysCallback callback) {
  unexportable_key_service_->DeleteKeysSlowlyAsync(
      key_ids, priority,
      base::BindOnce(&AdaptSizeType).Then(std::move(callback)));
}

void unexportable_keys::UnexportableKeyServiceProxyImpl::DeleteAllKeys(
    DeleteAllKeysCallback callback) {
  unexportable_key_service_->DeleteAllKeysSlowlyAsync(
      base::BindOnce(&AdaptSizeType).Then(std::move(callback)));
}

void UnexportableKeyServiceProxyImpl::OnSigningKeyCreated(
    NewSigningKeyCallback callback,
    ServiceErrorOr<UnexportableSigningKeyId> error_or_key_id) {
  std::move(callback).Run(PopulateNewSigningKeyData(
      *unexportable_key_service_, std::move(error_or_key_id)));
}

void UnexportableKeyServiceProxyImpl::OnAttestationKeyCreated(
    NewAttestationKeyCallback callback,
    ServiceErrorOr<UnexportableAttestationKeyId> error_or_key_id) {
  std::move(callback).Run(PopulateNewAttestationKeyData(
      *unexportable_key_service_, std::move(error_or_key_id)));
}

void UnexportableKeyServiceProxyImpl::OnGetAllKeysForGarbageCollection(
    GetAllKeysForGarbageCollectionCallback callback,
    ServiceErrorOr<std::vector<UnexportableSigningKeyId>> error_or_key_ids) {
  std::move(callback).Run(PopulateAllNewKeyData(*unexportable_key_service_,
                                                std::move(error_or_key_ids)));
}

}  // namespace unexportable_keys
