// Copyright 2024 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/on_device_translation/service/on_device_translation_service.h"

#include <memory>

#include "base/functional/bind.h"
#include "base/types/pass_key.h"
#include "components/on_device_translation/metrics.h"
#include "components/on_device_translation/public/mojom/on_device_translation_service.mojom.h"
#include "components/on_device_translation/public/mojom/translator.mojom.h"
#include "components/on_device_translation/service/translate_kit_client.h"
#include "mojo/public/cpp/bindings/self_owned_receiver.h"

namespace on_device_translation {
namespace {
// TranslateKitTranslator provides translation functionalities based on the
// Translator from TranslateKitClient.
class TranslateKitTranslator : public mojom::OnDeviceTranslator {
 public:
  explicit TranslateKitTranslator(std::string source_lang,
                                  std::string target_lang,
                                  TranslateKitClient::Translator* translator)
      : source_lang_(std::move(source_lang)),
        target_lang_(std::move(target_lang)),
        translator_(translator) {
    CHECK(translator_);
  }
  ~TranslateKitTranslator() override = default;
  // Not copyable.
  TranslateKitTranslator(const TranslateKitTranslator&) = delete;
  TranslateKitTranslator& operator=(const TranslateKitTranslator&) = delete;

  // `mojom::OnDeviceTranslator` overrides:
  void Translate(const std::string& input,
                 TranslateCallback translate_callback) override {
    RecordOnDeviceTranslationLength(source_lang_, target_lang_, input.size());
    CHECK(translator_);
    std::move(translate_callback).Run(translator_->Translate(input));
  }

  void SplitSentences(const std::string& input,
                      SplitSentencesCallback callback) override {
    CHECK(translator_);
    std::move(callback).Run(translator_->SplitSentences(input));
  }

 private:
  std::string source_lang_;
  std::string target_lang_;
  // Owned by the TranslateKitClient managed by OnDeviceTranslationService.
  // The lifetime must be longer than this instance.
  raw_ptr<TranslateKitClient::Translator> translator_;
};
}  // namespace

// static
std::unique_ptr<OnDeviceTranslationService>
OnDeviceTranslationService::CreateForTesting(
    mojo::PendingReceiver<mojom::OnDeviceTranslationService> receiver,
    std::unique_ptr<TranslateKitClient> client) {
  return std::make_unique<OnDeviceTranslationService>(
      std::move(receiver), std::move(client),
      base::PassKey<OnDeviceTranslationService>());
}

void OnDeviceTranslationService::OnDisconnect() {
  if (translators_.empty()) {
    receiver_.reset();
  }
}

OnDeviceTranslationService::OnDeviceTranslationService(
    mojo::PendingReceiver<mojom::OnDeviceTranslationService> receiver)
    : receiver_(this, std::move(receiver)), client_(TranslateKitClient::Get()) {
  translators_.set_disconnect_handler(base::BindRepeating(
      &OnDeviceTranslationService::OnDisconnect, base::Unretained(this)));
}

OnDeviceTranslationService::OnDeviceTranslationService(
    mojo::PendingReceiver<mojom::OnDeviceTranslationService> receiver,
    std::unique_ptr<TranslateKitClient> client,
    base::PassKey<OnDeviceTranslationService>)
    : receiver_(this, std::move(receiver)),
      owning_client_for_testing_(std::move(client)),
      client_(owning_client_for_testing_.get()) {
  translators_.set_disconnect_handler(base::BindRepeating(
      &OnDeviceTranslationService::OnDisconnect, base::Unretained(this)));
}

OnDeviceTranslationService::~OnDeviceTranslationService() = default;

void OnDeviceTranslationService::SetServiceConfig(
    mojom::OnDeviceTranslationServiceConfigPtr config) {
  client_->SetConfig(std::move(config));
}

void OnDeviceTranslationService::CreateTranslator(
    const std::string& source_lang,
    const std::string& target_lang,
    mojo::PendingReceiver<on_device_translation::mojom::OnDeviceTranslator>
        receiver,
    CreateTranslatorCallback create_translator_callback) {
  auto maybe_translator = client_->GetTranslator(source_lang, target_lang);
  if (!maybe_translator.has_value()) {
    std::move(create_translator_callback).Run(maybe_translator.error());
    return;
  }
  translators_.Add(std::make_unique<TranslateKitTranslator>(
                       source_lang, target_lang, maybe_translator.value()),
                   std::move(receiver));
  std::move(create_translator_callback)
      .Run(mojom::CreateTranslatorResult::kSuccess);
}

void OnDeviceTranslationService::CanTranslate(
    const std::string& source_lang,
    const std::string& target_lang,
    CanTranslateCallback can_translate_callback) {
  std::move(can_translate_callback)
      .Run(client_->CanTranslate(source_lang, target_lang));
}

}  // namespace on_device_translation
