// Copyright 2019 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/login/saml/public_saml_url_fetcher.h"

#include <memory>
#include <optional>
#include <string>
#include <utility>
#include <vector>

#include "base/check.h"
#include "base/check_deref.h"
#include "base/functional/bind.h"
#include "base/logging.h"
#include "chrome/browser/ash/policy/core/browser_policy_connector_ash.h"
#include "chrome/browser/ash/policy/core/device_local_account.h"
#include "chrome/browser/ash/settings/device_settings_service.h"
#include "chromeos/ash/components/install_attributes/install_attributes.h"
#include "chromeos/ash/components/settings/cros_settings.h"
#include "components/account_id/account_id.h"
#include "components/policy/core/common/cloud/cloud_policy_constants.h"
#include "components/policy/core/common/cloud/device_management_service.h"
#include "components/policy/core/common/cloud/dm_auth.h"
#include "components/policy/core/common/cloud/dmserver_job_configurations.h"
#include "device_management_backend.pb.h"
#include "services/network/public/cpp/shared_url_loader_factory.h"

namespace ash {
namespace {

namespace em = ::enterprise_management;

std::string GetAccountId(std::string user_id) {
  std::vector<policy::DeviceLocalAccount> device_local_accounts =
      policy::GetDeviceLocalAccounts(CrosSettings::Get());
  for (auto account : device_local_accounts) {
    if (account.user_id == user_id) {
      return account.account_id;
    }
  }
  return std::string();
}

}  // namespace

PublicSamlUrlFetcher::PublicSamlUrlFetcher(
    policy::BrowserPolicyConnectorAsh* browser_policy_connector_ash,
    scoped_refptr<network::SharedURLLoaderFactory> shared_url_loader_factory,
    AccountId account_id)
    : browser_policy_connector_ash_(CHECK_DEREF(browser_policy_connector_ash)),
      shared_url_loader_factory_(std::move(shared_url_loader_factory)),
      account_id_(GetAccountId(account_id.GetUserEmail())) {
  CHECK(shared_url_loader_factory_);
}

PublicSamlUrlFetcher::~PublicSamlUrlFetcher() = default;

std::string PublicSamlUrlFetcher::GetRedirectUrl() {
  return redirect_url_;
}

bool PublicSamlUrlFetcher::FetchSucceeded() {
  return fetch_succeeded_;
}

void PublicSamlUrlFetcher::Fetch(base::OnceClosure callback) {
  DCHECK(!callback_);
  callback_ = std::move(callback);
  policy::DeviceManagementService* service =
      browser_policy_connector_ash_->device_management_service();
  std::unique_ptr<policy::DMServerJobConfiguration> config = std::make_unique<
      policy::DMServerJobConfiguration>(
      service,
      policy::DeviceManagementService::JobConfiguration::TYPE_REQUEST_SAML_URL,
      browser_policy_connector_ash_->GetInstallAttributes()->GetDeviceId(),
      /*critical=*/false,
      policy::DMAuth::FromDMToken(
          DeviceSettingsService::Get()->policy_data()->request_token()),
      /*oauth_token=*/std::nullopt, shared_url_loader_factory_,
      base::BindOnce(&PublicSamlUrlFetcher::OnPublicSamlUrlReceived,
                     weak_ptr_factory_.GetWeakPtr()));

  em::PublicSamlUserRequest* saml_url_request =
      config->request()->mutable_public_saml_user_request();
  saml_url_request->set_account_id(account_id_);
  fetch_request_job_ = service->CreateJob(std::move(config));
}

void PublicSamlUrlFetcher::OnPublicSamlUrlReceived(
    policy::DMServerJobResult result) {
  VLOG(1) << "Public SAML url response received. DM Status: "
          << result.dm_status;
  fetch_request_job_.reset();
  std::string user_id;

  switch (result.dm_status) {
    case policy::DM_STATUS_SUCCESS: {
      if (!result.response.has_public_saml_user_response()) {
        LOG(WARNING) << "Invalid public SAML url response.";
        break;
      }
      const em::PublicSamlUserResponse& saml_url_response =
          result.response.public_saml_user_response();

      if (!saml_url_response.has_saml_parameters()) {
        LOG(WARNING) << "Invalid public SAML url response.";
        break;
      }
      // Fetch has succeeded.
      fetch_succeeded_ = true;
      const em::SamlParametersProto& saml_params =
          saml_url_response.saml_parameters();
      redirect_url_ = saml_params.auth_redirect_url();
      break;
    }
    default: {  // All other error cases
      LOG(ERROR) << "Fetching public SAML url failed. DM Status: "
                 << result.dm_status;
      break;
    }
  }
  std::move(callback_).Run();
}

}  // namespace ash
