// Copyright 2021 The Chromium Authors
// Use of this source code is governed by a BSD-style license that can be
// found in the LICENSE file.

#include "chrome/browser/ash/app_list/search/ranking/ftrl_ranker.h"

#include "chrome/browser/ash/app_list/search/chrome_search_result.h"
#include "chrome/browser/ash/app_list/search/util/ftrl_optimizer.h"

namespace app_list {

// FtrlRanker ------------------------------------------------------------

FtrlRanker::FtrlRanker(FtrlOptimizer::Params params, FtrlOptimizer::Proto proto)
    : ftrl_(std::make_unique<FtrlOptimizer>(std::move(proto), params)) {}

FtrlRanker::~FtrlRanker() = default;

void FtrlRanker::AddExpert(std::unique_ptr<Ranker> ranker) {
  rankers_.push_back(std::move(ranker));
}

void FtrlRanker::Start(const std::u16string& query,
                       const CategoriesList& categories) {
  for (auto& ranker : rankers_)
    ranker->Start(query, categories);

  ftrl_->Clear();
}

void FtrlRanker::Train(const LaunchData& launch) {
  for (auto& ranker : rankers_)
    ranker->Train(launch);

  ftrl_->Train(launch.id);
}

void FtrlRanker::UpdateResultRanks(ResultsMap& results, ProviderType provider) {
  const auto it = results.find(provider);
  DCHECK(it != results.end());
  const auto& new_results = it->second;

  // Create a vector of result ids.
  std::vector<std::string> ids;
  for (const auto& result : new_results)
    ids.push_back(result->id());

  // Create a vector of vectors for each expert's scores.
  std::vector<std::vector<double>> expert_scores;
  for (auto& ranker : rankers_)
    expert_scores.push_back(ranker->GetResultRanks(results, provider));

  // Get the final scores from the FTRL optimizer set them on the results.
  std::vector<double> result_scores =
      ftrl_->Score(std::move(ids), std::move(expert_scores));
  DCHECK_EQ(new_results.size(), result_scores.size());
  for (size_t i = 0; i < new_results.size(); ++i)
    new_results[i]->scoring().set_ftrl_result_score(result_scores[i]);
}


// ResultScoringShim ----------------------------------------------------------

ResultScoringShim::ResultScoringShim(ResultScoringShim::ScoringMember member)
    : member_(member) {}

std::vector<double> ResultScoringShim::GetResultRanks(const ResultsMap& results,
                                                      ProviderType provider) {
  const auto it = results.find(provider);
  if (it == results.end())
    return {};

  std::vector<double> scores;
  for (const auto& result : it->second) {
    double score;
    switch (member_) {
      case ResultScoringShim::ScoringMember::kNormalizedRelevance:
        score = result->scoring().normalized_relevance();
        break;
      case ResultScoringShim::ScoringMember::kMrfuResultScore:
        score = result->scoring().mrfu_result_score();
        break;
    }
    scores.push_back(score);
  }
  return scores;
}
}  // namespace app_list
