// 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/passage_embeddings/core/passage_embeddings_service_controller.h"

#include <algorithm>
#include <utility>
#include <vector>

#include "base/functional/bind.h"
#include "base/functional/callback_helpers.h"
#include "base/logging.h"
#include "base/metrics/histogram_functions.h"
#include "base/notreached.h"
#include "base/task/thread_pool.h"
#include "components/optimization_guide/core/optimization_guide_util.h"
#include "components/passage_embeddings/core/internal/scheduling_embedder.h"
#include "components/passage_embeddings/core/passage_embeddings_features.h"
#include "components/passage_embeddings/core/passage_embeddings_types.h"
#include "mojo/public/cpp/bindings/callback_helpers.h"
#include "services/passage_embeddings/public/mojom/passage_embeddings.mojom.h"

namespace passage_embeddings {

namespace {

mojom::PassageEmbeddingsLoadModelsParamsPtr MakeModelParams(
    const base::FilePath& embeddings_path,
    const base::FilePath& sp_path,
    uint32_t input_window_size) {
  auto params = mojom::PassageEmbeddingsLoadModelsParams::New();
  params->embeddings_model = base::File(
      embeddings_path, base::File::FLAG_OPEN | base::File::FLAG_READ);
  params->sp_model =
      base::File(sp_path, base::File::FLAG_OPEN | base::File::FLAG_READ);
  params->input_window_size = input_window_size;
  return params;
}

// Makes the parameters used to run the passage embedder.
mojom::PassageEmbedderParamsPtr MakeEmbedderParams(bool execute_for_gemma) {
  auto params = mojom::PassageEmbedderParams::New();
  params->execute_for_gemma = execute_for_gemma;
  params->user_initiated_priority_num_threads =
      kUserInitiatedPriorityNumThreads.Get();
  params->urgent_priority_num_threads = kUrgentPriorityNumThreads.Get();
  params->passive_priority_num_threads = kPassivePriorityNumThreads.Get();
  params->embedder_cache_size = kEmbedderCacheSize.Get();
  params->allow_gpu_execution = kAllowGpuExecution.Get();
  return params;
}

mojom::PassagePriority PassagePriorityToMojom(PassagePriority priority) {
  switch (priority) {
    case kUserInitiated:
      return mojom::PassagePriority::kUserInitiated;
    case kUrgent:
      return mojom::PassagePriority::kUrgent;
    case kPassive:
    case kLatent:
      return mojom::PassagePriority::kPassive;
  }
}

class ScopedEmbeddingsModelInfoStatusLogger {
 public:
  ScopedEmbeddingsModelInfoStatusLogger() = default;
  ~ScopedEmbeddingsModelInfoStatusLogger() {
    CHECK_NE(EmbeddingsModelInfoStatus::kUnknown, status_);
    base::UmaHistogramEnumeration(kModelInfoMetricName, status_);
  }

  void set_status(EmbeddingsModelInfoStatus status) { status_ = status; }

