// Copyright 2026 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/live_caption/translation_dispatcher_on_device.h"

#include "base/metrics/histogram_functions.h"
#include "base/strings/string_util.h"
#include "base/types/expected.h"
#include "components/on_device_translation/public/mojom/translator.mojom.h"
#include "components/on_device_translation/service_controller.h"
#include "components/on_device_translation/service_controller_manager.h"
#include "url/gurl.h"

namespace captions {

namespace {

constexpr char kOnDeviceTranslateErrorReasonHistogram[] =
    "Accessibility.LiveTranslate.OnDeviceTranslation.ErrorReason";

// Records the error reason metric to UMA `count` times. When translator
// creation fails, multiple pending callbacks queued during creation fail
// simultaneously, so `count` accounts for each failed translation request.
void RecordOnDeviceTranslationErrorReason(
    OnDeviceTranslationErrorReason error_reason,
    size_t count = 1) {
  constexpr auto kExclusiveMax =
      static_cast<int>(OnDeviceTranslationErrorReason::kMaxValue) + 1;

  base::HistogramBase* histogram = base::LinearHistogram::FactoryGet(
      kOnDeviceTranslateErrorReasonHistogram, 1, kExclusiveMax,
      static_cast<size_t>(kExclusiveMax + 1),
      base::HistogramBase::kUmaTargetedHistogramFlag);

  histogram->AddCount(static_cast<int>(error_reason), count);
}

// Maps a service controller CreateTranslatorError enum value to the
// corresponding UMA metric OnDeviceTranslationErrorReason enum value.
OnDeviceTranslationErrorReason ToOnDeviceTranslationErrorReason(
    on_device_translation::OnDeviceTranslationController::CreateTranslatorError
        error) {
  using Error = on_device_translation::OnDeviceTranslationController::
      CreateTranslatorError;
  switch (error) {
    case Error::kInvalidBinary:
      return OnDeviceTranslationErrorReason::kCreateTranslatorInvalidBinary;
    case Error::kInvalidFunctionPointer:
      return OnDeviceTranslationErrorReason::
          kCreateTranslatorInvalidFunctionPointer;
    case Error::kFailedToInitialize:
      return OnDeviceTranslationErrorReason::
          kCreateTranslatorFailedToInitialize;
    case Error::kFailedToCreateTranslator:
      return OnDeviceTranslationErrorReason::
          kCreateTranslatorFailedToCreateTranslator;
    case Error::kInvalidVersion:
      return OnDeviceTranslationErrorReason::kCreateTranslatorInvalidVersion;
    case Error::kServiceCrashed:
      return OnDeviceTranslationErrorReason::kCreateTranslatorServiceCrashed;
    case Error::kNotSupportedLanguage:
      return OnDeviceTranslationErrorReason::
          kCreateTranslatorNotSupportedLanguage;
    case Error::kExceedsServiceCountLimitation:
      return OnDeviceTranslationErrorReason::
          kCreateTranslatorExceedsServiceCountLimitation;
    case Error::kExceedsPendingTaskCountLimitation:
      return OnDeviceTranslationErrorReason::
          kCreateTranslatorExceedsPendingTaskCountLimitation;
  }
  return OnDeviceTranslationErrorReason::kCreateTranslatorUnknownError;
}

}  // namespace

using ::on_device_translation::OnDeviceTranslationController;

TranslationDispatcherOnDevice::TranslationDispatcherOnDevice() = default;

TranslationDispatcherOnDevice::TranslationDispatcherOnDevice(
    std::unique_ptr<OnDeviceTranslationController> translation_controller)
    : translation_controller_(std::move(translation_controller)) {}

TranslationDispatcherOnDevice::~TranslationDispatcherOnDevice() = default;

void TranslationDispatcherOnDevice::GetTranslation(
    std::string_view result,
    std::string_view source_language,
    std::string_view target_language,
    TranslateEventCallback callback) {
  if (translator_.is_bound() && source_language_ == source_language &&
      target_language_ == target_language) {
    translator_->Translate(
        std::string(result),
        base::BindOnce(&TranslationDispatcherOnDevice::OnTranslated,
                       weak_factory_.GetWeakPtr(), std::move(callback)));
    return;
  }

  if (creation_in_progress_) {
    pending_callbacks_.emplace_back(std::string(result), std::move(callback));
    return;
  }

  translator_.reset();
  creation_in_progress_ = true;

  translation_controller_->CanTranslate(
      std::string(source_language), std::string(target_language),
      base::BindOnce(&TranslationDispatcherOnDevice::OnCanTranslate,
                     weak_factory_.GetWeakPtr(), std::string(source_language),
                     std::string(target_language), std::string(result),
                     std::move(callback)));
}

void TranslationDispatcherOnDevice::OnCanTranslate(
    const std::string& source_language,
    const std::string& target_language,
    const std::string& result,
    TranslateEventCallback callback,
    OnDeviceTranslationController::CanTranslateResult can_translate_result) {
  switch (can_translate_result) {
    case OnDeviceTranslationController::CanTranslateResult::kReadily:
    case OnDeviceTranslationController::CanTranslateResult::
        kAfterDownloadLibraryNotReady:
    case OnDeviceTranslationController::CanTranslateResult::
        kAfterDownloadLanguagePackNotReady:
    case OnDeviceTranslationController::CanTranslateResult::
        kAfterDownloadLibraryAndLanguagePackNotReady:
      translation_controller_->CreateTranslator(
          source_language, target_language,
          base::BindOnce(&TranslationDispatcherOnDevice::OnTranslationCreated,
                         weak_factory_.GetWeakPtr(), source_language,
                         target_language, result, std::move(callback)));
      return;
    case OnDeviceTranslationController::CanTranslateResult::
        kNoNotSupportedLanguage:
      ResetCreationAndNotifyFailure(
          std::move(callback),
          OnDeviceTranslationErrorReason::kCanTranslateLanguageNotSupported);
      return;
    case OnDeviceTranslationController::CanTranslateResult::kNoServiceCrashed:
      ResetCreationAndNotifyFailure(
          std::move(callback),
          OnDeviceTranslationErrorReason::kCanTranslateServiceCrashed);
      return;
    case OnDeviceTranslationController::CanTranslateResult::
        kNoExceedsServiceCountLimitation:
      ResetCreationAndNotifyFailure(
          std::move(callback), OnDeviceTranslationErrorReason::
                                   kCanTranslateExceedsServiceCountLimitation);
      return;
  }
}

void TranslationDispatcherOnDevice::ResetCreationAndNotifyFailure(
    TranslateEventCallback callback,
    OnDeviceTranslationErrorReason error_reason) {
  RecordOnDeviceTranslationErrorReason(error_reason,
                                       1 + pending_callbacks_.size());
  creation_in_progress_ = false;
  std::move(callback).Run(base::unexpected("Failed to create translator"));
  for (auto& pending_callback : pending_callbacks_) {
    std::move(pending_callback.second)
        .Run(base::unexpected("Failed to create translator"));
  }
  pending_callbacks_.clear();
}

void TranslationDispatcherOnDevice::OnTranslationCreated(
    const std::string& source_language,
    const std::string& target_language,
    const std::string& result,
    TranslateEventCallback callback,
    base::expected<
        mojo::PendingRemote<on_device_translation::mojom::OnDeviceTranslator>,
        OnDeviceTranslationController::CreateTranslatorError> translator) {
  if (!translator.has_value()) {
    ResetCreationAndNotifyFailure(
        std::move(callback),
        ToOnDeviceTranslationErrorReason(translator.error()));
    return;
  }
  source_language_ = source_language;
  target_language_ = target_language;
  translator_.Bind(std::move(translator.value()));
  translator_->Translate(
      result, base::BindOnce(&TranslationDispatcherOnDevice::OnTranslated,
                             weak_factory_.GetWeakPtr(), std::move(callback)));

  for (auto& pending_callback : pending_callbacks_) {
    translator_->Translate(
        pending_callback.first,
        base::BindOnce(&TranslationDispatcherOnDevice::OnTranslated,
                       weak_factory_.GetWeakPtr(),
                       std::move(pending_callback.second)));
  }
  pending_callbacks_.clear();
}

void TranslationDispatcherOnDevice::OnTranslated(
    TranslateEventCallback callback,
    const std::optional<std::string>& translation) {
  if (!translation) {
    RecordOnDeviceTranslationErrorReason(
        OnDeviceTranslationErrorReason::kTranslationExecutionFailed);
    std::move(callback).Run(base::unexpected("Failed to get translation"));
    return;
  }
  std::move(callback).Run(base::ok(translation.value()));
}

}  // namespace captions
