// Copyright 2022 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/optimization_guide/core/delivery/prediction_model_store.h"

#include "base/command_line.h"
#include "base/files/file_enumerator.h"
#include "base/files/file_path.h"
#include "base/files/file_util.h"
#include "base/memory/ptr_util.h"
#include "base/metrics/histogram_functions.h"
#include "base/rand_util.h"
#include "base/strings/strcat.h"
#include "base/strings/string_number_conversions.h"
#include "base/task/thread_pool.h"
#include "base/trace_event/trace_event.h"
#include "base/uuid.h"
#include "components/optimization_guide/core/delivery/model_info.h"
#include "components/optimization_guide/core/delivery/model_store_metadata_entry.h"
#include "components/optimization_guide/core/delivery/model_util.h"
#include "components/optimization_guide/core/delivery/prediction_model_override.h"
#include "components/optimization_guide/core/optimization_guide_features.h"
#include "components/optimization_guide/core/optimization_guide_prefs.h"
#include "components/prefs/pref_service.h"

namespace optimization_guide {

const base::FilePath::CharType kOptimizationGuideModelStoreDirPrefix[] =
    FILE_PATH_LITERAL("optimization_guide_model_store");

namespace {

constexpr size_t kBytesPerMegabyte = 1024 * 1024;


// Parses the OptimizationTarget from the string.
proto::OptimizationTarget ParseOptimizationTargetFromString(
    const std::string& optimization_target_str) {
  int optimization_target;
  if (!base::StringToInt(optimization_target_str, &optimization_target)) {
    return proto::OPTIMIZATION_TARGET_UNKNOWN;
  }
  if (!proto::OptimizationTarget_IsValid(optimization_target)) {
    return proto::OPTIMIZATION_TARGET_UNKNOWN;
  }
  return static_cast<proto::OptimizationTarget>(optimization_target);
}

void RemoveInvalidModelDirs(const base::FilePath& base_store_dir,
                            std::set<base::FilePath> valid_model_dirs) {
  std::vector<base::FilePath> invalid_model_dirs;
  base::FileEnumerator enumerator(base_store_dir, /*recursive=*/false,
                                  base::FileEnumerator::DIRECTORIES);
  for (base::FilePath optimization_target_dir = enumerator.Next();
       !optimization_target_dir.empty();
       optimization_target_dir = enumerator.Next()) {
    proto::OptimizationTarget optimization_target =
        ParseOptimizationTargetFromString(
            optimization_target_dir.BaseName().AsUTF8Unsafe());
    if (optimization_target == proto::OPTIMIZATION_TARGET_UNKNOWN) {
      // Remove the unknown dirs within the model store dir. This can
      // potentially happen when the opt target is deprecated, and marked as
      // reserved.
      invalid_model_dirs.push_back(optimization_target_dir);
      RecordPredictionModelStoreModelRemovalVersionHistogram(
          proto::OPTIMIZATION_TARGET_UNKNOWN,
          PredictionModelStoreModelRemovalReason::kInconsistentModelDir);
      continue;
    }
    base::FileEnumerator model_cache_keys_enumerator(
        optimization_target_dir, false, base::FileEnumerator::DIRECTORIES);
    for (base::FilePath model_cache_key_dir =
             model_cache_keys_enumerator.Next();
         !model_cache_key_dir.empty();
         model_cache_key_dir = model_cache_keys_enumerator.Next()) {
      base::FileEnumerator models_enumerator(model_cache_key_dir,
                                             /*recursive=*/false,
                                             base::FileEnumerator::DIRECTORIES);
      for (base::FilePath model_dir = models_enumerator.Next();
           !model_dir.empty(); model_dir = models_enumerator.Next()) {
        DCHECK(model_dir.IsAbsolute());
        if (valid_model_dirs.find(ConvertToRelativePath(
                base_store_dir, model_dir)) == valid_model_dirs.end()) {
          invalid_model_dirs.push_back(model_dir);
          RecordPredictionModelStoreModelRemovalVersionHistogram(
              optimization_target,
              PredictionModelStoreModelRemovalReason::kInconsistentModelDir);
        }
      }
    }
  }
  // The invalid dirs can be removed immediately, since this is called at init.
  for (const auto& invalid_model_dir : invalid_model_dirs) {
    DCHECK(invalid_model_dir.IsAbsolute());
    base::DeletePathRecursively(invalid_model_dir);
  }
}

void RecordModelStorageMetrics(const base::FilePath& base_store_dir) {
  base::FileEnumerator enumerator(base_store_dir, false,
                                  base::FileEnumerator::DIRECTORIES);
  for (base::FilePath optimization_target_dir = enumerator.Next();
       !optimization_target_dir.empty();
       optimization_target_dir = enumerator.Next()) {
    proto::OptimizationTarget optimization_target =
        ParseOptimizationTargetFromString(
            optimization_target_dir.BaseName().AsUTF8Unsafe());
    if (optimization_target == proto::OPTIMIZATION_TARGET_UNKNOWN) {
      continue;
    }
    size_t total_models = 0;
    base::FileEnumerator models_enumerator(optimization_target_dir, false,
                                           base::FileEnumerator::DIRECTORIES);
    for (base::FilePath model_dir = models_enumerator.Next();
         !model_dir.empty(); model_dir = models_enumerator.Next()) {
      total_models++;
    }
    base::UmaHistogramCounts100(
        base::StrCat({"OptimizationGuide.PredictionModelStore.ModelCount.",
                      GetStringNameForOptimizationTarget(optimization_target)}),
        total_models);
    base::UmaHistogramMemoryMB(
        base::StrCat(
            {"OptimizationGuide.PredictionModelStore.TotalDirectorySize.",
             GetStringNameForOptimizationTarget(optimization_target)}),
        base::ComputeDirectorySize(optimization_target_dir) /
            kBytesPerMegabyte);
  }
}

}  // namespace

PredictionModelStore::PredictionModelStore(PrefService& local_state)
    : ledger_(local_state),
      background_task_runner_(base::ThreadPool::CreateSequencedTaskRunner(
          {base::MayBlock(), base::TaskPriority::BEST_EFFORT})) {}

PredictionModelStore::~PredictionModelStore() = default;

void PredictionModelStore::Initialize(const base::FilePath& base_store_dir) {
  DCHECK_CALLED_ON_VALID_SEQUENCE(sequence_checker_);
  DCHECK(!base_store_dir.empty());

  if (!background_task_runner_) {
    // In unit tests, to avoid leaking a task runner between test runs, the task
    // runner can be reset at the end of each test. In that case, we'll need to
    // recreate it here.
    background_task_runner_ = base::ThreadPool::CreateSequencedTaskRunner(
        {base::MayBlock(), base::TaskPriority::BEST_EFFORT});
  }

  // Should not be initialized already.
  DCHECK(base_store_dir_.empty());

  base_store_dir_ = base_store_dir;
  PurgeInactiveModels();

  // Clean up any model files that were slated for deletion in previous
  // sessions.
  CleanUpOldModelFiles();

  // crbug.com/404966596 - Removing invalid model dirs could race with unpacking
  // model overrides. For now, we just skip it if any model overrides were
  // specified.
  if (!base::CommandLine::ForCurrentProcess()->HasSwitch(
          kModelOverrideSwitch)) {
    background_task_runner_->PostTask(
        FROM_HERE, base::BindOnce(&RemoveInvalidModelDirs, base_store_dir_,
                                  ledger_.GetValidModelDirs()));
  }
  background_task_runner_->PostTask(
      FROM_HERE, base::BindOnce(&RecordModelStorageMetrics, base_store_dir_));
}

bool PredictionModelStore::HasModel(
    proto::OptimizationTarget optimization_target,
    const ClientCacheKey& model_cache_key) const {
  DCHECK_CALLED_ON_VALID_SEQUENCE(sequence_checker_);
  auto metadata =
      ledger_.GetEntryIfExists(optimization_target, model_cache_key);
  if (!metadata) {
    return false;
  }
  // Model dir should exist and be a relative path.
  return metadata->GetModelBaseDir() &&
         !metadata->GetModelBaseDir()->IsAbsolute();
}

bool PredictionModelStore::HasModelWithVersion(
    proto::OptimizationTarget optimization_target,
    const ClientCacheKey& model_cache_key,
    int64_t version) {
  DCHECK_CALLED_ON_VALID_SEQUENCE(sequence_checker_);
  auto metadata =
      ledger_.GetEntryIfExists(optimization_target, model_cache_key);
  if (!metadata) {
    return false;
  }
  if (!metadata->GetModelBaseDir() ||
      metadata->GetModelBaseDir()->IsAbsolute()) {
    // Model dir should exist and be a relative path.
    return false;
  }
  auto actual_version = metadata->GetVersion();
  if (!actual_version) {
    RemoveModel(optimization_target, model_cache_key,
                PredictionModelStoreModelRemovalReason::kModelVersionInvalid);
    return false;
  }
  return *actual_version == version;
}

void PredictionModelStore::LoadModel(
    proto::OptimizationTarget optimization_target,
    const ClientCacheKey& model_cache_key,
    scoped_refptr<base::SequencedTaskRunner> model_task_runner,
    PredictionModelLoadedCallback callback) {
  DCHECK_CALLED_ON_VALID_SEQUENCE(sequence_checker_);

  TRACE_EVENT("optimization_guide", "PredictionModelStore::LoadModel", "target",
              GetStringNameForOptimizationTarget(optimization_target));

  auto metadata =
      ledger_.GetEntryIfExists(optimization_target, model_cache_key);
  if (!metadata) {
    std::move(callback).Run(std::nullopt);
    return;
  }
  if (!metadata->GetKeepBeyondValidDuration() &&
      metadata->GetExpiryTime() <= base::Time::Now()) {
    RemoveModel(
        optimization_target, model_cache_key,
        PredictionModelStoreModelRemovalReason::kModelExpiredOnLoadModel);
    std::move(callback).Run(std::nullopt);
    return;
  }
  auto base_model_dir = metadata->GetModelBaseDir();
  if (!base_model_dir || base_model_dir->IsAbsolute()) {
    RemoveModel(optimization_target, model_cache_key,
                PredictionModelStoreModelRemovalReason::kInvalidModelDir);
    std::move(callback).Run(std::nullopt);
    return;
  }

  model_task_runner->PostTaskAndReplyWithResult(
      FROM_HERE,
      base::BindOnce(&LoadAndVerifyModelOffThread, optimization_target,
                     base_store_dir_.Append(*base_model_dir)),
      base::BindOnce(&PredictionModelStore::OnModelLoaded,
                     weak_ptr_factory_.GetWeakPtr(), optimization_target,
                     model_cache_key, std::move(callback)));
}

void PredictionModelStore::OnModelLoaded(
    proto::OptimizationTarget optimization_target,
    const ClientCacheKey& model_cache_key,
    PredictionModelLoadedCallback callback,
    std::optional<ModelInfo> model_info) {
  DCHECK_CALLED_ON_VALID_SEQUENCE(sequence_checker_);

  TRACE_EVENT("optimization_guide", "PredictionModelStore::OnModelLoaded",
              "target",
              GetStringNameForOptimizationTarget(optimization_target));

  if (!model_info) {
    RemoveModel(optimization_target, model_cache_key,
                PredictionModelStoreModelRemovalReason::kModelLoadFailed);
    std::move(callback).Run(std::nullopt);
    return;
  }
  std::move(callback).Run(std::move(model_info));
}

void PredictionModelStore::UpdateMetadataForExistingModel(
    proto::OptimizationTarget optimization_target,
    const ClientCacheKey& model_cache_key,
    const proto::ModelInfo& model_info) {
  DCHECK_CALLED_ON_VALID_SEQUENCE(sequence_checker_);
  DCHECK(model_info.has_version());
  DCHECK_EQ(optimization_target, model_info.optimization_target());

  if (!HasModel(optimization_target, model_cache_key)) {
    return;
  }

  ModelStoreMetadataEntryUpdater metadata =
      ledger_.UpdateEntry(optimization_target, model_cache_key);
  DCHECK(!metadata.entry().GetModelBaseDir()->IsAbsolute());
  metadata.SetVersion(model_info.version());
  if (model_info.has_valid_duration()) {
    metadata.SetExpiryTime(
        base::Time::Now() +
        base::Seconds(model_info.valid_duration().seconds()));
  }
  metadata.SetKeepBeyondValidDuration(model_info.keep_beyond_valid_duration());
}

void PredictionModelStore::UpdateModel(
    proto::OptimizationTarget optimization_target,
    const ClientCacheKey& model_cache_key,
    const proto::ModelInfo& model_info,
    const base::FilePath& base_model_dir,
    base::OnceClosure callback) {
  DCHECK_CALLED_ON_VALID_SEQUENCE(sequence_checker_);
  DCHECK(model_info.has_version());
  DCHECK_EQ(optimization_target, model_info.optimization_target());
  DCHECK(base_store_dir_.IsParent(base_model_dir));

  ModelStoreMetadataEntryUpdater metadata =
      ledger_.UpdateEntry(optimization_target, model_cache_key);
  metadata.SetVersion(model_info.version());
  metadata.SetExpiryTime(
      base::Time::Now() +
      (model_info.has_valid_duration()
           ? base::Seconds(model_info.valid_duration().seconds())
           : ModelStoreMetadataEntry::kDefaultStoredModelValidDuration));
  metadata.SetKeepBeyondValidDuration(model_info.keep_beyond_valid_duration());

  auto old_model_dir = metadata.entry().GetModelBaseDir();
  if (old_model_dir) {
    RecordPredictionModelStoreModelRemovalVersionHistogram(
        optimization_target,
        PredictionModelStoreModelRemovalReason::kNewModelUpdate);
    ScheduleModelDirRemoval(*old_model_dir);
  }
  metadata.SetModelBaseDir(
      ConvertToRelativePath(base_store_dir_, base_model_dir));

  background_task_runner_->PostTaskAndReplyWithResult(
      FROM_HERE,
      base::BindOnce(&CheckAllPathsExist,
                     GetModelFilePaths(model_info, base_model_dir)),
      base::BindOnce(&PredictionModelStore::OnModelUpdateVerified,
                     weak_ptr_factory_.GetWeakPtr(), optimization_target,
                     model_cache_key, std::move(callback)));
}

void PredictionModelStore::OnModelUpdateVerified(
    proto::OptimizationTarget optimization_target,
    const ClientCacheKey& model_cache_key,
    base::OnceClosure callback,
    bool model_paths_exist) {
  DCHECK_CALLED_ON_VALID_SEQUENCE(sequence_checker_);
  if (!model_paths_exist) {
    RemoveModel(optimization_target, model_cache_key,
                PredictionModelStoreModelRemovalReason::
                    kModelUpdateFilePathVerifyFailed);
  }
  std::move(callback).Run();
}

base::FilePath PredictionModelStore::GetBaseModelDirForModelCacheKey(
    proto::OptimizationTarget optimization_target,
    const ClientCacheKey& model_cache_key) {
  DCHECK_CALLED_ON_VALID_SEQUENCE(sequence_checker_);

  DCHECK(!base_store_dir_.empty());
  auto base_model_dir = base_store_dir_
                            .AppendASCII(base::NumberToString(
                                static_cast<int>(optimization_target)))
                            .AppendASCII(model_cache_key.hexhash);
  return base_model_dir.AppendASCII(
      base::HexEncode(base::RandBytesAsVector(8)));
}

void PredictionModelStore::UpdateModelCacheKeyMapping(
    proto::OptimizationTarget optimization_target,
    const ClientCacheKey& client_model_cache_key,
    const proto::ModelCacheKey& server_model_cache_key) {
  DCHECK_CALLED_ON_VALID_SEQUENCE(sequence_checker_);
  ledger_.UpdateModelCacheKeyMapping(
      optimization_target, client_model_cache_key, server_model_cache_key);
}

void PredictionModelStore::RemoveModel(
    proto::OptimizationTarget optimization_target,
    const ClientCacheKey& model_cache_key,
    PredictionModelStoreModelRemovalReason model_remove_reason) {
  DCHECK_CALLED_ON_VALID_SEQUENCE(sequence_checker_);

  RecordPredictionModelStoreModelRemovalVersionHistogram(optimization_target,
                                                         model_remove_reason);
  ModelStoreMetadataEntryUpdater metadata =
      ledger_.UpdateEntry(optimization_target, model_cache_key);
  auto base_model_dir = metadata.entry().GetModelBaseDir();
  if (base_model_dir) {
    ScheduleModelDirRemoval(*base_model_dir);
  }
  // Continue removing the metadata even if the model dirs does not exist.
  metadata.ClearMetadata();
}

void PredictionModelStore::ScheduleModelDirRemoval(
    const base::FilePath& base_model_dir) {
  // Backward compatibility: Model dirs were absolute in the earlier versions,
  // and it was only in experiment. The latest versions use relative paths.
  // Convert to absolute paths to save in the pref, since absolute dirs could
  // become non-existent if IOS Chrome upgrade changes the sandbox dirs.
  DCHECK_CALLED_ON_VALID_SEQUENCE(sequence_checker_);
  DCHECK(!base_model_dir.IsAbsolute() ||
         base_store_dir_.IsParent(base_model_dir));
  base::FilePath relative_model_dir =
      base_model_dir.IsAbsolute()
          ? ConvertToRelativePath(base_store_dir_, base_model_dir)
          : base_model_dir;
  ledger_.AddPathToDelete(relative_model_dir);
}

void PredictionModelStore::PurgeInactiveModels() {
  DCHECK_CALLED_ON_VALID_SEQUENCE(sequence_checker_);
  for (const auto& expired_model_dir : ledger_.PurgeAllInactiveMetadata()) {
    // Backward compatibility: Model dirs were absolute in the earlier versions,
    // and it was only in experiment. The latest versions use relative paths.
    DCHECK(!expired_model_dir.IsAbsolute() ||
           base_store_dir_.IsParent(expired_model_dir));
    base::FilePath absolute_model_dir =
        expired_model_dir.IsAbsolute()
            ? expired_model_dir
            : base_store_dir_.Append(expired_model_dir);
    // This is called at startup. So no need to schedule the deletion of the
    // model dirs, and instead can be deleted immediately.
    background_task_runner_->PostTask(
        FROM_HERE, base::GetDeletePathRecursivelyCallback(absolute_model_dir));
  }
}

void PredictionModelStore::CleanUpOldModelFiles() {
  DCHECK_CALLED_ON_VALID_SEQUENCE(sequence_checker_);
  for (const auto entry : ledger_.GetPathsToDelete()) {
    // Backward compatibility: Model dirs were absolute in the earlier versions.
    // The latest versions use relative paths.
    auto path_to_delete = StringToFilePath(entry.first);
    DCHECK(path_to_delete);
    DCHECK(!path_to_delete->IsAbsolute() ||
           base_store_dir_.IsParent(*path_to_delete));
    base::FilePath absolute_path_to_delete =
        path_to_delete->IsAbsolute() ? *path_to_delete
                                     : base_store_dir_.Append(*path_to_delete);
    background_task_runner_->PostTaskAndReplyWithResult(
        FROM_HERE,
        base::BindOnce(&base::DeletePathRecursively, absolute_path_to_delete),
        base::BindOnce(&PredictionModelStore::OnFilePathDeleted,
                       weak_ptr_factory_.GetWeakPtr(), *path_to_delete));
  }
}

void PredictionModelStore::OnFilePathDeleted(
    const base::FilePath& path_to_delete,
    bool success) {
  DCHECK_CALLED_ON_VALID_SEQUENCE(sequence_checker_);
  if (!success) {
    // Try to delete again later.
    return;
  }
  ledger_.RemovePathToDelete(path_to_delete);
}

base::FilePath PredictionModelStore::GetBaseStoreDirForTesting() const {
  DCHECK_CALLED_ON_VALID_SEQUENCE(sequence_checker_);
  return base_store_dir_;
}

}  // namespace optimization_guide