 private:
  EmbeddingsModelInfoStatus status_ = EmbeddingsModelInfoStatus::kUnknown;
};

}  // namespace

PassageEmbeddingsServiceController::PassageEmbeddingsServiceController(
    PassageEmbeddingsServiceLauncher& launcher,
    bool execute_for_gemma)
    : launcher_(launcher),
      embedder_(std::make_unique<SchedulingEmbedder>(
          /*embedder_metadata_provider=*/this,
          /*get_embeddings_callback=*/
          base::BindRepeating(
              &PassageEmbeddingsServiceController::GetEmbeddings,
              base::Unretained(this)),
          kSchedulerMaxJobs.Get(),
          kSchedulerMaxBatchSize.Get(),
          kUsePerformanceScenario.Get(),
          execute_for_gemma)),
      execute_for_gemma_(execute_for_gemma) {}

PassageEmbeddingsServiceController::~PassageEmbeddingsServiceController() =
    default;

bool PassageEmbeddingsServiceController::MaybeUpdateModelInfo(
    base::optional_ref<const optimization_guide::ModelInfo> model_info) {
  // Got the same version again. Do not run through rest of logic.
  if (model_info && model_version_ == model_info->version) {
    return true;
  }

  // Reset everything, so if the model info is invalid, the service controller
  // would stop accepting requests.
  embeddings_model_path_.clear();
  sp_model_path_.clear();
  model_metadata_ = std::nullopt;
  model_version_ = 0;
  ResetEmbedderRemote();

  ScopedEmbeddingsModelInfoStatusLogger logger;
  if (!model_info.has_value()) {
    logger.set_status(EmbeddingsModelInfoStatus::kEmpty);
    return false;
  }

  // The only additional file should be the sentencepiece model.
  const std::vector<base::FilePath>& additional_files =
      model_info->additional_files;
  if (additional_files.size() != 1u) {
    logger.set_status(EmbeddingsModelInfoStatus::kInvalidAdditionalFiles);
    return false;
  }

  // Check validity of model metadata.
  const std::optional<optimization_guide::proto::Any>& metadata =
      model_info->model_metadata;
  if (!metadata) {
    logger.set_status(EmbeddingsModelInfoStatus::kNoMetadata);
    return false;
  }
  std::optional<optimization_guide::proto::PassageEmbeddingsModelMetadata>
      embeddings_metadata = optimization_guide::ParsedAnyMetadata<
          optimization_guide::proto::PassageEmbeddingsModelMetadata>(*metadata);
  if (!embeddings_metadata) {
    logger.set_status(EmbeddingsModelInfoStatus::kInvalidMetadata);
    return false;
  }

  model_version_ = model_info->version;
  model_metadata_ = embeddings_metadata;
  embeddings_model_path_ = model_info->model_file_path;
  sp_model_path_ = additional_files[0];

  CHECK(IsModelAvailable());
  logger.set_status(EmbeddingsModelInfoStatus::kValid);
  observer_list_.Notify(&EmbedderMetadataObserver::EmbedderMetadataUpdated,
                        GetEmbedderMetadata());
  return true;
}

void PassageEmbeddingsServiceController::LoadModelsToService(
    base::WeakPtr<PassageEmbeddingsServiceController> embedder_remote_weak_ptr,
    mojo::PendingReceiver<mojom::PassageEmbedder> receiver,
    base::ElapsedTimer service_launch_timer,
    mojom::PassageEmbeddingsLoadModelsParamsPtr params) {
  if (!embedder_remote_weak_ptr || !service_remote_) {
    // Close the model files in a background thread.
    base::ThreadPool::PostTaskAndReply(
        FROM_HERE, {base::MayBlock()},
        base::DoNothingWithBoundArgs(std::move(params)),
        base::BindOnce(&PassageEmbeddingsServiceController::OnLoadModelsResult,
                       embedder_remote_weak_ptr_factory_.GetWeakPtr(),
                       std::move(service_launch_timer), /*success=*/false));
    return;
  }

  service_remote_->LoadModels(
      std::move(params), MakeEmbedderParams(execute_for_gemma_),
      std::move(receiver),
      base::BindOnce(&PassageEmbeddingsServiceController::OnLoadModelsResult,
                     embedder_remote_weak_ptr_factory_.GetWeakPtr(),
                     std::move(service_launch_timer)));
}

void PassageEmbeddingsServiceController::OnLoadModelsResult(
    base::ElapsedTimer service_launch_timer,
    bool success) {
  if (!success) {
    return;
  }

  if (!execute_for_gemma_) {
    base::UmaHistogramTimes("History.Embeddings.Embedder.LaunchDuration",
                            service_launch_timer.Elapsed());
  } else {
    base::UmaHistogramTimes("AI.SemanticEmbedder.LaunchDuration",
                            service_launch_timer.Elapsed());
  }
}

bool PassageEmbeddingsServiceController::IsModelAvailable() {
  return !sp_model_path_.empty() && !embeddings_model_path_.empty();
}

bool PassageEmbeddingsServiceController::EmbedderRunning() {
  return !pending_requests_.empty();
}

Embedder* PassageEmbeddingsServiceController::GetEmbedder() {
  return embedder_.get();
}

void PassageEmbeddingsServiceController::AddObserver(
    EmbedderMetadataObserver* observer) {
  if (IsModelAvailable()) {
    observer->EmbedderMetadataUpdated(GetEmbedderMetadata());
  }
  observer_list_.AddObserver(observer);
}

void PassageEmbeddingsServiceController::RemoveObserver(
    EmbedderMetadataObserver* observer) {
  observer_list_.RemoveObserver(observer);
}

void PassageEmbeddingsServiceController::GetEmbeddings(
    std::vector<std::string> passages,
    PassagePriority priority,
    GetEmbeddingsResultCallback callback) {
  if (passages.empty()) {
    std::move(callback).Run({}, ComputeEmbeddingsStatus::kSuccess);
    return;
  }

  if (!IsModelAvailable()) {
    VLOG(1) << "Missing model path: embeddings='" << embeddings_model_path_
            << "'; sp='" << sp_model_path_ << "'";
    std::move(callback).Run({}, ComputeEmbeddingsStatus::kModelUnavailable);
    return;
  }

  if (!embedder_remote_) {
    base::ElapsedTimer service_launch_timer;
    MaybeLaunchService();

    mojo::PendingReceiver<mojom::PassageEmbedder> receiver =
        embedder_remote_.BindNewPipeAndPassReceiver();
    // Unretained is safe because `this` owns `embedder_remote_`, which
    // synchronously calls the disconnect and idle handlers.
    embedder_remote_.set_disconnect_handler(
        base::BindOnce(&PassageEmbeddingsServiceController::ResetEmbedderRemote,
                       base::Unretained(this)));
    embedder_remote_.set_idle_handler(
        kEmbedderTimeout.Get(),
        base::BindRepeating(
            &PassageEmbeddingsServiceController::ResetEmbedderRemote,
            base::Unretained(this)));
    base::ThreadPool::PostTaskAndReplyWithResult(
        FROM_HERE, {base::MayBlock()},
        base::BindOnce(&MakeModelParams, embeddings_model_path_, sp_model_path_,
                       model_metadata_->input_window_size()),
        base::BindOnce(&PassageEmbeddingsServiceController::LoadModelsToService,
                       weak_ptr_factory_.GetWeakPtr(),
                       embedder_remote_weak_ptr_factory_.GetWeakPtr(),
                       std::move(receiver), std::move(service_launch_timer)));
  }

  pending_requests_.push_back(next_request_id_);
  base::ElapsedTimer generate_embeddings_timer;
  std::pair<GetEmbeddingsResultCallback, GetEmbeddingsResultCallback>
      callbacks = base::SplitOnceCallback(std::move(callback));
  embedder_remote_->GenerateEmbeddings(
      std::move(passages), PassagePriorityToMojom(priority),
      mojo::WrapCallbackWithDropHandler(
          base::BindOnce(&PassageEmbeddingsServiceController::OnGotEmbeddings,
                         weak_ptr_factory_.GetWeakPtr(), next_request_id_,
                         std::move(callbacks.first),
                         std::move(generate_embeddings_timer), priority),
          base::BindOnce(&PassageEmbeddingsServiceController::OnDisconnected,
                         weak_ptr_factory_.GetWeakPtr(), next_request_id_,
                         std::move(callbacks.second))));
  next_request_id_++;
}

EmbedderMetadata PassageEmbeddingsServiceController::GetEmbedderMetadata() {
  if (model_metadata_->score_threshold() > 0.0) {
    return EmbedderMetadata(model_version_, model_metadata_->output_size(),
                            model_metadata_->score_threshold());
  }

  return EmbedderMetadata(model_version_, model_metadata_->output_size());
}

void PassageEmbeddingsServiceController::ResetEmbedderRemote() {
  embedder_remote_.reset();
  embedder_remote_weak_ptr_factory_.InvalidateWeakPtrs();
}

void PassageEmbeddingsServiceController::OnGotEmbeddings(
    RequestId request_id,
    GetEmbeddingsResultCallback callback,
    base::ElapsedTimer generate_embeddings_timer,
    PassagePriority priority,
    std::vector<mojom::PassageEmbeddingsResultPtr> results) {
  // Mojo invokes the callbacks in the order in which `GenerateEmbeddings()` was
  // called.
  CHECK(!pending_requests_.empty());
  CHECK_EQ(pending_requests_.front(), request_id);
  pending_requests_.pop_front();

  ComputeEmbeddingsStatus status =
      results.empty() ? ComputeEmbeddingsStatus::kExecutionFailure
                      : ComputeEmbeddingsStatus::kSuccess;

  if (status == ComputeEmbeddingsStatus::kSuccess) {
    const base::TimeDelta duration = generate_embeddings_timer.Elapsed();
    if (execute_for_gemma_) {
      base::UmaHistogramTimes("AI.SemanticEmbedder.TaskDuration", duration);
    } else {
      base::UmaHistogramTimes("History.Embeddings.TaskDuration", duration);
      const char* priority_histogram = nullptr;
      switch (priority) {
        case kUserInitiated:
          priority_histogram = "History.Embeddings.TaskDuration.UserInitiated";
          break;

        case kUrgent:
          priority_histogram = "History.Embeddings.TaskDuration.Urgent";
          break;

        case kPassive:
          priority_histogram = "History.Embeddings.TaskDuration.Passive";
          break;

        default:
          priority_histogram = "History.Embeddings.TaskDuration.Other";
      }
      base::UmaHistogramTimes(priority_histogram, duration);
    }
  }

  // Run the callback last to prevent UAF if the callback destroys `this`.
  std::move(callback).Run(std::move(results), status);
}

void PassageEmbeddingsServiceController::OnDisconnected(
    RequestId request_id,
    GetEmbeddingsResultCallback callback) {
  // On disconnect, drop handlers are invoked in an undefined order, so we must
  // be able to remove arbitrary request IDs.
  auto it = std::ranges::find(pending_requests_, request_id);
  CHECK(it != pending_requests_.end());
  pending_requests_.erase(it);

  std::move(callback).Run(std::vector<mojom::PassageEmbeddingsResultPtr>(),
                          ComputeEmbeddingsStatus::kExecutionFailure);
}

void PassageEmbeddingsServiceController::MaybeLaunchService() {
  if (service_remote_.is_bound() || !launcher_->AllowedToLaunch()) {
    return;
  }
  auto receiver = service_remote_.BindNewPipeAndPassReceiver();
  service_remote_.set_disconnect_handler(
      base::BindOnce(&PassageEmbeddingsServiceController::ResetServiceRemote,
                     base::Unretained(this), /*is_idle=*/false));
  service_remote_.set_idle_handler(
      kEmbeddingsServiceTimeout.Get(),
      base::BindRepeating(
          &PassageEmbeddingsServiceController::ResetServiceRemote,
          base::Unretained(this), /*is_idle=*/true));
  launcher_->LaunchService(std::move(receiver));
}

void PassageEmbeddingsServiceController::ResetServiceRemote(bool is_idle) {
  service_remote_.reset();
  launcher_->OnServiceDisconnected(is_idle);
}
}  // namespace passage_embeddings
